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