1use std::sync::atomic::{AtomicBool, Ordering};
18use std::sync::{Arc, Mutex};
19use std::time::Duration;
20
21use axum::extract::{Path, State};
22use axum::http::header;
23use axum::http::StatusCode;
24use axum::response::{IntoResponse, Response};
25use axum::routing::{get, post};
26use axum::Json;
27use mongreldb_core::schema::{Schema, TypeId};
28use mongreldb_core::{CancellationReason, Database, Value};
29use mongreldb_query::{
30 CancelOutcome, CompactFinishedQuery, ExternalTableModule, ManagedQueryBatches, MongrelSession,
31 QueryId, RegisteredQueryGuard, RegisteredSqlQuery, SqlQueryOptions, SqlQueryPhase,
32 SqlQueryRegistry, SqlStreamCompletion,
33};
34use serde::{Deserialize, Serialize};
35use serde_json::json;
36use sha2::Digest;
37use zeroize::Zeroizing;
38
39mod admission;
40mod audit;
41pub mod cluster_admin;
42pub mod cluster_runtime;
43pub mod fragment_rpc;
44mod kit;
45mod metrics;
46pub mod native;
47pub mod oidc;
48mod pre_cancel;
49mod prepared;
50mod procedure;
51pub mod remote_embedding;
52mod sessions;
53mod sql_idempotency;
54mod sql_pages;
55pub mod vault_kms;
56
57#[doc(hidden)]
59pub fn fuzz_validate_sql_cursor(value: &str, owner: &str, key: &[u8; 32]) {
60 sql_pages::fuzz_validate_cursor(value, owner, key);
61}
62mod trigger;
63
64pub use sessions::{spawn_session_reaper, SessionStore};
65
66fn client_closed_request_status() -> StatusCode {
67 StatusCode::from_u16(499).unwrap_or(StatusCode::BAD_REQUEST)
68}
69
70fn cancellation_checkpoint_error(query: &RegisteredSqlQuery) -> mongreldb_query::MongrelQueryError {
71 query.checkpoint().err().unwrap_or_else(|| {
72 mongreldb_query::MongrelQueryError::InvalidQueryState(
73 "cancellation notification observed without a terminal checkpoint".into(),
74 )
75 })
76}
77
78fn status_for_error(e: &mongreldb_core::MongrelError) -> StatusCode {
86 use mongreldb_core::MongrelError;
87 match e {
88 MongrelError::AuthRequired | MongrelError::InvalidCredentials { .. } => {
89 StatusCode::UNAUTHORIZED
90 }
91 MongrelError::AuthNotRequired => StatusCode::BAD_REQUEST,
92 MongrelError::PermissionDenied { .. } => StatusCode::FORBIDDEN,
93 MongrelError::InvalidArgument(_) => StatusCode::CONFLICT,
94 MongrelError::Conflict(_) => StatusCode::CONFLICT,
95 MongrelError::Deadlock { .. } | MongrelError::SerializationFailure { .. } => {
98 StatusCode::CONFLICT
99 }
100 MongrelError::ReadOnlyReplica => StatusCode::CONFLICT,
101 MongrelError::NotFound(_) => StatusCode::NOT_FOUND,
102 MongrelError::DeadlineExceeded => StatusCode::GATEWAY_TIMEOUT,
103 MongrelError::WorkBudgetExceeded => StatusCode::TOO_MANY_REQUESTS,
104 MongrelError::ResourceLimitExceeded { .. } | MongrelError::Full(_) => {
107 StatusCode::SERVICE_UNAVAILABLE
108 }
109 MongrelError::Cancelled => client_closed_request_status(),
110 MongrelError::CursorStale(_) => StatusCode::CONFLICT,
111 MongrelError::CursorExpired => StatusCode::GONE,
112 _ => StatusCode::INTERNAL_SERVER_ERROR,
113 }
114}
115
116fn status_for_query_error(e: &mongreldb_query::MongrelQueryError) -> StatusCode {
119 use mongreldb_query::MongrelQueryError;
120 match e {
121 MongrelQueryError::Core(core) => status_for_error(core),
122 MongrelQueryError::DeadlineExceeded { .. } => StatusCode::GATEWAY_TIMEOUT,
123 MongrelQueryError::QueryCancelled { .. } => client_closed_request_status(),
124 MongrelQueryError::QueryIdConflict { .. } => StatusCode::CONFLICT,
125 MongrelQueryError::QueryRegistryFull => StatusCode::SERVICE_UNAVAILABLE,
126 MongrelQueryError::ResultLimitExceeded { .. } => StatusCode::PAYLOAD_TOO_LARGE,
127 MongrelQueryError::TransactionAborted => StatusCode::CONFLICT,
128 MongrelQueryError::NoSqlTransaction | MongrelQueryError::SavepointNotFound { .. } => {
129 StatusCode::CONFLICT
130 }
131 MongrelQueryError::CommitOutcome { .. } => StatusCode::CONFLICT,
132 MongrelQueryError::OutcomeUnknown { .. } => StatusCode::CONFLICT,
133 _ => StatusCode::INTERNAL_SERVER_ERROR,
134 }
135}
136
137#[cfg(test)]
138mod error_status_tests {
139 use super::*;
140
141 #[tokio::test]
146 async fn deadlock_and_serialization_failure_are_409_with_precise_taxonomy() {
147 use mongreldb_core::MongrelError;
148 use mongreldb_query::MongrelQueryError;
149
150 let cases = [
151 (
152 MongrelError::Deadlock {
153 victim: 7,
154 cycle: "7 → 3 → 7".into(),
155 },
156 "deadlock",
157 9,
158 ),
159 (
160 MongrelError::SerializationFailure {
161 message: "ssi certification failed".into(),
162 },
163 "serialization failure",
164 8,
165 ),
166 ];
167 for (core, category, category_code) in cases {
168 assert_eq!(status_for_error(&core), StatusCode::CONFLICT, "{core}");
169 let error = MongrelQueryError::Core(core);
170 assert_eq!(
171 status_for_query_error(&error),
172 StatusCode::CONFLICT,
173 "{error}"
174 );
175 let response = query_error_response(&error, None);
176 assert_eq!(response.status(), StatusCode::CONFLICT, "{error}");
177 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
178 .await
179 .unwrap();
180 let body: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
181 assert_eq!(body.pointer("/error/category").unwrap(), category, "{body}");
182 assert_eq!(
183 body.pointer("/error/category_code").unwrap(),
184 category_code,
185 "{body}"
186 );
187 }
188 }
189}
190
191struct OptionalPrincipal(Option<mongreldb_core::Principal>);
195
196impl<S> axum::extract::FromRequestParts<S> for OptionalPrincipal
197where
198 S: Send + Sync,
199{
200 type Rejection = std::convert::Infallible;
201
202 async fn from_request_parts(
203 parts: &mut axum::http::request::Parts,
204 _state: &S,
205 ) -> Result<Self, Self::Rejection> {
206 Ok(OptionalPrincipal(
207 parts.extensions.get::<mongreldb_core::Principal>().cloned(),
208 ))
209 }
210}
211
212struct AppState {
213 db: Arc<Database>,
214 idem: kit::IdempotencyStore,
215 external_modules: Vec<Arc<dyn ExternalTableModule>>,
216 auth_token: Option<String>,
217 user_auth: bool,
219 metrics: Arc<metrics::Metrics>,
221 audit: Arc<audit::AuditLog>,
223 sessions: Arc<sessions::SessionStore>,
226 ai_semaphore: Arc<tokio::sync::Semaphore>,
228 query_registry: Arc<SqlQueryRegistry>,
230 query_lifecycle: Mutex<()>,
232 pre_cancellations: pre_cancel::PreCancelStore,
234 sql_idempotency: Arc<sql_idempotency::SqlIdempotencyStore>,
236 sql_pages: sql_pages::SqlPageStore,
238 sql_semaphore: Arc<tokio::sync::Semaphore>,
241 sql_page_semaphore: Arc<tokio::sync::Semaphore>,
243 sql_page_default_timeout: std::time::Duration,
244 sql_page_max_timeout: std::time::Duration,
245 max_request_bytes: usize,
249 accepting_sql: Arc<AtomicBool>,
250 cursor_mac_key: CursorMacKey,
252 reloadable: Arc<ReloadableConfig>,
254 drain: Arc<DrainControl>,
256 scheduler: admission::SchedulerAdmission,
259 node_governor: std::sync::Mutex<mongreldb_core::NodeMemoryGovernor>,
261 ai_generations: std::sync::Mutex<mongreldb_core::AiIndexGenerationRegistry>,
263 multi_region: std::sync::Mutex<mongreldb_cluster::multi_region::MultiRegionPolicy>,
265 ops_jobs: std::sync::Mutex<mongreldb_core::OpsJobStore>,
267 resource_groups: mongreldb_core::ResourceGroupRegistry,
270 embedding_providers: mongreldb_core::EmbeddingProviderRegistry,
274 cluster_runtime: Option<cluster_runtime::ClusterRuntimeHandle>,
279}
280
281struct ReloadableDuration(std::sync::atomic::AtomicU64);
285
286impl ReloadableDuration {
287 fn new(value: std::time::Duration) -> Self {
288 Self(std::sync::atomic::AtomicU64::new(
289 value.as_millis().min(u128::from(u64::MAX)) as u64,
290 ))
291 }
292
293 fn get(&self) -> std::time::Duration {
294 std::time::Duration::from_millis(self.0.load(Ordering::Relaxed))
295 }
296
297 fn set(&self, value: std::time::Duration) {
298 self.0.store(
299 value.as_millis().min(u128::from(u64::MAX)) as u64,
300 Ordering::Relaxed,
301 );
302 }
303
304 fn as_millis(&self) -> u64 {
305 self.0.load(Ordering::Relaxed)
306 }
307}
308
309struct ReloadableUsize(std::sync::atomic::AtomicU64);
311
312impl ReloadableUsize {
313 fn new(value: usize) -> Self {
314 Self(std::sync::atomic::AtomicU64::new(
315 u64::try_from(value).unwrap_or(u64::MAX),
316 ))
317 }
318
319 fn get(&self) -> usize {
320 usize::try_from(self.0.load(Ordering::Relaxed)).unwrap_or(usize::MAX)
321 }
322
323 fn set(&self, value: usize) {
324 self.0
325 .store(u64::try_from(value).unwrap_or(u64::MAX), Ordering::Relaxed);
326 }
327}
328
329struct ReloadableConfig {
340 slow_query_threshold: ReloadableDuration,
342 sql_default_timeout: ReloadableDuration,
345 sql_max_timeout: ReloadableDuration,
347 sql_cancel_grace: ReloadableDuration,
350 sql_max_output_rows: ReloadableUsize,
352 sql_max_output_bytes: ReloadableUsize,
354}
355
356#[derive(Clone, Debug)]
359struct MutableConfigValues {
360 slow_query_threshold: std::time::Duration,
361 sql_default_timeout: std::time::Duration,
362 sql_max_timeout: std::time::Duration,
363 sql_cancel_grace: std::time::Duration,
364 sql_max_output_rows: usize,
365 sql_max_output_bytes: usize,
366 history_retention_epochs: u64,
367}
368
369#[derive(Clone, Debug, Default, Deserialize)]
373struct MutableConfigOverrides {
374 #[serde(default)]
375 slow_query_ms: Option<u64>,
376 #[serde(default)]
377 sql_default_timeout_ms: Option<u64>,
378 #[serde(default)]
379 sql_max_timeout_ms: Option<u64>,
380 #[serde(default)]
381 sql_cancel_grace_ms: Option<u64>,
382 #[serde(default)]
383 sql_max_output_rows: Option<u64>,
384 #[serde(default)]
385 sql_max_output_bytes: Option<u64>,
386 #[serde(default)]
387 history_retention_epochs: Option<u64>,
388}
389
390#[derive(Clone, Debug, Serialize)]
394pub struct ReloadReport {
395 pub slow_query_ms: u64,
396 pub sql_default_timeout_ms: u64,
397 pub sql_max_timeout_ms: u64,
398 pub sql_cancel_grace_ms: u64,
399 pub sql_max_output_rows: u64,
400 pub sql_max_output_bytes: u64,
401 pub history_retention_epochs: u64,
402}
403
404fn mutable_config_from_env() -> MutableConfigValues {
407 let sql_max_timeout = default_sql_max_timeout();
408 MutableConfigValues {
409 slow_query_threshold: metrics::slow_query_threshold(),
410 sql_default_timeout: default_sql_default_timeout().min(sql_max_timeout),
411 sql_max_timeout,
412 sql_cancel_grace: default_sql_cancel_grace(),
413 sql_max_output_rows: default_sql_max_output_rows(),
414 sql_max_output_bytes: default_sql_max_output_bytes(),
415 history_retention_epochs: default_history_retention_epochs(),
416 }
417}
418
419impl MutableConfigValues {
420 fn apply_overrides(&mut self, overrides: &MutableConfigOverrides) -> Result<(), &'static str> {
424 fn positive(value: Option<u64>) -> Result<Option<u64>, &'static str> {
425 match value {
426 Some(0) => Err("reload override values must be positive"),
427 other => Ok(other),
428 }
429 }
430 if let Some(value) = positive(overrides.slow_query_ms)? {
431 self.slow_query_threshold = std::time::Duration::from_millis(value);
432 }
433 if let Some(value) = positive(overrides.sql_max_timeout_ms)? {
434 self.sql_max_timeout = std::time::Duration::from_millis(value);
435 }
436 if let Some(value) = positive(overrides.sql_default_timeout_ms)? {
437 self.sql_default_timeout = std::time::Duration::from_millis(value);
438 }
439 if let Some(value) = positive(overrides.sql_cancel_grace_ms)? {
440 self.sql_cancel_grace = std::time::Duration::from_millis(value);
441 }
442 if let Some(value) = positive(overrides.sql_max_output_rows)? {
443 self.sql_max_output_rows = usize::try_from(value).unwrap_or(usize::MAX);
444 }
445 if let Some(value) = positive(overrides.sql_max_output_bytes)? {
446 self.sql_max_output_bytes = usize::try_from(value).unwrap_or(usize::MAX);
447 }
448 if let Some(value) = positive(overrides.history_retention_epochs)? {
449 self.history_retention_epochs = value;
450 }
451 Ok(())
452 }
453}
454
455fn apply_mutable_config(
459 reloadable: &ReloadableConfig,
460 db: &mongreldb_core::Database,
461 values: &MutableConfigValues,
462) -> mongreldb_core::Result<ReloadReport> {
463 reloadable.sql_max_timeout.set(values.sql_max_timeout);
464 reloadable
465 .sql_default_timeout
466 .set(values.sql_default_timeout.min(values.sql_max_timeout));
467 reloadable.sql_cancel_grace.set(values.sql_cancel_grace);
468 reloadable
469 .sql_max_output_rows
470 .set(values.sql_max_output_rows);
471 reloadable
472 .sql_max_output_bytes
473 .set(values.sql_max_output_bytes);
474 reloadable
475 .slow_query_threshold
476 .set(values.slow_query_threshold);
477 db.set_history_retention_epochs(values.history_retention_epochs)?;
478 Ok(ReloadReport {
479 slow_query_ms: reloadable.slow_query_threshold.as_millis(),
480 sql_default_timeout_ms: reloadable.sql_default_timeout.as_millis(),
481 sql_max_timeout_ms: reloadable.sql_max_timeout.as_millis(),
482 sql_cancel_grace_ms: reloadable.sql_cancel_grace.as_millis(),
483 sql_max_output_rows: reloadable.sql_max_output_rows.get() as u64,
484 sql_max_output_bytes: reloadable.sql_max_output_bytes.get() as u64,
485 history_retention_epochs: values.history_retention_epochs,
486 })
487}
488
489#[derive(Default)]
492struct DrainControl {
493 lock: tokio::sync::Mutex<()>,
495 record: std::sync::Mutex<DrainRecord>,
496}
497
498#[derive(Clone, Default, Serialize)]
499struct DrainRecord {
500 initiated: bool,
502 completed: bool,
504 deadline_ms: u64,
506 lifecycle: String,
508 detail: String,
510}
511
512#[derive(Default)]
513struct CursorMacKey {
514 key: Mutex<Option<[u8; 32]>>,
515}
516
517impl CursorMacKey {
518 fn get(&self) -> mongreldb_core::Result<[u8; 32]> {
519 let mut key = match self.key.lock() {
520 Ok(key) => key,
521 Err(poisoned) => poisoned.into_inner(),
522 };
523 if let Some(key) = *key {
524 return Ok(key);
525 }
526 let mut generated = [0u8; 32];
527 mongreldb_core::encryption::fill_random(&mut generated)?;
528 *key = Some(generated);
529 Ok(generated)
530 }
531}
532
533#[derive(Clone)]
536pub struct ServerControl {
537 query_registry: Arc<SqlQueryRegistry>,
538 sessions: Arc<sessions::SessionStore>,
539 accepting_sql: Arc<AtomicBool>,
540 cancel_grace: std::time::Duration,
541 metrics: Arc<metrics::Metrics>,
542 reloadable: Arc<ReloadableConfig>,
543 db: Arc<Database>,
544 audit: Arc<audit::AuditLog>,
545 sql_idempotency: Arc<sql_idempotency::SqlIdempotencyStore>,
546 sql_semaphore: Arc<tokio::sync::Semaphore>,
547 scheduler: admission::SchedulerAdmission,
548 sql_priority: u8,
549 cluster_runtime: Option<cluster_runtime::ClusterRuntimeHandle>,
551}
552
553impl ServerControl {
554 pub fn query_registry(&self) -> Arc<SqlQueryRegistry> {
556 Arc::clone(&self.query_registry)
557 }
558
559 pub fn native_runtime(
561 &self,
562 db: Arc<Database>,
563 sessions: Arc<SessionStore>,
564 ) -> native::NativeRuntime {
565 native::NativeRuntime::new(db, sessions, Arc::clone(&self.query_registry))
566 .with_sql_idempotency(Arc::clone(&self.sql_idempotency))
567 .with_sql_admission(
568 Arc::clone(&self.sql_semaphore),
569 self.scheduler.clone(),
570 self.sql_priority,
571 )
572 }
573
574 pub async fn shutdown(&self) -> usize {
575 self.accepting_sql.store(false, Ordering::Release);
576 self.query_registry
577 .cancel_all(CancellationReason::ServerShutdown);
578 let deadline = tokio::time::Instant::now() + self.cancel_grace;
579 while self.query_registry.active_count() > 0 && tokio::time::Instant::now() < deadline {
580 tokio::time::sleep(std::time::Duration::from_millis(5)).await;
581 }
582 self.sessions.close_all();
583 let stuck_queries = self.query_registry.active_statuses();
584 for status in &stuck_queries {
585 eprintln!(
586 "[sql-cancel-stuck] query_id={} phase={}",
587 status.query_id,
588 query_phase_name(status.phase)
589 );
590 }
591 let stuck = stuck_queries.len();
592 self.metrics.add_sql_stuck_after_cancel(stuck);
593 if let Some(runtime) = &self.cluster_runtime {
594 if let Err(error) = runtime.shutdown().await {
595 eprintln!("[cluster-runtime] shutdown error: {error}");
596 } else {
597 eprintln!("[cluster-runtime] shutdown complete");
598 }
599 }
600 stuck
601 }
602
603 pub fn cluster_runtime(&self) -> Option<&cluster_runtime::ClusterRuntimeHandle> {
605 self.cluster_runtime.as_ref()
606 }
607
608 pub fn reload_config(&self) -> Result<ReloadReport, String> {
614 let values = mutable_config_from_env();
615 let report = apply_mutable_config(&self.reloadable, &self.db, &values)
616 .map_err(|error| error.to_string())?;
617 self.audit.record(
618 "system",
619 "admin.reload.ok",
620 "configuration reload applied (source: SIGHUP)",
621 );
622 Ok(report)
623 }
624}
625
626pub fn build_app(db: Arc<Database>) -> axum::Router {
627 build_app_with_config(
628 db,
629 std::iter::empty::<Arc<dyn ExternalTableModule>>(),
630 None,
631 None,
632 )
633}
634
635pub fn build_app_with_external_modules(
636 db: Arc<Database>,
637 external_modules: impl IntoIterator<Item = Arc<dyn ExternalTableModule>>,
638) -> axum::Router {
639 build_app_with_config(db, external_modules, None, None)
640}
641
642pub fn build_app_with_config(
644 db: Arc<Database>,
645 external_modules: impl IntoIterator<Item = Arc<dyn ExternalTableModule>>,
646 auth_token: Option<String>,
647 max_connections: Option<usize>,
648) -> axum::Router {
649 build_app_full(db, external_modules, auth_token, max_connections, false)
650}
651
652pub fn build_app_full(
655 db: Arc<Database>,
656 external_modules: impl IntoIterator<Item = Arc<dyn ExternalTableModule>>,
657 auth_token: Option<String>,
658 max_connections: Option<usize>,
659 user_auth: bool,
660) -> axum::Router {
661 let sessions = Arc::new(sessions::SessionStore::new(
662 default_max_sessions(),
663 default_session_idle_timeout(),
664 ));
665 build_app_with_sessions(
666 db,
667 external_modules,
668 auth_token,
669 max_connections,
670 user_auth,
671 sessions,
672 )
673}
674
675pub fn build_app_with_sessions(
679 db: Arc<Database>,
680 external_modules: impl IntoIterator<Item = Arc<dyn ExternalTableModule>>,
681 auth_token: Option<String>,
682 max_connections: Option<usize>,
683 user_auth: bool,
684 sessions: Arc<sessions::SessionStore>,
685) -> axum::Router {
686 build_app_with_sessions_and_control(
687 db,
688 external_modules,
689 auth_token,
690 max_connections,
691 user_auth,
692 sessions,
693 )
694 .0
695}
696
697pub fn build_app_with_sessions_and_control(
699 db: Arc<Database>,
700 external_modules: impl IntoIterator<Item = Arc<dyn ExternalTableModule>>,
701 auth_token: Option<String>,
702 max_connections: Option<usize>,
703 user_auth: bool,
704 sessions: Arc<sessions::SessionStore>,
705) -> (axum::Router, ServerControl) {
706 build_app_with_sessions_control_and_cluster(
707 db,
708 external_modules,
709 auth_token,
710 max_connections,
711 user_auth,
712 sessions,
713 None,
714 )
715}
716
717pub fn build_app_with_sessions_control_and_cluster(
723 db: Arc<Database>,
724 external_modules: impl IntoIterator<Item = Arc<dyn ExternalTableModule>>,
725 auth_token: Option<String>,
726 max_connections: Option<usize>,
727 user_auth: bool,
728 sessions: Arc<sessions::SessionStore>,
729 cluster_runtime: Option<cluster_runtime::ClusterRuntimeHandle>,
730) -> (axum::Router, ServerControl) {
731 db.set_replication_wal_retention_segments(default_replication_wal_segments());
732 if let Err(error) = db.set_history_retention_epochs(default_history_retention_epochs()) {
733 eprintln!("[history] failed to configure retention: {error}");
734 }
735 let max_active_queries = default_sql_max_active_queries();
736 let query_registry = Arc::new(SqlQueryRegistry::new_with_limits(
737 max_active_queries,
738 default_sql_finished_detail_max_entries(),
739 default_sql_finished_detail_max_bytes(),
740 default_sql_finished_compact_max_entries(),
741 default_sql_finished_compact_max_bytes(),
742 default_sql_finished_query_ttl(),
743 ));
744 let accepting_sql = Arc::new(AtomicBool::new(true));
745 let sql_cancel_grace = default_sql_cancel_grace();
746 let sql_max_timeout = default_sql_max_timeout();
747 let sql_default_timeout = default_sql_default_timeout().min(sql_max_timeout);
748 let metrics = Arc::new(metrics::Metrics::default());
749 let audit = Arc::new(audit::AuditLog::new(8192));
750 let reloadable = Arc::new(ReloadableConfig {
751 slow_query_threshold: ReloadableDuration::new(metrics::slow_query_threshold()),
752 sql_default_timeout: ReloadableDuration::new(sql_default_timeout),
753 sql_max_timeout: ReloadableDuration::new(sql_max_timeout),
754 sql_cancel_grace: ReloadableDuration::new(sql_cancel_grace),
755 sql_max_output_rows: ReloadableUsize::new(default_sql_max_output_rows()),
756 sql_max_output_bytes: ReloadableUsize::new(default_sql_max_output_bytes()),
757 });
758 let (idempotency_root, idempotency_integrity) =
759 match sql_idempotency::IdempotencyIntegrity::for_database(&db) {
760 Ok((root, integrity)) => (root, Some(integrity)),
761 Err(error) => {
762 eprintln!("[idempotency] durable integrity key unavailable: {error}");
763 (db.durable_root(), None)
764 }
765 };
766 let sql_idempotency = Arc::new(sql_idempotency::SqlIdempotencyStore::new_with_integrity(
767 Arc::clone(&idempotency_root),
768 idempotency_integrity.clone(),
769 default_sql_idempotency_ttl(),
770 default_sql_idempotency_max_entries(),
771 ));
772 let resource_groups = mongreldb_core::ResourceGroupRegistry::with_defaults();
773 let scheduler = admission::SchedulerAdmission::from_resource_groups(&resource_groups);
774 let sql_semaphore = Arc::new(tokio::sync::Semaphore::new(default_sql_max_concurrent()));
775 let sql_priority = admission::priority_for_class(
776 &resource_groups,
777 mongreldb_core::WorkloadClass::InteractiveSql,
778 );
779 let server_control = ServerControl {
780 query_registry: Arc::clone(&query_registry),
781 sessions: Arc::clone(&sessions),
782 accepting_sql: Arc::clone(&accepting_sql),
783 cancel_grace: sql_cancel_grace,
784 metrics: Arc::clone(&metrics),
785 reloadable: Arc::clone(&reloadable),
786 db: Arc::clone(&db),
787 audit: Arc::clone(&audit),
788 sql_idempotency: Arc::clone(&sql_idempotency),
789 sql_semaphore: Arc::clone(&sql_semaphore),
790 scheduler: scheduler.clone(),
791 sql_priority,
792 cluster_runtime: cluster_runtime.clone(),
793 };
794 let node_memory_governor = db.memory_governor().clone();
797 let state = Arc::new(AppState {
798 idem: kit::IdempotencyStore::new_with_integrity(
799 idempotency_root,
800 idempotency_integrity,
801 default_sql_idempotency_ttl(),
802 default_sql_idempotency_max_entries(),
803 ),
804 db,
805 external_modules: external_modules.into_iter().collect(),
806 auth_token,
807 user_auth,
808 metrics,
809 audit,
810 sessions,
811 ai_semaphore: Arc::new(tokio::sync::Semaphore::new(default_ai_max_concurrent())),
812 query_registry,
813 query_lifecycle: Mutex::new(()),
814 pre_cancellations: pre_cancel::PreCancelStore::new(
815 default_sql_pre_cancel_ttl(),
816 default_sql_pre_cancel_max_entries(),
817 default_sql_pre_cancel_max_bytes(),
818 default_sql_pre_cancel_max_entries_per_owner(),
819 default_sql_pre_cancel_rate_window(),
820 default_sql_pre_cancel_rate_per_owner(),
821 ),
822 sql_idempotency,
823 sql_pages: sql_pages::SqlPageStore::new(
824 default_sql_page_ttl(),
825 default_sql_page_max_entries(),
826 default_sql_page_max_bytes(),
827 default_sql_page_max_entries_per_owner(),
828 ),
829 sql_semaphore,
830 sql_page_semaphore: Arc::new(tokio::sync::Semaphore::new(
831 default_sql_page_max_concurrent(),
832 )),
833 sql_page_default_timeout: default_sql_page_default_timeout()
834 .min(default_sql_page_max_timeout()),
835 sql_page_max_timeout: default_sql_page_max_timeout(),
836 max_request_bytes: default_max_request_bytes(),
837 accepting_sql,
838 cursor_mac_key: CursorMacKey::default(),
839 reloadable,
840 drain: Arc::new(DrainControl::default()),
841 scheduler,
844 node_governor: std::sync::Mutex::new(mongreldb_core::NodeMemoryGovernor::new(
845 node_memory_governor,
846 )),
847 ai_generations: std::sync::Mutex::new(mongreldb_core::AiIndexGenerationRegistry::new()),
848 multi_region: std::sync::Mutex::new(
849 mongreldb_cluster::multi_region::MultiRegionPolicy::default(),
850 ),
851 ops_jobs: std::sync::Mutex::new(mongreldb_core::OpsJobStore::new()),
852 resource_groups,
853 embedding_providers: mongreldb_core::EmbeddingProviderRegistry::new(),
854 cluster_runtime,
855 });
856 let router = axum::Router::new()
857 .route("/health", get(health))
858 .route("/build-info", get(build_info))
859 .route("/capabilities", get(capabilities))
860 .route(
861 "/history/retention",
862 get(history_retention).put(set_history_retention),
863 )
864 .route("/metrics", get(metrics_handler))
865 .route("/audit", get(audit_handler))
866 .route("/admin/drain", get(drain_status).post(admin_drain))
867 .route("/admin/reload", post(admin_reload))
868 .route("/admin/cluster/status", get(cluster_admin::status))
870 .route("/admin/cluster/node/drain", post(cluster_admin::drain))
871 .route("/admin/cluster/node/remove", post(cluster_admin::remove))
872 .route("/tables", get(list_tables).post(create_table))
873 .route("/tables/{name}", axum::routing::delete(drop_table))
874 .route("/tables/{name}/put", post(put_row))
875 .route("/tables/{name}/count", get(count))
876 .route("/tables/{name}/commit", post(commit))
877 .route("/sql", post(sql))
878 .route("/sql/continue", post(continue_sql_page))
879 .route("/queries/{query_id}", get(query_status))
880 .route("/queries/{query_id}/cancel", post(cancel_query))
881 .route("/txn", post(txn))
882 .route("/sessions", post(create_session))
883 .route("/sessions/{id}", axum::routing::delete(close_session))
884 .route("/sessions/{id}/prepare", post(prepare_statement))
885 .route("/sessions/{id}/execute", post(execute_statement))
886 .route(
887 "/sessions/{id}/statements/{name}",
888 axum::routing::delete(deallocate_statement),
889 )
890 .route("/procedures", get(procedure::list).post(procedure::create))
891 .route(
892 "/procedures/{name}",
893 get(procedure::describe)
894 .put(procedure::replace)
895 .delete(procedure::drop_procedure),
896 )
897 .route("/procedures/{name}/call", post(procedure::call))
898 .route("/triggers", get(trigger::list).post(trigger::create))
899 .route(
900 "/triggers/{name}",
901 get(trigger::describe)
902 .put(trigger::replace)
903 .delete(trigger::drop_trigger),
904 )
905 .route("/kit/schema", get(kit::schema_all))
907 .route("/kit/schema/{table}", get(kit::schema_one))
908 .route("/kit/txn", post(kit::kit_txn))
909 .route("/kit/query", post(kit::kit_query))
910 .route("/kit/retrieve", post(kit::kit_retrieve))
911 .route("/kit/ann_rerank", post(kit::kit_ann_rerank))
912 .route("/kit/ai/metrics", get(kit::kit_ai_metrics))
913 .route("/kit/set_similarity", post(kit::kit_set_similarity))
914 .route("/kit/search", post(kit::kit_search))
915 .route("/kit/create_table", post(kit::kit_create_table))
916 .route("/kit/procedures/{name}/call", post(procedure::kit_call))
917 .route("/compact", post(compact_all))
918 .route("/tables/{name}/compact", post(compact_table))
919 .route("/wal/stream", get(wal_stream))
920 .route("/replication/snapshot", get(replication_snapshot))
921 .route("/events", get(events_stream))
922 .with_state(state.clone());
923
924 let router = if state.auth_token.is_some() || state.user_auth || state.db.require_auth_enabled()
928 {
929 router.layer(axum::middleware::from_fn_with_state(
930 state.clone(),
931 auth_middleware,
932 ))
933 } else {
934 router
935 };
936
937 let router = router
941 .layer(axum::middleware::from_fn_with_state(
942 state.clone(),
943 request_body_limit_middleware,
944 ))
945 .layer(axum::extract::DefaultBodyLimit::max(
946 state.max_request_bytes,
947 ));
948
949 let router = if let Some(max) = max_connections {
951 router.layer(tower::limit::ConcurrencyLimitLayer::new(max))
952 } else {
953 router
954 };
955 (router, server_control)
956}
957
958async fn auth_middleware(
965 axum::extract::State(state): axum::extract::State<Arc<AppState>>,
966 mut req: axum::extract::Request,
967 next: axum::middleware::Next,
968) -> Result<axum::response::Response, axum::http::StatusCode> {
969 let header = req
970 .headers()
971 .get("authorization")
972 .and_then(|v| v.to_str().ok())
973 .unwrap_or("");
974
975 let mut attempted = String::new();
979 let mut fail_reason = "no credentials provided".to_string();
980
981 if let Some(token) = &state.auth_token {
983 if let Some(provided) = header.strip_prefix("Bearer ") {
984 attempted = "token".to_string();
985 if provided == token {
986 state
987 .audit
988 .record("token", "login.ok", "bearer token accepted");
989 return Ok(next.run(req).await);
990 }
991 fail_reason = "invalid bearer token".to_string();
992 }
993 }
994
995 if state.user_auth {
997 if let Some(encoded) = header.strip_prefix("Basic ") {
998 if let Ok(decoded) = base64_decode(encoded) {
999 let decoded = Zeroizing::new(decoded);
1000 if let Ok(creds) = std::str::from_utf8(&decoded) {
1001 if let Some((username, password)) = creds.split_once(':') {
1002 attempted = username.to_string();
1003 if let Ok(Some(principal)) =
1004 state.db.authenticate_principal(username, password)
1005 {
1006 let username = principal.username.clone();
1007 drop(decoded);
1008 state
1009 .audit
1010 .record(&username, "login.ok", "basic credentials accepted");
1011 req.extensions_mut().insert(principal);
1012 return Ok(next.run(req).await);
1013 }
1014 fail_reason = "invalid basic credentials".to_string();
1015 } else {
1016 fail_reason = "malformed basic credentials (no ':')".to_string();
1017 }
1018 } else {
1019 fail_reason = "malformed basic credentials (non-utf8)".to_string();
1020 }
1021 } else {
1022 fail_reason = "malformed basic credentials (bad base64)".to_string();
1023 }
1024 }
1025 }
1026
1027 let who = if attempted.is_empty() {
1028 "anonymous"
1029 } else {
1030 attempted.as_str()
1031 };
1032 state.audit.record(who, "login.fail", fail_reason);
1033 Err(axum::http::StatusCode::UNAUTHORIZED)
1034}
1035
1036fn base64_decode(input: &str) -> Result<Vec<u8>, ()> {
1038 const TABLE: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
1039 let input: Vec<u8> = input
1040 .bytes()
1041 .filter(|&b| b != b'\n' && b != b'\r' && b != b' ')
1042 .collect();
1043 let mut out = Vec::with_capacity(input.len() * 3 / 4);
1044 let mut buf = 0u32;
1045 let mut bits = 0u32;
1046 for &b in &input {
1047 if b == b'=' {
1048 break;
1049 }
1050 let val = TABLE.iter().position(|&t| t == b).ok_or(())? as u32;
1051 buf = (buf << 6) | val;
1052 bits += 6;
1053 if bits >= 8 {
1054 bits -= 8;
1055 out.push((buf >> bits) as u8);
1056 }
1057 }
1058 Ok(out)
1059}
1060
1061fn structured_category_error_response(
1067 http_status: StatusCode,
1068 code: &'static str,
1069 message: impl Into<String>,
1070 category: mongreldb_types::errors::ErrorCategory,
1071) -> Response {
1072 (
1073 http_status,
1074 Json(json!({
1075 "error": {
1076 "code": code,
1077 "message": message.into(),
1078 "category": category.to_string(),
1079 "category_code": category.code(),
1080 "retryable": category.is_retryable(),
1081 }
1082 })),
1083 )
1084 .into_response()
1085}
1086
1087async fn request_body_limit_middleware(
1092 axum::extract::State(state): axum::extract::State<Arc<AppState>>,
1093 req: axum::extract::Request,
1094 next: axum::middleware::Next,
1095) -> axum::response::Response {
1096 let declared = req
1097 .headers()
1098 .get(axum::http::header::CONTENT_LENGTH)
1099 .and_then(|value| value.to_str().ok())
1100 .and_then(|value| value.parse::<u64>().ok());
1101 if let Some(declared) = declared {
1102 if declared > state.max_request_bytes as u64 {
1103 return structured_category_error_response(
1104 StatusCode::PAYLOAD_TOO_LARGE,
1105 "REQUEST_BODY_TOO_LARGE",
1106 format!(
1107 "request body of {declared} bytes exceeds the server limit of {} bytes",
1108 state.max_request_bytes
1109 ),
1110 mongreldb_types::errors::ErrorCategory::ResourceExhausted,
1111 );
1112 }
1113 }
1114 next.run(req).await
1115}
1116
1117async fn wal_stream(
1124 axum::extract::State(state): axum::extract::State<Arc<AppState>>,
1125 OptionalPrincipal(principal): OptionalPrincipal,
1126 axum::extract::Query(params): axum::extract::Query<WalStreamParams>,
1127) -> Result<Response, StatusCode> {
1128 state
1129 .db
1130 .require_for(
1131 request_principal(&state, &principal).as_ref(),
1132 &mongreldb_core::Permission::Admin,
1133 )
1134 .map_err(|error| status_for_error(&error))?;
1135 let since = params.since.unwrap_or(0);
1136 let db = Arc::clone(&state.db);
1137 let batch = tokio::task::spawn_blocking(move || db.replication_batch_since(since))
1138 .await
1139 .map_err(|_e| StatusCode::INTERNAL_SERVER_ERROR)?;
1140 let batch = batch.map_err(|e| {
1141 eprintln!("wal_stream error: {e}");
1142 StatusCode::INTERNAL_SERVER_ERROR
1143 })?;
1144 if batch.requires_snapshot {
1145 let mut response = (
1146 StatusCode::CONFLICT,
1147 "replication snapshot required: WAL retention gap or spilled run",
1148 )
1149 .into_response();
1150 response.headers_mut().insert(
1151 "x-mongreldb-replication-status",
1152 "snapshot-required"
1153 .parse()
1154 .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?,
1155 );
1156 set_replication_headers(
1157 &mut response,
1158 &batch.source_id,
1159 batch.from_epoch,
1160 batch.current_epoch,
1161 batch.earliest_epoch,
1162 batch.commit_count,
1163 &batch.records_sha256,
1164 )?;
1165 return Ok(response);
1166 }
1167 let mut body = String::new();
1168 for record in &batch.records {
1169 let json = serde_json::to_string(record).map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
1170 body.push_str(&json);
1171 body.push('\n');
1172 }
1173 let mut response = (
1174 [
1175 (header::CONTENT_TYPE, "application/x-ndjson".to_string()),
1176 (header::CACHE_CONTROL, "no-cache".to_string()),
1177 ],
1178 body,
1179 )
1180 .into_response();
1181 set_replication_headers(
1182 &mut response,
1183 &batch.source_id,
1184 batch.from_epoch,
1185 batch.current_epoch,
1186 batch.earliest_epoch,
1187 batch.commit_count,
1188 &batch.records_sha256,
1189 )?;
1190 Ok(response)
1191}
1192
1193async fn replication_snapshot(
1194 axum::extract::State(state): axum::extract::State<Arc<AppState>>,
1195 OptionalPrincipal(principal): OptionalPrincipal,
1196) -> Response {
1197 if let Err(error) = state.db.require_for(
1198 request_principal(&state, &principal).as_ref(),
1199 &mongreldb_core::Permission::Admin,
1200 ) {
1201 return (status_for_error(&error), error.to_string()).into_response();
1202 }
1203 let db = Arc::clone(&state.db);
1204 let snapshot = match tokio::task::spawn_blocking(move || db.replication_snapshot()).await {
1205 Ok(Ok(snapshot)) => snapshot,
1206 Ok(Err(error)) => {
1207 return (StatusCode::INTERNAL_SERVER_ERROR, error.to_string()).into_response()
1208 }
1209 Err(error) => {
1210 return (StatusCode::INTERNAL_SERVER_ERROR, error.to_string()).into_response()
1211 }
1212 };
1213 let epoch = snapshot.epoch();
1214 match snapshot.encode() {
1217 Ok(bytes) => {
1218 let mut response =
1219 ([(header::CONTENT_TYPE, "application/octet-stream")], bytes).into_response();
1220 let Ok(value) = epoch.to_string().parse() else {
1221 return (
1222 StatusCode::INTERNAL_SERVER_ERROR,
1223 "invalid replication epoch response header",
1224 )
1225 .into_response();
1226 };
1227 response
1228 .headers_mut()
1229 .insert("x-mongreldb-current-epoch", value);
1230 let source_id = snapshot
1231 .source_id()
1232 .iter()
1233 .map(|byte| format!("{byte:02x}"))
1234 .collect::<String>();
1235 let Ok(source_id) = source_id.parse() else {
1236 return (
1237 StatusCode::INTERNAL_SERVER_ERROR,
1238 "invalid replication source response header",
1239 )
1240 .into_response();
1241 };
1242 response
1243 .headers_mut()
1244 .insert("x-mongreldb-source-id", source_id);
1245 response
1246 }
1247 Err(error) => (StatusCode::INTERNAL_SERVER_ERROR, error.to_string()).into_response(),
1248 }
1249}
1250
1251fn set_replication_headers(
1252 response: &mut Response,
1253 source_id: &[u8; 32],
1254 from: u64,
1255 current: u64,
1256 earliest: Option<u64>,
1257 commit_count: u64,
1258 records_sha256: &[u8; 32],
1259) -> Result<(), StatusCode> {
1260 let source_id = source_id
1261 .iter()
1262 .map(|byte| format!("{byte:02x}"))
1263 .collect::<String>();
1264 response.headers_mut().insert(
1265 "x-mongreldb-source-id",
1266 source_id
1267 .parse()
1268 .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?,
1269 );
1270 response.headers_mut().insert(
1271 "x-mongreldb-from-epoch",
1272 from.to_string()
1273 .parse()
1274 .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?,
1275 );
1276 response.headers_mut().insert(
1277 "x-mongreldb-current-epoch",
1278 current
1279 .to_string()
1280 .parse()
1281 .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?,
1282 );
1283 response.headers_mut().insert(
1284 "x-mongreldb-commit-count",
1285 commit_count
1286 .to_string()
1287 .parse()
1288 .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?,
1289 );
1290 let digest = records_sha256
1291 .iter()
1292 .map(|byte| format!("{byte:02x}"))
1293 .collect::<String>();
1294 response.headers_mut().insert(
1295 "x-mongreldb-records-sha256",
1296 digest
1297 .parse()
1298 .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?,
1299 );
1300 if let Some(earliest) = earliest {
1301 response.headers_mut().insert(
1302 "x-mongreldb-earliest-epoch",
1303 earliest
1304 .to_string()
1305 .parse()
1306 .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?,
1307 );
1308 }
1309 Ok(())
1310}
1311
1312#[derive(serde::Deserialize)]
1313struct WalStreamParams {
1314 since: Option<u64>,
1315}
1316
1317async fn events_stream(
1322 axum::extract::State(state): axum::extract::State<Arc<AppState>>,
1323 OptionalPrincipal(principal): OptionalPrincipal,
1324 headers: axum::http::HeaderMap,
1325) -> Result<Response, StatusCode> {
1326 use axum::response::sse::{Event, KeepAlive, Sse};
1327 use futures::Stream;
1328 use std::collections::VecDeque;
1329 use std::convert::Infallible;
1330
1331 state
1332 .db
1333 .require_for(
1334 request_principal(&state, &principal).as_ref(),
1335 &mongreldb_core::Permission::Admin,
1336 )
1337 .map_err(|error| status_for_error(&error))?;
1338
1339 struct State {
1340 db: Arc<Database>,
1341 receiver: tokio::sync::broadcast::Receiver<mongreldb_core::ChangeEvent>,
1342 change_wake: tokio::sync::broadcast::Receiver<()>,
1343 interval: tokio::time::Interval,
1344 pending: VecDeque<mongreldb_core::ChangeEvent>,
1345 last_id: Option<String>,
1346 poll_now: bool,
1347 done: bool,
1348 }
1349
1350 fn event(change: mongreldb_core::ChangeEvent) -> Event {
1351 let id = change.id.clone();
1352 let kind = if change.op == "notify" {
1353 "notify"
1354 } else {
1355 "change"
1356 };
1357 let mut event = Event::default().event(kind).data(
1358 serde_json::to_string(&change)
1359 .unwrap_or_else(|error| format!(r#"{{"error":"{error}"}}"#)),
1360 );
1361 if let Some(id) = id {
1362 event = event.id(id);
1363 }
1364 event
1365 }
1366
1367 let last_id = match headers.get("last-event-id") {
1368 Some(value) => Some(
1369 value
1370 .to_str()
1371 .map_err(|_| StatusCode::BAD_REQUEST)?
1372 .to_owned(),
1373 ),
1374 None => None,
1375 };
1376 let receiver = state.db.subscribe_changes();
1377 let change_wake = state.db.subscribe_change_commits();
1378 let db = Arc::clone(&state.db);
1379 let resume = last_id.clone();
1380 let initial = tokio::task::spawn_blocking(move || db.change_events_since(resume.as_deref()))
1381 .await
1382 .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?
1383 .map_err(|error| match error {
1384 mongreldb_core::MongrelError::InvalidArgument(_) => StatusCode::BAD_REQUEST,
1385 _ => StatusCode::INTERNAL_SERVER_ERROR,
1386 })?;
1387 if initial.gap {
1388 return Ok((
1389 StatusCode::CONFLICT,
1390 Json(json!({
1391 "error": "cdc retention gap",
1392 "earliest_epoch": initial.earliest_epoch,
1393 "current_epoch": initial.current_epoch,
1394 })),
1395 )
1396 .into_response());
1397 }
1398
1399 let stream: std::pin::Pin<Box<dyn Stream<Item = Result<Event, Infallible>> + Send>> = Box::pin(
1400 futures::stream::unfold(
1401 State {
1402 db: Arc::clone(&state.db),
1403 receiver,
1404 change_wake,
1405 interval: tokio::time::interval(std::time::Duration::from_millis(250)),
1406 pending: initial.events.into(),
1407 last_id,
1408 poll_now: false,
1409 done: false,
1410 },
1411 |mut stream| async move {
1412 if stream.done {
1413 return None;
1414 }
1415 loop {
1416 if stream.poll_now {
1417 stream.poll_now = false;
1418 let db = Arc::clone(&stream.db);
1419 let last_id = stream.last_id.clone();
1420 match tokio::task::spawn_blocking(move || {
1421 db.change_events_since(last_id.as_deref())
1422 })
1423 .await
1424 {
1425 Ok(Ok(batch)) if batch.gap => {
1426 stream.done = true;
1427 let gap = Event::default().event("gap").data(
1428 json!({
1429 "error": "cdc retention gap",
1430 "earliest_epoch": batch.earliest_epoch,
1431 "current_epoch": batch.current_epoch,
1432 })
1433 .to_string(),
1434 );
1435 return Some((Ok(gap), stream));
1436 }
1437 Ok(Ok(batch)) => stream.pending.extend(batch.events),
1438 Ok(Err(error)) => {
1439 stream.done = true;
1440 return Some((
1441 Ok(Event::default().event("error").data(error.to_string())),
1442 stream,
1443 ));
1444 }
1445 Err(error) => {
1446 stream.done = true;
1447 return Some((
1448 Ok(Event::default().event("error").data(error.to_string())),
1449 stream,
1450 ));
1451 }
1452 }
1453 }
1454 if let Some(change) = stream.pending.pop_front() {
1455 if let Some(id) = &change.id {
1456 stream.last_id = Some(id.clone());
1457 }
1458 return Some((Ok(event(change)), stream));
1459 }
1460 tokio::select! {
1461 received = stream.receiver.recv() => {
1462 match received {
1463 Ok(change) => return Some((Ok(event(change)), stream)),
1464 Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {},
1465 Err(tokio::sync::broadcast::error::RecvError::Closed) => {
1466 return None;
1467 }
1468 }
1469 }
1470 received = stream.change_wake.recv() => {
1471 match received {
1472 Ok(()) | Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
1473 stream.poll_now = true;
1474 }
1475 Err(tokio::sync::broadcast::error::RecvError::Closed) => {}
1476 }
1477 }
1478 _ = stream.interval.tick() => {
1479 stream.poll_now = true;
1480 }
1481 }
1482 }
1483 },
1484 ),
1485 );
1486
1487 Ok(Sse::new(stream)
1488 .keep_alive(
1489 KeepAlive::new()
1490 .interval(std::time::Duration::from_secs(15))
1491 .text("keep-alive"),
1492 )
1493 .into_response())
1494}
1495
1496pub fn spawn_auto_compactor(db: Arc<Database>) {
1501 if let Err(error) = std::thread::Builder::new()
1502 .name("mongreldb-auto-compact".into())
1503 .spawn(move || loop {
1504 std::thread::sleep(std::time::Duration::from_secs(30));
1505 for name in db.table_names() {
1506 let Ok(handle) = db.table(&name) else {
1507 continue;
1508 };
1509 let mut t = handle.lock();
1510 let before = t.run_count();
1511 match t.maybe_compact() {
1512 Ok(true) => {
1513 eprintln!(
1514 "[auto-compact] {name}: {} runs -> {}",
1515 before,
1516 t.run_count()
1517 );
1518 }
1519 Ok(false) => {}
1520 Err(e) => {
1521 eprintln!("[auto-compact] {name}: compaction failed: {e}");
1522 }
1523 }
1524 }
1525 })
1526 {
1527 eprintln!("[auto-compact] failed to start background thread: {error}");
1528 }
1529}
1530
1531async fn health() -> StatusCode {
1532 StatusCode::OK
1533}
1534
1535#[derive(Debug, Serialize)]
1536struct SqlCancellationCapabilities {
1537 version: u8,
1538 client_query_ids: bool,
1539 cancel_endpoint: bool,
1540 query_status: bool,
1541 pre_registration_cancel: bool,
1542 stream_disconnect_cancels: bool,
1543}
1544
1545#[derive(Debug, Serialize)]
1546struct SqlIdempotencyCapabilities {
1547 version: u8,
1548 durable_pre_execution_intent: bool,
1549 replay_committed_receipt: bool,
1550 indeterminate_never_reexecutes: bool,
1551}
1552
1553#[derive(Debug, Serialize)]
1554struct SqlPaginationCapabilities {
1555 version: u8,
1556 continuation_endpoint: &'static str,
1557 retained_snapshot: bool,
1558 projection_required: bool,
1559 byte_and_token_hints: bool,
1560}
1561
1562#[derive(Debug, Serialize)]
1563struct CapabilitiesResponse {
1564 sql_cancellation: SqlCancellationCapabilities,
1565 sql_idempotency: SqlIdempotencyCapabilities,
1566 sql_pagination: SqlPaginationCapabilities,
1567}
1568
1569async fn capabilities() -> Json<CapabilitiesResponse> {
1570 Json(CapabilitiesResponse {
1571 sql_cancellation: SqlCancellationCapabilities {
1572 version: 2,
1573 client_query_ids: true,
1574 cancel_endpoint: true,
1575 query_status: true,
1576 pre_registration_cancel: true,
1577 stream_disconnect_cancels: true,
1578 },
1579 sql_idempotency: SqlIdempotencyCapabilities {
1580 version: 1,
1581 durable_pre_execution_intent: true,
1582 replay_committed_receipt: true,
1583 indeterminate_never_reexecutes: true,
1584 },
1585 sql_pagination: SqlPaginationCapabilities {
1586 version: 1,
1587 continuation_endpoint: "/sql/continue",
1588 retained_snapshot: true,
1589 projection_required: true,
1590 byte_and_token_hints: true,
1591 },
1592 })
1593}
1594
1595async fn build_info() -> Json<mongreldb_query::BuildInfo> {
1596 Json(mongreldb_query::build_info())
1597}
1598
1599#[derive(Debug, Deserialize)]
1600struct HistoryRetentionRequest {
1601 #[serde(default)]
1602 history_retention_epochs: serde_json::Value,
1603}
1604
1605#[derive(Debug, Serialize)]
1606struct HistoryRetentionResponse {
1607 history_retention_epochs: u64,
1608 earliest_retained_epoch: u64,
1609}
1610
1611fn history_retention_response(db: &Database) -> HistoryRetentionResponse {
1612 HistoryRetentionResponse {
1613 history_retention_epochs: db.history_retention_epochs(),
1614 earliest_retained_epoch: db.earliest_retained_epoch().0,
1615 }
1616}
1617
1618async fn history_retention(
1620 State(state): State<Arc<AppState>>,
1621 OptionalPrincipal(principal): OptionalPrincipal,
1622) -> Response {
1623 if let Err(error) = state.db.require_for(
1624 request_principal(&state, &principal).as_ref(),
1625 &mongreldb_core::Permission::Admin,
1626 ) {
1627 return (status_for_error(&error), error.to_string()).into_response();
1628 }
1629 Json(history_retention_response(&state.db)).into_response()
1630}
1631
1632async fn set_history_retention(
1634 State(state): State<Arc<AppState>>,
1635 OptionalPrincipal(principal): OptionalPrincipal,
1636 Json(request): Json<HistoryRetentionRequest>,
1637) -> Response {
1638 if let Err(error) = state.db.require_for(
1639 request_principal(&state, &principal).as_ref(),
1640 &mongreldb_core::Permission::Admin,
1641 ) {
1642 return (status_for_error(&error), error.to_string()).into_response();
1643 }
1644 let Some(epochs) = request.history_retention_epochs.as_u64() else {
1645 return (
1646 StatusCode::BAD_REQUEST,
1647 Json(json!({"error": "history_retention_epochs must be a u64"})),
1648 )
1649 .into_response();
1650 };
1651 match state.db.set_history_retention_epochs(epochs) {
1652 Ok(()) => Json(history_retention_response(&state.db)).into_response(),
1653 Err(error) => (status_for_error(&error), error.to_string()).into_response(),
1654 }
1655}
1656
1657async fn audit_handler(
1662 State(state): State<Arc<AppState>>,
1663 OptionalPrincipal(principal): OptionalPrincipal,
1664) -> Response {
1665 if let Err(error) = state.db.require_for(
1666 request_principal(&state, &principal).as_ref(),
1667 &mongreldb_core::Permission::Admin,
1668 ) {
1669 return (status_for_error(&error), error.to_string()).into_response();
1670 }
1671 let recent = state.audit.recent();
1672 Json(recent).into_response()
1673}
1674
1675fn require_admin(
1680 state: &AppState,
1681 principal: &Option<mongreldb_core::Principal>,
1682 action: &str,
1683) -> Result<String, Box<Response>> {
1684 let owner = request_owner(state, principal);
1685 if let Err(error) = state.db.require_for(
1686 request_principal(state, principal).as_ref(),
1687 &mongreldb_core::Permission::Admin,
1688 ) {
1689 state
1690 .audit
1691 .record(owner, format!("{action}.fail"), "authorization failed");
1692 return Err(Box::new(
1693 (status_for_error(&error), error.to_string()).into_response(),
1694 ));
1695 }
1696 Ok(owner)
1697}
1698
1699fn require_writes_open(state: &AppState) -> Option<Response> {
1706 if state.accepting_sql.load(Ordering::Acquire)
1707 && state.db.lifecycle_state() == mongreldb_core::LifecycleState::Open
1708 {
1709 return None;
1710 }
1711 Some(
1712 (
1713 StatusCode::SERVICE_UNAVAILABLE,
1714 "server is draining; writes are closed",
1715 )
1716 .into_response(),
1717 )
1718}
1719
1720fn drain_status_json(state: &AppState) -> serde_json::Value {
1722 let record = state
1723 .drain
1724 .record
1725 .lock()
1726 .map(|record| record.clone())
1727 .unwrap_or_else(|poisoned| poisoned.into_inner().clone());
1728 json!({
1729 "lifecycle": state.db.lifecycle_state().to_string(),
1730 "accepting_sql": state.accepting_sql.load(Ordering::Acquire),
1731 "active_queries": state.query_registry.active_count(),
1732 "live_sessions": state.sessions.len(),
1733 "drain": record,
1734 })
1735}
1736
1737async fn drain_status(
1741 State(state): State<Arc<AppState>>,
1742 OptionalPrincipal(principal): OptionalPrincipal,
1743) -> Response {
1744 if let Err(response) = require_admin(&state, &principal, "admin.drain_status") {
1745 return *response;
1746 }
1747 Json(drain_status_json(&state)).into_response()
1748}
1749
1750async fn admin_drain(
1761 State(state): State<Arc<AppState>>,
1762 OptionalPrincipal(principal): OptionalPrincipal,
1763 body: Option<Json<DrainRequest>>,
1764) -> Response {
1765 let owner = match require_admin(&state, &principal, "admin.drain") {
1766 Ok(owner) => owner,
1767 Err(response) => return *response,
1768 };
1769 let deadline_ms = body
1770 .and_then(|Json(request)| request.drain_deadline_ms)
1771 .unwrap_or_else(default_drain_deadline_ms);
1772 if deadline_ms == 0 {
1773 return (
1774 StatusCode::BAD_REQUEST,
1775 Json(json!({"error": "drain_deadline_ms must be positive"})),
1776 )
1777 .into_response();
1778 }
1779 let _serialization = state.drain.lock.lock().await;
1780 if matches!(
1781 state.db.lifecycle_state(),
1782 mongreldb_core::LifecycleState::Closed
1783 ) {
1784 return Json(drain_status_json(&state)).into_response();
1786 }
1787 state.audit.record(
1788 owner.clone(),
1789 "admin.drain",
1790 format!("initiated deadline_ms={deadline_ms}"),
1791 );
1792 let deadline = std::time::Duration::from_millis(deadline_ms);
1793 let started = std::time::Instant::now();
1794 state.accepting_sql.store(false, Ordering::Release);
1796 state
1798 .query_registry
1799 .cancel_all(CancellationReason::ServerShutdown);
1800 let grace = state.reloadable.sql_cancel_grace.get().min(deadline);
1801 let grace_deadline = tokio::time::Instant::now() + grace;
1802 while state.query_registry.active_count() > 0 && tokio::time::Instant::now() < grace_deadline {
1803 tokio::time::sleep(std::time::Duration::from_millis(5)).await;
1804 }
1805 state.sessions.close_all();
1807 let remaining = deadline.saturating_sub(started.elapsed());
1809 let core = state.db.core();
1810 let result = tokio::task::spawn_blocking(move || core.shutdown(remaining)).await;
1811 let lifecycle = state.db.lifecycle_state().to_string();
1812 let (completed, detail) = match &result {
1813 Ok(Ok(())) => (true, format!("drain completed; core {lifecycle}")),
1814 Ok(Err(error)) => (false, format!("drain incomplete: {error}")),
1815 Err(error) => (false, format!("drain task failed: {error}")),
1816 };
1817 if completed {
1818 state.audit.record(
1819 owner,
1820 "admin.drain.ok",
1821 format!("deadline_ms={deadline_ms} lifecycle={lifecycle}"),
1822 );
1823 } else {
1824 state.audit.record(
1825 owner,
1826 "admin.drain.fail",
1827 format!("deadline_ms={deadline_ms} {detail}"),
1828 );
1829 }
1830 if let Ok(mut record) = state.drain.record.lock() {
1831 record.initiated = true;
1832 record.completed = completed;
1833 record.deadline_ms = deadline_ms;
1834 record.lifecycle = lifecycle;
1835 record.detail = detail;
1836 }
1837 let status = drain_status_json(&state);
1838 if completed {
1839 Json(status).into_response()
1840 } else {
1841 (StatusCode::CONFLICT, Json(status)).into_response()
1843 }
1844}
1845
1846#[derive(Deserialize)]
1847struct DrainRequest {
1848 #[serde(default)]
1849 drain_deadline_ms: Option<u64>,
1850}
1851
1852fn default_drain_deadline_ms() -> u64 {
1855 positive_env_u64("MONGRELDB_DRAIN_DEADLINE_MS", 30_000)
1856}
1857
1858async fn admin_reload(
1870 State(state): State<Arc<AppState>>,
1871 OptionalPrincipal(principal): OptionalPrincipal,
1872 body: Option<Json<MutableConfigOverrides>>,
1873) -> Response {
1874 let owner = match require_admin(&state, &principal, "admin.reload") {
1875 Ok(owner) => owner,
1876 Err(response) => return *response,
1877 };
1878 let mut values = mutable_config_from_env();
1879 if let Some(Json(overrides)) = body {
1880 if let Err(message) = values.apply_overrides(&overrides) {
1881 return (StatusCode::BAD_REQUEST, Json(json!({ "error": message }))).into_response();
1882 }
1883 }
1884 match apply_mutable_config(&state.reloadable, &state.db, &values) {
1885 Ok(report) => {
1886 state.audit.record(
1887 owner,
1888 "admin.reload.ok",
1889 "slow_query_ms sql_default_timeout_ms sql_max_timeout_ms sql_cancel_grace_ms sql_max_output_rows sql_max_output_bytes history_retention_epochs",
1890 );
1891 Json(json!({ "reloaded": true, "applied": report })).into_response()
1892 }
1893 Err(error) => {
1894 state
1895 .audit
1896 .record(owner, "admin.reload.fail", error.to_string());
1897 (status_for_error(&error), error.to_string()).into_response()
1898 }
1899 }
1900}
1901
1902fn default_max_sessions() -> usize {
1904 std::env::var("MONGRELBL_MAX_SESSIONS")
1905 .ok()
1906 .and_then(|v| v.parse().ok())
1907 .unwrap_or(256)
1908}
1909
1910fn default_session_idle_timeout() -> std::time::Duration {
1912 std::time::Duration::from_secs(
1913 std::env::var("MONGRELBL_SESSION_IDLE_TIMEOUT_SECS")
1914 .ok()
1915 .and_then(|v| v.parse().ok())
1916 .unwrap_or(300),
1917 )
1918}
1919
1920fn default_replication_wal_segments() -> usize {
1921 let replication = std::env::var("MONGRELDB_REPLICATION_WAL_SEGMENTS")
1922 .ok()
1923 .and_then(|value| value.parse().ok())
1924 .unwrap_or(16);
1925 let cdc = std::env::var("MONGRELDB_CDC_WAL_SEGMENTS")
1926 .ok()
1927 .and_then(|value| value.parse().ok())
1928 .unwrap_or(16);
1929 replication.max(cdc)
1930}
1931
1932fn default_history_retention_epochs() -> u64 {
1933 std::env::var("MONGRELDB_HISTORY_RETENTION_EPOCHS")
1934 .ok()
1935 .and_then(|value| value.parse().ok())
1936 .unwrap_or(1024)
1937}
1938
1939fn default_ai_max_concurrent() -> usize {
1940 std::env::var("MONGRELDB_AI_MAX_CONCURRENT")
1941 .ok()
1942 .and_then(|value| value.parse().ok())
1943 .filter(|value| *value > 0)
1944 .unwrap_or(4)
1945}
1946
1947fn positive_env_u64(name: &str, default: u64) -> u64 {
1948 std::env::var(name)
1949 .ok()
1950 .and_then(|value| value.parse().ok())
1951 .filter(|value| *value > 0)
1952 .unwrap_or(default)
1953}
1954
1955fn positive_env_usize(name: &str, default: usize) -> usize {
1956 std::env::var(name)
1957 .ok()
1958 .and_then(|value| value.parse().ok())
1959 .filter(|value| *value > 0)
1960 .unwrap_or(default)
1961}
1962
1963fn default_sql_default_timeout() -> std::time::Duration {
1964 std::time::Duration::from_millis(positive_env_u64("MONGRELDB_SQL_DEFAULT_TIMEOUT_MS", 30_000))
1965}
1966
1967fn default_sql_max_timeout() -> std::time::Duration {
1968 std::time::Duration::from_millis(positive_env_u64("MONGRELDB_SQL_MAX_TIMEOUT_MS", 300_000))
1969}
1970
1971fn default_sql_max_concurrent() -> usize {
1972 positive_env_usize(
1973 "MONGRELDB_SQL_MAX_CONCURRENT",
1974 std::thread::available_parallelism()
1975 .map(usize::from)
1976 .unwrap_or(4),
1977 )
1978}
1979
1980fn default_sql_max_active_queries() -> usize {
1981 positive_env_usize("MONGRELDB_SQL_MAX_ACTIVE_QUERIES", 1_024)
1982}
1983
1984fn default_sql_page_max_concurrent() -> usize {
1985 positive_env_usize("MONGRELDB_SQL_PAGE_MAX_CONCURRENT", 16)
1986}
1987
1988fn default_sql_page_default_timeout() -> std::time::Duration {
1989 std::time::Duration::from_millis(positive_env_u64(
1990 "MONGRELDB_SQL_PAGE_DEFAULT_TIMEOUT_MS",
1991 5_000,
1992 ))
1993}
1994
1995fn default_sql_page_max_timeout() -> std::time::Duration {
1996 std::time::Duration::from_millis(positive_env_u64(
1997 "MONGRELDB_SQL_PAGE_MAX_TIMEOUT_MS",
1998 30_000,
1999 ))
2000}
2001
2002fn default_sql_finished_query_ttl() -> std::time::Duration {
2003 std::time::Duration::from_secs(positive_env_u64(
2004 "MONGRELDB_SQL_FINISHED_QUERY_TTL_SECS",
2005 60,
2006 ))
2007}
2008
2009fn default_sql_finished_detail_max_entries() -> usize {
2010 positive_env_usize("MONGRELDB_SQL_FINISHED_DETAIL_MAX_ENTRIES", 2_048)
2011}
2012
2013fn default_sql_finished_detail_max_bytes() -> usize {
2014 positive_env_usize("MONGRELDB_SQL_FINISHED_DETAIL_MAX_BYTES", 8 * 1024 * 1024)
2015}
2016
2017fn default_sql_finished_compact_max_entries() -> usize {
2018 positive_env_usize("MONGRELDB_SQL_FINISHED_COMPACT_MAX_ENTRIES", 100_000)
2019}
2020
2021fn default_sql_finished_compact_max_bytes() -> usize {
2022 positive_env_usize("MONGRELDB_SQL_FINISHED_COMPACT_MAX_BYTES", 32 * 1024 * 1024)
2023}
2024
2025fn default_sql_pre_cancel_ttl() -> std::time::Duration {
2026 std::time::Duration::from_millis(positive_env_u64("MONGRELDB_SQL_PRE_CANCEL_TTL_MS", 15_000))
2027}
2028
2029fn default_sql_pre_cancel_max_entries() -> usize {
2030 positive_env_usize("MONGRELDB_SQL_PRE_CANCEL_MAX_ENTRIES", 2_048)
2031}
2032
2033fn default_sql_pre_cancel_max_bytes() -> usize {
2034 positive_env_usize("MONGRELDB_SQL_PRE_CANCEL_MAX_BYTES", 1024 * 1024)
2035}
2036
2037fn default_sql_pre_cancel_max_entries_per_owner() -> usize {
2038 positive_env_usize("MONGRELDB_SQL_PRE_CANCEL_MAX_PER_OWNER", 256)
2039}
2040
2041fn default_sql_pre_cancel_rate_window() -> std::time::Duration {
2042 std::time::Duration::from_millis(positive_env_u64(
2043 "MONGRELDB_SQL_PRE_CANCEL_RATE_WINDOW_MS",
2044 1_000,
2045 ))
2046}
2047
2048fn default_sql_pre_cancel_rate_per_owner() -> usize {
2049 positive_env_usize("MONGRELDB_SQL_PRE_CANCEL_RATE_PER_OWNER", 256)
2050}
2051
2052pub(crate) fn default_sql_idempotency_ttl() -> std::time::Duration {
2053 std::time::Duration::from_secs(positive_env_u64(
2054 "MONGRELDB_SQL_IDEMPOTENCY_TTL_SECS",
2055 86_400,
2056 ))
2057}
2058
2059pub(crate) fn default_sql_idempotency_max_entries() -> usize {
2060 positive_env_usize("MONGRELDB_SQL_IDEMPOTENCY_MAX_ENTRIES", 4_096)
2061}
2062
2063fn default_sql_page_ttl() -> std::time::Duration {
2064 std::time::Duration::from_secs(positive_env_u64("MONGRELDB_SQL_PAGE_TTL_SECS", 60))
2065}
2066
2067fn default_sql_page_max_entries() -> usize {
2068 positive_env_usize("MONGRELDB_SQL_PAGE_MAX_ENTRIES", 128)
2069}
2070
2071fn default_sql_page_max_bytes() -> usize {
2072 positive_env_usize("MONGRELDB_SQL_PAGE_MAX_RETAINED_BYTES", 128 * 1024 * 1024)
2073}
2074
2075fn default_sql_page_max_entries_per_owner() -> usize {
2076 positive_env_usize("MONGRELDB_SQL_PAGE_MAX_PER_OWNER", 16)
2077}
2078
2079fn default_sql_cancel_grace() -> std::time::Duration {
2080 std::time::Duration::from_millis(positive_env_u64("MONGRELDB_SQL_CANCEL_GRACE_MS", 1_000))
2081}
2082
2083fn default_sql_max_output_bytes() -> usize {
2084 positive_env_usize("MONGRELDB_SQL_MAX_OUTPUT_BYTES", 64 * 1024 * 1024)
2085}
2086
2087fn default_sql_max_output_rows() -> usize {
2088 positive_env_usize("MONGRELDB_SQL_MAX_OUTPUT_ROWS", 1_000_000)
2089}
2090
2091fn default_max_request_bytes() -> usize {
2095 positive_env_usize("MONGRELDB_MAX_REQUEST_BYTES", 2 * 1024 * 1024)
2096}
2097
2098fn request_owner(state: &AppState, principal: &Option<mongreldb_core::Principal>) -> String {
2101 if let Some(p) = principal {
2102 return format!("user:{}:{}", p.user_id, p.created_epoch);
2103 }
2104 if let Some(token) = state.auth_token.as_deref() {
2105 let mut digest = sha2::Sha256::new();
2106 digest.update(b"mongreldb-server-bearer-owner-v1\0");
2107 digest.update((token.len() as u64).to_le_bytes());
2108 digest.update(token.as_bytes());
2109 return format!("bearer:{}", sql_idempotency::hex(&digest.finalize()));
2110 }
2111 if state.user_auth || state.db.require_auth_enabled() {
2112 return "unauthenticated".into();
2113 }
2114 "anonymous".into()
2115}
2116
2117fn request_principal(
2118 state: &AppState,
2119 principal: &Option<mongreldb_core::Principal>,
2120) -> Option<mongreldb_core::Principal> {
2121 if let Some(principal) = principal {
2122 return state
2123 .db
2124 .resolve_current_principal(principal)
2125 .or_else(|| Some(principal.clone()));
2126 }
2127 state.auth_token.as_ref().and_then(|_| {
2128 if state.db.require_auth_enabled() {
2129 return state
2130 .db
2131 .principal_snapshot()
2132 .and_then(|principal| state.db.resolve_current_principal(&principal))
2133 .filter(|principal| principal.is_admin);
2134 }
2135 Some(mongreldb_core::Principal {
2136 user_id: 0,
2137 created_epoch: 0,
2138 username: "token".into(),
2139 is_admin: true,
2140 roles: Vec::new(),
2141 permissions: Vec::new(),
2142 })
2143 })
2144}
2145
2146fn current_request_principal(
2147 state: &AppState,
2148 principal: &Option<mongreldb_core::Principal>,
2149) -> Option<mongreldb_core::Principal> {
2150 if let Some(principal) = principal {
2151 return state.db.resolve_current_principal(principal);
2152 }
2153 request_principal(state, principal)
2154}
2155
2156fn request_identity_is_current(
2157 state: &AppState,
2158 principal: &Option<mongreldb_core::Principal>,
2159) -> bool {
2160 if principal.is_some() || state.auth_token.is_some() {
2161 return current_request_principal(state, principal).is_some();
2162 }
2163 !state.user_auth && !state.db.require_auth_enabled()
2164}
2165
2166fn protocol_identity(
2173 principal: &Option<mongreldb_core::Principal>,
2174) -> mongreldb_protocol::request::AuthenticatedIdentity {
2175 match principal {
2176 Some(principal) => mongreldb_protocol::request::AuthenticatedIdentity::CatalogUser {
2177 username: principal.username.clone(),
2178 user_id: principal.user_id,
2179 created_version: principal.created_epoch,
2180 },
2181 None => mongreldb_protocol::request::AuthenticatedIdentity::Credentialless,
2182 }
2183}
2184
2185async fn create_session(
2190 State(state): State<Arc<AppState>>,
2191 OptionalPrincipal(principal): OptionalPrincipal,
2192) -> Response {
2193 if !state.accepting_sql.load(Ordering::Acquire) {
2194 return (StatusCode::SERVICE_UNAVAILABLE, "server is shutting down").into_response();
2195 }
2196 if !request_identity_is_current(&state, &principal) {
2197 return StatusCode::UNAUTHORIZED.into_response();
2198 }
2199 let owner = request_owner(&state, &principal);
2200 let session = match MongrelSession::open_with_external_modules_as(
2201 Arc::clone(&state.db),
2202 state.external_modules.iter().cloned(),
2203 request_principal(&state, &principal),
2204 ) {
2205 Ok(session) => session.with_query_registry(Arc::clone(&state.query_registry)),
2206 Err(e) => return (status_for_query_error(&e), e.to_string()).into_response(),
2207 };
2208 match state
2209 .sessions
2210 .create_with_identity(session, owner.clone(), protocol_identity(&principal))
2211 {
2212 Some(token) => {
2213 state.audit.record(owner, "session.open", "session created");
2214 Json(json!({ "session_id": token })).into_response()
2215 }
2216 None => (
2217 StatusCode::SERVICE_UNAVAILABLE,
2218 "session limit reached; close an idle session or raise --max-sessions",
2219 )
2220 .into_response(),
2221 }
2222}
2223
2224async fn close_session(
2227 State(state): State<Arc<AppState>>,
2228 OptionalPrincipal(principal): OptionalPrincipal,
2229 Path(id): Path<String>,
2230) -> Response {
2231 if !request_identity_is_current(&state, &principal) {
2232 return StatusCode::NOT_FOUND.into_response();
2233 }
2234 let owner = request_owner(&state, &principal);
2235 if let Some(entry) = state.sessions.take_for_close(&id, &owner) {
2236 entry
2237 .session()
2238 .query_registry()
2239 .cancel_session(&id, CancellationReason::SessionClosed);
2240 let deadline = tokio::time::Instant::now() + state.reloadable.sql_cancel_grace.get();
2241 while entry.session().query_registry().active_for_session(&id) > 0
2242 && tokio::time::Instant::now() < deadline
2243 {
2244 tokio::time::sleep(std::time::Duration::from_millis(5)).await;
2245 }
2246 state.audit.record(owner, "session.close", "session closed");
2247 StatusCode::OK.into_response()
2248 } else {
2249 (
2250 StatusCode::NOT_FOUND,
2251 "session not found or not owned by caller",
2252 )
2253 .into_response()
2254 }
2255}
2256
2257async fn dispatch_buffered_sql_format(
2260 state: &AppState,
2261 format: Option<&str>,
2262 output: ManagedQueryBatches,
2263 query_id: QueryId,
2264 test_hook: Option<mongreldb_query::SqlTestHook>,
2265 output_limits: (usize, usize),
2266) -> Response {
2267 let format = format.unwrap_or("json").to_owned();
2268 let (max_rows, max_bytes) = output_limits;
2269 let serialization_batches = output.batches().to_vec();
2273 let serialization_query = output.query().clone();
2274 let serialized = tokio::task::spawn_blocking(move || {
2275 serialize_buffered_output(
2276 &format,
2277 &serialization_batches,
2278 &serialization_query,
2279 max_rows,
2280 max_bytes,
2281 test_hook.as_ref(),
2282 )
2283 })
2284 .await;
2285 let result = match serialized {
2286 Ok(result) => result,
2287 Err(error) => {
2288 output
2289 .query()
2290 .record_serialization_failure("SERIALIZATION_WORKER_FAILED");
2291 output.fail();
2292 state.metrics.inc_sql_errors();
2293 return terminal_server_error_response(
2294 state,
2295 query_id,
2296 StatusCode::INTERNAL_SERVER_ERROR,
2297 "SERIALIZATION_WORKER_FAILED",
2298 error.to_string(),
2299 );
2300 }
2301 };
2302 match result {
2303 Ok(serialized) => {
2304 if let Err(error) = output.try_complete() {
2305 state.metrics.inc_sql_errors();
2306 return tracked_query_error_response(state, &error, Some(query_id));
2307 }
2308 state.metrics.add_sql_output_bytes(serialized.bytes.len());
2309 let content_type = if serialized.arrow {
2310 "application/vnd.apache.arrow.file"
2311 } else {
2312 "application/json"
2313 };
2314 with_query_id(
2315 ([(header::CONTENT_TYPE, content_type)], serialized.bytes).into_response(),
2316 query_id,
2317 )
2318 }
2319 Err(BufferedSerializationError::Query(error)) => {
2320 output.fail();
2321 state.metrics.inc_sql_errors();
2322 tracked_query_error_response(state, &error, Some(query_id))
2323 }
2324 Err(BufferedSerializationError::Limit(message)) => {
2325 output.fail_result_limit();
2326 state.metrics.inc_sql_errors();
2327 terminal_server_error_response(
2328 state,
2329 query_id,
2330 StatusCode::PAYLOAD_TOO_LARGE,
2331 "RESULT_LIMIT_EXCEEDED",
2332 message,
2333 )
2334 }
2335 Err(BufferedSerializationError::Encoding(message)) => {
2336 output.fail_serialization();
2337 state.metrics.inc_sql_errors();
2338 terminal_server_error_response(
2339 state,
2340 query_id,
2341 StatusCode::INTERNAL_SERVER_ERROR,
2342 "SERIALIZATION_FAILED",
2343 message,
2344 )
2345 }
2346 }
2347}
2348
2349#[derive(Debug)]
2350enum PaginatedSerializationError {
2351 Query(mongreldb_query::MongrelQueryError),
2352 Limit(String),
2353 Projection(String),
2354 Encoding(String),
2355}
2356
2357struct SerializedPageRows {
2358 rows: Vec<serde_json::Value>,
2359 retained_bytes: usize,
2360}
2361
2362struct PaginatedJsonReader<'a> {
2363 cursor: std::io::Cursor<&'a [u8]>,
2364 query: &'a RegisteredSqlQuery,
2365 test_hook: Option<&'a mongreldb_query::SqlTestHook>,
2366}
2367
2368impl std::io::Read for PaginatedJsonReader<'_> {
2369 fn read(&mut self, buffer: &mut [u8]) -> std::io::Result<usize> {
2370 if let Some(hook) = self.test_hook {
2371 hook(mongreldb_query::SqlTestHookPoint::DuringPaginationDeserialization);
2372 }
2373 self.query
2374 .checkpoint()
2375 .map_err(|error| std::io::Error::other(error.to_string()))?;
2376 let length = buffer.len().min(64 * 1024);
2378 std::io::Read::read(&mut self.cursor, &mut buffer[..length])
2379 }
2380}
2381
2382const PAGINATED_VALUE_NODE_BYTES: usize = 64;
2383const PAGINATED_OBJECT_ENTRY_BYTES: usize = 128;
2384const PAGINATED_MEMORY_LIMIT_ERROR: &str = "SQL retained output memory limit exceeded";
2385
2386struct PaginatedDecodeBudget<'a> {
2387 used: usize,
2388 limit: usize,
2389 nodes: usize,
2390 exceeded: bool,
2391 query: &'a RegisteredSqlQuery,
2392 test_hook: Option<&'a mongreldb_query::SqlTestHook>,
2393}
2394
2395impl PaginatedDecodeBudget<'_> {
2396 fn begin_value<E: serde::de::Error>(&mut self) -> Result<(), E> {
2397 self.nodes = self.nodes.saturating_add(1);
2398 if self.nodes & 255 == 0 {
2399 if let Some(hook) = self.test_hook {
2400 hook(mongreldb_query::SqlTestHookPoint::DuringPaginationDeserialization);
2401 }
2402 self.query.checkpoint().map_err(E::custom)?;
2403 }
2404 self.charge::<E>(PAGINATED_VALUE_NODE_BYTES)
2405 }
2406
2407 fn charge<E: serde::de::Error>(&mut self, bytes: usize) -> Result<(), E> {
2408 let next = self.used.saturating_add(bytes);
2409 if next > self.limit {
2410 self.exceeded = true;
2411 return Err(E::custom(PAGINATED_MEMORY_LIMIT_ERROR));
2412 }
2413 self.used = next;
2414 Ok(())
2415 }
2416}
2417
2418struct BudgetedJsonValueSeed<'a, 'query> {
2419 budget: &'a mut PaginatedDecodeBudget<'query>,
2420}
2421
2422impl<'de> serde::de::DeserializeSeed<'de> for BudgetedJsonValueSeed<'_, '_> {
2423 type Value = serde_json::Value;
2424
2425 fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
2426 where
2427 D: serde::Deserializer<'de>,
2428 {
2429 self.budget.begin_value::<D::Error>()?;
2430 deserializer.deserialize_any(BudgetedJsonValueVisitor {
2431 budget: self.budget,
2432 })
2433 }
2434}
2435
2436struct BudgetedJsonValueVisitor<'a, 'query> {
2437 budget: &'a mut PaginatedDecodeBudget<'query>,
2438}
2439
2440impl<'de> serde::de::Visitor<'de> for BudgetedJsonValueVisitor<'_, '_> {
2441 type Value = serde_json::Value;
2442
2443 fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2444 formatter.write_str("a JSON value")
2445 }
2446
2447 fn visit_bool<E>(self, value: bool) -> Result<Self::Value, E> {
2448 Ok(serde_json::Value::Bool(value))
2449 }
2450
2451 fn visit_i64<E>(self, value: i64) -> Result<Self::Value, E> {
2452 Ok(serde_json::Value::Number(value.into()))
2453 }
2454
2455 fn visit_u64<E>(self, value: u64) -> Result<Self::Value, E> {
2456 Ok(serde_json::Value::Number(value.into()))
2457 }
2458
2459 fn visit_f64<E>(self, value: f64) -> Result<Self::Value, E>
2460 where
2461 E: serde::de::Error,
2462 {
2463 serde_json::Number::from_f64(value)
2464 .map(serde_json::Value::Number)
2465 .ok_or_else(|| E::custom("JSON number is not finite"))
2466 }
2467
2468 fn visit_str<E>(self, value: &str) -> Result<Self::Value, E>
2469 where
2470 E: serde::de::Error,
2471 {
2472 self.budget.charge::<E>(value.len())?;
2473 Ok(serde_json::Value::String(value.to_owned()))
2474 }
2475
2476 fn visit_borrowed_str<E>(self, value: &'de str) -> Result<Self::Value, E>
2477 where
2478 E: serde::de::Error,
2479 {
2480 self.visit_str(value)
2481 }
2482
2483 fn visit_string<E>(self, value: String) -> Result<Self::Value, E>
2484 where
2485 E: serde::de::Error,
2486 {
2487 self.budget.charge::<E>(value.capacity())?;
2488 Ok(serde_json::Value::String(value))
2489 }
2490
2491 fn visit_none<E>(self) -> Result<Self::Value, E> {
2492 Ok(serde_json::Value::Null)
2493 }
2494
2495 fn visit_unit<E>(self) -> Result<Self::Value, E> {
2496 Ok(serde_json::Value::Null)
2497 }
2498
2499 fn visit_some<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
2500 where
2501 D: serde::Deserializer<'de>,
2502 {
2503 deserializer.deserialize_any(self)
2504 }
2505
2506 fn visit_newtype_struct<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
2507 where
2508 D: serde::Deserializer<'de>,
2509 {
2510 deserializer.deserialize_any(self)
2511 }
2512
2513 fn visit_seq<A>(self, mut sequence: A) -> Result<Self::Value, A::Error>
2514 where
2515 A: serde::de::SeqAccess<'de>,
2516 {
2517 let mut values = Vec::new();
2518 while let Some(value) = sequence.next_element_seed(BudgetedJsonValueSeed {
2519 budget: self.budget,
2520 })? {
2521 values.push(value);
2522 }
2523 Ok(serde_json::Value::Array(values))
2524 }
2525
2526 fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
2527 where
2528 A: serde::de::MapAccess<'de>,
2529 {
2530 let mut values = serde_json::Map::new();
2531 while let Some(key) = map.next_key::<String>()? {
2532 self.budget
2533 .charge::<A::Error>(PAGINATED_OBJECT_ENTRY_BYTES.saturating_add(key.capacity()))?;
2534 let value = map.next_value_seed(BudgetedJsonValueSeed {
2535 budget: self.budget,
2536 })?;
2537 values.insert(key, value);
2538 }
2539 Ok(serde_json::Value::Object(values))
2540 }
2541}
2542
2543struct BudgetedJsonRowsSeed<'a, 'query> {
2544 budget: &'a mut PaginatedDecodeBudget<'query>,
2545}
2546
2547impl<'de> serde::de::DeserializeSeed<'de> for BudgetedJsonRowsSeed<'_, '_> {
2548 type Value = Vec<serde_json::Value>;
2549
2550 fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
2551 where
2552 D: serde::Deserializer<'de>,
2553 {
2554 deserializer.deserialize_seq(BudgetedJsonRowsVisitor {
2555 budget: self.budget,
2556 })
2557 }
2558}
2559
2560struct BudgetedJsonRowsVisitor<'a, 'query> {
2561 budget: &'a mut PaginatedDecodeBudget<'query>,
2562}
2563
2564impl<'de> serde::de::Visitor<'de> for BudgetedJsonRowsVisitor<'_, '_> {
2565 type Value = Vec<serde_json::Value>;
2566
2567 fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2568 formatter.write_str("a JSON array of SQL rows")
2569 }
2570
2571 fn visit_seq<A>(self, mut sequence: A) -> Result<Self::Value, A::Error>
2572 where
2573 A: serde::de::SeqAccess<'de>,
2574 {
2575 let mut rows = Vec::new();
2576 while let Some(row) = sequence.next_element_seed(BudgetedJsonValueSeed {
2577 budget: self.budget,
2578 })? {
2579 rows.push(row);
2580 }
2581 Ok(rows)
2582 }
2583}
2584
2585#[allow(clippy::too_many_arguments)]
2586async fn dispatch_paginated_sql(
2587 state: &AppState,
2588 output: ManagedQueryBatches,
2589 query_id: QueryId,
2590 owner: &str,
2591 pagination: ResolvedSqlPagination,
2592 output_limits: (usize, usize),
2593 test_hook: Option<mongreldb_query::SqlTestHook>,
2594 binding: sql_pages::SqlPageBinding,
2595) -> Response {
2596 let projection = pagination.projection;
2597 let serialization_projection = projection.clone();
2598 let serialization_batches = output.batches().to_vec();
2599 let serialization_query = output.query().clone();
2600 let response_test_hook = test_hook.clone();
2601 let retained_memory_limit = output_limits.1.min(state.sql_pages.max_retained_bytes());
2602 let serialized = tokio::task::spawn_blocking(move || {
2603 serialize_paginated_rows(
2604 &serialization_batches,
2605 &serialization_query,
2606 &serialization_projection,
2607 output_limits.0,
2608 output_limits.1,
2609 retained_memory_limit,
2610 test_hook.as_ref(),
2611 )
2612 })
2613 .await;
2614 let serialized = match serialized {
2615 Ok(result) => result,
2616 Err(error) => {
2617 output
2618 .query()
2619 .record_serialization_failure("SERIALIZATION_WORKER_FAILED");
2620 output.fail();
2621 state.metrics.inc_sql_errors();
2622 return terminal_server_error_response(
2623 state,
2624 query_id,
2625 StatusCode::INTERNAL_SERVER_ERROR,
2626 "SERIALIZATION_WORKER_FAILED",
2627 error.to_string(),
2628 );
2629 }
2630 };
2631 let serialized = match serialized {
2632 Ok(serialized) => serialized,
2633 Err(PaginatedSerializationError::Query(error)) => {
2634 output.fail();
2635 state.metrics.inc_sql_errors();
2636 return tracked_query_error_response(state, &error, Some(query_id));
2637 }
2638 Err(PaginatedSerializationError::Limit(message)) => {
2639 output.fail_result_limit();
2640 state.metrics.inc_sql_errors();
2641 return terminal_server_error_response(
2642 state,
2643 query_id,
2644 StatusCode::PAYLOAD_TOO_LARGE,
2645 "RESULT_LIMIT_EXCEEDED",
2646 message,
2647 );
2648 }
2649 Err(PaginatedSerializationError::Projection(message)) => {
2650 output.fail_with_error(
2651 "INVALID_SQL_PROJECTION",
2652 mongreldb_query::QueryTerminalErrorCategory::Execution,
2653 );
2654 state.metrics.inc_sql_errors();
2655 return terminal_server_error_response(
2656 state,
2657 query_id,
2658 StatusCode::BAD_REQUEST,
2659 "INVALID_SQL_PROJECTION",
2660 message,
2661 );
2662 }
2663 Err(PaginatedSerializationError::Encoding(message)) => {
2664 output.fail_serialization();
2665 state.metrics.inc_sql_errors();
2666 return terminal_server_error_response(
2667 state,
2668 query_id,
2669 StatusCode::INTERNAL_SERVER_ERROR,
2670 "SERIALIZATION_FAILED",
2671 message,
2672 );
2673 }
2674 };
2675 let retained_bytes = serialized.retained_bytes;
2676 let current_binding = sql_pages::SqlPageBinding {
2677 security_version: state.db.security_version(),
2678 catalog_epoch: state.db.catalog_snapshot().db_epoch,
2679 };
2680 if current_binding != binding {
2681 output.fail_with_error(
2682 "SQL_CURSOR_EXPIRED",
2683 mongreldb_query::QueryTerminalErrorCategory::Execution,
2684 );
2685 return with_query_id(
2686 sql_cursor_error_response(
2687 StatusCode::CONFLICT,
2688 "SQL_CURSOR_EXPIRED",
2689 "authorization or schema changed while retaining the SQL result",
2690 query_id,
2691 ),
2692 query_id,
2693 );
2694 }
2695 let retained = match state.sql_pages.insert(
2696 owner,
2697 serialized.rows,
2698 projection,
2699 pagination.limits,
2700 retained_bytes,
2701 binding,
2702 ) {
2703 Ok(retained) => retained,
2704 Err(sql_pages::InsertError::Full | sql_pages::InsertError::OwnerLimit) => {
2705 output.fail_with_error(
2706 "SQL_PAGE_STORE_FULL",
2707 mongreldb_query::QueryTerminalErrorCategory::Execution,
2708 );
2709 state.metrics.inc_sql_errors();
2710 return terminal_server_error_response(
2711 state,
2712 query_id,
2713 StatusCode::SERVICE_UNAVAILABLE,
2714 "SQL_PAGE_STORE_FULL",
2715 "retained SQL result capacity reached",
2716 );
2717 }
2718 Err(sql_pages::InsertError::EntropyUnavailable) => {
2719 output.fail_with_error(
2720 "ENTROPY_UNAVAILABLE",
2721 mongreldb_query::QueryTerminalErrorCategory::Execution,
2722 );
2723 state.metrics.inc_sql_errors();
2724 return terminal_server_error_response(
2725 state,
2726 query_id,
2727 StatusCode::INTERNAL_SERVER_ERROR,
2728 "ENTROPY_UNAVAILABLE",
2729 "OS CSPRNG unavailable",
2730 );
2731 }
2732 };
2733 let cursor_mac_key = match state.cursor_mac_key.get() {
2734 Ok(key) => key,
2735 Err(_) => {
2736 state.sql_pages.discard(&retained);
2737 output.fail_with_error(
2738 "ENTROPY_UNAVAILABLE",
2739 mongreldb_query::QueryTerminalErrorCategory::Execution,
2740 );
2741 state.metrics.inc_sql_errors();
2742 return terminal_server_error_response(
2743 state,
2744 query_id,
2745 StatusCode::INTERNAL_SERVER_ERROR,
2746 "ENTROPY_UNAVAILABLE",
2747 "OS CSPRNG unavailable",
2748 );
2749 }
2750 };
2751 let page = match sql_pages::SqlPageStore::first_page(&retained, &cursor_mac_key) {
2752 Ok(page) => page,
2753 Err(sql_pages::PageError::RowExceedsLimits) => {
2754 state.sql_pages.discard(&retained);
2755 output.fail_result_limit();
2756 state.metrics.inc_sql_errors();
2757 return terminal_server_error_response(
2758 state,
2759 query_id,
2760 StatusCode::PAYLOAD_TOO_LARGE,
2761 "RESULT_LIMIT_EXCEEDED",
2762 "one projected row exceeds the page byte or token limit",
2763 );
2764 }
2765 Err(sql_pages::PageError::OffsetInvalid) => {
2766 state.sql_pages.discard(&retained);
2767 output.fail_with_error(
2768 "INVALID_PAGE_OFFSET",
2769 mongreldb_query::QueryTerminalErrorCategory::Serialization,
2770 );
2771 state.metrics.inc_sql_errors();
2772 return terminal_server_error_response(
2773 state,
2774 query_id,
2775 StatusCode::INTERNAL_SERVER_ERROR,
2776 "INVALID_PAGE_OFFSET",
2777 "failed to create the first SQL result page",
2778 );
2779 }
2780 Err(sql_pages::PageError::Cancelled) => unreachable!("first-page rendering has no control"),
2781 };
2782 let discard_after_response = page.next_cursor.is_none();
2783 let page_byte_count = page.byte_count;
2784 let serialization_query = output.query().clone();
2787 let encoded_page = tokio::task::spawn_blocking(move || {
2788 serialize_sql_page_controlled(page, &serialization_query)
2789 })
2790 .await;
2791 let encoded_page = match encoded_page {
2792 Ok(Ok(encoded_page)) => encoded_page,
2793 Ok(Err(ControlledPageSerializationError::Query(error))) => {
2794 state.sql_pages.discard(&retained);
2795 output.fail();
2796 state.metrics.inc_sql_errors();
2797 return tracked_query_error_response(state, &error, Some(query_id));
2798 }
2799 Ok(Err(ControlledPageSerializationError::Encoding)) => {
2800 if let Err(error) = output.query().checkpoint() {
2804 state.sql_pages.discard(&retained);
2805 output.fail();
2806 state.metrics.inc_sql_errors();
2807 return tracked_query_error_response(state, &error, Some(query_id));
2808 }
2809 state.sql_pages.discard(&retained);
2810 output.fail_serialization();
2811 state.metrics.inc_sql_errors();
2812 return terminal_server_error_response(
2813 state,
2814 query_id,
2815 StatusCode::INTERNAL_SERVER_ERROR,
2816 "SERIALIZATION_FAILED",
2817 "failed to serialize the first SQL result page",
2818 );
2819 }
2820 Err(_) => {
2821 if let Err(error) = output.query().checkpoint() {
2822 state.sql_pages.discard(&retained);
2823 output.fail();
2824 state.metrics.inc_sql_errors();
2825 return tracked_query_error_response(state, &error, Some(query_id));
2826 }
2827 state.sql_pages.discard(&retained);
2828 output.fail_serialization();
2829 state.metrics.inc_sql_errors();
2830 return terminal_server_error_response(
2831 state,
2832 query_id,
2833 StatusCode::INTERNAL_SERVER_ERROR,
2834 "SERIALIZATION_WORKER_FAILED",
2835 "SQL page response serialization worker failed",
2836 );
2837 }
2838 };
2839 if let Some(hook) = response_test_hook {
2840 hook(mongreldb_query::SqlTestHookPoint::AfterPageResponseSerialization);
2841 }
2842 if let Err(error) = output.query().checkpoint() {
2843 state.sql_pages.discard(&retained);
2844 output.fail();
2847 state.metrics.inc_sql_errors();
2848 return tracked_query_error_response(state, &error, Some(query_id));
2849 }
2850 if let Err(error) = output.try_complete() {
2851 state.sql_pages.discard(&retained);
2852 state.metrics.inc_sql_errors();
2853 return tracked_query_error_response(state, &error, Some(query_id));
2854 }
2855 if discard_after_response {
2856 state.sql_pages.discard(&retained);
2857 }
2858 state.metrics.add_sql_output_bytes(page_byte_count);
2859 with_query_id(sql_page_response(encoded_page), query_id)
2860}
2861
2862fn serialize_paginated_rows(
2863 batches: &[arrow::record_batch::RecordBatch],
2864 query: &RegisteredSqlQuery,
2865 projection: &[String],
2866 max_rows: usize,
2867 max_bytes: usize,
2868 retained_memory_limit: usize,
2869 test_hook: Option<&mongreldb_query::SqlTestHook>,
2870) -> Result<SerializedPageRows, PaginatedSerializationError> {
2871 const ROW_CHECKPOINT_INTERVAL: usize = 256;
2872 if batches.is_empty() {
2873 let retained_bytes = sql_pages::accounted_bytes(2, &[], projection, || query.checkpoint())
2874 .map_err(PaginatedSerializationError::Query)?;
2875 if retained_bytes > retained_memory_limit {
2876 return Err(PaginatedSerializationError::Limit(
2877 PAGINATED_MEMORY_LIMIT_ERROR.into(),
2878 ));
2879 }
2880 return Ok(SerializedPageRows {
2881 rows: Vec::new(),
2882 retained_bytes,
2883 });
2884 }
2885 let fields = batches[0].schema();
2886 let mut indices = Vec::with_capacity(projection.len());
2887 for name in projection {
2888 let matches: Vec<_> = fields
2889 .fields()
2890 .iter()
2891 .enumerate()
2892 .filter_map(|(index, field)| (field.name() == name).then_some(index))
2893 .collect();
2894 if matches.len() != 1 {
2895 return Err(PaginatedSerializationError::Projection(format!(
2896 "projected output column {name:?} is missing or ambiguous"
2897 )));
2898 }
2899 indices.push(matches[0]);
2900 }
2901
2902 let mut writer_output = LimitedOutput::new(max_bytes);
2903 let mut rows = 0usize;
2904 let encoding = (|| {
2905 let mut writer = arrow::json::writer::ArrayWriter::new(&mut writer_output);
2906 for batch in batches {
2907 let batch = batch.project(&indices).map_err(|error| error.to_string())?;
2908 for offset in (0..batch.num_rows()).step_by(ROW_CHECKPOINT_INTERVAL) {
2909 if let Some(hook) = test_hook {
2910 hook(mongreldb_query::SqlTestHookPoint::BeforeSerializationBatch);
2911 }
2912 query.checkpoint().map_err(|error| error.to_string())?;
2913 let length = ROW_CHECKPOINT_INTERVAL.min(batch.num_rows() - offset);
2914 rows = rows.saturating_add(length);
2915 if rows > max_rows {
2916 return Err("SQL retained output row limit exceeded".into());
2917 }
2918 let slice = batch.slice(offset, length);
2919 writer
2920 .write_batches(&[&slice])
2921 .map_err(|error| error.to_string())?;
2922 }
2923 }
2924 writer.finish().map_err(|error| error.to_string())
2925 })();
2926 if let Err(error) = encoding {
2927 if let Err(query_error) = query.checkpoint() {
2928 return Err(PaginatedSerializationError::Query(query_error));
2929 }
2930 if writer_output.exceeded || rows > max_rows {
2931 return Err(PaginatedSerializationError::Limit(error));
2932 }
2933 return Err(PaginatedSerializationError::Encoding(error));
2934 }
2935 let bytes = writer_output.bytes.len();
2936 let retained_base = sql_pages::accounted_bytes(bytes, &[], projection, || query.checkpoint())
2937 .map_err(PaginatedSerializationError::Query)?;
2938 if retained_base > retained_memory_limit {
2939 return Err(PaginatedSerializationError::Limit(
2940 PAGINATED_MEMORY_LIMIT_ERROR.into(),
2941 ));
2942 }
2943 let reader = PaginatedJsonReader {
2944 cursor: std::io::Cursor::new(writer_output.bytes.as_slice()),
2945 query,
2946 test_hook,
2947 };
2948 let mut deserializer =
2949 serde_json::Deserializer::from_reader(std::io::BufReader::with_capacity(64 * 1024, reader));
2950 let mut budget = PaginatedDecodeBudget {
2951 used: retained_base,
2952 limit: retained_memory_limit,
2953 nodes: 0,
2954 exceeded: false,
2955 query,
2956 test_hook,
2957 };
2958 let rows = match serde::de::DeserializeSeed::deserialize(
2959 BudgetedJsonRowsSeed {
2960 budget: &mut budget,
2961 },
2962 &mut deserializer,
2963 ) {
2964 Ok(rows) => {
2965 if let Err(error) = deserializer.end() {
2966 return Err(PaginatedSerializationError::Encoding(error.to_string()));
2967 }
2968 rows
2969 }
2970 Err(error) => {
2971 if let Err(query_error) = query.checkpoint() {
2972 return Err(PaginatedSerializationError::Query(query_error));
2973 }
2974 if budget.exceeded {
2975 return Err(PaginatedSerializationError::Limit(
2976 PAGINATED_MEMORY_LIMIT_ERROR.into(),
2977 ));
2978 }
2979 return Err(PaginatedSerializationError::Encoding(error.to_string()));
2980 }
2981 };
2982 if let Some(hook) = test_hook {
2983 hook(mongreldb_query::SqlTestHookPoint::AfterSerialization);
2984 }
2985 let retained_bytes =
2986 sql_pages::accounted_bytes(bytes, &rows, projection, || query.checkpoint())
2987 .map_err(PaginatedSerializationError::Query)?;
2988 if retained_bytes > retained_memory_limit {
2989 return Err(PaginatedSerializationError::Limit(
2990 PAGINATED_MEMORY_LIMIT_ERROR.into(),
2991 ));
2992 }
2993 Ok(SerializedPageRows {
2994 rows,
2995 retained_bytes,
2996 })
2997}
2998
2999enum ControlledPageSerializationError {
3000 Query(mongreldb_query::MongrelQueryError),
3001 Encoding,
3002}
3003
3004struct ControlledPageWriter<'a> {
3005 bytes: Vec<u8>,
3006 query: &'a RegisteredSqlQuery,
3007}
3008
3009impl std::io::Write for ControlledPageWriter<'_> {
3010 fn write(&mut self, buffer: &[u8]) -> std::io::Result<usize> {
3011 self.query
3012 .checkpoint()
3013 .map_err(|error| std::io::Error::other(error.to_string()))?;
3014 self.bytes.extend_from_slice(buffer);
3015 Ok(buffer.len())
3016 }
3017
3018 fn flush(&mut self) -> std::io::Result<()> {
3019 Ok(())
3020 }
3021}
3022
3023fn serialize_sql_page_controlled(
3024 page: sql_pages::SqlPage,
3025 query: &RegisteredSqlQuery,
3026) -> Result<Vec<u8>, ControlledPageSerializationError> {
3027 query
3028 .checkpoint()
3029 .map_err(ControlledPageSerializationError::Query)?;
3030 let value = json!({
3031 "status": "completed",
3032 "rows": page.rows,
3033 "next_cursor": page.next_cursor,
3034 "page": {
3035 "offset": page.offset,
3036 "row_count": page.row_count,
3037 "total_rows": page.total_rows,
3038 "byte_count": page.byte_count,
3039 "estimated_tokens": page.estimated_tokens,
3040 "limits": page.limits,
3041 "projection": page.projection,
3042 "expires_at_ms": page.expires_at_ms,
3043 "snapshot": "retained_result",
3044 "token_estimate": "ceil(projected_json_bytes/4)",
3045 }
3046 });
3047 let mut writer = ControlledPageWriter {
3048 bytes: Vec::new(),
3049 query,
3050 };
3051 if let Err(error) = serde_json::to_writer(&mut writer, &value) {
3052 if let Err(query_error) = query.checkpoint() {
3055 return Err(ControlledPageSerializationError::Query(query_error));
3056 }
3057 let _ = error;
3058 return Err(ControlledPageSerializationError::Encoding);
3059 }
3060 query
3061 .checkpoint()
3062 .map_err(ControlledPageSerializationError::Query)?;
3063 Ok(writer.bytes)
3064}
3065
3066fn sql_page_response(body: Vec<u8>) -> Response {
3067 ([(header::CONTENT_TYPE, "application/json")], body).into_response()
3068}
3069
3070fn validate_stmt_name(name: &str) -> Result<(), String> {
3074 let mut chars = name.chars();
3075 match chars.next() {
3076 Some(c) if c.is_ascii_alphabetic() || c == '_' => {}
3077 _ => return Err("statement name must start with a letter or underscore".into()),
3078 }
3079 if !chars.all(|c| c.is_ascii_alphanumeric() || c == '_') {
3080 return Err("statement name may contain only letters, digits, or underscore".into());
3081 }
3082 Ok(())
3083}
3084
3085fn render_sql_literal(v: &serde_json::Value) -> Result<String, String> {
3090 match v {
3091 serde_json::Value::Null => Ok("NULL".into()),
3092 serde_json::Value::Bool(b) => {
3093 if *b {
3094 Ok("TRUE".into())
3095 } else {
3096 Ok("FALSE".into())
3097 }
3098 }
3099 serde_json::Value::Number(n) => Ok(n.to_string()),
3100 serde_json::Value::String(s) => {
3102 let mut out = String::with_capacity(s.len() + 2);
3103 out.push('\'');
3104 for c in s.chars() {
3105 if c == '\'' {
3106 out.push_str("''");
3107 } else {
3108 out.push(c);
3109 }
3110 }
3111 out.push('\'');
3112 Ok(out)
3113 }
3114 _ => Err("prepared-statement parameters must be scalar (null/bool/number/string)".into()),
3116 }
3117}
3118
3119fn validate_prepared_params(
3124 entry: &sessions::SessionEntry,
3125 name: &str,
3126 binding: &mut mongreldb_protocol::prepared::PreparedStatementBinding,
3127 params: &[serde_json::Value],
3128) -> Result<(), Box<Response>> {
3129 let provided = prepared::parameter_type_names(params);
3130 if binding.parameter_types.is_empty() {
3131 if !provided.is_empty() {
3132 binding.parameter_types = provided;
3133 entry.insert_prepared_binding(name.to_owned(), binding.clone());
3134 }
3135 return Ok(());
3136 }
3137 if binding.parameter_types != provided {
3138 return Err(Box::new(structured_category_error_response(
3139 StatusCode::BAD_REQUEST,
3140 "PREPARED_PARAMETER_MISMATCH",
3141 format!(
3142 "prepared statement {name:?} expects parameters of types {:?}, got {provided:?}",
3143 binding.parameter_types
3144 ),
3145 mongreldb_types::errors::ErrorCategory::ClusterVersionMismatch,
3146 )));
3147 }
3148 Ok(())
3149}
3150
3151#[derive(Deserialize)]
3152struct PrepareRequest {
3153 name: String,
3154 sql: String,
3155 #[serde(default)]
3159 param_types: Option<Vec<String>>,
3160 #[serde(default)]
3161 query_id: Option<QueryId>,
3162 #[serde(default)]
3163 timeout_ms: Option<u64>,
3164}
3165
3166async fn prepare_statement(
3175 State(state): State<Arc<AppState>>,
3176 OptionalPrincipal(principal): OptionalPrincipal,
3177 Path(id): Path<String>,
3178 headers: axum::http::HeaderMap,
3179 Json(req): Json<PrepareRequest>,
3180) -> Response {
3181 if !state.accepting_sql.load(Ordering::Acquire) {
3182 return (StatusCode::SERVICE_UNAVAILABLE, "server is shutting down").into_response();
3183 }
3184 if !request_identity_is_current(&state, &principal) {
3185 return StatusCode::NOT_FOUND.into_response();
3186 }
3187 if let Err(msg) = validate_stmt_name(&req.name) {
3188 return (StatusCode::BAD_REQUEST, msg).into_response();
3189 }
3190 if let Some(param_types) = req.param_types.as_deref() {
3191 if let Err(msg) = prepared::validate_parameter_type_names(param_types) {
3192 return (StatusCode::BAD_REQUEST, msg).into_response();
3193 }
3194 }
3195 let owner = request_owner(&state, &principal);
3196 let Some(entry) = state.sessions.get(&id, &owner) else {
3197 return (
3198 StatusCode::NOT_FOUND,
3199 "session not found or not owned by caller",
3200 )
3201 .into_response();
3202 };
3203 let (options, query_id) = match resolve_query_options(
3204 &state,
3205 &headers,
3206 req.query_id,
3207 req.timeout_ms,
3208 owner,
3209 Some(id),
3210 ) {
3211 Ok(options) => options,
3212 Err(response) => return *response,
3213 };
3214 let query = match register_controlled_query(&state, &entry.session(), options) {
3215 Ok(query) => query,
3216 Err(error) => return tracked_query_error_response(&state, &error, Some(query_id)),
3217 };
3218 let registration = RegisteredQueryGuard::new(query);
3219 if mongreldb_query::contains_boolean_ai_predicate(&req.sql) {
3220 registration.fail();
3221 return with_query_id(remote_boolean_ai_error(), query_id);
3222 }
3223 let _sql_permit = match acquire_sql_permit(&state, &entry.session(), registration.query()).await
3224 {
3225 Ok(permit) => permit,
3226 Err(error) => return tracked_query_error_response(&state, &error, Some(query_id)),
3227 };
3228 let _guard = tokio::select! {
3229 guard = entry.lock.lock() => guard,
3230 _ = registration.query().control().cancelled() => {
3231 return tracked_query_error_response(
3232 &state,
3233 &cancellation_checkpoint_error(registration.query()),
3234 Some(query_id),
3235 );
3236 }
3237 };
3238 if entry.is_closed() {
3239 return with_query_id(
3240 (StatusCode::NOT_FOUND, "session no longer available").into_response(),
3241 query_id,
3242 );
3243 }
3244 entry.touch();
3245 let sql = format!("PREPARE {} AS {}", req.name, req.sql);
3246 let query = registration.into_query();
3247 match entry.session().run_with_query(&sql, query).await {
3248 Ok(_) => {
3249 let catalog_state = prepared::CatalogState::capture(&state.db);
3251 let statement_id = entry.allocate_statement_id();
3252 entry.insert_prepared_binding(
3253 req.name.clone(),
3254 prepared::build_binding(
3255 statement_id,
3256 req.sql.clone(),
3257 req.param_types.clone().unwrap_or_default(),
3258 &catalog_state,
3259 ),
3260 );
3261 with_query_id(
3262 Json(json!({ "prepared": req.name, "statement_id": statement_id.get() }))
3263 .into_response(),
3264 query_id,
3265 )
3266 }
3267 Err(error) => tracked_query_error_response(&state, &error, Some(query_id)),
3268 }
3269}
3270
3271#[derive(Deserialize)]
3272struct ExecuteRequest {
3273 name: String,
3274 params: Vec<serde_json::Value>,
3275 #[serde(default)]
3276 format: Option<String>,
3277 #[serde(default)]
3278 query_id: Option<QueryId>,
3279 #[serde(default)]
3280 timeout_ms: Option<u64>,
3281}
3282
3283async fn execute_statement(
3287 State(state): State<Arc<AppState>>,
3288 OptionalPrincipal(principal): OptionalPrincipal,
3289 Path(id): Path<String>,
3290 headers: axum::http::HeaderMap,
3291 Json(req): Json<ExecuteRequest>,
3292) -> Response {
3293 if !state.accepting_sql.load(Ordering::Acquire) {
3294 return (StatusCode::SERVICE_UNAVAILABLE, "server is shutting down").into_response();
3295 }
3296 if !request_identity_is_current(&state, &principal) {
3297 return StatusCode::NOT_FOUND.into_response();
3298 }
3299 if let Err(msg) = validate_stmt_name(&req.name) {
3300 return (StatusCode::BAD_REQUEST, msg).into_response();
3301 }
3302 let owner = request_owner(&state, &principal);
3303 let Some(entry) = state.sessions.get(&id, &owner) else {
3304 return (
3305 StatusCode::NOT_FOUND,
3306 "session not found or not owned by caller",
3307 )
3308 .into_response();
3309 };
3310 let (options, query_id) = match resolve_query_options(
3311 &state,
3312 &headers,
3313 req.query_id,
3314 req.timeout_ms,
3315 owner,
3316 Some(id),
3317 ) {
3318 Ok(options) => options,
3319 Err(response) => return *response,
3320 };
3321 let query = match register_controlled_query(&state, &entry.session(), options) {
3322 Ok(query) => query,
3323 Err(error) => return tracked_query_error_response(&state, &error, Some(query_id)),
3324 };
3325 let registration = RegisteredQueryGuard::new(query);
3326 let sql_permit = match acquire_sql_permit(&state, &entry.session(), registration.query()).await
3327 {
3328 Ok(permit) => permit,
3329 Err(error) => return tracked_query_error_response(&state, &error, Some(query_id)),
3330 };
3331 let _guard = tokio::select! {
3332 guard = entry.lock.lock() => guard,
3333 _ = registration.query().control().cancelled() => {
3334 return tracked_query_error_response(
3335 &state,
3336 &cancellation_checkpoint_error(registration.query()),
3337 Some(query_id),
3338 );
3339 }
3340 };
3341 if entry.is_closed() {
3342 return with_query_id(
3343 (StatusCode::NOT_FOUND, "session no longer available").into_response(),
3344 query_id,
3345 );
3346 }
3347 entry.touch();
3348 if let Some((statement_id, mut binding)) = entry.prepared_binding(&req.name) {
3352 if let Err(response) =
3353 validate_prepared_params(&entry, &req.name, &mut binding, &req.params)
3354 {
3355 registration.fail();
3356 state.metrics.inc_sql_errors();
3357 return with_query_id(*response, query_id);
3358 }
3359 let catalog_state = prepared::CatalogState::capture(&state.db);
3360 if !catalog_state.is_compatible(&binding) {
3361 state.audit.record(
3366 request_owner(&state, &principal),
3367 "prepared.invalidate",
3368 format!(
3369 "statement {:?} invalidated by a catalog/schema change",
3370 req.name
3371 ),
3372 );
3373 let deallocate = format!("DEALLOCATE {}", req.name);
3374 let _ = entry.session().run(&deallocate).await;
3375 let replan_error: Option<String> = 'replan: {
3376 if entry.session().staged_sql_operation_count().is_some() {
3379 Some(
3380 "the session has an open transaction; COMMIT or ROLLBACK, then re-prepare"
3381 .to_owned(),
3382 )
3383 } else {
3384 let fresh = match MongrelSession::open_with_external_modules_as(
3391 Arc::clone(&state.db),
3392 state.external_modules.iter().cloned(),
3393 request_principal(&state, &principal),
3394 ) {
3395 Ok(session) => {
3396 session.with_query_registry(Arc::clone(&state.query_registry))
3397 }
3398 Err(error) => break 'replan Some(error.to_string()),
3399 };
3400 fresh.set_test_hook(entry.session().sql_test_hook());
3402 let replan = format!("PREPARE {} AS {}", req.name, binding.sql);
3403 match fresh.run(&replan).await {
3404 Ok(_) => {
3405 entry.replace_session(fresh);
3406 let rebound = prepared::build_binding(
3407 statement_id,
3408 binding.sql.clone(),
3409 binding.parameter_types.clone(),
3410 &prepared::CatalogState::capture(&state.db),
3411 );
3412 entry.insert_prepared_binding(req.name.clone(), rebound);
3413 None
3414 }
3415 Err(error) => Some(error.to_string()),
3416 }
3417 }
3418 };
3419 match &replan_error {
3420 Some(_) => state.audit.record(
3421 request_owner(&state, &principal),
3422 "prepared.replan.fail",
3423 format!(
3424 "statement {:?} could not be replanned; re-prepare required",
3425 req.name
3426 ),
3427 ),
3428 None => state.audit.record(
3429 request_owner(&state, &principal),
3430 "prepared.replan.ok",
3431 format!(
3432 "statement {:?} replanned against the current catalog",
3433 req.name
3434 ),
3435 ),
3436 }
3437 if let Some(error) = replan_error {
3438 entry.remove_prepared_binding(&req.name);
3439 registration.fail();
3440 state.metrics.inc_sql_errors();
3441 return with_query_id(
3442 structured_category_error_response(
3443 StatusCode::CONFLICT,
3444 "SCHEMA_VERSION_MISMATCH",
3445 format!(
3446 "prepared statement {:?} was invalidated by a catalog/schema change and could not be replanned: {error}; re-prepare the statement",
3447 req.name
3448 ),
3449 mongreldb_types::errors::ErrorCategory::SchemaVersionMismatch,
3450 ),
3451 query_id,
3452 );
3453 }
3454 }
3455 }
3456 state.metrics.inc_sql_queries();
3457 let literals: Vec<String> = match req
3458 .params
3459 .iter()
3460 .map(render_sql_literal)
3461 .collect::<Result<_, _>>()
3462 {
3463 Ok(v) => v,
3464 Err(msg) => {
3465 state.metrics.inc_sql_errors();
3466 return (StatusCode::BAD_REQUEST, msg).into_response();
3467 }
3468 };
3469 let sql = format!("EXECUTE {}({})", req.name, literals.join(", "));
3470 let start = std::time::Instant::now();
3471 let result = if req.format.as_deref() == Some("arrow-stream") {
3472 let query = registration.into_query();
3473 match entry
3474 .session()
3475 .run_stream_with_query_for_serialization(&sql, query)
3476 .await
3477 {
3478 Ok((stream, completion)) => Ok(sql_arrow_stream_response_controlled(
3479 stream,
3480 completion,
3481 sql_permit,
3482 (
3483 state.reloadable.sql_max_output_rows.get(),
3484 state.reloadable.sql_max_output_bytes.get(),
3485 ),
3486 &state,
3487 query_id,
3488 entry.session().sql_test_hook(),
3489 )),
3490 Err(error) => Err(error),
3491 }
3492 } else {
3493 let query = registration.into_query();
3494 match entry
3495 .session()
3496 .run_with_query_for_serialization_with_limits(
3497 &sql,
3498 query,
3499 mongreldb_query::SqlCollectionLimits::new(
3500 state.reloadable.sql_max_output_rows.get(),
3501 state.reloadable.sql_max_output_bytes.get(),
3502 ),
3503 )
3504 .await
3505 {
3506 Ok(output) => Ok(dispatch_buffered_sql_format(
3507 &state,
3508 req.format.as_deref(),
3509 output,
3510 query_id,
3511 entry.session().sql_test_hook(),
3512 (
3513 state.reloadable.sql_max_output_rows.get(),
3514 state.reloadable.sql_max_output_bytes.get(),
3515 ),
3516 )
3517 .await),
3518 Err(error) => Err(error),
3519 }
3520 };
3521 let elapsed = start.elapsed();
3522 if elapsed >= state.reloadable.slow_query_threshold.get() {
3523 state.metrics.inc_slow_queries();
3524 eprintln!(
3525 "[slow-query] {}\u{00b5}s \u{2014} EXECUTE {}",
3526 elapsed.as_micros(),
3527 req.name
3528 );
3529 }
3530 match result {
3531 Ok(response) => with_query_id(response, query_id),
3532 Err(e) => {
3533 state.metrics.inc_sql_errors();
3534 let msg = format!("{e}");
3536 let status = if msg.contains("does not exist") {
3537 StatusCode::NOT_FOUND
3538 } else {
3539 status_for_query_error(&e)
3540 };
3541 if status == status_for_query_error(&e) {
3542 tracked_query_error_response(&state, &e, Some(query_id))
3543 } else {
3544 with_query_id(
3545 (status, format!("{msg} ({}µs)", elapsed.as_micros())).into_response(),
3546 query_id,
3547 )
3548 }
3549 }
3550 }
3551}
3552
3553async fn deallocate_statement(
3556 State(state): State<Arc<AppState>>,
3557 OptionalPrincipal(principal): OptionalPrincipal,
3558 Path((id, name)): Path<(String, String)>,
3559 headers: axum::http::HeaderMap,
3560) -> Response {
3561 if !state.accepting_sql.load(Ordering::Acquire) {
3562 return (StatusCode::SERVICE_UNAVAILABLE, "server is shutting down").into_response();
3563 }
3564 if !request_identity_is_current(&state, &principal) {
3565 return StatusCode::NOT_FOUND.into_response();
3566 }
3567 if let Err(msg) = validate_stmt_name(&name) {
3568 return (StatusCode::BAD_REQUEST, msg).into_response();
3569 }
3570 let owner = request_owner(&state, &principal);
3571 let Some(entry) = state.sessions.get(&id, &owner) else {
3572 return (
3573 StatusCode::NOT_FOUND,
3574 "session not found or not owned by caller",
3575 )
3576 .into_response();
3577 };
3578 let (options, query_id) =
3579 match resolve_query_options(&state, &headers, None, None, owner, Some(id)) {
3580 Ok(options) => options,
3581 Err(response) => return *response,
3582 };
3583 let query = match register_controlled_query(&state, &entry.session(), options) {
3584 Ok(query) => query,
3585 Err(error) => return tracked_query_error_response(&state, &error, Some(query_id)),
3586 };
3587 let registration = RegisteredQueryGuard::new(query);
3588 let _sql_permit = match acquire_sql_permit(&state, &entry.session(), registration.query()).await
3589 {
3590 Ok(permit) => permit,
3591 Err(error) => return tracked_query_error_response(&state, &error, Some(query_id)),
3592 };
3593 let _guard = tokio::select! {
3594 guard = entry.lock.lock() => guard,
3595 _ = registration.query().control().cancelled() => {
3596 return tracked_query_error_response(
3597 &state,
3598 &cancellation_checkpoint_error(registration.query()),
3599 Some(query_id),
3600 );
3601 }
3602 };
3603 if entry.is_closed() {
3604 return with_query_id(
3605 (StatusCode::NOT_FOUND, "session no longer available").into_response(),
3606 query_id,
3607 );
3608 }
3609 entry.touch();
3610 state.metrics.inc_sql_queries();
3611 let sql = format!("DEALLOCATE {name}");
3612 let start = std::time::Instant::now();
3613 let result = entry
3614 .session()
3615 .run_with_query(&sql, registration.into_query())
3616 .await;
3617 let elapsed = start.elapsed();
3618 if elapsed >= state.reloadable.slow_query_threshold.get() {
3619 state.metrics.inc_slow_queries();
3620 eprintln!(
3621 "[slow-query] {}\u{00b5}s query_id={} operation=DEALLOCATE",
3622 elapsed.as_micros(),
3623 query_id,
3624 );
3625 }
3626 match result {
3627 Ok(_) => {
3628 entry.remove_prepared_binding(&name);
3629 with_query_id(
3630 Json(json!({ "deallocated": name })).into_response(),
3631 query_id,
3632 )
3633 }
3634 Err(error) => {
3635 state.metrics.inc_sql_errors();
3636 tracked_query_error_response(&state, &error, Some(query_id))
3637 }
3638 }
3639}
3640
3641async fn metrics_handler(
3644 State(state): State<Arc<AppState>>,
3645 OptionalPrincipal(principal): OptionalPrincipal,
3646) -> Response {
3647 if let Err(error) = state.db.require_for(
3648 request_principal(&state, &principal).as_ref(),
3649 &mongreldb_core::Permission::Admin,
3650 ) {
3651 return (status_for_error(&error), error.to_string()).into_response();
3652 }
3653 let body = state.metrics.prometheus_text(
3654 state.db.table_names().len(),
3655 state.query_registry.stats(),
3656 (
3657 state.pre_cancellations.len(),
3658 state.pre_cancellations.approximate_bytes(),
3659 ),
3660 );
3661 (
3662 [(
3663 header::CONTENT_TYPE,
3664 "text/plain; version=0.0.4; charset=utf-8".to_string(),
3665 )],
3666 body,
3667 )
3668 .into_response()
3669}
3670
3671async fn compact_all(
3673 State(state): State<Arc<AppState>>,
3674 OptionalPrincipal(principal): OptionalPrincipal,
3675) -> (StatusCode, Json<serde_json::Value>) {
3676 if let Err(error) = state.db.require_for(
3677 request_principal(&state, &principal).as_ref(),
3678 &mongreldb_core::Permission::Ddl,
3679 ) {
3680 return (
3681 status_for_error(&error),
3682 Json(json!({ "status": "error", "message": error.to_string() })),
3683 );
3684 }
3685 match state.db.compact() {
3686 Ok((compacted, skipped)) => (
3687 StatusCode::OK,
3688 Json(json!({
3689 "status": "ok",
3690 "compacted": compacted,
3691 "skipped": skipped,
3692 })),
3693 ),
3694 Err(e) => (
3695 StatusCode::INTERNAL_SERVER_ERROR,
3696 Json(json!({ "status": "error", "message": format!("{e}") })),
3697 ),
3698 }
3699}
3700
3701async fn compact_table(
3703 State(state): State<Arc<AppState>>,
3704 OptionalPrincipal(principal): OptionalPrincipal,
3705 Path(name): Path<String>,
3706) -> (StatusCode, Json<serde_json::Value>) {
3707 if let Err(error) = state.db.require_for(
3708 request_principal(&state, &principal).as_ref(),
3709 &mongreldb_core::Permission::Ddl,
3710 ) {
3711 return (
3712 status_for_error(&error),
3713 Json(json!({ "status": "error", "table": name, "message": error.to_string() })),
3714 );
3715 }
3716 match state.db.compact_table(&name) {
3717 Ok(true) => (
3718 StatusCode::OK,
3719 Json(json!({ "status": "compacted", "table": name })),
3720 ),
3721 Ok(false) => (
3722 StatusCode::OK,
3723 Json(json!({ "status": "skipped", "table": name, "reason": "fewer than 2 runs" })),
3724 ),
3725 Err(e) => (
3726 StatusCode::INTERNAL_SERVER_ERROR,
3727 Json(json!({ "status": "error", "table": name, "message": format!("{e}") })),
3728 ),
3729 }
3730}
3731
3732#[derive(Deserialize)]
3733struct CreateTableRequest {
3734 name: String,
3735 columns: Vec<ColumnDefJson>,
3736}
3737
3738#[derive(Deserialize)]
3739struct ColumnDefJson {
3740 id: u16,
3741 name: String,
3742 ty: String,
3743 primary_key: bool,
3744 #[serde(default)]
3745 nullable: bool,
3746}
3747
3748async fn create_table(
3749 State(state): State<Arc<AppState>>,
3750 OptionalPrincipal(principal): OptionalPrincipal,
3751 Json(req): Json<CreateTableRequest>,
3752) -> Response {
3753 if let Some(response) = require_writes_open(&state) {
3754 return response;
3755 }
3756 if let Err(error) = state.db.require_for(
3757 request_principal(&state, &principal).as_ref(),
3758 &mongreldb_core::Permission::Ddl,
3759 ) {
3760 return (status_for_error(&error), error.to_string()).into_response();
3761 }
3762 let mut columns = Vec::new();
3763 for c in &req.columns {
3764 let ty = match c.ty.as_str() {
3765 "int64" | "bigint" => TypeId::Int64,
3766 "float64" | "double" => TypeId::Float64,
3767 "bytes" | "varchar" | "text" => TypeId::Bytes,
3768 "bool" => TypeId::Bool,
3769 other => {
3770 return (StatusCode::BAD_REQUEST, format!("unknown type: {other}")).into_response()
3771 }
3772 };
3773 let mut flags = mongreldb_core::schema::ColumnFlags::empty();
3774 if c.primary_key {
3775 flags = flags.with(mongreldb_core::schema::ColumnFlags::PRIMARY_KEY);
3776 }
3777 if c.nullable {
3778 flags = flags.with(mongreldb_core::schema::ColumnFlags::NULLABLE);
3779 }
3780 columns.push(mongreldb_core::schema::ColumnDef {
3781 id: c.id,
3782 name: c.name.clone(),
3783 ty,
3784 flags,
3785 default_value: None,
3786 embedding_source: None,
3787 });
3788 }
3789 let schema = Schema {
3790 schema_id: 0,
3791 columns,
3792 indexes: vec![],
3793 colocation: vec![],
3794 constraints: Default::default(),
3795 clustered: false,
3796 };
3797 if let Err(msg) = validate_table_name(&req.name) {
3798 return (StatusCode::BAD_REQUEST, msg).into_response();
3799 }
3800 match state.db.create_table(&req.name, schema) {
3801 Ok(id) => Json(json!({
3802 "table_id": id,
3803 "table_id_text": id.to_string()
3804 }))
3805 .into_response(),
3806 Err(error) => crate::kit::durable_core_error_response(&error)
3807 .unwrap_or_else(|| (status_for_error(&error), error.to_string()).into_response()),
3808 }
3809}
3810
3811async fn list_tables(
3812 State(state): State<Arc<AppState>>,
3813 OptionalPrincipal(principal): OptionalPrincipal,
3814) -> Json<Vec<String>> {
3815 let principal = request_principal(&state, &principal);
3816 Json(
3817 state
3818 .db
3819 .table_names()
3820 .into_iter()
3821 .filter(|table| {
3822 state
3823 .db
3824 .select_column_ids_for(table, principal.as_ref())
3825 .is_ok()
3826 })
3827 .collect(),
3828 )
3829}
3830
3831async fn drop_table(
3832 State(state): State<Arc<AppState>>,
3833 OptionalPrincipal(principal): OptionalPrincipal,
3834 Path(name): Path<String>,
3835) -> Response {
3836 if let Some(response) = require_writes_open(&state) {
3837 return response;
3838 }
3839 if let Err(error) = state.db.require_for(
3840 request_principal(&state, &principal).as_ref(),
3841 &mongreldb_core::Permission::Ddl,
3842 ) {
3843 return (status_for_error(&error), error.to_string()).into_response();
3844 }
3845 match state.db.drop_table_with_epoch(&name) {
3846 Ok(epoch) => Json(json!({
3847 "status": "committed",
3848 "epoch": epoch.0,
3849 "epoch_text": epoch.0.to_string()
3850 }))
3851 .into_response(),
3852 Err(error) => crate::kit::durable_core_error_response(&error)
3853 .unwrap_or_else(|| (status_for_error(&error), error.to_string()).into_response()),
3854 }
3855}
3856
3857#[derive(Deserialize)]
3858struct PutRequest {
3859 row: Vec<serde_json::Value>,
3860}
3861
3862pub(crate) fn json_to_value(v: &serde_json::Value, expected: &TypeId) -> Value {
3863 match (v, expected) {
3864 (serde_json::Value::Number(n), TypeId::Float64) => {
3865 n.as_f64().map(Value::Float64).unwrap_or(Value::Null)
3866 }
3867 (serde_json::Value::Number(n), TypeId::Int64) => {
3868 n.as_i64().map(Value::Int64).unwrap_or(Value::Null)
3869 }
3870 (serde_json::Value::String(s), TypeId::Bytes) => Value::Bytes(s.as_bytes().to_vec()),
3871 (serde_json::Value::String(s), TypeId::Enum { variants }) => {
3872 if variants.iter().any(|v| v == s) {
3873 Value::Bytes(s.as_bytes().to_vec())
3874 } else {
3875 Value::Null
3876 }
3877 }
3878 (serde_json::Value::Bool(b), TypeId::Bool) => Value::Bool(*b),
3879 (serde_json::Value::Array(arr), TypeId::Embedding { dim }) => {
3882 if arr.len() as u32 != *dim {
3883 return Value::Null;
3884 }
3885 let vec: Option<Vec<f32>> =
3886 arr.iter().map(|el| el.as_f64().map(|f| f as f32)).collect();
3887 vec.map(Value::Embedding).unwrap_or(Value::Null)
3888 }
3889 (serde_json::Value::Null, _) => Value::Null,
3890 (serde_json::Value::Number(n), _) => {
3892 if let Some(i) = n.as_i64() {
3893 Value::Int64(i)
3894 } else if let Some(f) = n.as_f64() {
3895 Value::Float64(f)
3896 } else {
3897 Value::Null
3898 }
3899 }
3900 (serde_json::Value::String(s), _) => Value::Bytes(s.as_bytes().to_vec()),
3901 (serde_json::Value::Bool(b), _) => Value::Bool(*b),
3902 _ => Value::Null,
3903 }
3904}
3905
3906fn legacy_json_to_value(value: &serde_json::Value, expected: &TypeId) -> Result<Value, String> {
3907 if value.is_null() {
3908 return Ok(Value::Null);
3909 }
3910 match expected {
3911 TypeId::Bool => value
3912 .as_bool()
3913 .map(Value::Bool)
3914 .ok_or_else(|| "expected a boolean".into()),
3915 TypeId::Int8
3916 | TypeId::Int16
3917 | TypeId::Int32
3918 | TypeId::Int64
3919 | TypeId::UInt8
3920 | TypeId::UInt16
3921 | TypeId::UInt32
3922 | TypeId::UInt64
3923 | TypeId::TimestampNanos
3924 | TypeId::Date32
3925 | TypeId::Date64
3926 | TypeId::Time64 => value
3927 .as_i64()
3928 .map(Value::Int64)
3929 .ok_or_else(|| "expected a signed 64-bit integer".into()),
3930 TypeId::Float32 | TypeId::Float64 => value
3931 .as_f64()
3932 .filter(|value| value.is_finite())
3933 .map(Value::Float64)
3934 .ok_or_else(|| "expected a finite number".into()),
3935 TypeId::Bytes => match value.as_str() {
3936 Some(value) => Ok(Value::Bytes(value.as_bytes().to_vec())),
3937 None => decode_tagged_hex(value, "bytes").map(Value::Bytes),
3938 },
3939 TypeId::Enum { variants } => {
3940 let bytes = match value.as_str() {
3941 Some(value) => value.as_bytes().to_vec(),
3942 None => decode_tagged_hex(value, "bytes")?,
3943 };
3944 let value =
3945 std::str::from_utf8(&bytes).map_err(|_| "enum variant is not UTF-8".to_string())?;
3946 if !variants.iter().any(|variant| variant == value) {
3947 return Err("expected a declared enum variant".into());
3948 }
3949 Ok(Value::Bytes(bytes))
3950 }
3951 TypeId::Embedding { dim } => {
3952 let values = value
3953 .as_array()
3954 .ok_or_else(|| "expected an embedding array".to_string())?;
3955 if values.len() != *dim as usize {
3956 return Err(format!("expected an embedding with {dim} values"));
3957 }
3958 values
3959 .iter()
3960 .map(|value| {
3961 value
3962 .as_f64()
3963 .map(|value| value as f32)
3964 .filter(|value| value.is_finite())
3965 .ok_or_else(|| "embedding values must be finite numbers".to_string())
3966 })
3967 .collect::<Result<Vec<_>, _>>()
3968 .map(Value::Embedding)
3969 }
3970 TypeId::Decimal128 { .. } => {
3971 let object = exact_tagged_object(value, "decimal", &["unscaled"])?;
3972 let text = object["unscaled"]
3973 .as_str()
3974 .ok_or_else(|| "decimal unscaled value must be a string".to_string())?;
3975 let value = text
3976 .parse::<i128>()
3977 .map_err(|_| "decimal unscaled value is invalid".to_string())?;
3978 if value.to_string() != text {
3979 return Err("decimal unscaled value is not canonical".into());
3980 }
3981 Ok(Value::Decimal(value))
3982 }
3983 TypeId::Interval => {
3984 let object = exact_tagged_object(value, "interval", &["months", "days", "nanos"])?;
3985 let months = canonical_i64(&object["months"], "interval months")?;
3986 let days = canonical_i64(&object["days"], "interval days")?
3987 .try_into()
3988 .map_err(|_| "interval days is outside i32 range".to_string())?;
3989 let nanos = canonical_i64(&object["nanos"], "interval nanos")?;
3990 Ok(Value::Interval {
3991 months,
3992 days,
3993 nanos,
3994 })
3995 }
3996 TypeId::Uuid => {
3997 let bytes = decode_tagged_hex(value, "uuid")?;
3998 let bytes: [u8; 16] = bytes
3999 .try_into()
4000 .map_err(|_| "UUID must contain exactly 16 bytes".to_string())?;
4001 Ok(Value::Uuid(bytes))
4002 }
4003 TypeId::Json => {
4004 let bytes = decode_tagged_hex(value, "json")?;
4005 std::str::from_utf8(&bytes).map_err(|_| "JSON value is not UTF-8".to_string())?;
4006 serde_json::from_slice::<serde_json::Value>(&bytes)
4007 .map_err(|error| format!("JSON value is invalid: {error}"))?;
4008 Ok(Value::Json(bytes))
4009 }
4010 TypeId::Array { .. } => Err("legacy put does not support array columns".into()),
4011 }
4012}
4013
4014fn exact_tagged_object<'a>(
4015 value: &'a serde_json::Value,
4016 expected_kind: &str,
4017 fields: &[&str],
4018) -> Result<&'a serde_json::Map<String, serde_json::Value>, String> {
4019 let object = value
4020 .as_object()
4021 .ok_or_else(|| format!("expected tagged {expected_kind} value"))?;
4022 if object.len() != fields.len() + 1
4023 || object
4024 .get("$mongreldb_type")
4025 .and_then(|value| value.as_str())
4026 != Some(expected_kind)
4027 || fields.iter().any(|field| !object.contains_key(*field))
4028 {
4029 return Err(format!("invalid tagged {expected_kind} value"));
4030 }
4031 Ok(object)
4032}
4033
4034fn decode_tagged_hex(value: &serde_json::Value, expected_kind: &str) -> Result<Vec<u8>, String> {
4035 let object = exact_tagged_object(value, expected_kind, &["hex"])?;
4036 let encoded = object["hex"]
4037 .as_str()
4038 .ok_or_else(|| format!("tagged {expected_kind} hex must be a string"))?;
4039 if encoded.len() % 2 != 0 {
4040 return Err(format!("tagged {expected_kind} hex has odd length"));
4041 }
4042 encoded
4043 .as_bytes()
4044 .chunks_exact(2)
4045 .map(|pair| {
4046 let high = hex_nibble(pair[0])?;
4047 let low = hex_nibble(pair[1])?;
4048 Ok((high << 4) | low)
4049 })
4050 .collect()
4051}
4052
4053fn hex_nibble(value: u8) -> Result<u8, String> {
4054 match value {
4055 b'0'..=b'9' => Ok(value - b'0'),
4056 b'a'..=b'f' => Ok(value - b'a' + 10),
4057 _ => Err("hex value must use lowercase ASCII digits".into()),
4058 }
4059}
4060
4061fn canonical_i64(value: &serde_json::Value, field: &str) -> Result<i64, String> {
4062 let text = value
4063 .as_str()
4064 .ok_or_else(|| format!("{field} must be a string"))?;
4065 let value = text
4066 .parse::<i64>()
4067 .map_err(|_| format!("{field} is invalid"))?;
4068 if value.to_string() != text {
4069 return Err(format!("{field} is not canonical"));
4070 }
4071 Ok(value)
4072}
4073
4074#[cfg(test)]
4075mod legacy_wire_tests {
4076 use super::*;
4077
4078 #[test]
4079 fn typed_values_are_exact_and_malformed_values_fail_closed() {
4080 assert_eq!(
4081 legacy_json_to_value(
4082 &json!({"$mongreldb_type": "bytes", "hex": "00ff61"}),
4083 &TypeId::Bytes,
4084 )
4085 .unwrap(),
4086 Value::Bytes(vec![0, 0xff, b'a'])
4087 );
4088 for (value, ty) in [
4089 (
4090 json!({"$mongreldb_type": "bytes", "hex": "00FF"}),
4091 TypeId::Bytes,
4092 ),
4093 (
4094 json!({"$mongreldb_type": "decimal", "unscaled": "01"}),
4095 TypeId::Decimal128 {
4096 precision: 38,
4097 scale: 0,
4098 },
4099 ),
4100 (
4101 json!({"$mongreldb_type": "uuid", "hex": "00"}),
4102 TypeId::Uuid,
4103 ),
4104 (
4105 json!({"$mongreldb_type": "json", "hex": "7b"}),
4106 TypeId::Json,
4107 ),
4108 (json!([1.0]), TypeId::Embedding { dim: 2 }),
4109 (json!([1e100, 1.0]), TypeId::Embedding { dim: 2 }),
4110 ] {
4111 assert!(legacy_json_to_value(&value, &ty).is_err(), "{value}");
4112 }
4113 }
4114}
4115
4116fn parse_cells(
4119 row: &[serde_json::Value],
4120 schema: &mongreldb_core::schema::Schema,
4121) -> Result<Vec<(u16, Value)>, String> {
4122 if row.len() & 1 != 0 {
4123 return Err("row must be an even-length array of [col_id, value] pairs".into());
4124 }
4125 let mut out = Vec::with_capacity(row.len() / 2);
4126 let mut seen = std::collections::HashSet::new();
4127 for chunk in row.chunks(2) {
4128 let col_id = chunk[0]
4129 .as_u64()
4130 .and_then(|value| u16::try_from(value).ok())
4131 .ok_or("column id must be an unsigned 16-bit integer")?;
4132 if !seen.insert(col_id) {
4133 return Err(format!("duplicate column id {col_id}"));
4134 }
4135 let expected = schema
4136 .columns
4137 .iter()
4138 .find(|c| c.id == col_id)
4139 .map(|c| c.ty.clone())
4140 .ok_or_else(|| format!("unknown column id {col_id}"))?;
4141 let val = legacy_json_to_value(&chunk[1], &expected)?;
4142 out.push((col_id, val));
4143 }
4144 Ok(out)
4145}
4146
4147pub(crate) fn validate_table_name(name: &str) -> Result<(), String> {
4149 if name.is_empty() {
4150 return Err("table name must not be empty".into());
4151 }
4152 if name.contains('/') || name.contains('\\') || name.contains('\0') {
4153 return Err("table name contains invalid characters".into());
4154 }
4155 Ok(())
4156}
4157
4158async fn put_row(
4159 State(state): State<Arc<AppState>>,
4160 OptionalPrincipal(principal): OptionalPrincipal,
4161 Path(name): Path<String>,
4162 Json(req): Json<PutRequest>,
4163) -> Response {
4164 if let Some(response) = require_writes_open(&state) {
4165 return response;
4166 }
4167 let handle = match state.db.table(&name) {
4168 Ok(h) => h,
4169 Err(e) => return (StatusCode::NOT_FOUND, e.to_string()).into_response(),
4170 };
4171 let schema = handle.lock().schema().clone();
4172 let row = match parse_cells(&req.row, &schema) {
4173 Ok(r) => r,
4174 Err(msg) => return (StatusCode::BAD_REQUEST, msg).into_response(),
4175 };
4176 state.metrics.inc_puts();
4177 let principal = request_principal(&state, &principal);
4178 match state.db.put_for(&name, row, principal.as_ref()) {
4179 Ok(rid) => Json(json!({ "row_id": rid.0.to_string() })).into_response(),
4180 Err(e) => (status_for_error(&e), e.to_string()).into_response(),
4181 }
4182}
4183
4184async fn count(
4185 State(state): State<Arc<AppState>>,
4186 OptionalPrincipal(principal): OptionalPrincipal,
4187 Path(name): Path<String>,
4188) -> Response {
4189 let principal = request_principal(&state, &principal);
4190 match state.db.count_for(&name, principal.as_ref()) {
4191 Ok(count) => Json(json!({ "count": count })).into_response(),
4192 Err(error) => (status_for_error(&error), error.to_string()).into_response(),
4193 }
4194}
4195
4196async fn commit(
4197 State(state): State<Arc<AppState>>,
4198 OptionalPrincipal(principal): OptionalPrincipal,
4199 Path(name): Path<String>,
4200) -> Response {
4201 if let Some(response) = require_writes_open(&state) {
4202 return response;
4203 }
4204 if let Err(error) = state.db.require_for(
4205 request_principal(&state, &principal).as_ref(),
4206 &mongreldb_core::Permission::Update {
4207 table: name.clone(),
4208 },
4209 ) {
4210 return (status_for_error(&error), error.to_string()).into_response();
4211 }
4212 let handle = match state.db.table(&name) {
4213 Ok(h) => h,
4214 Err(e) => return (StatusCode::NOT_FOUND, e.to_string()).into_response(),
4215 };
4216 let mut g = handle.lock();
4217 state.metrics.inc_commits();
4218 match g.commit() {
4219 Ok(epoch) => Json(json!({
4220 "epoch": epoch.0,
4221 "epoch_text": epoch.0.to_string()
4222 }))
4223 .into_response(),
4224 Err(error) => crate::kit::durable_core_error_response(&error)
4225 .unwrap_or_else(|| (status_for_error(&error), error.to_string()).into_response()),
4226 }
4227}
4228
4229#[derive(Deserialize)]
4230struct SqlRequest {
4231 sql: String,
4232 #[serde(default)]
4235 format: Option<String>,
4236 #[serde(default)]
4238 query_id: Option<QueryId>,
4239 #[serde(default)]
4240 timeout_ms: Option<u64>,
4241 #[serde(default)]
4242 max_output_rows: Option<u64>,
4243 #[serde(default)]
4244 max_output_bytes: Option<u64>,
4245 #[serde(default)]
4246 idempotency_key: Option<String>,
4247 #[serde(default)]
4248 pagination: Option<SqlPaginationRequest>,
4249}
4250
4251#[derive(Clone, Debug, Deserialize, Serialize)]
4252struct SqlPaginationRequest {
4253 page_size_rows: u64,
4254 projection: Vec<String>,
4255 #[serde(default)]
4256 max_page_bytes: Option<u64>,
4257 #[serde(default)]
4258 max_page_tokens: Option<u64>,
4259}
4260
4261#[derive(Clone)]
4262struct ResolvedSqlPagination {
4263 projection: Vec<String>,
4264 limits: sql_pages::SqlPageLimits,
4265}
4266
4267struct ResolvedSqlRequest {
4268 request: SqlRequest,
4269 output_limits: (usize, usize),
4270 idempotency: Option<sql_idempotency::SqlIdempotencyExecution>,
4271 pagination: Option<ResolvedSqlPagination>,
4272}
4273
4274fn query_error_response(
4275 error: &mongreldb_query::MongrelQueryError,
4276 query_id: Option<QueryId>,
4277) -> Response {
4278 query_error_response_with_status(error, query_id, None)
4279}
4280
4281fn query_error_category(
4286 error: &mongreldb_query::MongrelQueryError,
4287) -> mongreldb_types::errors::ErrorCategory {
4288 use mongreldb_query::MongrelQueryError;
4289 use mongreldb_types::errors::ErrorCategory;
4290 match error {
4291 MongrelQueryError::Core(error) => error.category(),
4292 MongrelQueryError::QueryCancelled { .. } => ErrorCategory::Cancelled,
4293 MongrelQueryError::DeadlineExceeded { .. } => ErrorCategory::DeadlineExceeded,
4294 MongrelQueryError::CommitOutcome { committed, .. } => {
4298 if *committed {
4299 ErrorCategory::CommitOutcomeUnknown
4300 } else {
4301 ErrorCategory::TransactionAborted
4302 }
4303 }
4304 MongrelQueryError::OutcomeUnknown { .. } => ErrorCategory::CommitOutcomeUnknown,
4305 MongrelQueryError::TransactionAborted => ErrorCategory::TransactionAborted,
4306 MongrelQueryError::NoSqlTransaction => ErrorCategory::TransactionAborted,
4308 MongrelQueryError::SavepointNotFound { .. } => ErrorCategory::StaleMetadata,
4311 MongrelQueryError::QueryRegistryFull | MongrelQueryError::ResultLimitExceeded { .. } => {
4312 ErrorCategory::ResourceExhausted
4313 }
4314 MongrelQueryError::QueryIdConflict { .. }
4318 | MongrelQueryError::InvalidQueryState(_)
4319 | MongrelQueryError::Arrow(_)
4320 | MongrelQueryError::DataFusion(_) => ErrorCategory::ClusterVersionMismatch,
4321 MongrelQueryError::Schema(_) => ErrorCategory::SchemaVersionMismatch,
4322 _ => ErrorCategory::ReplicaUnavailable,
4325 }
4326}
4327
4328fn query_error_response_with_status(
4329 error: &mongreldb_query::MongrelQueryError,
4330 query_id: Option<QueryId>,
4331 status: Option<&mongreldb_query::QueryStatus>,
4332) -> Response {
4333 use mongreldb_query::MongrelQueryError;
4334 let (base_code, id) = match error {
4335 MongrelQueryError::QueryCancelled { query_id, .. } => ("QUERY_CANCELLED", Some(*query_id)),
4336 MongrelQueryError::DeadlineExceeded { query_id, .. } => {
4337 ("DEADLINE_EXCEEDED", Some(*query_id))
4338 }
4339 MongrelQueryError::QueryIdConflict { query_id } => ("QUERY_ID_CONFLICT", Some(*query_id)),
4340 MongrelQueryError::QueryRegistryFull => ("QUERY_REGISTRY_FULL", query_id),
4341 MongrelQueryError::ResultLimitExceeded { query_id, .. } => {
4342 ("RESULT_LIMIT_EXCEEDED", Some(*query_id))
4343 }
4344 MongrelQueryError::TransactionAborted => ("TRANSACTION_ABORTED", query_id),
4345 MongrelQueryError::NoSqlTransaction => ("NO_SQL_TRANSACTION", query_id),
4346 MongrelQueryError::SavepointNotFound { .. } => ("SAVEPOINT_NOT_FOUND", query_id),
4347 MongrelQueryError::CommitOutcome { query_id, .. } => ("COMMIT_OUTCOME", Some(*query_id)),
4348 MongrelQueryError::OutcomeUnknown { query_id, .. } => {
4349 ("QUERY_OUTCOME_UNKNOWN", Some(*query_id))
4350 }
4351 _ => ("QUERY_FAILED", query_id),
4352 };
4353 let (
4354 error_committed,
4355 error_committed_statements,
4356 error_last_commit_epoch,
4357 error_first_commit_statement_index,
4358 error_last_commit_statement_index,
4359 ) = match error {
4360 MongrelQueryError::QueryCancelled {
4361 committed,
4362 committed_statements,
4363 last_commit_epoch,
4364 first_commit_statement_index,
4365 last_commit_statement_index,
4366 ..
4367 }
4368 | MongrelQueryError::DeadlineExceeded {
4369 committed,
4370 committed_statements,
4371 last_commit_epoch,
4372 first_commit_statement_index,
4373 last_commit_statement_index,
4374 ..
4375 }
4376 | MongrelQueryError::ResultLimitExceeded {
4377 committed,
4378 committed_statements,
4379 last_commit_epoch,
4380 first_commit_statement_index,
4381 last_commit_statement_index,
4382 ..
4383 } => (
4384 *committed,
4385 *committed_statements,
4386 *last_commit_epoch,
4387 *first_commit_statement_index,
4388 *last_commit_statement_index,
4389 ),
4390 MongrelQueryError::CommitOutcome {
4391 committed,
4392 committed_statements,
4393 last_commit_epoch,
4394 first_commit_statement_index,
4395 last_commit_statement_index,
4396 ..
4397 } => (
4398 *committed,
4399 *committed_statements,
4400 *last_commit_epoch,
4401 *first_commit_statement_index,
4402 *last_commit_statement_index,
4403 ),
4404 _ => (false, 0, None, None, None),
4405 };
4406 let committed = status.map_or_else(
4407 || error_committed,
4408 |status| status.durable_outcome.committed,
4409 );
4410 let outcome_unknown = matches!(error, MongrelQueryError::OutcomeUnknown { .. })
4411 || status.is_some_and(|status| status.outcome_unknown);
4412 let code = status
4413 .and_then(|status| {
4414 status
4415 .terminal_error
4416 .as_ref()
4417 .map(|error| error.code.as_str())
4418 })
4419 .unwrap_or(match (base_code, committed) {
4420 ("QUERY_CANCELLED", true) => "QUERY_CANCELLED_AFTER_COMMIT",
4421 ("DEADLINE_EXCEEDED", true) => "DEADLINE_AFTER_COMMIT",
4422 _ => base_code,
4423 });
4424 let response_status = if outcome_unknown {
4425 "outcome_unknown"
4426 } else {
4427 status
4428 .and_then(mongreldb_query::QueryStatus::terminal_state)
4429 .map(terminal_state_name)
4430 .unwrap_or_else(|| match (error, committed) {
4431 (MongrelQueryError::QueryCancelled { .. }, true) => "cancelled_after_commit",
4432 (MongrelQueryError::QueryCancelled { .. }, false) => "cancelled_before_commit",
4433 (MongrelQueryError::DeadlineExceeded { .. }, true) => "deadline_after_commit",
4434 (MongrelQueryError::DeadlineExceeded { .. }, false) => "deadline_before_commit",
4435 (_, true) => "committed_with_error",
4436 _ => "failed_before_commit",
4437 })
4438 };
4439 let (completed_statements, statement_index) = status.map_or_else(
4440 || match error {
4441 MongrelQueryError::QueryCancelled {
4442 completed_statements,
4443 cancelled_statement_index,
4444 ..
4445 }
4446 | MongrelQueryError::DeadlineExceeded {
4447 completed_statements,
4448 cancelled_statement_index,
4449 ..
4450 } => (*completed_statements, *cancelled_statement_index),
4451 MongrelQueryError::ResultLimitExceeded {
4452 completed_statements,
4453 statement_index,
4454 ..
4455 } => (*completed_statements, *statement_index),
4456 MongrelQueryError::CommitOutcome {
4457 completed_statements,
4458 statement_index,
4459 ..
4460 } => (*completed_statements, *statement_index),
4461 _ => (0, 0),
4462 },
4463 |status| (status.completed_statements, status.statement_index),
4464 );
4465 let committed_statements = status.map_or(error_committed_statements, |status| {
4466 status.durable_outcome.committed_statements
4467 });
4468 let last_commit_epoch = status.map_or(error_last_commit_epoch, |status| {
4469 status.durable_outcome.last_commit_epoch
4470 });
4471 let first_commit_statement_index = status
4472 .map_or(error_first_commit_statement_index, |status| {
4473 status.durable_outcome.first_commit_statement_index
4474 });
4475 let last_commit_statement_index = status.map_or(error_last_commit_statement_index, |status| {
4476 status.durable_outcome.last_commit_statement_index
4477 });
4478 let cancellation_reason = status
4479 .map(|status| status.cancellation_reason)
4480 .or(match error {
4481 MongrelQueryError::QueryCancelled { reason, .. } => Some(*reason),
4482 MongrelQueryError::DeadlineExceeded { .. } => Some(CancellationReason::Deadline),
4483 _ => None,
4484 })
4485 .map(cancellation_reason_name);
4486 let cancel_outcome = match error {
4487 MongrelQueryError::QueryCancelled { .. } | MongrelQueryError::DeadlineExceeded { .. } => {
4488 Some("accepted")
4489 }
4490 _ => status.and_then(query_cancel_outcome),
4491 };
4492 let outcome = if outcome_unknown {
4493 json!({
4494 "committed": null,
4495 "committed_statements": null,
4496 "last_commit_epoch": null,
4497 "last_commit_epoch_text": null,
4498 "first_commit_statement_index": null,
4499 "last_commit_statement_index": null,
4500 "completed_statements": null,
4501 "statement_index": null,
4502 "serialization": "unknown",
4503 })
4504 } else {
4505 status.map_or_else(
4506 || {
4507 json!({
4508 "committed": committed,
4509 "committed_statements": committed_statements,
4510 "last_commit_epoch": last_commit_epoch,
4511 "last_commit_epoch_text": epoch_text(last_commit_epoch),
4512 "first_commit_statement_index": first_commit_statement_index,
4513 "last_commit_statement_index": last_commit_statement_index,
4514 "completed_statements": completed_statements,
4515 "statement_index": statement_index,
4516 "serialization": "unknown",
4517 })
4518 },
4519 |status| query_outcome_json(Some(status)),
4520 )
4521 };
4522 let http_status = match code {
4523 "QUERY_CANCELLED_AFTER_COMMIT" | "DEADLINE_AFTER_COMMIT" => StatusCode::CONFLICT,
4524 "QUERY_CANCELLED" => client_closed_request_status(),
4525 "DEADLINE_EXCEEDED" => StatusCode::GATEWAY_TIMEOUT,
4526 _ => status_for_query_error(error),
4527 };
4528 let category = query_error_category(error);
4532 let mut response = (
4533 http_status,
4534 Json(json!({
4535 "query_id": id.map(|value| value.to_string()),
4536 "status": response_status,
4537 "terminal_state": response_status,
4538 "committed": (!outcome_unknown).then_some(committed),
4539 "committed_statements": (!outcome_unknown).then_some(committed_statements),
4540 "last_commit_epoch": (!outcome_unknown).then_some(last_commit_epoch).flatten(),
4541 "last_commit_epoch_text": (!outcome_unknown).then_some(epoch_text(last_commit_epoch)).flatten(),
4542 "first_commit_statement_index": (!outcome_unknown).then_some(first_commit_statement_index).flatten(),
4543 "last_commit_statement_index": (!outcome_unknown).then_some(last_commit_statement_index).flatten(),
4544 "completed_statements": (!outcome_unknown).then_some(completed_statements),
4545 "statement_index": (!outcome_unknown).then_some(statement_index),
4546 "cancel_outcome": cancel_outcome,
4547 "cancellation_reason": cancellation_reason,
4548 "retryable": matches!(error, MongrelQueryError::QueryRegistryFull),
4549 "server_state": status.map(|status| query_phase_name(status.phase)),
4550 "outcome": outcome,
4551 "error": {
4552 "code": code,
4553 "message": error.to_string(),
4554 "category": category.to_string(),
4555 "category_code": category.code(),
4556 "query_id": id.map(|value| value.to_string()),
4557 "committed": (!outcome_unknown).then_some(committed),
4558 "retryable": matches!(error, MongrelQueryError::QueryRegistryFull),
4559 }
4560 })),
4561 )
4562 .into_response();
4563 if let Some(id) = id {
4564 add_query_id_header(&mut response, id);
4565 }
4566 response
4567}
4568
4569fn record_query_error(metrics: &metrics::Metrics, error: &mongreldb_query::MongrelQueryError) {
4570 match error {
4571 mongreldb_query::MongrelQueryError::QueryCancelled { reason, .. } => {
4572 metrics.inc_sql_cancelled(*reason)
4573 }
4574 mongreldb_query::MongrelQueryError::DeadlineExceeded { .. } => {
4575 metrics.inc_sql_deadline_exceeded();
4576 metrics.inc_sql_cancelled(CancellationReason::Deadline);
4577 }
4578 _ => {}
4579 }
4580}
4581
4582fn tracked_query_error_response(
4583 state: &AppState,
4584 error: &mongreldb_query::MongrelQueryError,
4585 query_id: Option<QueryId>,
4586) -> Response {
4587 record_query_error(&state.metrics, error);
4588 if let mongreldb_query::MongrelQueryError::QueryCancelled { query_id, .. } = error {
4589 if let Some(requested_at) = state
4590 .query_registry
4591 .status(*query_id)
4592 .and_then(|status| status.cancel_requested_at)
4593 {
4594 state
4595 .metrics
4596 .observe_sql_cancel_latency(requested_at.elapsed());
4597 }
4598 }
4599 let status = if matches!(
4600 error,
4601 mongreldb_query::MongrelQueryError::QueryIdConflict { .. }
4602 | mongreldb_query::MongrelQueryError::QueryRegistryFull
4603 ) {
4604 None
4605 } else {
4606 query_id
4607 .or(match error {
4608 mongreldb_query::MongrelQueryError::QueryCancelled { query_id, .. }
4609 | mongreldb_query::MongrelQueryError::DeadlineExceeded { query_id, .. }
4610 | mongreldb_query::MongrelQueryError::CommitOutcome { query_id, .. }
4611 | mongreldb_query::MongrelQueryError::OutcomeUnknown { query_id, .. } => {
4612 Some(*query_id)
4613 }
4614 _ => None,
4615 })
4616 .and_then(|query_id| state.query_registry.status(query_id))
4617 };
4618 query_error_response_with_status(error, query_id, status.as_ref())
4619}
4620
4621fn add_query_id_header(response: &mut Response, query_id: QueryId) {
4622 if let Ok(value) = axum::http::HeaderValue::from_str(&query_id.to_string()) {
4623 response.headers_mut().insert("x-mongreldb-query-id", value);
4624 }
4625}
4626
4627fn with_query_id(mut response: Response, query_id: QueryId) -> Response {
4628 add_query_id_header(&mut response, query_id);
4629 response
4630}
4631
4632fn bad_query_control_request(message: impl Into<String>, query_id: Option<QueryId>) -> Response {
4633 let mut response = (
4634 StatusCode::BAD_REQUEST,
4635 Json(json!({
4636 "query_id": query_id.map(|value| value.to_string()),
4637 "status": "failed_before_commit",
4638 "terminal_state": "failed_before_commit",
4639 "committed": false,
4640 "committed_statements": 0,
4641 "last_commit_epoch": null,
4642 "last_commit_epoch_text": null,
4643 "first_commit_statement_index": null,
4644 "last_commit_statement_index": null,
4645 "completed_statements": 0,
4646 "statement_index": 0,
4647 "cancel_outcome": null,
4648 "cancellation_reason": null,
4649 "retryable": false,
4650 "server_state": "failed",
4651 "outcome": {
4652 "committed": false,
4653 "committed_statements": 0,
4654 "last_commit_epoch": null,
4655 "last_commit_epoch_text": null,
4656 "first_commit_statement_index": null,
4657 "last_commit_statement_index": null,
4658 "completed_statements": 0,
4659 "statement_index": 0,
4660 "serialization": "not_started",
4661 },
4662 "error": {
4663 "code": "INVALID_QUERY_OPTIONS",
4664 "message": message.into(),
4665 "query_id": query_id.map(|value| value.to_string()),
4666 "committed": false,
4667 "retryable": false,
4668 }
4669 })),
4670 )
4671 .into_response();
4672 if let Some(query_id) = query_id {
4673 add_query_id_header(&mut response, query_id);
4674 }
4675 response
4676}
4677
4678fn resolve_query_options(
4679 state: &AppState,
4680 headers: &axum::http::HeaderMap,
4681 body_query_id: Option<QueryId>,
4682 body_timeout_ms: Option<u64>,
4683 owner: String,
4684 session_id: Option<String>,
4685) -> std::result::Result<(SqlQueryOptions, QueryId), Box<Response>> {
4686 let query_id = match body_query_id {
4687 Some(query_id) => query_id,
4688 None => match headers.get("x-mongreldb-query-id") {
4689 Some(value) => {
4690 let value = value.to_str().map_err(|_| {
4691 Box::new(bad_query_control_request(
4692 "X-MongrelDB-Query-ID is not valid text",
4693 None,
4694 ))
4695 })?;
4696 value
4697 .parse()
4698 .map_err(|error: mongreldb_query::MongrelQueryError| {
4699 Box::new(bad_query_control_request(error.to_string(), None))
4700 })?
4701 }
4702 None => {
4703 QueryId::random().map_err(|error| Box::new(query_error_response(&error, None)))?
4704 }
4705 },
4706 };
4707 let timeout_ms = match body_timeout_ms {
4708 Some(timeout_ms) => timeout_ms,
4709 None => match headers.get("x-mongreldb-timeout-ms") {
4710 Some(value) => value
4711 .to_str()
4712 .ok()
4713 .and_then(|value| value.parse::<u64>().ok())
4714 .ok_or_else(|| {
4715 Box::new(bad_query_control_request(
4716 "X-MongrelDB-Timeout-Ms must be a positive integer",
4717 Some(query_id),
4718 ))
4719 })?,
4720 None => state.reloadable.sql_default_timeout.as_millis(),
4721 },
4722 };
4723 if timeout_ms == 0 {
4724 return Err(Box::new(bad_query_control_request(
4725 "timeout_ms must be positive",
4726 Some(query_id),
4727 )));
4728 }
4729 let timeout = std::time::Duration::from_millis(timeout_ms);
4730 if timeout > state.reloadable.sql_max_timeout.get() {
4731 return Err(Box::new(bad_query_control_request(
4732 format!(
4733 "timeout_ms exceeds server maximum of {}",
4734 state.reloadable.sql_max_timeout.as_millis()
4735 ),
4736 Some(query_id),
4737 )));
4738 }
4739 Ok((
4740 SqlQueryOptions {
4741 query_id: Some(query_id),
4742 timeout: Some(timeout),
4743 owner: Some(owner),
4744 session_id,
4745 parent_control: None,
4746 },
4747 query_id,
4748 ))
4749}
4750
4751fn resolve_sql_output_limits(
4752 state: &AppState,
4753 request: &SqlRequest,
4754 query_id: QueryId,
4755) -> std::result::Result<(usize, usize), Box<Response>> {
4756 fn resolve(
4757 requested: Option<u64>,
4758 configured: usize,
4759 name: &str,
4760 query_id: QueryId,
4761 ) -> std::result::Result<usize, Box<Response>> {
4762 if requested == Some(0) {
4763 return Err(Box::new(bad_query_control_request(
4764 format!("{name} must be positive"),
4765 Some(query_id),
4766 )));
4767 }
4768 let requested = requested
4769 .and_then(|value| usize::try_from(value).ok())
4770 .unwrap_or(usize::MAX);
4771 Ok(requested.min(configured))
4772 }
4773
4774 Ok((
4775 resolve(
4776 request.max_output_rows,
4777 state.reloadable.sql_max_output_rows.get(),
4778 "max_output_rows",
4779 query_id,
4780 )?,
4781 resolve(
4782 request.max_output_bytes,
4783 state.reloadable.sql_max_output_bytes.get(),
4784 "max_output_bytes",
4785 query_id,
4786 )?,
4787 ))
4788}
4789
4790fn resolve_sql_pagination(
4791 headers: &axum::http::HeaderMap,
4792 request: &SqlRequest,
4793 output_limits: (usize, usize),
4794 registration: RegisteredQueryGuard,
4795 query_id: QueryId,
4796) -> Result<(RegisteredQueryGuard, Option<ResolvedSqlPagination>), Box<Response>> {
4797 let Some(pagination) = request.pagination.as_ref() else {
4798 return Ok((registration, None));
4799 };
4800 if requested_sql_idempotency_key(headers, request)
4801 .ok()
4802 .flatten()
4803 .is_some()
4804 {
4805 return Err(Box::new(registered_sql_error_response(
4806 registration,
4807 query_id,
4808 StatusCode::BAD_REQUEST,
4809 "INCOMPATIBLE_SQL_CONTROLS",
4810 "idempotency_key cannot be combined with SQL pagination",
4811 false,
4812 )));
4813 }
4814 if request
4815 .format
4816 .as_deref()
4817 .is_some_and(|format| format != "json")
4818 {
4819 return Err(Box::new(registered_sql_error_response(
4820 registration,
4821 query_id,
4822 StatusCode::BAD_REQUEST,
4823 "PAGINATION_REQUIRES_JSON",
4824 "SQL pagination supports JSON responses only",
4825 false,
4826 )));
4827 }
4828 registration.query().set_sql_metadata(&request.sql);
4829 if !mongreldb_query::is_single_read_only_query(&request.sql) {
4830 return Err(Box::new(registered_sql_error_response(
4831 registration,
4832 query_id,
4833 StatusCode::BAD_REQUEST,
4834 "PAGINATION_REQUIRES_SINGLE_READ_QUERY",
4835 "SQL pagination accepts exactly one read-only query statement",
4836 false,
4837 )));
4838 }
4839 let page_size = match usize::try_from(pagination.page_size_rows) {
4840 Ok(0) | Err(_) => {
4841 return Err(Box::new(registered_sql_error_response(
4842 registration,
4843 query_id,
4844 StatusCode::BAD_REQUEST,
4845 "INVALID_PAGINATION_OPTIONS",
4846 "pagination.page_size_rows must be positive",
4847 false,
4848 )))
4849 }
4850 Ok(value) => value.min(output_limits.0),
4851 };
4852 if pagination.projection.is_empty() || pagination.projection.len() > 128 {
4853 return Err(Box::new(registered_sql_error_response(
4854 registration,
4855 query_id,
4856 StatusCode::BAD_REQUEST,
4857 "INVALID_SQL_PROJECTION",
4858 "pagination.projection must contain between 1 and 128 output column names",
4859 false,
4860 )));
4861 }
4862 let mut seen = std::collections::HashSet::new();
4863 let metadata_bytes = pagination
4864 .projection
4865 .iter()
4866 .map(String::len)
4867 .fold(0usize, usize::saturating_add);
4868 if metadata_bytes > 16 * 1024
4869 || pagination.projection.iter().any(|column| {
4870 column.is_empty()
4871 || column == "*"
4872 || column.len() > 256
4873 || !seen.insert(column.as_str())
4874 })
4875 {
4876 return Err(Box::new(registered_sql_error_response(
4877 registration,
4878 query_id,
4879 StatusCode::BAD_REQUEST,
4880 "INVALID_SQL_PROJECTION",
4881 "pagination.projection requires unique explicit output names of at most 256 bytes",
4882 false,
4883 )));
4884 }
4885 let max_page_bytes = match pagination.max_page_bytes {
4886 Some(0) => {
4887 return Err(Box::new(registered_sql_error_response(
4888 registration,
4889 query_id,
4890 StatusCode::BAD_REQUEST,
4891 "INVALID_PAGINATION_OPTIONS",
4892 "pagination.max_page_bytes must be positive",
4893 false,
4894 )))
4895 }
4896 Some(value) => usize::try_from(value)
4897 .unwrap_or(usize::MAX)
4898 .min(output_limits.1),
4899 None => output_limits.1.min(1024 * 1024),
4900 };
4901 let token_cap = (output_limits.1.saturating_add(3) / 4).max(1);
4902 let max_page_tokens = match pagination.max_page_tokens {
4903 Some(0) => {
4904 return Err(Box::new(registered_sql_error_response(
4905 registration,
4906 query_id,
4907 StatusCode::BAD_REQUEST,
4908 "INVALID_PAGINATION_OPTIONS",
4909 "pagination.max_page_tokens must be positive",
4910 false,
4911 )))
4912 }
4913 Some(value) => usize::try_from(value).unwrap_or(usize::MAX).min(token_cap),
4914 None => (max_page_bytes.saturating_add(3) / 4).max(1),
4915 };
4916 Ok((
4917 registration,
4918 Some(ResolvedSqlPagination {
4919 projection: pagination.projection.clone(),
4920 limits: sql_pages::SqlPageLimits {
4921 rows: page_size,
4922 bytes: max_page_bytes,
4923 tokens: max_page_tokens,
4924 },
4925 }),
4926 ))
4927}
4928
4929fn requested_sql_idempotency_key(
4930 headers: &axum::http::HeaderMap,
4931 request: &SqlRequest,
4932) -> Result<Option<String>, &'static str> {
4933 let header = match headers.get("idempotency-key") {
4934 Some(value) => Some(
4935 value
4936 .to_str()
4937 .map_err(|_| "Idempotency-Key must be valid UTF-8")?,
4938 ),
4939 None => None,
4940 };
4941 match (request.idempotency_key.as_deref(), header) {
4942 (Some(body), Some(header)) if body != header => {
4943 Err("body idempotency_key and Idempotency-Key header must match")
4944 }
4945 (Some(body), _) => Ok(Some(body.to_owned())),
4946 (None, Some(header)) => Ok(Some(header.to_owned())),
4947 (None, None) => Ok(None),
4948 }
4949}
4950
4951fn sql_idempotency_binding(
4952 request: &SqlRequest,
4953 output_limits: (usize, usize),
4954 session_id: Option<&str>,
4955 expires_after_ms: u64,
4956) -> Result<sql_idempotency::SqlIdempotencyBinding, serde_json::Error> {
4957 let request_semantics = serde_json::to_vec(&json!({
4958 "format": request.format.as_deref().unwrap_or("json"),
4959 "max_output_rows": output_limits.0,
4960 "max_output_bytes": output_limits.1,
4961 "pagination": request.pagination.as_ref(),
4962 }))?;
4963 let session_semantics = session_id.map_or_else(
4964 || b"ephemeral".to_vec(),
4965 |session_id| {
4966 let mut semantics = b"session\0".to_vec();
4967 semantics.extend_from_slice(session_id.as_bytes());
4968 semantics
4969 },
4970 );
4971 Ok(sql_idempotency::SqlIdempotencyBinding {
4972 sql_fingerprint: mongreldb_query::normalized_sql_fingerprint(&request.sql),
4973 parameter_hash: sql_idempotency::hash(b"[]"),
4976 request_semantics_hash: sql_idempotency::hash(&request_semantics),
4977 session_semantics_hash: sql_idempotency::hash(&session_semantics),
4978 expires_after_ms,
4979 })
4980}
4981
4982struct SqlIdempotencyContext<'a> {
4983 headers: &'a axum::http::HeaderMap,
4984 request: &'a SqlRequest,
4985 output_limits: (usize, usize),
4986 owner: &'a str,
4987 session_id: Option<&'a str>,
4988 session_in_transaction: bool,
4989 query_id: QueryId,
4990}
4991
4992async fn begin_sql_idempotency(
4993 state: &AppState,
4994 context: SqlIdempotencyContext<'_>,
4995 registration: RegisteredQueryGuard,
4996) -> Result<
4997 (
4998 RegisteredQueryGuard,
4999 Option<sql_idempotency::SqlIdempotencyExecution>,
5000 ),
5001 Response,
5002> {
5003 let SqlIdempotencyContext {
5004 headers,
5005 request,
5006 output_limits,
5007 owner,
5008 session_id,
5009 session_in_transaction,
5010 query_id,
5011 } = context;
5012 let key = match requested_sql_idempotency_key(headers, request) {
5013 Ok(key) => key,
5014 Err(message) => {
5015 return Err(registered_sql_error_response(
5016 registration,
5017 query_id,
5018 StatusCode::BAD_REQUEST,
5019 "INVALID_IDEMPOTENCY_KEY",
5020 message,
5021 false,
5022 ))
5023 }
5024 };
5025 let Some(key) = key else {
5026 return Ok((registration, None));
5027 };
5028 match mongreldb_query::classify_sql_idempotency(&request.sql) {
5029 mongreldb_query::SqlIdempotencyClass::ReadOnly
5030 | mongreldb_query::SqlIdempotencyClass::Unsupported => {
5031 return Err(registered_sql_error_response(
5032 registration,
5033 query_id,
5034 StatusCode::BAD_REQUEST,
5035 "IDEMPOTENCY_REQUIRES_SINGLE_WRITE",
5036 "idempotency_key accepts one non-transaction SQL write statement",
5037 false,
5038 ));
5039 }
5040 mongreldb_query::SqlIdempotencyClass::SingleWrite => {}
5041 }
5042 if let Err(message) = sql_idempotency::SqlIdempotencyStore::validate_key(&key) {
5043 return Err(registered_sql_error_response(
5044 registration,
5045 query_id,
5046 StatusCode::BAD_REQUEST,
5047 "INVALID_IDEMPOTENCY_KEY",
5048 message,
5049 false,
5050 ));
5051 }
5052 if request
5053 .format
5054 .as_deref()
5055 .is_some_and(|format| format != "json")
5056 {
5057 return Err(registered_sql_error_response(
5058 registration,
5059 query_id,
5060 StatusCode::BAD_REQUEST,
5061 "IDEMPOTENCY_REQUIRES_JSON",
5062 "SQL idempotency supports buffered JSON responses only",
5063 false,
5064 ));
5065 }
5066 if session_in_transaction {
5067 return Err(registered_sql_error_response(
5068 registration,
5069 query_id,
5070 StatusCode::CONFLICT,
5071 "IDEMPOTENCY_UNSUPPORTED_IN_TRANSACTION",
5072 "SQL idempotency cannot be used inside an open session transaction",
5073 false,
5074 ));
5075 }
5076 registration.query().set_sql_metadata(&request.sql);
5077 let binding = match sql_idempotency_binding(
5078 request,
5079 output_limits,
5080 session_id,
5081 state.sql_idempotency.expires_after_ms(),
5082 ) {
5083 Ok(binding) => binding,
5084 Err(_) => {
5085 return Err(registered_sql_error_response(
5086 registration,
5087 query_id,
5088 StatusCode::INTERNAL_SERVER_ERROR,
5089 "SERIALIZATION_FAILED",
5090 "failed to serialize SQL idempotency request semantics",
5091 false,
5092 ))
5093 }
5094 };
5095 let begin = tokio::select! {
5096 begin = state.sql_idempotency.begin(owner, &key, binding) => begin,
5097 _ = registration.query().control().cancelled() => {
5098 return Err(tracked_query_error_response(
5099 state,
5100 &cancellation_checkpoint_error(registration.query()),
5101 Some(query_id),
5102 ));
5103 }
5104 };
5105 match begin {
5106 sql_idempotency::BeginResult::Execute(execution) => Ok((registration, Some(execution))),
5107 sql_idempotency::BeginResult::Replay {
5108 receipt,
5109 expires_at_ms,
5110 } => match restore_idempotency_replay(registration, &receipt) {
5111 Ok(()) => Err(sql_idempotency_receipt_response(
5112 query_id,
5113 &receipt,
5114 true,
5115 expires_at_ms,
5116 true,
5117 )),
5118 Err(error) => Err(tracked_query_error_response(state, &error, Some(query_id))),
5119 },
5120 sql_idempotency::BeginResult::Mismatch => Err(registered_sql_error_response(
5121 registration,
5122 query_id,
5123 StatusCode::CONFLICT,
5124 "IDEMPOTENCY_KEY_REUSE_MISMATCH",
5125 "idempotency key was already used with different SQL or request semantics",
5126 false,
5127 )),
5128 sql_idempotency::BeginResult::Indeterminate { created_at_ms } => Err(
5129 sql_idempotency_indeterminate_response(registration, query_id, created_at_ms),
5130 ),
5131 sql_idempotency::BeginResult::Full => Err(registered_sql_error_response(
5132 registration,
5133 query_id,
5134 StatusCode::SERVICE_UNAVAILABLE,
5135 "IDEMPOTENCY_STORE_FULL",
5136 "SQL idempotency receipt store is full",
5137 true,
5138 )),
5139 sql_idempotency::BeginResult::Unavailable(_reason) => Err(registered_sql_error_response(
5140 registration,
5141 query_id,
5142 StatusCode::SERVICE_UNAVAILABLE,
5143 "IDEMPOTENCY_STORE_UNAVAILABLE",
5144 "could not durably reserve the SQL idempotency key",
5145 true,
5146 )),
5147 }
5148}
5149
5150fn restore_idempotency_replay(
5151 registration: RegisteredQueryGuard,
5152 receipt: &sql_idempotency::SqlDurableReceipt,
5153) -> mongreldb_query::Result<()> {
5154 use mongreldb_query::{
5155 DurableOutcome, QueryTerminalError, QueryTerminalErrorCategory, QueryTerminalState,
5156 SerializationOutcome,
5157 };
5158
5159 let invalid_receipt = |field: &str| {
5160 mongreldb_query::MongrelQueryError::InvalidQueryState(format!(
5161 "durable SQL idempotency receipt has invalid {field}"
5162 ))
5163 };
5164 let terminal_state = match receipt.status.as_str() {
5165 "completed" => QueryTerminalState::Completed,
5166 "failed_before_commit" => QueryTerminalState::FailedBeforeCommit,
5167 "cancelled_before_commit" => QueryTerminalState::CancelledBeforeCommit,
5168 "deadline_before_commit" => QueryTerminalState::DeadlineBeforeCommit,
5169 "committed" => QueryTerminalState::Committed,
5170 "committed_with_error" => QueryTerminalState::CommittedWithError,
5171 "partially_committed" => QueryTerminalState::PartiallyCommitted,
5172 "cancelled_after_commit" => QueryTerminalState::CancelledAfterCommit,
5173 "deadline_after_commit" => QueryTerminalState::DeadlineAfterCommit,
5174 _ => return Err(invalid_receipt("terminal state")),
5175 };
5176 let serialization = match receipt.outcome.serialization.as_str() {
5177 "not_started" => SerializationOutcome::NotStarted,
5178 "in_progress" => SerializationOutcome::InProgress,
5179 "succeeded" => SerializationOutcome::Succeeded,
5180 "failed" => SerializationOutcome::Failed,
5181 _ => return Err(invalid_receipt("serialization state")),
5182 };
5183 let terminal_error = match receipt.terminal_error.as_ref() {
5184 Some(error) => Some(QueryTerminalError {
5185 code: error.code.clone(),
5186 category: match error.category.as_str() {
5187 "cancellation" => QueryTerminalErrorCategory::Cancellation,
5188 "deadline" => QueryTerminalErrorCategory::Deadline,
5189 "result_limit" => QueryTerminalErrorCategory::ResultLimit,
5190 "serialization" => QueryTerminalErrorCategory::Serialization,
5191 "execution" => QueryTerminalErrorCategory::Execution,
5192 _ => return Err(invalid_receipt("terminal error category")),
5193 },
5194 }),
5195 None => None,
5196 };
5197 let cancellation_reason = CancellationReason::from_protocol_str(&receipt.cancellation_reason)
5198 .ok_or_else(|| invalid_receipt("cancellation reason"))?;
5199 let phase = match receipt.server_state.as_str() {
5200 "completed" => SqlQueryPhase::Completed,
5201 "cancelled" => SqlQueryPhase::Cancelled,
5202 "failed" => SqlQueryPhase::Failed,
5203 _ => return Err(invalid_receipt("server state")),
5204 };
5205 let query = registration.into_query();
5206 query.restore_replayed_outcome(
5207 DurableOutcome {
5208 committed: receipt.outcome.committed,
5209 committed_statements: receipt.outcome.committed_statements,
5210 last_commit_epoch: receipt.outcome.last_commit_epoch,
5211 first_commit_statement_index: receipt.outcome.first_commit_statement_index,
5212 last_commit_statement_index: receipt.outcome.last_commit_statement_index,
5213 commit_ts: receipt
5217 .commit_receipt
5218 .as_ref()
5219 .map(sql_idempotency::SqlCommitReceipt::commit_ts),
5220 },
5221 receipt.outcome.completed_statements,
5222 receipt.outcome.statement_index,
5223 serialization,
5224 terminal_error,
5225 terminal_state,
5226 cancellation_reason,
5227 phase,
5228 );
5229 query.try_complete()
5230}
5231
5232fn sql_idempotency_indeterminate_response(
5233 registration: RegisteredQueryGuard,
5234 query_id: QueryId,
5235 created_at_ms: Option<u64>,
5236) -> Response {
5237 registration.query().mark_outcome_unknown();
5238 registration.fail();
5239 with_query_id(
5240 (
5241 StatusCode::CONFLICT,
5242 Json(json!({
5243 "query_id": query_id.to_string(),
5244 "status": "outcome_unknown",
5245 "terminal_state": "outcome_unknown",
5246 "committed": null,
5247 "committed_statements": null,
5248 "last_commit_epoch": null,
5249 "last_commit_epoch_text": null,
5250 "first_commit_statement_index": null,
5251 "last_commit_statement_index": null,
5252 "completed_statements": null,
5253 "statement_index": null,
5254 "cancel_outcome": null,
5255 "cancellation_reason": null,
5256 "retryable": false,
5257 "server_state": "failed",
5258 "idempotency_replayed": true,
5259 "idempotency_intent_created_at_ms": created_at_ms,
5260 "outcome": {
5261 "committed": null,
5262 "committed_statements": null,
5263 "last_commit_epoch": null,
5264 "last_commit_epoch_text": null,
5265 "first_commit_statement_index": null,
5266 "last_commit_statement_index": null,
5267 "completed_statements": null,
5268 "statement_index": null,
5269 "serialization": "unknown",
5270 },
5271 "error": {
5272 "code": "QUERY_OUTCOME_UNKNOWN",
5273 "message": "a durable write intent exists without a durable receipt; the SQL was not re-executed",
5274 "query_id": query_id.to_string(),
5275 "committed": null,
5276 "retryable": false,
5277 }
5278 })),
5279 )
5280 .into_response(),
5281 query_id,
5282 )
5283}
5284
5285fn registered_sql_error_response(
5286 registration: RegisteredQueryGuard,
5287 query_id: QueryId,
5288 status: StatusCode,
5289 code: &'static str,
5290 message: impl Into<String>,
5291 retryable: bool,
5292) -> Response {
5293 let message = message.into();
5294 registration
5295 .query()
5296 .record_terminal_error(code, mongreldb_query::QueryTerminalErrorCategory::Execution);
5297 registration.fail();
5298 with_query_id(
5299 (
5300 status,
5301 Json(json!({
5302 "query_id": query_id.to_string(),
5303 "status": "failed_before_commit",
5304 "terminal_state": "failed_before_commit",
5305 "committed": false,
5306 "committed_statements": 0,
5307 "last_commit_epoch": null,
5308 "last_commit_epoch_text": null,
5309 "first_commit_statement_index": null,
5310 "last_commit_statement_index": null,
5311 "completed_statements": 0,
5312 "statement_index": 0,
5313 "cancel_outcome": null,
5314 "cancellation_reason": null,
5315 "retryable": retryable,
5316 "server_state": "failed",
5317 "outcome": {
5318 "committed": false,
5319 "committed_statements": 0,
5320 "last_commit_epoch": null,
5321 "last_commit_epoch_text": null,
5322 "first_commit_statement_index": null,
5323 "last_commit_statement_index": null,
5324 "completed_statements": 0,
5325 "statement_index": 0,
5326 "serialization": "not_started",
5327 },
5328 "error": {
5329 "code": code,
5330 "message": message,
5331 "query_id": query_id.to_string(),
5332 "committed": false,
5333 "retryable": retryable,
5334 }
5335 })),
5336 )
5337 .into_response(),
5338 query_id,
5339 )
5340}
5341
5342fn register_controlled_query(
5343 state: &AppState,
5344 session: &MongrelSession,
5345 options: SqlQueryOptions,
5346) -> std::result::Result<RegisteredSqlQuery, mongreldb_query::MongrelQueryError> {
5347 let query_id = options.query_id.ok_or_else(|| {
5348 mongreldb_query::MongrelQueryError::InvalidQueryState(
5349 "server query registration requires a query id".into(),
5350 )
5351 })?;
5352 let owner = options.owner.clone().unwrap_or_default();
5353 let session_id = options.session_id.clone();
5354 let _lifecycle = state
5355 .query_lifecycle
5356 .lock()
5357 .unwrap_or_else(|error| error.into_inner());
5358 let pre_cancel_reason = match state.pre_cancellations.lookup_for_registration(
5359 query_id,
5360 &owner,
5361 session_id.as_deref(),
5362 ) {
5363 pre_cancel::RegistrationLookup::NoReservation => None,
5364 pre_cancel::RegistrationLookup::Matching(reason) => Some(reason),
5365 pre_cancel::RegistrationLookup::ReservedByAnotherIdentity => {
5366 return Err(mongreldb_query::MongrelQueryError::QueryIdConflict { query_id });
5367 }
5368 };
5369 let query = session.register_query(options)?;
5370 let Some(reason) = pre_cancel_reason else {
5371 return Ok(query);
5372 };
5373 state
5374 .pre_cancellations
5375 .take(query_id, &owner, session_id.as_deref());
5376 query.request_cancel(reason);
5377 let error = query.checkpoint().err().unwrap_or_else(|| {
5378 mongreldb_query::MongrelQueryError::InvalidQueryState(format!(
5379 "pre-cancelled query {query_id} remained runnable"
5380 ))
5381 });
5382 query.fail();
5383 Err(error)
5384}
5385
5386async fn acquire_sql_permit(
5396 state: &AppState,
5397 session: &MongrelSession,
5398 query: &RegisteredSqlQuery,
5399) -> std::result::Result<admission::SqlAdmissionGuard, mongreldb_query::MongrelQueryError> {
5400 refresh_node_pressure(state);
5401
5402 session.fire_test_hook(mongreldb_query::SqlTestHookPoint::WaitingForSqlPermit);
5403
5404 let permit = tokio::select! {
5406 permit = Arc::clone(&state.sql_semaphore).acquire_owned() => permit.map_err(|_| {
5407 mongreldb_query::MongrelQueryError::InvalidQueryState(
5408 "SQL admission semaphore closed".into(),
5409 )
5410 })?,
5411 _ = query.control().cancelled() => {
5412 return Err(cancellation_checkpoint_error(query));
5413 }
5414 };
5415
5416 let class = mongreldb_core::WorkloadClass::InteractiveSql;
5417 let priority = admission::priority_for_class(&state.resource_groups, class);
5418 let types_query_id = mongreldb_types::ids::QueryId::from_bytes(*query.id().as_bytes());
5419
5420 let work = match state
5422 .scheduler
5423 .submit_and_wait(
5424 admission::AdmitRequest {
5425 tenant: "default",
5426 class,
5427 priority,
5428 deadline: None,
5429 query_id: Some(types_query_id),
5430 tag: "sql",
5431 },
5432 query.control().cancelled(),
5433 )
5434 .await
5435 {
5436 Ok(work) => work,
5437 Err(admission::AdmitError::Rejected(error)) => {
5438 return Err(admission::scheduler_error_to_query(error));
5439 }
5440 Err(admission::AdmitError::Cancelled) => {
5441 return Err(cancellation_checkpoint_error(query));
5442 }
5443 Err(admission::AdmitError::PressureRejected { resource }) => {
5444 return Err(admission::admit_error_to_query(
5445 admission::AdmitError::PressureRejected { resource },
5446 ));
5447 }
5448 };
5449
5450 Ok(admission::SqlAdmissionGuard::new(permit, work))
5451}
5452
5453fn refresh_node_pressure(state: &AppState) {
5456 let Ok(mut governor) = state.node_governor.lock() else {
5457 return;
5458 };
5459 let db_gov = state.db.memory_governor();
5460 let ai_capacity = default_ai_max_concurrent();
5461 let inputs = admission::build_pressure_inputs(&admission::PressureInputSources {
5462 db_reserved_bytes: db_gov.total_used(),
5463 db_max_bytes: db_gov.max_bytes(),
5464 node_configured_max_bytes: governor.governor.max_bytes(),
5465 tablet_reserved_bytes: governor.tablet_reserved_bytes(),
5466 ai_capacity,
5467 ai_available: state.ai_semaphore.available_permits(),
5468 process_rss_bytes: admission::process_rss_bytes(),
5469 });
5470 admission::refresh_pressure(&mut governor, &inputs, &state.scheduler, Some(db_gov));
5471}
5472
5473fn caller_may_manage_query(
5474 state: &AppState,
5475 principal: &Option<mongreldb_core::Principal>,
5476 owner: Option<&str>,
5477) -> bool {
5478 let current = current_request_principal(state, principal);
5479 if (principal.is_some()
5480 || state.auth_token.is_some()
5481 || state.user_auth
5482 || state.db.require_auth_enabled())
5483 && current.is_none()
5484 {
5485 return false;
5486 }
5487 current.is_some_and(|principal| principal.is_admin)
5488 || owner == Some(request_owner(state, principal).as_str())
5489}
5490
5491fn query_phase_name(phase: SqlQueryPhase) -> &'static str {
5492 match phase {
5493 SqlQueryPhase::Queued => "queued",
5494 SqlQueryPhase::Planning => "planning",
5495 SqlQueryPhase::Executing => "executing",
5496 SqlQueryPhase::Streaming => "streaming",
5497 SqlQueryPhase::Serializing => "serializing",
5498 SqlQueryPhase::CommitCritical => "commit_critical",
5499 SqlQueryPhase::Cancelling => "cancelling",
5500 SqlQueryPhase::Completed => "completed",
5501 SqlQueryPhase::Failed => "failed",
5502 SqlQueryPhase::Cancelled => "cancelled",
5503 }
5504}
5505
5506fn commit_fence_outcome_name(outcome: mongreldb_query::CommitFenceOutcome) -> &'static str {
5507 match outcome {
5508 mongreldb_query::CommitFenceOutcome::NotReached => "not_reached",
5509 mongreldb_query::CommitFenceOutcome::CancelWon => "cancel_won",
5510 mongreldb_query::CommitFenceOutcome::CommitWon => "commit_won",
5511 }
5512}
5513
5514fn terminal_state_name(state: mongreldb_query::QueryTerminalState) -> &'static str {
5515 use mongreldb_query::QueryTerminalState;
5516 match state {
5517 QueryTerminalState::OutcomeUnknown => "outcome_unknown",
5518 QueryTerminalState::Completed => "completed",
5519 QueryTerminalState::FailedBeforeCommit => "failed_before_commit",
5520 QueryTerminalState::CancelledBeforeCommit => "cancelled_before_commit",
5521 QueryTerminalState::DeadlineBeforeCommit => "deadline_before_commit",
5522 QueryTerminalState::Committed => "committed",
5523 QueryTerminalState::CommittedWithError => "committed_with_error",
5524 QueryTerminalState::PartiallyCommitted => "partially_committed",
5525 QueryTerminalState::CancelledAfterCommit => "cancelled_after_commit",
5526 QueryTerminalState::DeadlineAfterCommit => "deadline_after_commit",
5527 }
5528}
5529
5530fn serialization_outcome_name(outcome: mongreldb_query::SerializationOutcome) -> &'static str {
5531 use mongreldb_query::SerializationOutcome;
5532 match outcome {
5533 SerializationOutcome::NotStarted => "not_started",
5534 SerializationOutcome::InProgress => "in_progress",
5535 SerializationOutcome::Succeeded => "succeeded",
5536 SerializationOutcome::Failed => "failed",
5537 }
5538}
5539
5540fn terminal_error_category_name(
5541 category: mongreldb_query::QueryTerminalErrorCategory,
5542) -> &'static str {
5543 use mongreldb_query::QueryTerminalErrorCategory;
5544 match category {
5545 QueryTerminalErrorCategory::Cancellation => "cancellation",
5546 QueryTerminalErrorCategory::Deadline => "deadline",
5547 QueryTerminalErrorCategory::ResultLimit => "result_limit",
5548 QueryTerminalErrorCategory::Serialization => "serialization",
5549 QueryTerminalErrorCategory::Execution => "execution",
5550 }
5551}
5552
5553fn terminal_error_retryable(error: Option<&mongreldb_query::QueryTerminalError>) -> bool {
5554 error.is_some_and(|error| {
5555 matches!(
5556 error.code.as_str(),
5557 "IDEMPOTENCY_STORE_FULL" | "IDEMPOTENCY_STORE_UNAVAILABLE"
5558 )
5559 })
5560}
5561
5562fn epoch_text(epoch: Option<u64>) -> Option<String> {
5563 epoch.map(|epoch| epoch.to_string())
5564}
5565
5566fn cancellation_reason_name(reason: CancellationReason) -> &'static str {
5567 reason.as_str()
5568}
5569
5570fn query_cancel_outcome(status: &mongreldb_query::QueryStatus) -> Option<&'static str> {
5571 match status.phase {
5572 SqlQueryPhase::CommitCritical => Some("too_late"),
5573 SqlQueryPhase::Completed | SqlQueryPhase::Failed | SqlQueryPhase::Cancelled => {
5574 Some("already_finished")
5575 }
5576 SqlQueryPhase::Cancelling => Some("accepted"),
5577 _ => None,
5578 }
5579}
5580
5581fn query_outcome_json(status: Option<&mongreldb_query::QueryStatus>) -> serde_json::Value {
5582 let Some(status) = status else {
5583 return json!({
5584 "committed": false,
5585 "committed_statements": 0,
5586 "last_commit_epoch": null,
5587 "last_commit_epoch_text": null,
5588 "first_commit_statement_index": null,
5589 "last_commit_statement_index": null,
5590 "completed_statements": 0,
5591 "statement_index": 0,
5592 "serialization": "not_started",
5593 });
5594 };
5595 if status.outcome_unknown {
5596 return json!({
5597 "committed": null,
5598 "committed_statements": null,
5599 "last_commit_epoch": null,
5600 "last_commit_epoch_text": null,
5601 "first_commit_statement_index": null,
5602 "last_commit_statement_index": null,
5603 "completed_statements": null,
5604 "statement_index": null,
5605 "serialization": "unknown",
5606 });
5607 }
5608 json!({
5609 "committed": status.durable_outcome.committed,
5610 "committed_statements": (!status.outcome_unknown).then_some(status.durable_outcome.committed_statements),
5611 "last_commit_epoch": (!status.outcome_unknown).then_some(status.durable_outcome.last_commit_epoch).flatten(),
5612 "last_commit_epoch_text": (!status.outcome_unknown).then_some(epoch_text(status.durable_outcome.last_commit_epoch)).flatten(),
5613 "first_commit_statement_index": (!status.outcome_unknown).then_some(status.durable_outcome.first_commit_statement_index).flatten(),
5614 "last_commit_statement_index": (!status.outcome_unknown).then_some(status.durable_outcome.last_commit_statement_index).flatten(),
5615 "completed_statements": (!status.outcome_unknown).then_some(status.completed_statements),
5616 "statement_index": (!status.outcome_unknown).then_some(status.statement_index),
5617 "serialization": serialization_outcome_name(status.serialization_outcome),
5618 })
5619}
5620
5621fn sql_terminal_idempotency_receipt(
5622 status: &mongreldb_query::QueryStatus,
5623) -> Option<sql_idempotency::SqlDurableReceipt> {
5624 if status.outcome_unknown {
5625 return None;
5626 }
5627 let terminal_state = status.terminal_state()?;
5628 if !status.durable_outcome.committed
5629 && terminal_state != mongreldb_query::QueryTerminalState::Completed
5630 {
5631 return None;
5632 }
5633 Some(sql_idempotency::SqlDurableReceipt {
5634 original_query_id: status.query_id.to_string(),
5635 status: status
5636 .terminal_state()
5637 .map(terminal_state_name)
5638 .unwrap_or("committed")
5639 .to_owned(),
5640 server_state: query_phase_name(status.phase).to_owned(),
5641 cancellation_reason: cancellation_reason_name(status.cancellation_reason).to_owned(),
5642 outcome: sql_idempotency::SqlReceiptOutcome {
5643 committed: status.durable_outcome.committed,
5644 committed_statements: status.durable_outcome.committed_statements,
5645 last_commit_epoch: status.durable_outcome.last_commit_epoch,
5646 last_commit_epoch_text: epoch_text(status.durable_outcome.last_commit_epoch),
5647 first_commit_statement_index: status.durable_outcome.first_commit_statement_index,
5648 last_commit_statement_index: status.durable_outcome.last_commit_statement_index,
5649 completed_statements: status.completed_statements,
5650 statement_index: status.statement_index,
5651 serialization: serialization_outcome_name(status.serialization_outcome).to_owned(),
5652 },
5653 terminal_error: status.terminal_error.as_ref().map(|error| {
5654 sql_idempotency::SqlReceiptTerminalError {
5655 code: error.code.clone(),
5656 category: terminal_error_category_name(error.category).to_owned(),
5657 }
5658 }),
5659 commit_receipt: None,
5660 })
5661}
5662
5663fn sql_idempotency_receipt_response(
5664 query_id: QueryId,
5665 receipt: &sql_idempotency::SqlDurableReceipt,
5666 replayed: bool,
5667 expires_at_ms: u64,
5668 persisted: bool,
5669) -> Response {
5670 let mut body = json!({
5671 "query_id": query_id.to_string(),
5672 "original_query_id": receipt.original_query_id,
5673 "status": receipt.status,
5674 "terminal_state": receipt.status,
5675 "server_state": receipt.server_state,
5676 "cancel_outcome": "already_finished",
5677 "cancellation_reason": receipt.cancellation_reason,
5678 "committed": receipt.outcome.committed,
5679 "committed_statements": receipt.outcome.committed_statements,
5680 "last_commit_epoch": receipt.outcome.last_commit_epoch,
5681 "last_commit_epoch_text": receipt.outcome.last_commit_epoch_text.as_deref(),
5682 "first_commit_statement_index": receipt.outcome.first_commit_statement_index,
5683 "last_commit_statement_index": receipt.outcome.last_commit_statement_index,
5684 "completed_statements": receipt.outcome.completed_statements,
5685 "statement_index": receipt.outcome.statement_index,
5686 "retryable": false,
5687 "idempotency_replayed": replayed,
5688 "idempotency_persisted": persisted,
5689 "idempotency_expires_at_ms": expires_at_ms,
5690 "outcome": receipt.outcome,
5691 "terminal_error": receipt.terminal_error,
5692 });
5693 if let Some(commit_receipt) = &receipt.commit_receipt {
5698 body["commit_receipt"] = json!(commit_receipt);
5699 }
5700 let mut response = Json(body).into_response();
5701 response.headers_mut().insert(
5702 "idempotency-replayed",
5703 axum::http::HeaderValue::from_static(if replayed { "true" } else { "false" }),
5704 );
5705 response.headers_mut().insert(
5706 "idempotency-persisted",
5707 axum::http::HeaderValue::from_static(if persisted { "true" } else { "false" }),
5708 );
5709 if let Ok(value) = axum::http::HeaderValue::from_str(&receipt.original_query_id) {
5710 response
5711 .headers_mut()
5712 .insert("x-mongreldb-original-query-id", value);
5713 }
5714 with_query_id(response, query_id)
5715}
5716
5717fn terminal_server_error_response(
5718 state: &AppState,
5719 query_id: QueryId,
5720 http_status: StatusCode,
5721 base_code: &'static str,
5722 message: impl Into<String>,
5723) -> Response {
5724 let status = state.query_registry.status(query_id);
5725 let committed = status
5726 .as_ref()
5727 .is_some_and(|status| status.durable_outcome.committed);
5728 let code = if committed && base_code.starts_with("SERIALIZATION_") {
5729 "SERIALIZATION_FAILED_AFTER_COMMIT"
5730 } else {
5731 base_code
5732 };
5733 let category = {
5736 use mongreldb_types::errors::ErrorCategory;
5737 match code {
5738 "RESULT_LIMIT_EXCEEDED" | "SQL_PAGE_STORE_FULL" | "ENTROPY_UNAVAILABLE" => {
5739 ErrorCategory::ResourceExhausted
5740 }
5741 "INVALID_SQL_PROJECTION" | "INVALID_PAGE_OFFSET" => {
5745 ErrorCategory::ClusterVersionMismatch
5746 }
5747 _ if code.starts_with("SERIALIZATION_") => ErrorCategory::ClusterVersionMismatch,
5748 _ => ErrorCategory::ReplicaUnavailable,
5749 }
5750 };
5751 let response_status = status
5752 .as_ref()
5753 .and_then(mongreldb_query::QueryStatus::terminal_state)
5754 .map(terminal_state_name)
5755 .unwrap_or(if committed {
5756 "committed_with_error"
5757 } else {
5758 "failed_before_commit"
5759 });
5760 let outcome = query_outcome_json(status.as_ref());
5761 with_query_id(
5762 (
5763 http_status,
5764 Json(json!({
5765 "query_id": query_id.to_string(),
5766 "status": response_status,
5767 "terminal_state": response_status,
5768 "committed": committed,
5769 "committed_statements": status.as_ref().map_or(0, |status| status.durable_outcome.committed_statements),
5770 "last_commit_epoch": status.as_ref().and_then(|status| status.durable_outcome.last_commit_epoch),
5771 "last_commit_epoch_text": epoch_text(status.as_ref().and_then(|status| status.durable_outcome.last_commit_epoch)),
5772 "first_commit_statement_index": status.as_ref().and_then(|status| status.durable_outcome.first_commit_statement_index),
5773 "last_commit_statement_index": status.as_ref().and_then(|status| status.durable_outcome.last_commit_statement_index),
5774 "completed_statements": status.as_ref().map_or(0, |status| status.completed_statements),
5775 "statement_index": status.as_ref().map_or(0, |status| status.statement_index),
5776 "cancel_outcome": null,
5777 "cancellation_reason": status.as_ref().map(|status| cancellation_reason_name(status.cancellation_reason)),
5778 "retryable": false,
5779 "server_state": status.as_ref().map(|status| query_phase_name(status.phase)),
5780 "outcome": outcome,
5781 "error": {
5782 "code": code,
5783 "message": message.into(),
5784 "category": category.to_string(),
5785 "category_code": category.code(),
5786 "query_id": query_id.to_string(),
5787 "committed": committed,
5788 "retryable": false,
5789 }
5790 })),
5791 )
5792 .into_response(),
5793 query_id,
5794 )
5795}
5796
5797fn query_not_found_response(query_id: Option<QueryId>) -> Response {
5798 let mut response = (
5799 StatusCode::NOT_FOUND,
5800 Json(json!({
5801 "query_id": query_id.map(|value| value.to_string()),
5802 "status": "unknown",
5803 "terminal_state": null,
5804 "committed": null,
5805 "committed_statements": null,
5806 "last_commit_epoch": null,
5807 "last_commit_epoch_text": null,
5808 "first_commit_statement_index": null,
5809 "last_commit_statement_index": null,
5810 "completed_statements": null,
5811 "statement_index": null,
5812 "cancel_outcome": "not_found",
5813 "cancellation_reason": null,
5814 "retryable": false,
5815 "server_state": "not_found",
5816 "outcome": {
5817 "committed": null,
5818 "committed_statements": null,
5819 "last_commit_epoch": null,
5820 "last_commit_epoch_text": null,
5821 "first_commit_statement_index": null,
5822 "last_commit_statement_index": null,
5823 "completed_statements": null,
5824 "statement_index": null,
5825 "serialization": "unknown",
5826 },
5827 "error": {
5828 "code": "QUERY_NOT_FOUND",
5829 "message": "query not found",
5830 "query_id": query_id.map(|value| value.to_string()),
5831 "committed": null,
5832 "retryable": false,
5833 }
5834 })),
5835 )
5836 .into_response();
5837 if let Some(query_id) = query_id {
5838 add_query_id_header(&mut response, query_id);
5839 }
5840 response
5841}
5842
5843fn query_session_header(
5844 headers: &axum::http::HeaderMap,
5845 query_id: Option<QueryId>,
5846) -> std::result::Result<Option<String>, Box<Response>> {
5847 match headers.get("x-session-id") {
5848 Some(value) => match value.to_str() {
5849 Ok(value) if value.len() <= 256 => Ok(Some(value.to_owned())),
5850 _ => Err(Box::new(bad_query_control_request(
5851 "X-Session-ID must be valid text no longer than 256 bytes",
5852 query_id,
5853 ))),
5854 },
5855 None => Ok(None),
5856 }
5857}
5858
5859fn pre_cancelled_query_response(
5860 query_id: QueryId,
5861 reason: CancellationReason,
5862 status: StatusCode,
5863) -> Response {
5864 with_query_id(
5865 (
5866 status,
5867 Json(json!({
5868 "query_id": query_id.to_string(),
5869 "status": "cancelled_before_start",
5870 "terminal_state": "cancelled_before_start",
5871 "state": "pre_cancelled",
5872 "server_state": "pre_cancelled",
5873 "cancel_outcome": "pre_cancelled",
5874 "committed": false,
5875 "committed_statements": 0,
5876 "last_commit_epoch": null,
5877 "last_commit_epoch_text": null,
5878 "first_commit_statement_index": null,
5879 "last_commit_statement_index": null,
5880 "completed_statements": 0,
5881 "statement_index": 0,
5882 "cancellation_reason": cancellation_reason_name(reason),
5883 "outcome": {
5884 "committed": false,
5885 "committed_statements": 0,
5886 "last_commit_epoch": null,
5887 "last_commit_epoch_text": null,
5888 "first_commit_statement_index": null,
5889 "last_commit_statement_index": null,
5890 "completed_statements": 0,
5891 "statement_index": 0,
5892 "serialization": "not_started",
5893 },
5894 "terminal_error": {
5895 "code": "QUERY_CANCELLED",
5896 "category": "cancellation",
5897 },
5898 "retryable": false,
5899 })),
5900 )
5901 .into_response(),
5902 query_id,
5903 )
5904}
5905
5906fn compact_finished_query_response(status: &CompactFinishedQuery) -> Response {
5907 let query_id = status.query_id;
5908 let durable = &status.durable_outcome;
5909 let outcome_unknown =
5910 status.terminal_state == mongreldb_query::QueryTerminalState::OutcomeUnknown;
5911 let terminal_error = status.terminal_error.as_ref().map(|error| {
5912 json!({
5913 "code": error.code,
5914 "category": terminal_error_category_name(error.category),
5915 })
5916 });
5917 with_query_id(
5918 Json(json!({
5919 "detail": "compact",
5920 "query_id": query_id.to_string(),
5921 "status": terminal_state_name(status.terminal_state),
5922 "terminal_state": terminal_state_name(status.terminal_state),
5923 "state": query_phase_name(status.phase),
5924 "server_state": query_phase_name(status.phase),
5925 "cancel_outcome": "already_finished",
5926 "code": "QUERY_ALREADY_FINISHED",
5927 "committed": (!outcome_unknown).then_some(durable.committed),
5928 "committed_statements": (!outcome_unknown).then_some(durable.committed_statements),
5929 "last_commit_epoch": (!outcome_unknown).then_some(durable.last_commit_epoch).flatten(),
5930 "last_commit_epoch_text": (!outcome_unknown).then_some(epoch_text(durable.last_commit_epoch)).flatten(),
5931 "first_commit_statement_index": (!outcome_unknown).then_some(durable.first_commit_statement_index).flatten(),
5932 "last_commit_statement_index": (!outcome_unknown).then_some(durable.last_commit_statement_index).flatten(),
5933 "completed_statements": (!outcome_unknown).then_some(status.completed_statements),
5934 "statement_index": (!outcome_unknown).then_some(status.statement_index),
5935 "cancellation_reason": cancellation_reason_name(status.cancellation_reason),
5936 "outcome": {
5937 "committed": (!outcome_unknown).then_some(durable.committed),
5938 "committed_statements": (!outcome_unknown).then_some(durable.committed_statements),
5939 "last_commit_epoch": (!outcome_unknown).then_some(durable.last_commit_epoch).flatten(),
5940 "last_commit_epoch_text": (!outcome_unknown).then_some(epoch_text(durable.last_commit_epoch)).flatten(),
5941 "first_commit_statement_index": (!outcome_unknown).then_some(durable.first_commit_statement_index).flatten(),
5942 "last_commit_statement_index": (!outcome_unknown).then_some(durable.last_commit_statement_index).flatten(),
5943 "completed_statements": (!outcome_unknown).then_some(status.completed_statements),
5944 "statement_index": (!outcome_unknown).then_some(status.statement_index),
5945 "serialization": serialization_outcome_name(status.serialization_outcome),
5946 },
5947 "terminal_error": terminal_error,
5948 "retryable": terminal_error_retryable(status.terminal_error.as_ref()),
5949 }))
5950 .into_response(),
5951 query_id,
5952 )
5953}
5954
5955async fn query_status(
5956 State(state): State<Arc<AppState>>,
5957 OptionalPrincipal(principal): OptionalPrincipal,
5958 Path(query_id): Path<String>,
5959 headers: axum::http::HeaderMap,
5960) -> Response {
5961 let Ok(query_id) = query_id.parse::<QueryId>() else {
5962 return query_not_found_response(None);
5963 };
5964 if !request_identity_is_current(&state, &principal) {
5965 return query_not_found_response(Some(query_id));
5966 }
5967 let requested_session = match query_session_header(&headers, Some(query_id)) {
5968 Ok(session_id) => session_id,
5969 Err(response) => return *response,
5970 };
5971 let owner = request_owner(&state, &principal);
5972 let is_admin =
5973 current_request_principal(&state, &principal).is_some_and(|principal| principal.is_admin);
5974 let _lifecycle = state
5975 .query_lifecycle
5976 .lock()
5977 .unwrap_or_else(|error| error.into_inner());
5978 let Some(status) = state.query_registry.status(query_id) else {
5979 if let Some(finished) = state.query_registry.compact_finished_status(query_id) {
5980 if !caller_may_manage_query(&state, &principal, finished.owner.as_deref())
5981 || requested_session
5982 .as_deref()
5983 .is_some_and(|session| finished.session_id.as_deref() != Some(session))
5984 {
5985 return query_not_found_response(Some(query_id));
5986 }
5987 return compact_finished_query_response(&finished);
5988 }
5989 let reason = state
5990 .pre_cancellations
5991 .reason(query_id, &owner, requested_session.as_deref())
5992 .or_else(|| {
5993 is_admin.then(|| match requested_session.as_deref() {
5994 Some(session_id) => state
5995 .pre_cancellations
5996 .reason_for_query_in_session(query_id, session_id),
5997 None => state.pre_cancellations.reason_for_query(query_id),
5998 })?
5999 });
6000 if let Some(reason) = reason {
6001 return pre_cancelled_query_response(query_id, reason, StatusCode::OK);
6002 }
6003 return query_not_found_response(Some(query_id));
6004 };
6005 if !caller_may_manage_query(&state, &principal, status.owner.as_deref())
6006 || requested_session
6007 .as_deref()
6008 .is_some_and(|session| status.session_id.as_deref() != Some(session))
6009 {
6010 return query_not_found_response(Some(query_id));
6011 }
6012 let terminal_status = status.terminal_state().map(terminal_state_name);
6013 let terminal_error = status.terminal_error.as_ref().map(|error| {
6014 json!({
6015 "code": error.code,
6016 "category": terminal_error_category_name(error.category),
6017 })
6018 });
6019 let retryable = terminal_error_retryable(status.terminal_error.as_ref());
6020 let cancel_outcome = query_cancel_outcome(&status);
6021 let outcome = query_outcome_json(Some(&status));
6022 let response = Json(json!({
6023 "query_id": query_id.to_string(),
6024 "status": terminal_status.unwrap_or(if status.durable_outcome.committed {
6025 "committed"
6026 } else {
6027 "running"
6028 }),
6029 "terminal_state": terminal_status,
6030 "state": query_phase_name(status.phase),
6031 "server_state": query_phase_name(status.phase),
6032 "started_ms_ago": status.started_at.elapsed().as_millis(),
6033 "deadline_ms_remaining": status.deadline.map(|deadline| {
6034 deadline.saturating_duration_since(std::time::Instant::now()).as_millis()
6035 }),
6036 "session_id": status.session_id,
6037 "operation": status.operation,
6038 "committed": (!status.outcome_unknown).then_some(status.committed),
6039 "committed_statements": (!status.outcome_unknown).then_some(status.durable_outcome.committed_statements),
6040 "last_commit_epoch": (!status.outcome_unknown).then_some(status.durable_outcome.last_commit_epoch).flatten(),
6041 "last_commit_epoch_text": (!status.outcome_unknown).then_some(epoch_text(status.durable_outcome.last_commit_epoch)).flatten(),
6042 "first_commit_statement_index": (!status.outcome_unknown).then_some(status.durable_outcome.first_commit_statement_index).flatten(),
6043 "last_commit_statement_index": (!status.outcome_unknown).then_some(status.durable_outcome.last_commit_statement_index).flatten(),
6044 "cancellation_reason": cancellation_reason_name(status.cancellation_reason),
6045 "completed_statements": (!status.outcome_unknown).then_some(status.completed_statements),
6046 "statement_index": (!status.outcome_unknown).then_some(status.statement_index),
6047 "cancel_outcome": cancel_outcome,
6048 "retryable": retryable,
6049 "outcome": outcome,
6050 "terminal_error": terminal_error,
6051 "trace": {
6052 "queue_duration_us": status.queue_duration.as_micros(),
6053 "planning_duration_us": status.planning_duration.as_micros(),
6054 "execution_duration_us": status.execution_duration.as_micros(),
6055 "serialization_duration_us": status.serialization_duration.as_micros(),
6056 "cancel_requested_phase": status.cancel_requested_phase.map(query_phase_name),
6057 "cancel_observed_phase": status.cancel_observed_phase.map(query_phase_name),
6058 "commit_fence_outcome": commit_fence_outcome_name(status.commit_fence_outcome),
6059 },
6060 }))
6061 .into_response();
6062 with_query_id(response, query_id)
6063}
6064
6065async fn cancel_query(
6066 State(state): State<Arc<AppState>>,
6067 OptionalPrincipal(principal): OptionalPrincipal,
6068 Path(query_id): Path<String>,
6069 headers: axum::http::HeaderMap,
6070) -> Response {
6071 let Ok(query_id) = query_id.parse::<QueryId>() else {
6072 return query_not_found_response(None);
6073 };
6074 if !request_identity_is_current(&state, &principal) {
6075 return query_not_found_response(Some(query_id));
6076 }
6077 let requested_session = match query_session_header(&headers, Some(query_id)) {
6078 Ok(session_id) => session_id,
6079 Err(response) => return *response,
6080 };
6081 let owner = request_owner(&state, &principal);
6082 let _lifecycle = state
6083 .query_lifecycle
6084 .lock()
6085 .unwrap_or_else(|error| error.into_inner());
6086 state.metrics.inc_sql_cancel_requests();
6087 let Some(status) = state.query_registry.status(query_id) else {
6088 if let Some(finished) = state.query_registry.compact_finished_status(query_id) {
6089 if !caller_may_manage_query(&state, &principal, finished.owner.as_deref())
6090 || requested_session
6091 .as_deref()
6092 .is_some_and(|session| finished.session_id.as_deref() != Some(session))
6093 {
6094 return query_not_found_response(Some(query_id));
6095 }
6096 return compact_finished_query_response(&finished);
6097 }
6098 return match state.pre_cancellations.insert(
6099 query_id,
6100 &owner,
6101 requested_session.as_deref(),
6102 CancellationReason::ClientRequest,
6103 ) {
6104 Ok(()) => {
6105 state.metrics.inc_sql_commit_cancel_winner_cancel();
6106 pre_cancelled_query_response(
6107 query_id,
6108 CancellationReason::ClientRequest,
6109 StatusCode::ACCEPTED,
6110 )
6111 }
6112 Err(pre_cancel::InsertError::MetadataTooLarge) => bad_query_control_request(
6113 "query owner or session metadata exceeds 256 bytes",
6114 Some(query_id),
6115 ),
6116 Err(
6117 pre_cancel::InsertError::Full
6118 | pre_cancel::InsertError::OwnerLimit
6119 | pre_cancel::InsertError::RateLimited,
6120 ) => with_query_id(
6121 (
6122 StatusCode::TOO_MANY_REQUESTS,
6123 Json(json!({
6124 "query_id": query_id.to_string(),
6125 "status": "failed_before_commit",
6126 "terminal_state": "failed_before_commit",
6127 "server_state": "failed",
6128 "cancel_outcome": null,
6129 "cancellation_reason": null,
6130 "committed": false,
6131 "committed_statements": 0,
6132 "last_commit_epoch": null,
6133 "last_commit_epoch_text": null,
6134 "first_commit_statement_index": null,
6135 "last_commit_statement_index": null,
6136 "completed_statements": 0,
6137 "statement_index": 0,
6138 "retryable": true,
6139 "outcome": {
6140 "committed": false,
6141 "committed_statements": 0,
6142 "last_commit_epoch": null,
6143 "last_commit_epoch_text": null,
6144 "first_commit_statement_index": null,
6145 "last_commit_statement_index": null,
6146 "completed_statements": 0,
6147 "statement_index": 0,
6148 "serialization": "not_started",
6149 },
6150 "error": {
6151 "code": "QUERY_REGISTRY_FULL",
6152 "message": "pre-registration cancellation limit reached",
6153 "query_id": query_id.to_string(),
6154 "committed": false,
6155 "retryable": true,
6156 }
6157 })),
6158 )
6159 .into_response(),
6160 query_id,
6161 ),
6162 };
6163 };
6164 if !caller_may_manage_query(&state, &principal, status.owner.as_deref())
6165 || requested_session
6166 .as_deref()
6167 .is_some_and(|session| status.session_id.as_deref() != Some(session))
6168 {
6169 return query_not_found_response(Some(query_id));
6170 }
6171 let (http_status, mut body) = match state.query_registry.cancel(query_id) {
6172 CancelOutcome::Accepted => (
6173 {
6174 state.metrics.inc_sql_commit_cancel_winner_cancel();
6175 StatusCode::ACCEPTED
6176 },
6177 json!({
6178 "query_id": query_id.to_string(),
6179 "state": "cancellation_requested",
6180 "cancel_outcome": "accepted",
6181 }),
6182 ),
6183 CancelOutcome::AlreadyCancelling => (
6184 StatusCode::OK,
6185 json!({
6186 "query_id": query_id.to_string(),
6187 "state": "cancelling",
6188 "cancel_outcome": "already_cancelling",
6189 }),
6190 ),
6191 CancelOutcome::TooLate => (
6192 {
6193 state.metrics.inc_sql_commit_cancel_winner_commit();
6194 StatusCode::CONFLICT
6195 },
6196 json!({
6197 "query_id": query_id.to_string(),
6198 "state": "commit_critical",
6199 "cancel_outcome": "too_late",
6200 "committed": status.durable_outcome.committed,
6201 "outcome": query_outcome_json(Some(&status)),
6202 "retryable": false,
6203 "error": {
6204 "code": "CANCEL_TOO_LATE",
6205 "message": "the query has entered its durable commit phase",
6206 "committed": status.durable_outcome.committed,
6207 "retryable": false,
6208 }
6209 }),
6210 ),
6211 CancelOutcome::AlreadyFinished => (
6212 StatusCode::OK,
6213 json!({
6214 "query_id": query_id.to_string(),
6215 "state": "finished",
6216 "status": status.terminal_state().map(terminal_state_name),
6217 "cancel_outcome": "already_finished",
6218 "code": "QUERY_ALREADY_FINISHED",
6219 "committed": status.durable_outcome.committed,
6220 "outcome": query_outcome_json(Some(&status)),
6221 "retryable": false,
6222 }),
6223 ),
6224 CancelOutcome::NotFound => return query_not_found_response(Some(query_id)),
6225 };
6226 let status = state.query_registry.status(query_id).unwrap_or(status);
6227 let response_status = status.terminal_state().map(terminal_state_name).unwrap_or(
6228 if status.durable_outcome.committed {
6229 "committed"
6230 } else {
6231 "running"
6232 },
6233 );
6234 if let Some(body) = body.as_object_mut() {
6235 body.insert("status".into(), json!(response_status));
6236 body.insert(
6237 "terminal_state".into(),
6238 json!(status.terminal_state().map(terminal_state_name)),
6239 );
6240 body.insert(
6241 "committed".into(),
6242 json!((!status.outcome_unknown).then_some(status.durable_outcome.committed)),
6243 );
6244 body.insert(
6245 "committed_statements".into(),
6246 json!((!status.outcome_unknown).then_some(status.durable_outcome.committed_statements)),
6247 );
6248 body.insert(
6249 "last_commit_epoch".into(),
6250 json!((!status.outcome_unknown)
6251 .then_some(status.durable_outcome.last_commit_epoch)
6252 .flatten()),
6253 );
6254 body.insert(
6255 "last_commit_epoch_text".into(),
6256 json!((!status.outcome_unknown)
6257 .then_some(epoch_text(status.durable_outcome.last_commit_epoch))
6258 .flatten()),
6259 );
6260 body.insert(
6261 "first_commit_statement_index".into(),
6262 json!((!status.outcome_unknown)
6263 .then_some(status.durable_outcome.first_commit_statement_index)
6264 .flatten()),
6265 );
6266 body.insert(
6267 "last_commit_statement_index".into(),
6268 json!((!status.outcome_unknown)
6269 .then_some(status.durable_outcome.last_commit_statement_index)
6270 .flatten()),
6271 );
6272 body.insert(
6273 "completed_statements".into(),
6274 json!((!status.outcome_unknown).then_some(status.completed_statements)),
6275 );
6276 body.insert(
6277 "statement_index".into(),
6278 json!((!status.outcome_unknown).then_some(status.statement_index)),
6279 );
6280 body.insert(
6281 "cancellation_reason".into(),
6282 json!(cancellation_reason_name(status.cancellation_reason)),
6283 );
6284 body.insert("retryable".into(), json!(false));
6285 body.insert("server_state".into(), json!(query_phase_name(status.phase)));
6286 body.insert("outcome".into(), query_outcome_json(Some(&status)));
6287 }
6288 with_query_id((http_status, Json(body)).into_response(), query_id)
6289}
6290
6291#[derive(Deserialize)]
6292#[serde(deny_unknown_fields)]
6293struct SqlContinuationRequest {
6294 cursor: String,
6295 #[serde(default)]
6296 operation_id: Option<QueryId>,
6297 #[serde(default)]
6298 timeout_ms: Option<u64>,
6299}
6300
6301fn register_page_operation(
6302 state: &AppState,
6303 options: SqlQueryOptions,
6304) -> mongreldb_query::Result<RegisteredSqlQuery> {
6305 let query_id = options.query_id.ok_or_else(|| {
6306 mongreldb_query::MongrelQueryError::InvalidQueryState(
6307 "page operation registration requires an operation id".into(),
6308 )
6309 })?;
6310 let owner = options.owner.clone().unwrap_or_default();
6311 let session_id = options.session_id.clone();
6312 let _lifecycle = state
6313 .query_lifecycle
6314 .lock()
6315 .unwrap_or_else(|error| error.into_inner());
6316 let reason = match state.pre_cancellations.lookup_for_registration(
6317 query_id,
6318 &owner,
6319 session_id.as_deref(),
6320 ) {
6321 pre_cancel::RegistrationLookup::NoReservation => None,
6322 pre_cancel::RegistrationLookup::Matching(reason) => Some(reason),
6323 pre_cancel::RegistrationLookup::ReservedByAnotherIdentity => {
6324 return Err(mongreldb_query::MongrelQueryError::QueryIdConflict { query_id });
6325 }
6326 };
6327 let query = state.query_registry.register(options)?;
6328 if let Some(reason) = reason {
6329 state
6330 .pre_cancellations
6331 .take(query_id, &owner, session_id.as_deref());
6332 query.request_cancel(reason);
6333 let error = cancellation_checkpoint_error(&query);
6334 query.fail();
6335 return Err(error);
6336 }
6337 Ok(query)
6338}
6339
6340async fn continue_sql_page(
6341 State(state): State<Arc<AppState>>,
6342 OptionalPrincipal(principal): OptionalPrincipal,
6343 headers: axum::http::HeaderMap,
6344 Json(request): Json<SqlContinuationRequest>,
6345) -> Response {
6346 let query_id = match request.operation_id {
6347 Some(query_id) => query_id,
6348 None => match QueryId::random() {
6349 Ok(query_id) => query_id,
6350 Err(error) => return query_error_response(&error, None),
6351 },
6352 };
6353 if !request_identity_is_current(&state, &principal) {
6354 return with_query_id(
6355 sql_cursor_error_response(
6356 StatusCode::NOT_FOUND,
6357 "SQL_CURSOR_NOT_FOUND",
6358 "SQL continuation result is unavailable",
6359 query_id,
6360 ),
6361 query_id,
6362 );
6363 }
6364 let owner = request_owner(&state, &principal);
6365 let session_id = match query_session_header(&headers, Some(query_id)) {
6366 Ok(session_id) => session_id,
6367 Err(response) => return *response,
6368 };
6369 let timeout_ms = request.timeout_ms.unwrap_or_else(|| {
6370 state
6371 .sql_page_default_timeout
6372 .as_millis()
6373 .min(u128::from(u64::MAX)) as u64
6374 });
6375 if timeout_ms == 0 || Duration::from_millis(timeout_ms) > state.sql_page_max_timeout {
6376 return bad_query_control_request(
6377 format!(
6378 "timeout_ms must be positive and no greater than {}",
6379 state.sql_page_max_timeout.as_millis()
6380 ),
6381 Some(query_id),
6382 );
6383 }
6384 let query = match register_page_operation(
6385 &state,
6386 SqlQueryOptions {
6387 query_id: Some(query_id),
6388 timeout: Some(Duration::from_millis(timeout_ms)),
6389 owner: Some(owner.clone()),
6390 session_id,
6391 parent_control: None,
6392 },
6393 ) {
6394 Ok(query) => query,
6395 Err(error) => return tracked_query_error_response(&state, &error, Some(query_id)),
6396 };
6397 query.set_sql_metadata("CONTINUE SQL PAGE");
6398 let _permit = match tokio::select! {
6399 permit = Arc::clone(&state.sql_page_semaphore).acquire_owned() => permit.map_err(|_| {
6400 mongreldb_query::MongrelQueryError::InvalidQueryState(
6401 "SQL page admission semaphore closed".into(),
6402 )
6403 }),
6404 _ = query.control().cancelled() => Err(cancellation_checkpoint_error(&query)),
6405 } {
6406 Ok(permit) => permit,
6407 Err(error) => {
6408 query.fail();
6409 return tracked_query_error_response(&state, &error, Some(query_id));
6410 }
6411 };
6412 if let Err(error) = query.transition(SqlQueryPhase::Queued, SqlQueryPhase::Executing) {
6413 query.fail();
6414 return tracked_query_error_response(&state, &error, Some(query_id));
6415 }
6416 let fail = |status, code, message| {
6417 query.record_terminal_error(code, mongreldb_query::QueryTerminalErrorCategory::Execution);
6418 query.fail();
6419 with_query_id(
6420 sql_cursor_error_response(status, code, message, query_id),
6421 query_id,
6422 )
6423 };
6424 if request.cursor.is_empty() || request.cursor.len() > 2_048 {
6425 return fail(
6426 StatusCode::BAD_REQUEST,
6427 "INVALID_SQL_CURSOR",
6428 "invalid SQL continuation cursor",
6429 );
6430 }
6431 if let Err(error) = query.checkpoint() {
6432 query.fail();
6433 return tracked_query_error_response(&state, &error, Some(query_id));
6434 }
6435 let cursor_mac_key = match state.cursor_mac_key.get() {
6436 Ok(key) => key,
6437 Err(_) => {
6438 return fail(
6439 StatusCode::INTERNAL_SERVER_ERROR,
6440 "ENTROPY_UNAVAILABLE",
6441 "OS CSPRNG unavailable",
6442 );
6443 }
6444 };
6445 match state.sql_pages.continue_page_with_control(
6446 &request.cursor,
6447 &owner,
6448 &cursor_mac_key,
6449 sql_pages::SqlPageBinding {
6450 security_version: state.db.security_version(),
6451 catalog_epoch: state.db.catalog_snapshot().db_epoch,
6452 },
6453 &query,
6454 ) {
6455 Ok(page) => {
6456 let page_byte_count = page.byte_count;
6457 if let Err(error) = query.begin_serialization() {
6458 query.fail();
6459 return tracked_query_error_response(&state, &error, Some(query_id));
6460 }
6461 let serialization_query = query.clone();
6462 match tokio::task::spawn_blocking(move || {
6463 serialize_sql_page_controlled(page, &serialization_query)
6464 })
6465 .await
6466 {
6467 Ok(Ok(body)) => {
6468 if let Err(error) = query.try_complete() {
6469 return tracked_query_error_response(&state, &error, Some(query_id));
6470 }
6471 state.metrics.add_sql_output_bytes(page_byte_count);
6472 with_query_id(sql_page_response(body), query_id)
6473 }
6474 Ok(Err(ControlledPageSerializationError::Query(error))) => {
6475 query.fail();
6476 tracked_query_error_response(&state, &error, Some(query_id))
6477 }
6478 Ok(Err(ControlledPageSerializationError::Encoding)) => {
6479 if let Err(error) = query.checkpoint() {
6480 query.fail();
6481 return tracked_query_error_response(&state, &error, Some(query_id));
6482 }
6483 fail(
6484 StatusCode::INTERNAL_SERVER_ERROR,
6485 "SERIALIZATION_FAILED",
6486 "failed to serialize SQL continuation page",
6487 )
6488 }
6489 Err(_) => {
6490 if let Err(error) = query.checkpoint() {
6491 query.fail();
6492 return tracked_query_error_response(&state, &error, Some(query_id));
6493 }
6494 fail(
6495 StatusCode::INTERNAL_SERVER_ERROR,
6496 "SERIALIZATION_WORKER_FAILED",
6497 "SQL continuation serialization worker failed",
6498 )
6499 }
6500 }
6501 }
6502 Err(sql_pages::CursorError::Cancelled) => {
6503 let error = cancellation_checkpoint_error(&query);
6504 query.fail();
6505 tracked_query_error_response(&state, &error, Some(query_id))
6506 }
6507 Err(sql_pages::CursorError::Invalid) => fail(
6508 StatusCode::BAD_REQUEST,
6509 "INVALID_SQL_CURSOR",
6510 "invalid SQL continuation cursor",
6511 ),
6512 Err(sql_pages::CursorError::Expired) => fail(
6513 StatusCode::GONE,
6514 "SQL_CURSOR_EXPIRED",
6515 "SQL continuation cursor expired",
6516 ),
6517 Err(sql_pages::CursorError::NotFound) => fail(
6518 StatusCode::NOT_FOUND,
6519 "SQL_CURSOR_NOT_FOUND",
6520 "SQL continuation result is unavailable",
6521 ),
6522 Err(sql_pages::CursorError::PageLimit) => fail(
6523 StatusCode::PAYLOAD_TOO_LARGE,
6524 "RESULT_LIMIT_EXCEEDED",
6525 "one projected row exceeds the page byte or token limit",
6526 ),
6527 }
6528}
6529
6530fn sql_cursor_error_response(
6531 status: StatusCode,
6532 code: &'static str,
6533 message: &'static str,
6534 query_id: QueryId,
6535) -> Response {
6536 let category = {
6539 use mongreldb_types::errors::ErrorCategory;
6540 match code {
6541 "SQL_CURSOR_NOT_FOUND" | "SQL_CURSOR_EXPIRED" => ErrorCategory::StaleMetadata,
6544 "RESULT_LIMIT_EXCEEDED" | "ENTROPY_UNAVAILABLE" => ErrorCategory::ResourceExhausted,
6545 _ => ErrorCategory::ClusterVersionMismatch,
6546 }
6547 };
6548 (
6549 status,
6550 Json(json!({
6551 "query_id": query_id.to_string(),
6552 "status": "failed_before_commit",
6553 "terminal_state": "failed_before_commit",
6554 "server_state": "failed",
6555 "committed": false,
6556 "committed_statements": 0,
6557 "last_commit_epoch": null,
6558 "last_commit_epoch_text": null,
6559 "first_commit_statement_index": null,
6560 "last_commit_statement_index": null,
6561 "completed_statements": 0,
6562 "statement_index": 0,
6563 "cancel_outcome": "already_finished",
6564 "cancellation_reason": "none",
6565 "retryable": false,
6566 "outcome": {
6567 "committed": false,
6568 "committed_statements": 0,
6569 "last_commit_epoch": null,
6570 "last_commit_epoch_text": null,
6571 "first_commit_statement_index": null,
6572 "last_commit_statement_index": null,
6573 "completed_statements": 0,
6574 "statement_index": 0,
6575 "serialization": "not_started",
6576 },
6577 "error": {
6578 "code": code,
6579 "message": message,
6580 "category": category.to_string(),
6581 "category_code": category.code(),
6582 "query_id": query_id.to_string(),
6583 "committed": false,
6584 "retryable": false,
6585 }
6586 })),
6587 )
6588 .into_response()
6589}
6590
6591async fn sql(
6592 State(state): State<Arc<AppState>>,
6593 OptionalPrincipal(principal): OptionalPrincipal,
6594 headers: axum::http::HeaderMap,
6595 Json(req): Json<SqlRequest>,
6596) -> Response {
6597 if !state.accepting_sql.load(Ordering::Acquire) {
6598 return (StatusCode::SERVICE_UNAVAILABLE, "server is shutting down").into_response();
6599 }
6600 if !request_identity_is_current(&state, &principal) {
6601 return StatusCode::UNAUTHORIZED.into_response();
6602 }
6603 let session_id = match query_session_header(&headers, None) {
6608 Ok(session_id) => session_id,
6609 Err(response) => return *response,
6610 };
6611
6612 let owner = request_owner(&state, &principal);
6613 if let Some(sid) = session_id {
6614 let Some(entry) = state.sessions.get(&sid, &owner) else {
6615 return (
6616 StatusCode::NOT_FOUND,
6617 "session not found or not owned by caller",
6618 )
6619 .into_response();
6620 };
6621 let (options, query_id) = match resolve_query_options(
6622 &state,
6623 &headers,
6624 req.query_id,
6625 req.timeout_ms,
6626 owner.clone(),
6627 Some(sid.clone()),
6628 ) {
6629 Ok(options) => options,
6630 Err(response) => return *response,
6631 };
6632 let query = match register_controlled_query(&state, &entry.session(), options) {
6633 Ok(query) => query,
6634 Err(error) => return tracked_query_error_response(&state, &error, Some(query_id)),
6635 };
6636 let registration = RegisteredQueryGuard::new(query);
6637 if mongreldb_query::contains_boolean_ai_predicate(&req.sql) {
6638 registration.fail();
6639 return with_query_id(remote_boolean_ai_error(), query_id);
6640 }
6641 let output_limits = match resolve_sql_output_limits(&state, &req, query_id) {
6642 Ok(limits) => limits,
6643 Err(response) => {
6644 registration.fail();
6645 return *response;
6646 }
6647 };
6648 let (registration, pagination) =
6649 match resolve_sql_pagination(&headers, &req, output_limits, registration, query_id) {
6650 Ok(resolved) => resolved,
6651 Err(response) => return *response,
6652 };
6653 let sql_permit =
6654 match acquire_sql_permit(&state, &entry.session(), registration.query()).await {
6655 Ok(permit) => permit,
6656 Err(error) => {
6657 return tracked_query_error_response(&state, &error, Some(query_id));
6658 }
6659 };
6660 let _guard = tokio::select! {
6662 guard = entry.lock.lock() => guard,
6663 _ = registration.query().control().cancelled() => {
6664 return tracked_query_error_response(
6665 &state,
6666 &cancellation_checkpoint_error(registration.query()),
6667 Some(query_id),
6668 );
6669 }
6670 };
6671 if entry.is_closed() {
6674 return (StatusCode::NOT_FOUND, "session no longer available").into_response();
6675 }
6676 if req.idempotency_key.is_some() || headers.contains_key("idempotency-key") {
6677 entry
6678 .session()
6679 .fire_test_hook(mongreldb_query::SqlTestHookPoint::BeforeServerIdempotencyCheck);
6680 }
6681 let (registration, idempotency) = match begin_sql_idempotency(
6682 &state,
6683 SqlIdempotencyContext {
6684 headers: &headers,
6685 request: &req,
6686 output_limits,
6687 owner: &owner,
6688 session_id: Some(&sid),
6689 session_in_transaction: entry.session().staged_sql_operation_count().is_some(),
6690 query_id,
6691 },
6692 registration,
6693 )
6694 .await
6695 {
6696 Ok(resolved) => resolved,
6697 Err(response) => return response,
6698 };
6699 entry.touch();
6700 let query = registration.into_query();
6701 let (response, idempotent_commit_ts) = execute_sql(
6702 &state,
6703 &principal,
6704 &entry.session(),
6705 ResolvedSqlRequest {
6706 request: req,
6707 output_limits,
6708 idempotency,
6709 pagination,
6710 },
6711 query,
6712 query_id,
6713 sql_permit,
6714 )
6715 .await;
6716 let durable = state
6728 .query_registry
6729 .status(query_id)
6730 .map(|status| status.durable_outcome);
6731 entry.sync_record_after_request(match durable {
6732 Some(outcome) if outcome.committed => Some(
6733 idempotent_commit_ts
6734 .or(outcome.commit_ts)
6735 .or_else(|| {
6736 outcome.last_commit_epoch.and_then(|epoch| {
6737 state.db.commit_ts_for_epoch(mongreldb_core::Epoch(epoch))
6738 })
6739 })
6740 .unwrap_or_else(|| ryw_commit_timestamp(&state.db)),
6741 ),
6742 _ => None,
6743 });
6744 response
6745 } else {
6746 let session = match MongrelSession::open_with_external_modules_as(
6747 Arc::clone(&state.db),
6748 state.external_modules.iter().cloned(),
6749 request_principal(&state, &principal),
6750 ) {
6751 Ok(session) => session.with_query_registry(Arc::clone(&state.query_registry)),
6752 Err(e) => return (status_for_query_error(&e), e.to_string()).into_response(),
6753 };
6754 let (options, query_id) = match resolve_query_options(
6755 &state,
6756 &headers,
6757 req.query_id,
6758 req.timeout_ms,
6759 owner.clone(),
6760 None,
6761 ) {
6762 Ok(options) => options,
6763 Err(response) => return *response,
6764 };
6765 let query = match register_controlled_query(&state, &session, options) {
6766 Ok(query) => query,
6767 Err(error) => return tracked_query_error_response(&state, &error, Some(query_id)),
6768 };
6769 let registration = RegisteredQueryGuard::new(query);
6770 if mongreldb_query::contains_boolean_ai_predicate(&req.sql) {
6771 registration.fail();
6772 return with_query_id(remote_boolean_ai_error(), query_id);
6773 }
6774 let output_limits = match resolve_sql_output_limits(&state, &req, query_id) {
6775 Ok(limits) => limits,
6776 Err(response) => {
6777 registration.fail();
6778 return *response;
6779 }
6780 };
6781 let (registration, pagination) =
6782 match resolve_sql_pagination(&headers, &req, output_limits, registration, query_id) {
6783 Ok(resolved) => resolved,
6784 Err(response) => return *response,
6785 };
6786 let (registration, idempotency) = match begin_sql_idempotency(
6787 &state,
6788 SqlIdempotencyContext {
6789 headers: &headers,
6790 request: &req,
6791 output_limits,
6792 owner: &owner,
6793 session_id: None,
6794 session_in_transaction: false,
6795 query_id,
6796 },
6797 registration,
6798 )
6799 .await
6800 {
6801 Ok(resolved) => resolved,
6802 Err(response) => return response,
6803 };
6804 let sql_permit = match acquire_sql_permit(&state, &session, registration.query()).await {
6805 Ok(permit) => permit,
6806 Err(error) => {
6807 if let Some(idempotency) = idempotency {
6808 idempotency.abort();
6809 }
6810 return tracked_query_error_response(&state, &error, Some(query_id));
6811 }
6812 };
6813 let query = registration.into_query();
6814 execute_sql(
6815 &state,
6816 &principal,
6817 &session,
6818 ResolvedSqlRequest {
6819 request: req,
6820 output_limits,
6821 idempotency,
6822 pagination,
6823 },
6824 query,
6825 query_id,
6826 sql_permit,
6827 )
6828 .await
6829 .0
6830 }
6831}
6832
6833fn ryw_commit_timestamp(db: &mongreldb_core::Database) -> mongreldb_types::hlc::HlcTimestamp {
6842 if let Some(ts) = db.begin().read_ts() {
6843 return ts;
6844 }
6845 mongreldb_types::hlc::HlcTimestamp {
6846 physical_micros: sessions::now_unix_micros(),
6847 logical: 0,
6848 node_tiebreaker: 0,
6849 }
6850}
6851
6852async fn execute_sql(
6861 state: &AppState,
6862 principal: &Option<mongreldb_core::Principal>,
6863 session: &MongrelSession,
6864 request: ResolvedSqlRequest,
6865 query: RegisteredSqlQuery,
6866 query_id: QueryId,
6867 sql_permit: admission::SqlAdmissionGuard,
6868) -> (Response, Option<mongreldb_types::hlc::HlcTimestamp>) {
6869 let ResolvedSqlRequest {
6870 request: req,
6871 output_limits,
6872 idempotency,
6873 pagination,
6874 } = request;
6875 if let Some(response) = cluster_admin::try_admin_sql(state, principal, &req.sql).await {
6879 drop(sql_permit);
6880 return (response, None);
6881 }
6882 let idempotency_query = idempotency.as_ref().map(|_| query.clone());
6885 state.metrics.inc_sql_queries();
6886 let audited = audit::is_audited_sql(&req.sql);
6887 let actor = request_owner(state, principal);
6888 let page_binding = sql_pages::SqlPageBinding {
6889 security_version: state.db.security_version(),
6890 catalog_epoch: state.db.catalog_snapshot().db_epoch,
6891 };
6892 let start = std::time::Instant::now();
6893 let result = if let Some(pagination) = pagination {
6899 match session
6900 .run_with_query_for_serialization_with_limits(
6901 &req.sql,
6902 query,
6903 mongreldb_query::SqlCollectionLimits::new(output_limits.0, output_limits.1),
6904 )
6905 .await
6906 {
6907 Ok(output) => Ok(dispatch_paginated_sql(
6908 state,
6909 output,
6910 query_id,
6911 &actor,
6912 pagination,
6913 output_limits,
6914 session.sql_test_hook(),
6915 page_binding,
6916 )
6917 .await),
6918 Err(error) => Err(error),
6919 }
6920 } else if req.format.as_deref() == Some("arrow-stream") {
6921 match session
6922 .run_stream_with_query_for_serialization(&req.sql, query)
6923 .await
6924 {
6925 Ok((stream, completion)) => Ok(sql_arrow_stream_response_controlled(
6926 stream,
6927 completion,
6928 sql_permit,
6929 output_limits,
6930 state,
6931 query_id,
6932 session.sql_test_hook(),
6933 )),
6934 Err(error) => Err(error),
6935 }
6936 } else {
6937 match session
6938 .run_with_query_for_serialization_with_limits(
6939 &req.sql,
6940 query,
6941 mongreldb_query::SqlCollectionLimits::new(output_limits.0, output_limits.1),
6942 )
6943 .await
6944 {
6945 Ok(output) => Ok(dispatch_buffered_sql_format(
6946 state,
6947 req.format.as_deref(),
6948 output,
6949 query_id,
6950 session.sql_test_hook(),
6951 output_limits,
6952 )
6953 .await),
6954 Err(error) => Err(error),
6955 }
6956 };
6957 let elapsed = start.elapsed();
6958 if elapsed >= state.reloadable.slow_query_threshold.get() {
6961 state.metrics.inc_slow_queries();
6962 eprintln!(
6963 "[slow-query] {}\u{00b5}s query_id={} operation={}",
6964 elapsed.as_micros(),
6965 query_id,
6966 safe_sql_operation(&req.sql)
6967 );
6968 }
6969 if audited {
6972 let (action, detail) = audit::redacted_ddl_detail(&req.sql, result.is_ok());
6973 state.audit.record(actor, action, detail);
6974 }
6975 let response = match result {
6976 Ok(response) => with_query_id(response, query_id),
6977 Err(e) => {
6978 state.metrics.inc_sql_errors();
6979 tracked_query_error_response(state, &e, Some(query_id))
6980 }
6981 };
6982 let Some(idempotency) = idempotency else {
6983 return (response, None);
6984 };
6985 let status = idempotency_query.map(|query| query.status());
6986 if let Some(receipt) = status.as_ref().and_then(sql_terminal_idempotency_receipt) {
6987 let mut receipt = receipt;
6988 if receipt.outcome.committed {
6995 let db = Arc::clone(&state.db);
6996 let owner = idempotency.owner().to_owned();
6997 let key = idempotency.key().to_owned();
6998 let binding = idempotency.binding().clone();
6999 let ttl = idempotency.ttl();
7000 receipt.commit_receipt = tokio::task::spawn_blocking(move || {
7001 sql_idempotency::record_core_idempotency_commit(&db, &owner, &key, &binding, ttl)
7002 })
7003 .await
7004 .unwrap_or_else(|error| {
7005 eprintln!("[idempotency] core ledger record task failed: {error}");
7006 None
7007 });
7008 }
7009 let commit_ts = receipt
7010 .commit_receipt
7011 .as_ref()
7012 .map(sql_idempotency::SqlCommitReceipt::commit_ts);
7013 let (expires_at_ms, persisted) = idempotency.commit(receipt.clone());
7014 return (
7015 sql_idempotency_receipt_response(query_id, &receipt, false, expires_at_ms, persisted),
7016 commit_ts,
7017 );
7018 }
7019 if status.as_ref().is_some_and(can_abort_idempotency_intent) {
7020 idempotency.abort();
7021 }
7022 (response, None)
7023}
7024
7025fn can_abort_idempotency_intent(status: &mongreldb_query::QueryStatus) -> bool {
7026 !status.outcome_unknown
7027 && !status.durable_outcome.committed
7028 && matches!(
7029 status.terminal_state(),
7030 Some(
7031 mongreldb_query::QueryTerminalState::FailedBeforeCommit
7032 | mongreldb_query::QueryTerminalState::CancelledBeforeCommit
7033 | mongreldb_query::QueryTerminalState::DeadlineBeforeCommit
7034 )
7035 )
7036}
7037
7038fn safe_sql_operation(sql: &str) -> String {
7039 sql.split_whitespace()
7040 .next()
7041 .unwrap_or("UNKNOWN")
7042 .chars()
7043 .filter(|character| character.is_ascii_alphabetic())
7044 .take(16)
7045 .collect::<String>()
7046 .to_ascii_uppercase()
7047}
7048
7049fn remote_boolean_ai_error() -> Response {
7050 (
7051 StatusCode::BAD_REQUEST,
7052 "Boolean ANN/Sparse SQL is disabled remotely; use scored SQL functions",
7053 )
7054 .into_response()
7055}
7056
7057#[derive(Debug)]
7058struct SerializedOutput {
7059 bytes: Vec<u8>,
7060 arrow: bool,
7061}
7062
7063#[derive(Debug)]
7064enum BufferedSerializationError {
7065 Query(mongreldb_query::MongrelQueryError),
7066 Limit(String),
7067 Encoding(String),
7068}
7069
7070struct LimitedOutput {
7071 bytes: Vec<u8>,
7072 max_bytes: usize,
7073 exceeded: bool,
7074}
7075
7076impl LimitedOutput {
7077 fn new(max_bytes: usize) -> Self {
7078 Self {
7079 bytes: Vec::new(),
7080 max_bytes,
7081 exceeded: false,
7082 }
7083 }
7084}
7085
7086impl std::io::Write for LimitedOutput {
7087 fn write(&mut self, bytes: &[u8]) -> std::io::Result<usize> {
7088 if self.bytes.len().saturating_add(bytes.len()) > self.max_bytes {
7089 self.exceeded = true;
7090 return Err(std::io::Error::other("SQL output byte limit exceeded"));
7091 }
7092 self.bytes.extend_from_slice(bytes);
7093 Ok(bytes.len())
7094 }
7095
7096 fn flush(&mut self) -> std::io::Result<()> {
7097 Ok(())
7098 }
7099}
7100
7101fn serialize_buffered_output(
7102 format: &str,
7103 batches: &[arrow::record_batch::RecordBatch],
7104 query: &RegisteredSqlQuery,
7105 max_rows: usize,
7106 max_bytes: usize,
7107 test_hook: Option<&mongreldb_query::SqlTestHook>,
7108) -> std::result::Result<SerializedOutput, BufferedSerializationError> {
7109 const ROW_CHECKPOINT_INTERVAL: usize = 256;
7110 let mut rows = 0usize;
7111 let mut writer_output = LimitedOutput::new(max_bytes);
7112
7113 if format == "arrow" {
7114 if batches.is_empty() {
7115 if let Some(hook) = test_hook {
7116 hook(mongreldb_query::SqlTestHookPoint::AfterSerialization);
7117 }
7118 return Ok(SerializedOutput {
7119 bytes: Vec::new(),
7120 arrow: true,
7121 });
7122 }
7123 let schema = batches[0].schema();
7124 let encoding_result = (|| {
7125 let mut writer =
7126 arrow::ipc::writer::FileWriter::try_new(&mut writer_output, schema.as_ref())
7127 .map_err(|error| error.to_string())?;
7128 for batch in batches {
7129 for offset in (0..batch.num_rows()).step_by(ROW_CHECKPOINT_INTERVAL) {
7130 if let Some(hook) = test_hook {
7131 hook(mongreldb_query::SqlTestHookPoint::BeforeSerializationBatch);
7132 }
7133 query.checkpoint().map_err(|error| error.to_string())?;
7134 let length = ROW_CHECKPOINT_INTERVAL.min(batch.num_rows() - offset);
7135 rows = rows.saturating_add(length);
7136 if rows > max_rows {
7137 return Err("SQL output row limit exceeded".into());
7138 }
7139 writer
7140 .write(&batch.slice(offset, length))
7141 .map_err(|error| error.to_string())?;
7142 }
7143 }
7144 writer.finish().map_err(|error| error.to_string())
7145 })();
7146 if let Err(error) = encoding_result {
7147 if let Err(query_error) = query.checkpoint() {
7148 return Err(BufferedSerializationError::Query(query_error));
7149 }
7150 if writer_output.exceeded || rows > max_rows {
7151 return Err(BufferedSerializationError::Limit(error));
7152 }
7153 return Err(BufferedSerializationError::Encoding(error));
7154 }
7155 if let Some(hook) = test_hook {
7156 hook(mongreldb_query::SqlTestHookPoint::AfterSerialization);
7157 }
7158 return Ok(SerializedOutput {
7159 bytes: writer_output.bytes,
7160 arrow: true,
7161 });
7162 }
7163
7164 let encoding_result = (|| {
7165 let mut writer = arrow::json::writer::ArrayWriter::new(&mut writer_output);
7166 for batch in batches {
7167 for offset in (0..batch.num_rows()).step_by(ROW_CHECKPOINT_INTERVAL) {
7168 if let Some(hook) = test_hook {
7169 hook(mongreldb_query::SqlTestHookPoint::BeforeSerializationBatch);
7170 }
7171 query.checkpoint().map_err(|error| error.to_string())?;
7172 let length = ROW_CHECKPOINT_INTERVAL.min(batch.num_rows() - offset);
7173 rows = rows.saturating_add(length);
7174 if rows > max_rows {
7175 return Err("SQL output row limit exceeded".into());
7176 }
7177 let slice = batch.slice(offset, length);
7178 writer
7179 .write_batches(&[&slice])
7180 .map_err(|error| error.to_string())?;
7181 }
7182 }
7183 writer.finish().map_err(|error| error.to_string())
7184 })();
7185 if let Err(error) = encoding_result {
7186 if let Err(query_error) = query.checkpoint() {
7187 return Err(BufferedSerializationError::Query(query_error));
7188 }
7189 if writer_output.exceeded || rows > max_rows {
7190 return Err(BufferedSerializationError::Limit(error));
7191 }
7192 return Err(BufferedSerializationError::Encoding(error));
7193 }
7194 if let Some(hook) = test_hook {
7195 hook(mongreldb_query::SqlTestHookPoint::AfterSerialization);
7196 }
7197 Ok(SerializedOutput {
7198 bytes: writer_output.bytes,
7199 arrow: false,
7200 })
7201}
7202
7203#[cfg(test)]
7206fn sql_arrow_stream_response(batches: mongreldb_query::MongrelRecordBatchStream) -> Response {
7207 use futures::{stream, StreamExt};
7208
7209 const STREAM_CT: &str = "application/vnd.apache.arrow.stream";
7210
7211 let schema = batches.schema();
7212 let mut writer = match arrow::ipc::writer::StreamWriter::try_new(Vec::new(), schema.as_ref()) {
7213 Ok(w) => w,
7214 Err(e) => {
7215 return (
7216 StatusCode::INTERNAL_SERVER_ERROR,
7217 format!("arrow stream init error: {e}"),
7218 )
7219 .into_response()
7220 }
7221 };
7222 let schema_chunk: Vec<u8> = std::mem::take(writer.get_mut());
7225 let batch_stream = stream::unfold(
7226 (batches, Some(writer)),
7227 |(mut batches, writer)| async move {
7228 let mut writer = writer?;
7229 match batches.next().await {
7230 Some(Ok(batch)) => match writer.write(&batch) {
7231 Ok(()) => {
7232 let chunk = std::mem::take(writer.get_mut());
7233 Some((Ok(chunk), (batches, Some(writer))))
7234 }
7235 Err(error) => Some((Err(std::io::Error::other(error)), (batches, None))),
7236 },
7237 Some(Err(error)) => Some((Err(std::io::Error::other(error)), (batches, None))),
7238 None => match writer.finish() {
7239 Ok(()) => {
7240 let chunk = std::mem::take(writer.get_mut());
7241 Some((Ok(chunk), (batches, None)))
7242 }
7243 Err(error) => Some((Err(std::io::Error::other(error)), (batches, None))),
7244 },
7245 }
7246 },
7247 );
7248
7249 let schema_item: Result<Vec<u8>, std::io::Error> = Ok(schema_chunk);
7252 let full = stream::iter([schema_item]).chain(batch_stream);
7253 let body = axum::body::Body::from_stream(full);
7254 ([(header::CONTENT_TYPE, STREAM_CT)], body).into_response()
7255}
7256
7257fn sql_arrow_stream_response_controlled(
7258 batches: mongreldb_query::MongrelRecordBatchStream,
7259 completion: SqlStreamCompletion,
7260 sql_permit: admission::SqlAdmissionGuard,
7261 limits: (usize, usize),
7262 state: &AppState,
7263 query_id: QueryId,
7264 test_hook: Option<mongreldb_query::SqlTestHook>,
7265) -> Response {
7266 use futures::{stream, StreamExt};
7267
7268 const STREAM_CT: &str = "application/vnd.apache.arrow.stream";
7269 let (max_rows, max_bytes) = limits;
7270 let schema = batches.schema();
7271 let mut writer = match arrow::ipc::writer::StreamWriter::try_new(Vec::new(), schema.as_ref()) {
7272 Ok(writer) => writer,
7273 Err(error) => {
7274 completion.fail_serialization();
7275 drop(batches);
7276 return terminal_server_error_response(
7277 state,
7278 query_id,
7279 StatusCode::INTERNAL_SERVER_ERROR,
7280 "SERIALIZATION_FAILED",
7281 format!("arrow stream init error: {error}"),
7282 );
7283 }
7284 };
7285 let schema_chunk = std::mem::take(writer.get_mut());
7286 if schema_chunk.len() > max_bytes {
7287 completion.fail_result_limit();
7288 drop(batches);
7289 return terminal_server_error_response(
7290 state,
7291 query_id,
7292 StatusCode::PAYLOAD_TOO_LARGE,
7293 "RESULT_LIMIT_EXCEEDED",
7294 "SQL output byte limit exceeded",
7295 );
7296 }
7297 let metrics = Arc::clone(&state.metrics);
7298 metrics.add_sql_output_bytes(schema_chunk.len());
7299 let batch_stream = stream::unfold(
7300 (
7301 batches,
7302 Some(writer),
7303 completion,
7304 Some(sql_permit),
7305 0usize,
7306 schema_chunk.len(),
7307 metrics,
7308 ),
7309 move |(mut batches, writer, completion, permit, rows, bytes, metrics)| {
7310 let test_hook = test_hook.clone();
7311 async move {
7312 let mut writer = writer?;
7313 match batches.next().await {
7314 Some(Ok(batch)) => {
7315 let next_rows = rows.saturating_add(batch.num_rows());
7316 if next_rows > max_rows {
7317 completion.fail_result_limit();
7318 return Some((
7319 Err(std::io::Error::other("SQL output row limit exceeded")),
7320 (batches, None, completion, permit, next_rows, bytes, metrics),
7321 ));
7322 }
7323 match writer.write(&batch) {
7324 Ok(()) => {
7325 let chunk = std::mem::take(writer.get_mut());
7326 let next_bytes = bytes.saturating_add(chunk.len());
7327 if next_bytes > max_bytes {
7328 completion.fail_result_limit();
7329 return Some((
7330 Err(std::io::Error::other(
7331 "SQL output byte limit exceeded",
7332 )),
7333 (
7334 batches, None, completion, permit, next_rows,
7335 next_bytes, metrics,
7336 ),
7337 ));
7338 }
7339 metrics.add_sql_output_bytes(chunk.len());
7340 Some((
7341 Ok(chunk),
7342 (
7343 batches,
7344 Some(writer),
7345 completion,
7346 permit,
7347 next_rows,
7348 next_bytes,
7349 metrics,
7350 ),
7351 ))
7352 }
7353 Err(error) => {
7354 completion.fail_serialization();
7355 Some((
7356 Err(std::io::Error::other(error)),
7357 (batches, None, completion, permit, rows, bytes, metrics),
7358 ))
7359 }
7360 }
7361 }
7362 Some(Err(error)) => Some((
7363 Err(std::io::Error::other(error)),
7364 (batches, None, completion, permit, rows, bytes, metrics),
7365 )),
7366 None => match writer.finish() {
7367 Ok(()) => {
7368 let chunk = std::mem::take(writer.get_mut());
7369 let next_bytes = bytes.saturating_add(chunk.len());
7370 if next_bytes > max_bytes {
7371 completion.fail_result_limit();
7372 return Some((
7373 Err(std::io::Error::other("SQL output byte limit exceeded")),
7374 (batches, None, completion, permit, rows, next_bytes, metrics),
7375 ));
7376 }
7377 if let Some(hook) = test_hook {
7378 hook(mongreldb_query::SqlTestHookPoint::AfterSerialization);
7379 }
7380 match completion.try_complete() {
7381 Ok(()) => {
7382 metrics.add_sql_output_bytes(chunk.len());
7383 Some((
7384 Ok(chunk),
7385 (
7386 batches, None, completion, permit, rows, next_bytes,
7387 metrics,
7388 ),
7389 ))
7390 }
7391 Err(error) => {
7392 metrics.inc_sql_errors();
7393 Some((
7394 Err(std::io::Error::other(error.to_string())),
7395 (batches, None, completion, permit, rows, bytes, metrics),
7396 ))
7397 }
7398 }
7399 }
7400 Err(error) => {
7401 completion.fail_serialization();
7402 Some((
7403 Err(std::io::Error::other(error)),
7404 (batches, None, completion, permit, rows, bytes, metrics),
7405 ))
7406 }
7407 },
7408 }
7409 }
7410 },
7411 );
7412 let schema_item: Result<Vec<u8>, std::io::Error> = Ok(schema_chunk);
7413 let body = axum::body::Body::from_stream(stream::iter([schema_item]).chain(batch_stream));
7414 ([(header::CONTENT_TYPE, STREAM_CT)], body).into_response()
7415}
7416
7417#[derive(Deserialize)]
7418struct TxnOp {
7419 table: String,
7420 op: String,
7421 cells: Option<Vec<serde_json::Value>>,
7422 row_id: Option<u64>,
7423}
7424
7425#[derive(Deserialize)]
7426struct TxnRequest {
7427 ops: Vec<TxnOp>,
7428}
7429
7430async fn txn(
7431 State(state): State<Arc<AppState>>,
7432 OptionalPrincipal(principal): OptionalPrincipal,
7433 Json(req): Json<TxnRequest>,
7434) -> Response {
7435 if let Some(response) = require_writes_open(&state) {
7436 return response;
7437 }
7438 let mut parsed: Vec<(String, TxnAction)> = Vec::with_capacity(req.ops.len());
7442 for op in &req.ops {
7443 match op.op.as_str() {
7444 "put" => {
7445 let cells_json = match op.cells.as_ref() {
7446 Some(c) if !c.is_empty() => c,
7447 _ => {
7448 return (StatusCode::BAD_REQUEST, "put op requires non-empty cells")
7449 .into_response()
7450 }
7451 };
7452 let handle = match state.db.table(&op.table) {
7453 Ok(h) => h,
7454 Err(e) => return (StatusCode::NOT_FOUND, e.to_string()).into_response(),
7455 };
7456 let schema = handle.lock().schema().clone();
7457 let cells = match parse_cells(cells_json, &schema) {
7458 Ok(c) => c,
7459 Err(msg) => return (StatusCode::BAD_REQUEST, msg).into_response(),
7460 };
7461 parsed.push((op.table.clone(), TxnAction::Put(cells)));
7462 }
7463 "delete" => {
7464 let rid = match op.row_id {
7465 Some(r) => r,
7466 None => {
7467 return (StatusCode::BAD_REQUEST, "delete op requires row_id")
7468 .into_response()
7469 }
7470 };
7471 parsed.push((op.table.clone(), TxnAction::Delete(rid)));
7472 }
7473 other => {
7474 return (StatusCode::BAD_REQUEST, format!("unknown op: {other}")).into_response()
7475 }
7476 }
7477 }
7478
7479 state.metrics.inc_txns();
7480 let mut transaction = state.db.begin_as(request_principal(&state, &principal));
7481 let result = (|| {
7482 for (table, action) in &parsed {
7483 match action {
7484 TxnAction::Put(cells) => {
7485 transaction.put(table, cells.clone())?;
7486 }
7487 TxnAction::Delete(rid) => {
7488 transaction.delete(table, mongreldb_core::RowId(*rid))?;
7489 }
7490 }
7491 }
7492 transaction.commit()
7493 })();
7494 match result {
7495 Ok(epoch) => Json(json!({
7496 "status": "committed",
7497 "epoch": epoch.0,
7498 "epoch_text": epoch.0.to_string()
7499 }))
7500 .into_response(),
7501 Err(error) => crate::kit::durable_core_error_response(&error)
7502 .unwrap_or_else(|| (status_for_error(&error), error.to_string()).into_response()),
7503 }
7504}
7505
7506enum TxnAction {
7507 Put(Vec<(u16, Value)>),
7508 Delete(u64),
7509}
7510
7511#[cfg(test)]
7512mod auth_tests {
7513 use super::*;
7514 use mongreldb_core::Database;
7515 use tempfile::tempdir;
7516
7517 #[test]
7518 fn slow_query_operation_does_not_include_literals() {
7519 let sql = "CREATE USER alice PASSWORD 'never-log-this'";
7520 let operation = safe_sql_operation(sql);
7521 assert_eq!(operation, "CREATE");
7522 assert!(!operation.contains("never-log-this"));
7523 }
7524
7525 #[tokio::test]
7526 async fn auth_rejects_missing_token() {
7527 let dir = tempdir().unwrap();
7528 let db = Arc::new(Database::create(dir.path()).unwrap());
7529 let app = build_app_with_config(db, std::iter::empty(), Some("secret".into()), None);
7530 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
7531 let addr = listener.local_addr().unwrap();
7532 tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
7533 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
7535
7536 let client = reqwest::Client::new();
7537 let resp = client
7538 .get(format!("http://{addr}/health"))
7539 .send()
7540 .await
7541 .unwrap();
7542 assert_eq!(resp.status(), 401);
7543 }
7544
7545 #[tokio::test]
7546 async fn auth_accepts_valid_token() {
7547 let dir = tempdir().unwrap();
7548 let db = Arc::new(Database::create(dir.path()).unwrap());
7549 let app = build_app_with_config(db, std::iter::empty(), Some("secret".into()), None);
7550 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
7551 let addr = listener.local_addr().unwrap();
7552 tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
7553 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
7554
7555 let client = reqwest::Client::new();
7556 let resp = client
7557 .get(format!("http://{addr}/health"))
7558 .header("Authorization", "Bearer secret")
7559 .send()
7560 .await
7561 .unwrap();
7562 assert_eq!(resp.status(), 200);
7563 }
7564
7565 #[tokio::test]
7566 async fn no_auth_when_token_unset() {
7567 let dir = tempdir().unwrap();
7568 let db = Arc::new(Database::create(dir.path()).unwrap());
7569 let app = build_app_with_config(db, std::iter::empty(), None, None);
7570 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
7571 let addr = listener.local_addr().unwrap();
7572 tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
7573 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
7574
7575 let client = reqwest::Client::new();
7576 let resp = client
7577 .get(format!("http://{addr}/health"))
7578 .send()
7579 .await
7580 .unwrap();
7581 assert_eq!(resp.status(), 200);
7582 }
7583
7584 #[tokio::test]
7585 async fn capabilities_advertise_sql_cancellation_v2() {
7586 let dir = tempdir().unwrap();
7587 let db = Arc::new(Database::create(dir.path()).unwrap());
7588 let app = build_app_with_config(db, std::iter::empty(), None, None);
7589 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
7590 let addr = listener.local_addr().unwrap();
7591 tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
7592
7593 let body: serde_json::Value = reqwest::Client::new()
7594 .get(format!("http://{addr}/capabilities"))
7595 .send()
7596 .await
7597 .unwrap()
7598 .json()
7599 .await
7600 .unwrap();
7601 assert_eq!(body["sql_cancellation"]["version"], 2);
7602 assert_eq!(body["sql_cancellation"]["client_query_ids"], true);
7603 assert_eq!(body["sql_cancellation"]["cancel_endpoint"], true);
7604 assert_eq!(body["sql_cancellation"]["query_status"], true);
7605 assert_eq!(body["sql_cancellation"]["pre_registration_cancel"], true);
7606 assert_eq!(body["sql_cancellation"]["stream_disconnect_cancels"], true);
7607 assert_eq!(body["sql_idempotency"]["version"], 1);
7608 assert_eq!(
7609 body["sql_idempotency"]["indeterminate_never_reexecutes"],
7610 true
7611 );
7612 assert_eq!(body["sql_pagination"]["version"], 1);
7613 assert_eq!(
7614 body["sql_pagination"]["continuation_endpoint"],
7615 "/sql/continue"
7616 );
7617 }
7618}
7619
7620#[cfg(test)]
7621mod query_response_tests {
7622 use super::*;
7623
7624 #[test]
7625 fn cancellation_reason_names_are_stable_snake_case() {
7626 assert_eq!(cancellation_reason_name(CancellationReason::None), "none");
7627 assert_eq!(
7628 cancellation_reason_name(CancellationReason::ClientRequest),
7629 "client_request"
7630 );
7631 assert_eq!(
7632 cancellation_reason_name(CancellationReason::ClientDisconnected),
7633 "client_disconnected"
7634 );
7635 assert_eq!(
7636 cancellation_reason_name(CancellationReason::SessionClosed),
7637 "session_closed"
7638 );
7639 assert_eq!(
7640 cancellation_reason_name(CancellationReason::ServerShutdown),
7641 "server_shutdown"
7642 );
7643 assert_eq!(
7644 cancellation_reason_name(CancellationReason::Deadline),
7645 "deadline"
7646 );
7647 }
7648
7649 #[test]
7650 fn unknown_outcome_never_proves_idempotency_intent_safe_to_abort() {
7651 let registry = Arc::new(SqlQueryRegistry::default());
7652 let unknown_id: QueryId = "102132435465768798a9bacbdcedfe0f".parse().unwrap();
7653 let unknown = registry
7654 .register(SqlQueryOptions {
7655 query_id: Some(unknown_id),
7656 ..SqlQueryOptions::default()
7657 })
7658 .unwrap();
7659 unknown.mark_outcome_unknown();
7660 unknown.fail();
7661 assert!(!can_abort_idempotency_intent(
7662 ®istry.status(unknown_id).unwrap()
7663 ));
7664
7665 let failed_id: QueryId = "2031425364758697a8b9cadbecfd0e1f".parse().unwrap();
7666 let failed = registry
7667 .register(SqlQueryOptions {
7668 query_id: Some(failed_id),
7669 ..SqlQueryOptions::default()
7670 })
7671 .unwrap();
7672 failed.fail();
7673 assert!(can_abort_idempotency_intent(
7674 ®istry.status(failed_id).unwrap()
7675 ));
7676 }
7677
7678 #[test]
7679 fn unknown_outcome_never_becomes_durable_receipt() {
7680 let registry = Arc::new(SqlQueryRegistry::default());
7681 let query = registry.register(SqlQueryOptions::default()).unwrap();
7682 query.record_commit(0, 42);
7683 query.mark_outcome_unknown();
7684 query.fail();
7685 let status = registry.status(query.id()).unwrap();
7686 assert!(status.durable_outcome.committed);
7687 assert!(status.outcome_unknown);
7688 assert!(sql_terminal_idempotency_receipt(&status).is_none());
7689 }
7690
7691 #[test]
7692 fn successful_noop_write_becomes_noncommitting_receipt() {
7693 let registry = Arc::new(SqlQueryRegistry::default());
7694 let query = registry.register(SqlQueryOptions::default()).unwrap();
7695 query
7696 .transition(SqlQueryPhase::Queued, SqlQueryPhase::Executing)
7697 .unwrap();
7698 query.try_complete().unwrap();
7699 let status = registry.status(query.id()).unwrap();
7700 assert!(!status.durable_outcome.committed);
7701 assert!(!can_abort_idempotency_intent(&status));
7702 let receipt = sql_terminal_idempotency_receipt(&status).unwrap();
7703 assert_eq!(receipt.status, "completed");
7704 assert!(!receipt.outcome.committed);
7705 assert_eq!(receipt.outcome.committed_statements, 0);
7706 assert_eq!(receipt.outcome.last_commit_epoch, None);
7707 }
7708
7709 #[test]
7710 fn conflicting_idempotency_key_sources_are_rejected() {
7711 let request = SqlRequest {
7712 sql: "INSERT INTO items VALUES (1)".into(),
7713 format: None,
7714 query_id: None,
7715 timeout_ms: None,
7716 max_output_rows: None,
7717 max_output_bytes: None,
7718 idempotency_key: Some("body-key".into()),
7719 pagination: None,
7720 };
7721 let mut headers = axum::http::HeaderMap::new();
7722 headers.insert("idempotency-key", "header-key".parse().unwrap());
7723 assert_eq!(
7724 requested_sql_idempotency_key(&headers, &request),
7725 Err("body idempotency_key and Idempotency-Key header must match")
7726 );
7727
7728 headers.insert("idempotency-key", "body-key".parse().unwrap());
7729 assert_eq!(
7730 requested_sql_idempotency_key(&headers, &request),
7731 Ok(Some("body-key".into()))
7732 );
7733 }
7734
7735 #[test]
7736 fn idempotency_binding_includes_pagination_semantics() {
7737 let mut request = SqlRequest {
7738 sql: "INSERT INTO items VALUES (1)".into(),
7739 format: None,
7740 query_id: None,
7741 timeout_ms: None,
7742 max_output_rows: None,
7743 max_output_bytes: None,
7744 idempotency_key: Some("key".into()),
7745 pagination: None,
7746 };
7747 let unpaged = sql_idempotency_binding(&request, (100, 1_024), None, 60_000).unwrap();
7748 request.pagination = Some(SqlPaginationRequest {
7749 page_size_rows: 10,
7750 projection: vec!["id".into()],
7751 max_page_bytes: Some(512),
7752 max_page_tokens: Some(128),
7753 });
7754 let paged = sql_idempotency_binding(&request, (100, 1_024), None, 60_000).unwrap();
7755 assert_ne!(unpaged.request_semantics_hash, paged.request_semantics_hash);
7756 }
7757
7758 #[test]
7759 fn paginated_decode_stops_nested_heap_amplification_at_budget() {
7760 let registry = Arc::new(SqlQueryRegistry::default());
7761 let query = registry.register(SqlQueryOptions::default()).unwrap();
7762 let json = format!("[[{}]]", vec!["null"; 10_000].join(","));
7763 let mut deserializer = serde_json::Deserializer::from_slice(json.as_bytes());
7764 let mut budget = PaginatedDecodeBudget {
7765 used: 0,
7766 limit: 4 * 1024,
7767 nodes: 0,
7768 exceeded: false,
7769 query: &query,
7770 test_hook: None,
7771 };
7772 let error = serde::de::DeserializeSeed::deserialize(
7773 BudgetedJsonRowsSeed {
7774 budget: &mut budget,
7775 },
7776 &mut deserializer,
7777 )
7778 .unwrap_err();
7779 assert!(error.to_string().contains(PAGINATED_MEMORY_LIMIT_ERROR));
7780 assert!(budget.exceeded);
7781 assert!(budget.nodes < 100, "decoded {} nodes", budget.nodes);
7782 query.fail();
7783 }
7784
7785 #[tokio::test]
7786 async fn unknown_outcome_response_never_claims_no_commit() {
7787 let registry = Arc::new(SqlQueryRegistry::default());
7788 let query = registry.register(SqlQueryOptions::default()).unwrap();
7789 let query_id = query.id();
7790 query
7791 .transition(SqlQueryPhase::Queued, SqlQueryPhase::Executing)
7792 .unwrap();
7793 let error = query.outcome_unknown_error("fenced maintenance failed");
7794 query.fail();
7795 let status = registry.status(query_id).unwrap();
7796 assert_eq!(
7797 status.terminal_state(),
7798 Some(mongreldb_query::QueryTerminalState::OutcomeUnknown)
7799 );
7800
7801 let response = query_error_response_with_status(&error, Some(query_id), Some(&status));
7802 assert_eq!(response.status(), StatusCode::CONFLICT);
7803 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
7804 .await
7805 .unwrap();
7806 let body: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
7807 assert_eq!(body["status"], "outcome_unknown");
7808 assert_eq!(body["error"]["code"], "QUERY_OUTCOME_UNKNOWN");
7809 assert!(body["committed"].is_null());
7810 assert!(body["committed_statements"].is_null());
7811 assert!(body["last_commit_epoch"].is_null());
7812 assert!(body["completed_statements"].is_null());
7813 assert!(body["statement_index"].is_null());
7814 assert!(body["outcome"]["committed"].is_null());
7815 assert!(body["error"]["committed"].is_null());
7816 }
7817
7818 #[tokio::test]
7819 async fn retryable_idempotency_error_survives_terminal_status() {
7820 let registry = Arc::new(SqlQueryRegistry::default());
7821 let query = registry.register(SqlQueryOptions::default()).unwrap();
7822 let query_id = query.id();
7823 let response = registered_sql_error_response(
7824 RegisteredQueryGuard::new(query),
7825 query_id,
7826 StatusCode::SERVICE_UNAVAILABLE,
7827 "IDEMPOTENCY_STORE_UNAVAILABLE",
7828 "could not durably reserve the SQL idempotency key",
7829 true,
7830 );
7831 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
7832 .await
7833 .unwrap();
7834 let body: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
7835 assert_eq!(body["retryable"], true);
7836 assert_eq!(body["error"]["retryable"], true);
7837
7838 let status = registry.status(query_id).unwrap();
7839 assert_eq!(
7840 status.terminal_error.as_ref().unwrap().code,
7841 "IDEMPOTENCY_STORE_UNAVAILABLE"
7842 );
7843 assert!(terminal_error_retryable(status.terminal_error.as_ref()));
7844 }
7845
7846 #[tokio::test]
7847 async fn committed_idempotency_terminal_error_is_a_receipt() {
7848 let query_id: QueryId = "00112233445566778899aabbccddeeff".parse().unwrap();
7849 let receipt = sql_idempotency::SqlDurableReceipt {
7850 original_query_id: query_id.to_string(),
7851 status: "committed_with_error".into(),
7852 server_state: "failed".into(),
7853 cancellation_reason: "client_disconnected".into(),
7854 outcome: sql_idempotency::SqlReceiptOutcome {
7855 committed: true,
7856 committed_statements: 1,
7857 last_commit_epoch: Some(42),
7858 last_commit_epoch_text: Some("42".into()),
7859 first_commit_statement_index: Some(0),
7860 last_commit_statement_index: Some(0),
7861 completed_statements: 1,
7862 statement_index: 0,
7863 serialization: "failed".into(),
7864 },
7865 terminal_error: Some(sql_idempotency::SqlReceiptTerminalError {
7866 code: "SERIALIZATION_FAILED_AFTER_COMMIT".into(),
7867 category: "serialization".into(),
7868 }),
7869 commit_receipt: None,
7870 };
7871 let response = sql_idempotency_receipt_response(query_id, &receipt, false, 99, true);
7872 assert_eq!(response.status(), StatusCode::OK);
7873 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
7874 .await
7875 .unwrap();
7876 let body: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
7877 assert_eq!(body["status"], "committed_with_error");
7878 assert_eq!(body["server_state"], "failed");
7879 assert_eq!(body["cancellation_reason"], "client_disconnected");
7880 assert_eq!(body["first_commit_statement_index"], 0);
7881 assert_eq!(body["last_commit_statement_index"], 0);
7882 assert_eq!(
7883 body["terminal_error"]["code"],
7884 "SERIALIZATION_FAILED_AFTER_COMMIT"
7885 );
7886 }
7887
7888 #[test]
7889 fn durable_replay_restores_terminal_status_parity() {
7890 let registry = Arc::new(SqlQueryRegistry::default());
7891 let query_id: QueryId = "11223344556677889900aabbccddeeff".parse().unwrap();
7892 let query = registry
7893 .register(SqlQueryOptions {
7894 query_id: Some(query_id),
7895 ..SqlQueryOptions::default()
7896 })
7897 .unwrap();
7898 let receipt = sql_idempotency::SqlDurableReceipt {
7899 original_query_id: "00112233445566778899aabbccddeeff".into(),
7900 status: "cancelled_after_commit".into(),
7901 server_state: "cancelled".into(),
7902 cancellation_reason: "client_disconnected".into(),
7903 outcome: sql_idempotency::SqlReceiptOutcome {
7904 committed: true,
7905 committed_statements: 1,
7906 last_commit_epoch: Some(42),
7907 last_commit_epoch_text: Some("42".into()),
7908 first_commit_statement_index: Some(0),
7909 last_commit_statement_index: Some(0),
7910 completed_statements: 1,
7911 statement_index: 1,
7912 serialization: "failed".into(),
7913 },
7914 terminal_error: Some(sql_idempotency::SqlReceiptTerminalError {
7915 code: "QUERY_CANCELLED_AFTER_COMMIT".into(),
7916 category: "cancellation".into(),
7917 }),
7918 commit_receipt: None,
7919 };
7920 restore_idempotency_replay(RegisteredQueryGuard::new(query), &receipt).unwrap();
7921 let status = registry.status(query_id).unwrap();
7922 assert_eq!(status.phase, SqlQueryPhase::Cancelled);
7923 assert_eq!(
7924 status.terminal_state(),
7925 Some(mongreldb_query::QueryTerminalState::CancelledAfterCommit)
7926 );
7927 assert_eq!(
7928 status.cancellation_reason,
7929 CancellationReason::ClientDisconnected
7930 );
7931 assert_eq!(status.durable_outcome.committed_statements, 1);
7932 assert_eq!(status.durable_outcome.last_commit_epoch, Some(42));
7933 assert_eq!(
7934 status.terminal_error.unwrap().code,
7935 "QUERY_CANCELLED_AFTER_COMMIT"
7936 );
7937 }
7938
7939 #[test]
7940 fn durable_replay_rejects_invalid_authenticated_state() {
7941 let registry = Arc::new(SqlQueryRegistry::default());
7942 let query = registry.register(SqlQueryOptions::default()).unwrap();
7943 let receipt = sql_idempotency::SqlDurableReceipt {
7944 original_query_id: query.id().to_string(),
7945 status: "invented_terminal_state".into(),
7946 server_state: "completed".into(),
7947 cancellation_reason: "none".into(),
7948 outcome: sql_idempotency::SqlReceiptOutcome {
7949 committed: true,
7950 committed_statements: 1,
7951 last_commit_epoch: Some(42),
7952 last_commit_epoch_text: Some("42".into()),
7953 first_commit_statement_index: Some(0),
7954 last_commit_statement_index: Some(0),
7955 completed_statements: 1,
7956 statement_index: 0,
7957 serialization: "succeeded".into(),
7958 },
7959 terminal_error: None,
7960 commit_receipt: None,
7961 };
7962 let error =
7963 restore_idempotency_replay(RegisteredQueryGuard::new(query), &receipt).unwrap_err();
7964 assert!(error.to_string().contains("invalid terminal state"));
7965 }
7966
7967 #[test]
7968 fn durable_replay_cancel_wins_before_receipt_response() {
7969 let registry = Arc::new(SqlQueryRegistry::default());
7970 let query_id: QueryId = "22334455667788990011aabbccddeeff".parse().unwrap();
7971 let query = registry
7972 .register(SqlQueryOptions {
7973 query_id: Some(query_id),
7974 ..SqlQueryOptions::default()
7975 })
7976 .unwrap();
7977 assert_eq!(
7978 query.request_cancel(CancellationReason::ClientRequest),
7979 CancelOutcome::Accepted
7980 );
7981 let receipt = sql_idempotency::SqlDurableReceipt {
7982 original_query_id: "00112233445566778899aabbccddeeff".into(),
7983 status: "completed".into(),
7984 server_state: "completed".into(),
7985 cancellation_reason: "none".into(),
7986 outcome: sql_idempotency::SqlReceiptOutcome {
7987 committed: true,
7988 committed_statements: 1,
7989 last_commit_epoch: Some(42),
7990 last_commit_epoch_text: Some("42".into()),
7991 first_commit_statement_index: Some(0),
7992 last_commit_statement_index: Some(0),
7993 completed_statements: 1,
7994 statement_index: 0,
7995 serialization: "succeeded".into(),
7996 },
7997 terminal_error: None,
7998 commit_receipt: None,
7999 };
8000 let error = restore_idempotency_replay(RegisteredQueryGuard::new(query), &receipt)
8001 .expect_err("accepted cancellation must suppress replay success");
8002 assert!(matches!(
8003 error,
8004 mongreldb_query::MongrelQueryError::QueryCancelled { .. }
8005 ));
8006 let status = registry.status(query_id).unwrap();
8007 assert_eq!(status.phase, SqlQueryPhase::Cancelled);
8008 assert!(status.durable_outcome.committed);
8009 }
8010
8011 #[test]
8012 fn direct_query_handle_preserves_receipt_after_tombstone_eviction() {
8013 let registry = Arc::new(SqlQueryRegistry::new(
8014 1,
8015 1,
8016 usize::MAX,
8017 std::time::Duration::from_secs(60),
8018 ));
8019 let first = registry.register(SqlQueryOptions::default()).unwrap();
8020 first
8021 .transition(SqlQueryPhase::Queued, SqlQueryPhase::Executing)
8022 .unwrap();
8023 first.record_commit(0, 42);
8024 first.complete_current_statement();
8025 first.try_complete().unwrap();
8026
8027 let second = registry.register(SqlQueryOptions::default()).unwrap();
8028 second
8029 .transition(SqlQueryPhase::Queued, SqlQueryPhase::Executing)
8030 .unwrap();
8031 second.try_complete().unwrap();
8032
8033 assert!(registry.status(first.id()).is_none());
8034 let receipt = sql_terminal_idempotency_receipt(&first.status()).unwrap();
8035 assert_eq!(receipt.outcome.committed_statements, 1);
8036 assert_eq!(receipt.outcome.last_commit_epoch, Some(42));
8037 }
8038
8039 #[test]
8040 fn cancellation_checkpoint_mismatch_returns_typed_error() {
8041 let registry = Arc::new(SqlQueryRegistry::default());
8042 let query = registry.register(SqlQueryOptions::default()).unwrap();
8043 assert!(matches!(
8044 cancellation_checkpoint_error(&query),
8045 mongreldb_query::MongrelQueryError::InvalidQueryState(_)
8046 ));
8047 assert_eq!(
8048 query.request_cancel(CancellationReason::ClientRequest),
8049 CancelOutcome::Accepted
8050 );
8051 assert!(matches!(
8052 cancellation_checkpoint_error(&query),
8053 mongreldb_query::MongrelQueryError::QueryCancelled { .. }
8054 ));
8055 }
8056
8057 #[tokio::test]
8058 async fn cancellation_after_commit_reports_durable_outcome() {
8059 let registry = Arc::new(SqlQueryRegistry::default());
8060 let query = registry.register(SqlQueryOptions::default()).unwrap();
8061 let query_id = query.id();
8062 query
8063 .transition(SqlQueryPhase::Queued, SqlQueryPhase::Executing)
8064 .unwrap();
8065 query.record_commit(0, 42);
8066 assert_eq!(
8067 query.request_cancel(CancellationReason::ClientRequest),
8068 CancelOutcome::Accepted
8069 );
8070 let error = query.checkpoint().unwrap_err();
8071 query.fail();
8072 let status = registry.status(query_id).unwrap();
8073
8074 let response = query_error_response_with_status(&error, Some(query_id), Some(&status));
8075 assert_eq!(response.status(), StatusCode::CONFLICT);
8076 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
8077 .await
8078 .unwrap();
8079 let body: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
8080 assert_eq!(body["status"], "cancelled_after_commit");
8081 assert_eq!(body["error"]["code"], "QUERY_CANCELLED_AFTER_COMMIT");
8082 assert_eq!(body["committed"], true);
8083 assert_eq!(body["outcome"]["committed_statements"], 1);
8084 assert_eq!(body["outcome"]["last_commit_epoch"], 42);
8085 assert_eq!(body["outcome"]["last_commit_epoch_text"], "42");
8086 assert_eq!(body["first_commit_statement_index"], 0);
8087 assert_eq!(body["last_commit_statement_index"], 0);
8088 assert_eq!(body["outcome"]["first_commit_statement_index"], 0);
8089 assert_eq!(body["outcome"]["last_commit_statement_index"], 0);
8090
8091 let response = query_error_response_with_status(&error, Some(query_id), None);
8092 assert_eq!(response.status(), StatusCode::CONFLICT);
8093 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
8094 .await
8095 .unwrap();
8096 let body: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
8097 assert_eq!(body["status"], "cancelled_after_commit");
8098 assert_eq!(body["error"]["code"], "QUERY_CANCELLED_AFTER_COMMIT");
8099 assert_eq!(body["committed"], true);
8100 assert_eq!(body["committed_statements"], 1);
8101 assert_eq!(body["last_commit_epoch"], 42);
8102 assert_eq!(body["last_commit_epoch_text"], "42");
8103 assert_eq!(body["first_commit_statement_index"], 0);
8104 assert_eq!(body["last_commit_statement_index"], 0);
8105 assert_eq!(body["outcome"]["committed"], true);
8106 assert_eq!(body["outcome"]["committed_statements"], 1);
8107 assert_eq!(body["outcome"]["last_commit_epoch"], 42);
8108 assert_eq!(body["outcome"]["last_commit_epoch_text"], "42");
8109 assert_eq!(body["outcome"]["first_commit_statement_index"], 0);
8110 assert_eq!(body["outcome"]["last_commit_statement_index"], 0);
8111 }
8112
8113 #[tokio::test]
8114 async fn commit_outcome_fallback_preserves_exact_progress() {
8115 let query_id: QueryId = "33445566778899001122aabbccddeeff".parse().unwrap();
8116 let error = mongreldb_query::MongrelQueryError::CommitOutcome {
8117 query_id,
8118 committed: true,
8119 committed_statements: 3,
8120 last_commit_epoch: Some(77),
8121 first_commit_statement_index: Some(1),
8122 last_commit_statement_index: Some(4),
8123 completed_statements: 4,
8124 statement_index: 5,
8125 message: "durable outcome retained".into(),
8126 };
8127 let response = query_error_response_with_status(&error, Some(query_id), None);
8128 assert_eq!(response.status(), StatusCode::CONFLICT);
8129 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
8130 .await
8131 .unwrap();
8132 let body: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
8133 assert_eq!(body["status"], "committed_with_error");
8134 assert_eq!(body["committed"], true);
8135 assert_eq!(body["committed_statements"], 3);
8136 assert_eq!(body["last_commit_epoch"], 77);
8137 assert_eq!(body["last_commit_epoch_text"], "77");
8138 assert_eq!(body["first_commit_statement_index"], 1);
8139 assert_eq!(body["last_commit_statement_index"], 4);
8140 assert_eq!(body["completed_statements"], 4);
8141 assert_eq!(body["statement_index"], 5);
8142 assert_eq!(body["outcome"]["committed_statements"], 3);
8143 assert_eq!(body["outcome"]["last_commit_epoch"], 77);
8144 assert_eq!(body["outcome"]["first_commit_statement_index"], 1);
8145 assert_eq!(body["outcome"]["last_commit_statement_index"], 4);
8146 assert_eq!(body["outcome"]["completed_statements"], 4);
8147 assert_eq!(body["outcome"]["statement_index"], 5);
8148 }
8149
8150 #[test]
8151 fn status_cancel_outcome_matches_cancel_endpoint_state() {
8152 let registry = Arc::new(SqlQueryRegistry::default());
8153 let commit = registry.register(SqlQueryOptions::default()).unwrap();
8154 commit
8155 .transition(SqlQueryPhase::Queued, SqlQueryPhase::Executing)
8156 .unwrap();
8157 commit.enter_commit_critical().unwrap();
8158 assert_eq!(
8159 query_cancel_outcome(®istry.status(commit.id()).unwrap()),
8160 Some("too_late")
8161 );
8162
8163 let completed = registry.register(SqlQueryOptions::default()).unwrap();
8164 completed
8165 .transition(SqlQueryPhase::Queued, SqlQueryPhase::Executing)
8166 .unwrap();
8167 completed.try_complete().unwrap();
8168 assert_eq!(
8169 query_cancel_outcome(®istry.status(completed.id()).unwrap()),
8170 Some("already_finished")
8171 );
8172 }
8173}
8174
8175#[cfg(test)]
8176mod wal_stream_tests {
8177 use super::*;
8178 use mongreldb_client::ReplicationFollower;
8179 use mongreldb_core::Database;
8180 use tempfile::tempdir;
8181
8182 #[tokio::test]
8183 async fn wal_stream_returns_records_after_commit() {
8184 let dir = tempdir().unwrap();
8185 let db = Arc::new(Database::create(dir.path()).unwrap());
8186 let table_schema = mongreldb_core::schema::Schema {
8187 schema_id: 1,
8188 columns: vec![mongreldb_core::schema::ColumnDef {
8189 id: 1,
8190 name: "id".into(),
8191 ty: TypeId::Int64,
8192 flags: mongreldb_core::schema::ColumnFlags::empty()
8193 .with(mongreldb_core::schema::ColumnFlags::PRIMARY_KEY),
8194 default_value: None,
8195 embedding_source: None,
8196 }],
8197 indexes: vec![],
8198 colocation: vec![],
8199 constraints: Default::default(),
8200 clustered: false,
8201 };
8202 db.create_table("items", table_schema).unwrap();
8203 let handle = db.table("items").unwrap();
8205 handle.lock().put(vec![(1, Value::Int64(1))]).unwrap();
8206 handle.lock().flush().unwrap();
8207
8208 let app = build_app(db);
8209 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
8210 let addr = listener.local_addr().unwrap();
8211 tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
8212 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
8213
8214 let resp = reqwest::get(format!("http://{addr}/wal/stream"))
8215 .await
8216 .unwrap();
8217 assert_eq!(resp.status(), 200);
8218 let body = resp.text().await.unwrap();
8219 assert!(!body.is_empty(), "wal_stream should return records");
8221 assert!(body.contains("seq"), "response should contain seq field");
8222 }
8223
8224 #[tokio::test]
8225 async fn follower_bootstraps_and_applies_incremental_commit() {
8226 let leader_dir = tempdir().unwrap();
8227 let follower_dir = tempdir().unwrap();
8228 let follower_path = follower_dir.path().join("copy");
8229 let db = Arc::new(Database::create(leader_dir.path()).unwrap());
8230 db.create_table(
8231 "items",
8232 mongreldb_core::schema::Schema {
8233 schema_id: 1,
8234 columns: vec![mongreldb_core::schema::ColumnDef {
8235 id: 1,
8236 name: "id".into(),
8237 ty: TypeId::Int64,
8238 flags: mongreldb_core::schema::ColumnFlags::empty()
8239 .with(mongreldb_core::schema::ColumnFlags::PRIMARY_KEY),
8240 default_value: None,
8241 embedding_source: None,
8242 }],
8243 indexes: vec![],
8244 colocation: vec![],
8245 constraints: Default::default(),
8246 clustered: false,
8247 },
8248 )
8249 .unwrap();
8250 let handle = db.table("items").unwrap();
8251 handle.lock().put(vec![(1, Value::Int64(1))]).unwrap();
8252 handle.lock().commit().unwrap();
8253
8254 let app = build_app_with_config(
8255 Arc::clone(&db),
8256 std::iter::empty::<Arc<dyn ExternalTableModule>>(),
8257 Some("replication-secret".into()),
8258 None,
8259 );
8260 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
8261 let addr = listener.local_addr().unwrap();
8262 tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
8263 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
8264
8265 let leader_url = format!("http://{addr}");
8266 let first_path = follower_path.clone();
8267 let (mut follower, initial) = tokio::task::spawn_blocking(move || {
8268 let mut follower = ReplicationFollower::new(&leader_url, first_path)
8269 .unwrap()
8270 .with_bearer_token("replication-secret");
8271 let applied = follower.sync().unwrap();
8272 (follower, applied)
8273 })
8274 .await
8275 .unwrap();
8276 assert_eq!(initial, 0);
8277
8278 handle.lock().put(vec![(1, Value::Int64(2))]).unwrap();
8279 handle.lock().commit().unwrap();
8280 let applied = tokio::task::spawn_blocking(move || {
8281 let count = follower.sync().unwrap();
8282 (follower, count)
8283 })
8284 .await
8285 .unwrap();
8286 follower = applied.0;
8287 assert!(applied.1 > 0);
8288 assert!(follower.last_epoch() > 0);
8289
8290 let replica = Database::open(&follower_path).unwrap();
8291 assert_eq!(replica.table("items").unwrap().lock().count(), 2);
8292 drop(replica);
8293
8294 db.set_spill_threshold(1);
8295 db.transaction(|txn| {
8296 txn.put("items", vec![(1, Value::Int64(3))])?;
8297 Ok(())
8298 })
8299 .unwrap();
8300 let (follower_after_bootstrap, applied) = tokio::task::spawn_blocking(move || {
8301 let count = follower.sync().unwrap();
8302 (follower, count)
8303 })
8304 .await
8305 .unwrap();
8306 assert_eq!(applied, 0, "spilled run should trigger safe rebootstrap");
8307 assert!(follower_after_bootstrap.last_epoch() > 0);
8308 let replica = Database::open(&follower_path).unwrap();
8309 assert_eq!(replica.table("items").unwrap().lock().count(), 3);
8310 }
8311}
8312
8313#[cfg(test)]
8314mod metrics_tests {
8315 use super::*;
8316 use mongreldb_core::Database;
8317 use tempfile::tempdir;
8318
8319 async fn setup() -> (tempfile::TempDir, std::net::SocketAddr) {
8322 let dir = tempdir().unwrap();
8323 let db = Arc::new(Database::create(dir.path()).unwrap());
8324 let table_schema = mongreldb_core::schema::Schema {
8325 schema_id: 1,
8326 columns: vec![mongreldb_core::schema::ColumnDef {
8327 id: 1,
8328 name: "id".into(),
8329 ty: TypeId::Int64,
8330 flags: mongreldb_core::schema::ColumnFlags::empty()
8331 .with(mongreldb_core::schema::ColumnFlags::PRIMARY_KEY),
8332 default_value: None,
8333 embedding_source: None,
8334 }],
8335 indexes: vec![],
8336 colocation: vec![],
8337 constraints: Default::default(),
8338 clustered: false,
8339 };
8340 db.create_table("items", table_schema).unwrap();
8341 let app = build_app(db);
8342 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
8343 let addr = listener.local_addr().unwrap();
8344 tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
8345 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
8346 (dir, addr)
8347 }
8348
8349 #[tokio::test]
8350 async fn metrics_endpoint_returns_prometheus_text() {
8351 let (_dir, addr) = setup().await;
8352 let client = reqwest::Client::new();
8353
8354 let _ = client
8356 .post(format!("http://{addr}/tables/items/put"))
8357 .json(&json!({ "row": [1, 1] }))
8358 .send()
8359 .await
8360 .unwrap();
8361 let _ = client
8362 .post(format!("http://{addr}/sql"))
8363 .json(&json!({ "sql": "SELECT count(*) FROM items" }))
8364 .send()
8365 .await
8366 .unwrap();
8367
8368 let resp = client
8369 .get(format!("http://{addr}/metrics"))
8370 .send()
8371 .await
8372 .unwrap();
8373 assert_eq!(resp.status(), 200);
8374 let ct = resp
8375 .headers()
8376 .get("content-type")
8377 .and_then(|v| v.to_str().ok())
8378 .unwrap_or_default()
8379 .to_string();
8380 assert!(
8381 ct.contains("text/plain"),
8382 "content-type is prometheus text: {ct}"
8383 );
8384 let body = resp.text().await.unwrap();
8385 assert!(body.contains("# TYPE mongreldb_sql_queries_total counter"));
8387 assert!(body.contains("# TYPE mongreldb_puts_total counter"));
8388 assert!(body.contains("# TYPE mongreldb_tables gauge"));
8389 assert!(
8391 body.contains("mongreldb_sql_queries_total 1"),
8392 "sql_queries counter should reflect the /sql call: {body}"
8393 );
8394 assert!(
8395 body.contains("mongreldb_puts_total 1"),
8396 "puts counter should reflect the put call: {body}"
8397 );
8398 assert!(body.contains("mongreldb_tables 1"));
8400 }
8401
8402 #[tokio::test]
8403 async fn metrics_error_counter_increments_on_bad_sql() {
8404 let (_dir, addr) = setup().await;
8405 let client = reqwest::Client::new();
8406 let _ = client
8408 .post(format!("http://{addr}/sql"))
8409 .json(&json!({ "sql": "SELECT * FROM does_not_exist" }))
8410 .send()
8411 .await
8412 .unwrap();
8413 let body = client
8414 .get(format!("http://{addr}/metrics"))
8415 .send()
8416 .await
8417 .unwrap()
8418 .text()
8419 .await
8420 .unwrap();
8421 assert!(
8422 body.contains("mongreldb_sql_errors_total 1"),
8423 "sql_errors should increment on a failed query: {body}"
8424 );
8425 }
8426
8427 #[tokio::test]
8428 async fn arrow_stream_returns_ipc_stream_bytes() {
8429 let (_dir, addr) = setup().await;
8430 let client = reqwest::Client::new();
8431 for i in 1..=3 {
8434 let resp = client
8435 .post(format!("http://{addr}/tables/items/put"))
8436 .json(&json!({ "row": [1, i] }))
8437 .send()
8438 .await
8439 .unwrap();
8440 assert_eq!(resp.status(), 200, "put should succeed");
8441 }
8442 let _ = client
8443 .post(format!("http://{addr}/tables/items/commit"))
8444 .send()
8445 .await
8446 .unwrap();
8447 let count_body = client
8449 .get(format!("http://{addr}/tables/items/count"))
8450 .send()
8451 .await
8452 .unwrap()
8453 .text()
8454 .await
8455 .unwrap();
8456 assert!(
8457 count_body.contains("\"count\":3"),
8458 "expected 3 visible rows, got: {count_body}"
8459 );
8460 let resp = client
8461 .post(format!("http://{addr}/sql"))
8462 .json(&json!({ "sql": "SELECT count(*) FROM items", "format": "arrow-stream" }))
8463 .send()
8464 .await
8465 .unwrap();
8466 assert_eq!(resp.status(), 200, "streaming query should succeed");
8467 let ct = resp
8468 .headers()
8469 .get("content-type")
8470 .and_then(|v| v.to_str().ok())
8471 .unwrap_or_default()
8472 .to_string();
8473 assert!(
8474 ct.contains("application/vnd.apache.arrow.stream"),
8475 "content-type should be the arrow stream format: {ct}"
8476 );
8477 let bytes = resp.bytes().await.unwrap();
8478 assert!(
8482 !bytes.is_empty(),
8483 "arrow stream body should contain schema + batch + EOS"
8484 );
8485 assert!(
8486 bytes.starts_with(&0xFFFFFFFFu32.to_le_bytes()),
8487 "arrow stream must begin with the IPC continuation marker"
8488 );
8489 assert!(
8490 bytes.ends_with(&[0u8, 0, 0, 0]),
8491 "arrow stream should end with the EOS marker (trailing zero length)"
8492 );
8493 }
8494}
8495
8496#[cfg(test)]
8497mod streaming_tests {
8498 use super::*;
8499 use arrow::array::Int64Array;
8500 use arrow::datatypes::{DataType, Field, Schema};
8501 use arrow::record_batch::RecordBatch;
8502 use datafusion::common::DataFusionError;
8503 use datafusion::physical_plan::stream::RecordBatchStreamAdapter;
8504 use futures::StreamExt;
8505 use std::sync::Arc as StdArc;
8506
8507 fn batch_stream(batches: Vec<RecordBatch>) -> mongreldb_query::MongrelRecordBatchStream {
8508 let schema = batches
8509 .first()
8510 .map(RecordBatch::schema)
8511 .unwrap_or_else(|| StdArc::new(Schema::empty()));
8512 let batches =
8513 futures::stream::iter(batches.into_iter().map(Ok::<RecordBatch, DataFusionError>));
8514 Box::pin(RecordBatchStreamAdapter::new(schema, batches))
8515 }
8516
8517 #[tokio::test]
8522 async fn arrow_stream_serializes_multiple_batches_roundtrip() {
8523 let schema = StdArc::new(Schema::new(vec![Field::new("v", DataType::Int64, false)]));
8524 let b1 = RecordBatch::try_new(
8525 schema.clone(),
8526 vec![StdArc::new(Int64Array::from(vec![1, 2]))],
8527 )
8528 .unwrap();
8529 let b2 = RecordBatch::try_new(
8530 schema.clone(),
8531 vec![StdArc::new(Int64Array::from(vec![3, 4, 5]))],
8532 )
8533 .unwrap();
8534
8535 let resp = sql_arrow_stream_response(batch_stream(vec![b1, b2]));
8536 let bytes = axum::body::to_bytes(resp.into_body(), 1 << 20)
8537 .await
8538 .unwrap();
8539
8540 assert!(bytes.starts_with(&0xFFFFFFFFu32.to_le_bytes()));
8542
8543 let slice: &[u8] = bytes.as_ref();
8546 let mut reader = arrow::ipc::reader::StreamReader::try_new(slice, None).unwrap();
8547 let mut total_rows = 0;
8548 for batch in reader.by_ref() {
8549 let batch = batch.expect("each IPC message should decode");
8550 total_rows += batch.num_rows();
8551 }
8552 assert_eq!(
8553 total_rows, 5,
8554 "all rows should round-trip through the stream"
8555 );
8556 }
8557
8558 #[tokio::test]
8559 async fn arrow_stream_emits_schema_before_first_batch() {
8560 let schema = StdArc::new(Schema::new(vec![Field::new("v", DataType::Int64, false)]));
8561 let pending = futures::stream::pending::<Result<RecordBatch, DataFusionError>>();
8562 let batches = Box::pin(RecordBatchStreamAdapter::new(schema, pending));
8563 let mut body = sql_arrow_stream_response(batches)
8564 .into_body()
8565 .into_data_stream();
8566
8567 let chunk = tokio::time::timeout(std::time::Duration::from_millis(100), body.next())
8568 .await
8569 .expect("schema chunk should not wait for a query batch")
8570 .unwrap()
8571 .unwrap();
8572 assert!(chunk.starts_with(&0xFFFFFFFFu32.to_le_bytes()));
8573 }
8574
8575 #[tokio::test]
8576 async fn arrow_stream_empty_query_is_valid_ipc() {
8577 let resp = sql_arrow_stream_response(batch_stream(Vec::new()));
8578 let bytes = axum::body::to_bytes(resp.into_body(), 1 << 20)
8579 .await
8580 .unwrap();
8581 let slice: &[u8] = bytes.as_ref();
8582 let reader = arrow::ipc::reader::StreamReader::try_new(slice, None).unwrap();
8583 assert_eq!(reader.count(), 0);
8584 }
8585
8586 #[tokio::test]
8587 async fn buffered_output_limits_are_typed() {
8588 let dir = tempfile::tempdir().unwrap();
8589 let db = StdArc::new(mongreldb_core::Database::create(dir.path()).unwrap());
8590 let session = MongrelSession::open(db).unwrap();
8591 let query = session.register_query(SqlQueryOptions::default()).unwrap();
8592 let output = session
8593 .run_with_query_for_serialization("SELECT 1", query)
8594 .await
8595 .unwrap();
8596
8597 let row_error =
8598 serialize_buffered_output("json", output.batches(), output.query(), 0, 1024, None)
8599 .unwrap_err();
8600 assert!(matches!(row_error, BufferedSerializationError::Limit(_)));
8601 let byte_error =
8602 serialize_buffered_output("json", output.batches(), output.query(), 10, 1, None)
8603 .unwrap_err();
8604 assert!(matches!(byte_error, BufferedSerializationError::Limit(_)));
8605 output.fail();
8606 }
8607}
8608
8609#[cfg(test)]
8610mod audit_tests {
8611 use super::*;
8612 use mongreldb_core::Database;
8613 use tempfile::tempdir;
8614
8615 async fn auth_setup(password: &str) -> std::net::SocketAddr {
8617 let dir = tempdir().unwrap();
8618 let db = Arc::new(Database::create(dir.path()).unwrap());
8619 db.create_user("alice", password).unwrap();
8620 db.set_user_admin("alice", true).unwrap();
8621 let app = build_app_full(
8622 db,
8623 std::iter::empty::<Arc<dyn ExternalTableModule>>(),
8624 None,
8625 None,
8626 true,
8627 );
8628 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
8629 let addr = listener.local_addr().unwrap();
8630 tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
8631 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
8632 addr
8633 }
8634
8635 #[tokio::test]
8636 async fn audit_records_login_success_and_failure() {
8637 let addr = auth_setup("s3cret").await;
8638 let client = reqwest::Client::new();
8639
8640 let resp = client
8642 .get(format!("http://{addr}/health"))
8643 .header("Authorization", basic("alice", "s3cret"))
8644 .send()
8645 .await
8646 .unwrap();
8647 assert_eq!(resp.status(), 200);
8648
8649 let resp = client
8651 .get(format!("http://{addr}/health"))
8652 .header("Authorization", basic("alice", "wrong"))
8653 .send()
8654 .await
8655 .unwrap();
8656 assert_eq!(resp.status(), 401);
8657
8658 let body = client
8659 .get(format!("http://{addr}/audit"))
8660 .header("Authorization", basic("alice", "s3cret"))
8661 .send()
8662 .await
8663 .unwrap()
8664 .text()
8665 .await
8666 .unwrap();
8667 assert!(
8668 body.contains("\"action\":\"login.ok\""),
8669 "audit should record the successful login: {body}"
8670 );
8671 assert!(
8672 body.contains("\"action\":\"login.fail\""),
8673 "audit should record the failed login: {body}"
8674 );
8675 assert!(
8676 body.contains("\"principal\":\"alice\""),
8677 "audit should attribute events to alice: {body}"
8678 );
8679 }
8680
8681 #[tokio::test]
8682 async fn audit_records_ddl_sql() {
8683 let dir = tempdir().unwrap();
8684 let db = Arc::new(Database::create(dir.path()).unwrap());
8685 let app = build_app(db);
8686 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
8687 let addr = listener.local_addr().unwrap();
8688 tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
8689 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
8690
8691 let client = reqwest::Client::new();
8692 let _ = client
8693 .post(format!("http://{addr}/sql"))
8694 .json(&json!({ "sql": "CREATE TABLE t (id BIGINT PRIMARY KEY)" }))
8695 .send()
8696 .await
8697 .unwrap();
8698 let _ = client
8700 .post(format!("http://{addr}/sql"))
8701 .json(&json!({ "sql": "SELECT 1" }))
8702 .send()
8703 .await
8704 .unwrap();
8705
8706 let body = client
8707 .get(format!("http://{addr}/audit"))
8708 .send()
8709 .await
8710 .unwrap()
8711 .text()
8712 .await
8713 .unwrap();
8714 assert!(
8715 body.contains("\"action\":\"ddl.ok\""),
8716 "audit should record the successful DDL statement: {body}"
8717 );
8718 assert!(
8719 body.contains("CREATE TABLE"),
8720 "audit detail should carry the DDL snippet: {body}"
8721 );
8722 assert!(
8724 !body.contains("SELECT 1"),
8725 "non-DDL reads should not be audited: {body}"
8726 );
8727 }
8728
8729 #[tokio::test]
8730 async fn audit_redacts_credential_passwords() {
8731 let dir = tempdir().unwrap();
8732 let db = Arc::new(Database::create(dir.path()).unwrap());
8733 let app = build_app(db);
8734 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
8735 let addr = listener.local_addr().unwrap();
8736 tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
8737 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
8738
8739 let client = reqwest::Client::new();
8740 let _ = client
8743 .post(format!("http://{addr}/sql"))
8744 .json(&json!({ "sql": "CREATE USER alice WITH PASSWORD 'topsecret'" }))
8745 .send()
8746 .await
8747 .unwrap();
8748
8749 let body = client
8750 .get(format!("http://{addr}/audit"))
8751 .send()
8752 .await
8753 .unwrap()
8754 .text()
8755 .await
8756 .unwrap();
8757 assert!(
8758 !body.contains("topsecret"),
8759 "password must never appear in the audit log: {body}"
8760 );
8761 assert!(
8762 body.contains("redacted credential statement"),
8763 "credential DDL should be recorded as redacted: {body}"
8764 );
8765 }
8766
8767 fn basic(user: &str, pass: &str) -> String {
8768 let raw = format!("{user}:{pass}");
8769 format!("Basic {}", base64_encode(raw.as_bytes()))
8770 }
8771
8772 fn base64_encode(input: &[u8]) -> String {
8773 const TABLE: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
8774 let mut out = String::new();
8775 let mut buf = 0u32;
8776 let mut bits = 0u32;
8777 for &b in input {
8778 buf = (buf << 8) | b as u32;
8779 bits += 8;
8780 while bits >= 6 {
8781 bits -= 6;
8782 out.push(TABLE[((buf >> bits) & 0x3F) as usize] as char);
8783 }
8784 }
8785 if bits > 0 {
8786 out.push(TABLE[((buf << (6 - bits)) & 0x3F) as usize] as char);
8787 }
8788 while !out.len().is_multiple_of(4) {
8789 out.push('=');
8790 }
8791 out
8792 }
8793}
8794
8795#[cfg(test)]
8796mod session_tests {
8797 use super::*;
8798 use mongreldb_core::Database;
8799 use tempfile::tempdir;
8800
8801 async fn setup() -> (tempfile::TempDir, std::net::SocketAddr) {
8805 let dir = tempdir().unwrap();
8806 let db = Arc::new(Database::create(dir.path()).unwrap());
8807 let table_schema = mongreldb_core::schema::Schema {
8808 schema_id: 1,
8809 columns: vec![mongreldb_core::schema::ColumnDef {
8810 id: 1,
8811 name: "id".into(),
8812 ty: TypeId::Int64,
8813 flags: mongreldb_core::schema::ColumnFlags::empty()
8814 .with(mongreldb_core::schema::ColumnFlags::PRIMARY_KEY),
8815 default_value: None,
8816 embedding_source: None,
8817 }],
8818 indexes: vec![],
8819 colocation: vec![],
8820 constraints: Default::default(),
8821 clustered: false,
8822 };
8823 db.create_table("items", table_schema).unwrap();
8824 let app = build_app(db);
8825 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
8826 let addr = listener.local_addr().unwrap();
8827 tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
8828 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
8829 (dir, addr)
8830 }
8831
8832 async fn open_session(client: &reqwest::Client, addr: &std::net::SocketAddr) -> String {
8834 let resp = client
8835 .post(format!("http://{addr}/sessions"))
8836 .send()
8837 .await
8838 .unwrap();
8839 assert_eq!(resp.status(), 200);
8840 resp.json::<serde_json::Value>()
8841 .await
8842 .unwrap()
8843 .get("session_id")
8844 .unwrap()
8845 .as_str()
8846 .unwrap()
8847 .to_string()
8848 }
8849
8850 async fn sql_on(
8851 client: &reqwest::Client,
8852 addr: &std::net::SocketAddr,
8853 session: &str,
8854 sql: &str,
8855 ) -> reqwest::Response {
8856 client
8857 .post(format!("http://{addr}/sql"))
8858 .header("X-Session-ID", session)
8859 .json(&json!({ "sql": sql }))
8860 .send()
8861 .await
8862 .unwrap()
8863 }
8864
8865 async fn count_items(client: &reqwest::Client, addr: &std::net::SocketAddr) -> u64 {
8866 client
8867 .get(format!("http://{addr}/tables/items/count"))
8868 .send()
8869 .await
8870 .unwrap()
8871 .json::<serde_json::Value>()
8872 .await
8873 .unwrap()
8874 .get("count")
8875 .unwrap()
8876 .as_u64()
8877 .unwrap()
8878 }
8879
8880 #[tokio::test]
8881 async fn cross_request_transaction_commits() {
8882 let (_dir, addr) = setup().await;
8883 let client = reqwest::Client::new();
8884 let session = open_session(&client, &addr).await;
8885
8886 let r = sql_on(&client, &addr, &session, "BEGIN").await;
8888 assert_eq!(r.status(), 200);
8889 let r = sql_on(
8890 &client,
8891 &addr,
8892 &session,
8893 "INSERT INTO items (id) VALUES (1)",
8894 )
8895 .await;
8896 assert_eq!(r.status(), 200, "INSERT should stage successfully");
8897 assert_eq!(count_items(&client, &addr).await, 0);
8899 let r = sql_on(&client, &addr, &session, "COMMIT").await;
8900 assert_eq!(r.status(), 200);
8901
8902 assert_eq!(count_items(&client, &addr).await, 1);
8904 }
8905
8906 #[tokio::test]
8907 async fn cross_request_transaction_rolls_back() {
8908 let (_dir, addr) = setup().await;
8909 let client = reqwest::Client::new();
8910 let session = open_session(&client, &addr).await;
8911
8912 sql_on(&client, &addr, &session, "BEGIN").await;
8913 sql_on(
8914 &client,
8915 &addr,
8916 &session,
8917 "INSERT INTO items (id) VALUES (5)",
8918 )
8919 .await;
8920 let r = sql_on(&client, &addr, &session, "ROLLBACK").await;
8922 assert_eq!(r.status(), 200);
8923 assert_eq!(
8924 count_items(&client, &addr).await,
8925 0,
8926 "rollback discards staged writes"
8927 );
8928 }
8929
8930 #[tokio::test]
8931 async fn unknown_session_id_is_404() {
8932 let (_dir, addr) = setup().await;
8933 let client = reqwest::Client::new();
8934 let resp = client
8935 .post(format!("http://{addr}/sql"))
8936 .header("X-Session-ID", "does-not-exist")
8937 .json(&json!({ "sql": "SELECT 1" }))
8938 .send()
8939 .await
8940 .unwrap();
8941 assert_eq!(resp.status(), 404);
8942 }
8943
8944 #[tokio::test]
8945 async fn invalid_session_headers_do_not_autocommit() {
8946 let (_dir, addr) = setup().await;
8947 let client = reqwest::Client::new();
8948
8949 let resp = client
8950 .post(format!("http://{addr}/sql"))
8951 .header(
8952 "X-Session-ID",
8953 reqwest::header::HeaderValue::from_bytes(&[0xff]).unwrap(),
8954 )
8955 .json(&json!({ "sql": "INSERT INTO items (id) VALUES (1)" }))
8956 .send()
8957 .await
8958 .unwrap();
8959 assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
8960
8961 let resp = client
8962 .post(format!("http://{addr}/sql"))
8963 .header("X-Session-ID", "x".repeat(257))
8964 .json(&json!({ "sql": "INSERT INTO items (id) VALUES (2)" }))
8965 .send()
8966 .await
8967 .unwrap();
8968 assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
8969 assert_eq!(count_items(&client, &addr).await, 0);
8970 }
8971
8972 #[tokio::test]
8973 async fn close_session_ends_cross_request_state() {
8974 let (_dir, addr) = setup().await;
8975 let client = reqwest::Client::new();
8976 let session = open_session(&client, &addr).await;
8977
8978 sql_on(&client, &addr, &session, "BEGIN").await;
8980 let r = client
8981 .delete(format!("http://{addr}/sessions/{session}"))
8982 .send()
8983 .await
8984 .unwrap();
8985 assert_eq!(r.status(), 200);
8986
8987 let resp = sql_on(&client, &addr, &session, "COMMIT").await;
8989 assert_eq!(resp.status(), 404, "closed session is no longer usable");
8990 }
8991
8992 #[tokio::test]
8993 async fn no_session_header_uses_fresh_ephemeral_session() {
8994 let (_dir, addr) = setup().await;
8997 let client = reqwest::Client::new();
8998 let resp = client
8999 .post(format!("http://{addr}/sql"))
9000 .json(&json!({ "sql": "INSERT INTO items (id) VALUES (42)" }))
9001 .send()
9002 .await
9003 .unwrap();
9004 assert_eq!(resp.status(), 200);
9005 assert_eq!(count_items(&client, &addr).await, 1);
9006 }
9007
9008 #[tokio::test]
9009 async fn prepared_statement_prepare_execute_and_reuse() {
9010 let (_dir, addr) = setup().await;
9011 let client = reqwest::Client::new();
9012 let session = open_session(&client, &addr).await;
9013 sql_on(
9014 &client,
9015 &addr,
9016 &session,
9017 "INSERT INTO items (id) VALUES (1), (2), (3), (4)",
9018 )
9019 .await;
9020
9021 let resp = client
9023 .post(format!("http://{addr}/sessions/{session}/prepare"))
9024 .json(&json!({"name":"gt","sql":"SELECT id FROM items WHERE id > $1"}))
9025 .send()
9026 .await
9027 .unwrap();
9028 assert_eq!(resp.status(), 200);
9029
9030 let resp = client
9032 .post(format!("http://{addr}/sessions/{session}/execute"))
9033 .json(&json!({"name":"gt","params":[2]}))
9034 .send()
9035 .await
9036 .unwrap();
9037 assert_eq!(resp.status(), 200);
9038 let body = resp.json::<serde_json::Value>().await.unwrap();
9039 let arr = body
9040 .as_array()
9041 .expect("execute returns a JSON array of rows");
9042 assert_eq!(arr.len(), 2, "ids > 2 are {{3,4}}: {body}");
9043
9044 let resp = client
9046 .post(format!("http://{addr}/sessions/{session}/execute"))
9047 .json(&json!({"name":"gt","params":[3]}))
9048 .send()
9049 .await
9050 .unwrap();
9051 let body = resp.json::<serde_json::Value>().await.unwrap();
9052 assert_eq!(body.as_array().unwrap().len(), 1, "ids > 3 is {{4}}");
9053 }
9054
9055 #[tokio::test]
9056 async fn prepared_statement_deallocate_then_execute_fails() {
9057 let (_dir, addr) = setup().await;
9058 let client = reqwest::Client::new();
9059 let session = open_session(&client, &addr).await;
9060 let _ = client
9061 .post(format!("http://{addr}/sessions/{session}/prepare"))
9062 .json(&json!({"name":"p","sql":"SELECT $1"}))
9063 .send()
9064 .await
9065 .unwrap();
9066 let deallocate_query_id = "dadadadadadadadadadadadadadadada";
9067 let resp = client
9068 .delete(format!("http://{addr}/sessions/{session}/statements/p"))
9069 .header("X-MongrelDB-Query-ID", deallocate_query_id)
9070 .header("X-MongrelDB-Timeout-Ms", "10000")
9071 .send()
9072 .await
9073 .unwrap();
9074 assert_eq!(resp.status(), 200);
9075 assert_eq!(
9076 resp.headers()
9077 .get("X-MongrelDB-Query-ID")
9078 .unwrap()
9079 .to_str()
9080 .unwrap(),
9081 deallocate_query_id
9082 );
9083 let status = client
9084 .get(format!("http://{addr}/queries/{deallocate_query_id}"))
9085 .send()
9086 .await
9087 .unwrap()
9088 .json::<serde_json::Value>()
9089 .await
9090 .unwrap();
9091 assert_eq!(status["state"], "completed");
9092 assert_eq!(status["operation"], "DEALLOCATE");
9093 let resp = client
9095 .post(format!("http://{addr}/sessions/{session}/execute"))
9096 .json(&json!({"name":"p","params":[1]}))
9097 .send()
9098 .await
9099 .unwrap();
9100 assert_ne!(resp.status(), 200, "execute after DEALLOCATE must fail");
9101 }
9102
9103 #[tokio::test]
9104 async fn prepared_statement_deallocate_honors_pre_registration_cancel() {
9105 let (_dir, addr) = setup().await;
9106 let client = reqwest::Client::new();
9107 let session = open_session(&client, &addr).await;
9108 let prepared = client
9109 .post(format!("http://{addr}/sessions/{session}/prepare"))
9110 .json(&json!({"name":"p","sql":"SELECT $1"}))
9111 .send()
9112 .await
9113 .unwrap();
9114 assert_eq!(prepared.status(), StatusCode::OK);
9115
9116 let query_id = "dbdbdbdbdbdbdbdbdbdbdbdbdbdbdbdb";
9117 let cancel = client
9118 .post(format!("http://{addr}/queries/{query_id}/cancel"))
9119 .header("X-Session-ID", &session)
9120 .send()
9121 .await
9122 .unwrap();
9123 assert_eq!(cancel.status(), StatusCode::ACCEPTED);
9124 let deallocate = client
9125 .delete(format!("http://{addr}/sessions/{session}/statements/p"))
9126 .header("X-MongrelDB-Query-ID", query_id)
9127 .send()
9128 .await
9129 .unwrap();
9130 assert_eq!(deallocate.status().as_u16(), 499);
9131 assert_eq!(
9132 deallocate.json::<serde_json::Value>().await.unwrap()["error"]["code"],
9133 "QUERY_CANCELLED"
9134 );
9135
9136 let execute = client
9137 .post(format!("http://{addr}/sessions/{session}/execute"))
9138 .json(&json!({"name":"p","params":[1]}))
9139 .send()
9140 .await
9141 .unwrap();
9142 assert_eq!(execute.status(), StatusCode::OK);
9143 }
9144
9145 #[tokio::test]
9146 async fn prepared_statement_rejects_bad_name() {
9147 let (_dir, addr) = setup().await;
9148 let client = reqwest::Client::new();
9149 let session = open_session(&client, &addr).await;
9150 let resp = client
9151 .post(format!("http://{addr}/sessions/{session}/prepare"))
9152 .json(&json!({"name":"1bad","sql":"SELECT 1"}))
9153 .send()
9154 .await
9155 .unwrap();
9156 assert_eq!(
9157 resp.status(),
9158 400,
9159 "statement name starting with a digit must be rejected"
9160 );
9161 }
9162}