1use std::collections::HashMap;
5use std::pin::Pin;
6use std::str::FromStr;
7use std::sync::{Arc, Mutex, RwLock};
8use std::time::{Duration, SystemTime, UNIX_EPOCH};
9
10use arrow::ipc::writer::StreamWriter;
11use futures::{Stream, StreamExt};
12use mongreldb_core::{
13 Database, JwksCache, JwtValidationConfig, MongrelError, Principal, ServiceToken,
14 ServiceTokenRegistry, WorkloadClass,
15};
16use mongreldb_protocol::native;
17use mongreldb_protocol::request::AuthenticatedIdentity;
18use mongreldb_protocol::validate_native_context;
19use mongreldb_query::{
20 CancelOutcome, ManagedQueryBatches, MongrelQueryError, MongrelSession, QueryId,
21 SerializationOutcome, SqlQueryOptions, SqlQueryPhase, SqlQueryRegistry,
22};
23use mongreldb_types::errors::{ErrorCategory, RetryClass};
24use mongreldb_types::ids::TransactionId;
25use prost::Message;
26use tonic::{Code, Request, Response, Status};
27
28use crate::prepared;
29use crate::sessions::{SessionEntry, SessionStore};
30
31const AUTH_TOKEN_TTL: Duration = Duration::from_secs(300);
32const MAX_AUTH_TOKENS: usize = 4_096;
33
34enum NativeIdempotency {
35 None,
36 Execute(crate::sql_idempotency::SqlIdempotencyExecution),
37 Replay(crate::sql_idempotency::SqlDurableReceipt),
38}
39
40#[derive(Clone)]
41struct AuthGrant {
42 principal: Option<Principal>,
43 identity: AuthenticatedIdentity,
44 expires_unix_micros: u64,
45}
46
47struct PendingScram {
48 username: String,
49 session: mongreldb_core::ScramServerSession,
50 fake_user: bool,
51 expires_unix_micros: u64,
52}
53
54#[derive(Clone)]
55struct NativeOidc {
56 validation: JwtValidationConfig,
57 cache: Arc<JwksCache<crate::oidc::HttpsJwksProvider>>,
58}
59
60#[derive(Clone, Default)]
62pub struct NativeExternalAuth {
63 service_tokens: Arc<RwLock<ServiceTokenRegistry>>,
64 oidc: Option<NativeOidc>,
65}
66
67impl NativeExternalAuth {
68 pub fn new() -> Self {
69 Self::default()
70 }
71
72 pub fn upsert_service_token(&self, token: ServiceToken) -> Result<(), MongrelError> {
73 self.service_tokens
74 .write()
75 .map_err(|_| MongrelError::Other("service-token registry poisoned".into()))?
76 .upsert(token);
77 Ok(())
78 }
79
80 pub fn with_oidc(
81 mut self,
82 validation: JwtValidationConfig,
83 provider: crate::oidc::HttpsJwksProvider,
84 ) -> Self {
85 self.oidc = Some(NativeOidc {
86 validation,
87 cache: Arc::new(JwksCache::new(provider)),
88 });
89 self
90 }
91}
92
93#[derive(Clone)]
95pub struct NativeRuntime {
96 db: Arc<Database>,
97 sessions: Arc<SessionStore>,
98 query_registry: Arc<SqlQueryRegistry>,
99 sql_idempotency: Option<Arc<crate::sql_idempotency::SqlIdempotencyStore>>,
100 auth_grants: Arc<Mutex<HashMap<Vec<u8>, AuthGrant>>>,
101 scram_exchanges: Arc<Mutex<HashMap<Vec<u8>, PendingScram>>>,
102 external_auth: Option<NativeExternalAuth>,
103 sql_semaphore: Arc<tokio::sync::Semaphore>,
104 scheduler: crate::admission::SchedulerAdmission,
105 node_admission: crate::admission::NodeAdmissionController,
107 sql_priority: u8,
108}
109
110impl NativeRuntime {
111 pub fn new(
112 db: Arc<Database>,
113 sessions: Arc<SessionStore>,
114 query_registry: Arc<SqlQueryRegistry>,
115 ) -> Self {
116 let groups = mongreldb_core::ResourceGroupRegistry::with_defaults();
117 let memory = mongreldb_core::MemoryGovernor::new(mongreldb_core::GovernorConfig::new(
118 512 * 1024 * 1024,
119 ))
120 .expect("default native node memory governor");
121 let node_admission = crate::admission::NodeAdmissionController::new(&groups, memory);
122 let scheduler = node_admission.scheduler().clone();
123 Self {
124 db,
125 sessions,
126 query_registry,
127 sql_idempotency: None,
128 auth_grants: Arc::new(Mutex::new(HashMap::new())),
129 scram_exchanges: Arc::new(Mutex::new(HashMap::new())),
130 external_auth: None,
131 sql_semaphore: Arc::new(tokio::sync::Semaphore::new(
132 crate::default_sql_max_concurrent(),
133 )),
134 scheduler,
135 node_admission,
136 sql_priority: crate::admission::priority_for_class(
137 &groups,
138 WorkloadClass::InteractiveSql,
139 ),
140 }
141 }
142
143 pub(crate) fn with_sql_admission(
144 mut self,
145 semaphore: Arc<tokio::sync::Semaphore>,
146 scheduler: crate::admission::SchedulerAdmission,
147 priority: u8,
148 ) -> Self {
149 self.sql_semaphore = semaphore;
150 self.scheduler = scheduler;
151 self.sql_priority = priority;
152 self
153 }
154
155 pub(crate) fn with_node_admission(
157 mut self,
158 node_admission: crate::admission::NodeAdmissionController,
159 ) -> Self {
160 self.scheduler = node_admission.scheduler().clone();
161 self.node_admission = node_admission;
162 self
163 }
164
165 pub(crate) fn with_sql_idempotency(
166 mut self,
167 store: Arc<crate::sql_idempotency::SqlIdempotencyStore>,
168 ) -> Self {
169 self.sql_idempotency = Some(store);
170 self
171 }
172
173 pub fn with_external_auth(mut self, auth: NativeExternalAuth) -> Self {
174 self.external_auth = Some(auth);
175 self
176 }
177
178 fn issue_auth_grant(
179 &self,
180 principal: Option<Principal>,
181 identity: AuthenticatedIdentity,
182 ) -> Result<Vec<u8>, Status> {
183 let token = mongreldb_types::ids::QueryId::new_random()
184 .as_bytes()
185 .to_vec();
186 let now = now_unix_micros();
187 let mut grants = self
188 .auth_grants
189 .lock()
190 .map_err(|_| Status::internal("auth grant store poisoned"))?;
191 grants.retain(|_, grant| grant.expires_unix_micros > now);
192 if grants.len() >= MAX_AUTH_TOKENS {
193 return Err(Status::resource_exhausted(
194 "too many pending authentications",
195 ));
196 }
197 grants.insert(
198 token.clone(),
199 AuthGrant {
200 principal,
201 identity,
202 expires_unix_micros: now.saturating_add(AUTH_TOKEN_TTL.as_micros() as u64),
203 },
204 );
205 Ok(token)
206 }
207
208 fn consume_auth_grant(&self, token: &[u8]) -> Result<AuthGrant, Status> {
209 let grant = self
210 .auth_grants
211 .lock()
212 .map_err(|_| Status::internal("auth grant store poisoned"))?
213 .remove(token)
214 .ok_or_else(|| Status::unauthenticated("invalid or expired auth token"))?;
215 if grant.expires_unix_micros <= now_unix_micros() {
216 return Err(Status::unauthenticated("invalid or expired auth token"));
217 }
218 if let Some(principal) = &grant.principal {
219 let current = self
220 .db
221 .resolve_current_principal(principal)
222 .ok_or_else(|| Status::unauthenticated("principal was revoked"))?;
223 if current.user_id != principal.user_id
224 || current.created_epoch != principal.created_epoch
225 {
226 return Err(Status::unauthenticated("principal was revoked"));
227 }
228 }
229 Ok(grant)
230 }
231
232 fn session(
233 &self,
234 bytes: &[u8],
235 required_scope: &str,
236 ) -> Result<(String, Arc<SessionEntry>), Status> {
237 let token = id_hex(bytes, "session id")?;
238 let entry = self
239 .sessions
240 .get_by_token(&token)
241 .ok_or_else(|| Status::not_found("session not found"))?;
242 if let AuthenticatedIdentity::ExternalPrincipal { scopes, .. } =
243 &entry.protocol_record().principal
244 {
245 if !scopes
246 .iter()
247 .any(|scope| scope == "*" || scope == required_scope)
248 {
249 return Err(Status::permission_denied(format!(
250 "service principal lacks {required_scope} scope"
251 )));
252 }
253 }
254 Ok((token, entry))
255 }
256
257 fn session_principal(&self, entry: &SessionEntry) -> Result<Option<Principal>, Status> {
258 match entry.protocol_record().principal {
259 AuthenticatedIdentity::Credentialless => Ok(None),
260 AuthenticatedIdentity::CatalogUser {
261 username,
262 user_id,
263 created_version,
264 } => {
265 let principal = self
266 .db
267 .resolve_principal(&username)
268 .ok_or_else(|| Status::unauthenticated("principal was revoked"))?;
269 if principal.user_id != user_id || principal.created_epoch != created_version {
270 return Err(Status::unauthenticated("principal was revoked"));
271 }
272 Ok(Some(principal))
273 }
274 AuthenticatedIdentity::ServicePrincipal { .. } => Err(Status::permission_denied(
275 "internal service principal cannot use native catalog",
276 )),
277 AuthenticatedIdentity::ExternalPrincipal {
278 username,
279 user_id,
280 created_version,
281 ..
282 } => {
283 let principal = self
284 .db
285 .resolve_principal(&username)
286 .ok_or_else(|| Status::unauthenticated("principal was revoked"))?;
287 if principal.user_id != user_id || principal.created_epoch != created_version {
288 return Err(Status::unauthenticated("principal was revoked"));
289 }
290 Ok(Some(principal))
291 }
292 }
293 }
294
295 async fn execute_sql(
296 &self,
297 entry: Arc<SessionEntry>,
298 session_id: String,
299 query_id: QueryId,
300 sql: String,
301 context: Option<&native::RequestContext>,
302 ) -> Result<(Vec<u8>, ManagedQueryBatches), Status> {
303 let options = SqlQueryOptions {
304 query_id: Some(query_id),
305 timeout: request_timeout(context)?,
306 owner: Some(entry.owner.clone()),
307 session_id: Some(session_id),
308 parent_control: None,
309 };
310 let session = entry.session();
311 let query = session.register_query(options).map_err(query_status)?;
312 let permit = tokio::select! {
313 permit = Arc::clone(&self.sql_semaphore).acquire_owned() => {
314 permit.map_err(|_| Status::unavailable("native SQL admission closed"))?
315 }
316 _ = query.control().cancelled() => {
317 let error = query.checkpoint().err().unwrap_or_else(|| {
318 MongrelQueryError::InvalidQueryState(
319 "cancelled native SQL query remained runnable".into(),
320 )
321 });
322 query.fail();
323 return Err(query_status(error));
324 }
325 };
326 let types_query_id = mongreldb_types::ids::QueryId::from_bytes(*query.id().as_bytes());
327 const NATIVE_SQL_PARENT_BUDGET: u64 = 16 * 1024 * 1024;
329 let parent = match self
330 .node_admission
331 .admit_parent(
332 crate::admission::AdmitRequest {
333 tenant: &entry.owner,
334 class: WorkloadClass::InteractiveSql,
335 priority: self.sql_priority,
336 deadline: request_timeout(context)?,
337 query_id: Some(types_query_id),
338 tag: "native-sql",
339 },
340 mongreldb_core::MemoryClass::QueryExecution,
341 NATIVE_SQL_PARENT_BUDGET,
342 query.control().cancelled(),
343 )
344 .await
345 {
346 Ok(parent) => parent,
347 Err(error) => {
348 query.fail();
349 return Err(query_status(crate::admission::admit_error_to_query(error)));
350 }
351 };
352 let _admission = crate::admission::SqlAdmissionGuard::new(permit, parent);
353 let batches = session
354 .run_with_query_for_serialization(&sql, query)
355 .await
356 .map_err(query_status)?;
357 Ok((query_id.as_bytes().to_vec(), batches))
358 }
359
360 async fn begin_idempotency(
361 &self,
362 request: &native::ExecuteRequest,
363 entry: &SessionEntry,
364 session_id: &str,
365 sql: &str,
366 ) -> Result<NativeIdempotency, Status> {
367 let key = request
368 .context
369 .as_ref()
370 .map(|context| context.idempotency_key.as_str())
371 .unwrap_or_default();
372 if key.is_empty() {
373 return Ok(NativeIdempotency::None);
374 }
375 if !matches!(
376 request.command,
377 Some(native::execute_request::Command::Sql(_))
378 ) {
379 return Err(Status::invalid_argument(
380 "idempotency key requires one direct SQL write",
381 ));
382 }
383 if mongreldb_query::classify_sql_idempotency(sql)
384 != mongreldb_query::SqlIdempotencyClass::SingleWrite
385 {
386 return Err(Status::invalid_argument(
387 "idempotency key requires one direct SQL write",
388 ));
389 }
390 crate::sql_idempotency::SqlIdempotencyStore::validate_key(key)
391 .map_err(Status::invalid_argument)?;
392 if entry.session().staged_sql_operation_count().is_some() {
393 return Err(Status::failed_precondition(
394 "idempotency is unavailable inside an open transaction",
395 ));
396 }
397 let store = self.sql_idempotency.as_ref().ok_or_else(|| {
398 Status::unavailable("durable native SQL idempotency is not configured")
399 })?;
400 let parameters = request
401 .parameters
402 .iter()
403 .flatten()
404 .copied()
405 .collect::<Vec<_>>();
406 let binding = crate::sql_idempotency::SqlIdempotencyBinding {
407 sql_fingerprint: mongreldb_query::normalized_sql_fingerprint(sql),
408 parameter_hash: crate::sql_idempotency::hash(¶meters),
409 request_semantics_hash: crate::sql_idempotency::hash(b"native-arrow-ipc-v1"),
410 session_semantics_hash: crate::sql_idempotency::hash(session_id.as_bytes()),
411 expires_after_ms: store.expires_after_ms(),
412 };
413 match store.begin(&entry.owner, key, binding).await {
414 crate::sql_idempotency::BeginResult::Execute(execution) => {
415 Ok(NativeIdempotency::Execute(execution))
416 }
417 crate::sql_idempotency::BeginResult::Replay { receipt, .. } => {
418 Ok(NativeIdempotency::Replay(receipt))
419 }
420 crate::sql_idempotency::BeginResult::Mismatch => Err(structured_status(
421 Code::AlreadyExists,
422 ErrorCategory::TransactionConflict,
423 "idempotency key was reused with different request semantics",
424 )),
425 crate::sql_idempotency::BeginResult::Indeterminate { .. } => Err(structured_status(
426 Code::Unknown,
427 ErrorCategory::CommitOutcomeUnknown,
428 "idempotency outcome is indeterminate",
429 )),
430 crate::sql_idempotency::BeginResult::Full => Err(structured_status(
431 Code::ResourceExhausted,
432 ErrorCategory::ResourceExhausted,
433 "idempotency store is full",
434 )),
435 crate::sql_idempotency::BeginResult::Unavailable(_reason) => Err(structured_status(
436 Code::Unavailable,
437 ErrorCategory::ReplicaUnavailable,
438 "idempotency store is unavailable",
439 )),
440 }
441 }
442
443 fn finish_idempotency(
444 &self,
445 execution: crate::sql_idempotency::SqlIdempotencyExecution,
446 query_id: QueryId,
447 ) {
448 let Some(status) = self.query_registry.status(query_id) else {
449 return;
450 };
451 if let Some(mut receipt) = super::sql_terminal_idempotency_receipt(&status) {
452 receipt.commit_receipt = crate::sql_idempotency::record_core_idempotency_commit(
453 &self.db,
454 execution.owner(),
455 execution.key(),
456 execution.binding(),
457 execution.ttl(),
458 );
459 execution.commit(receipt);
460 } else if super::can_abort_idempotency_intent(&status) {
461 execution.abort();
462 }
463 }
464
465 async fn resolve_sql(
466 &self,
467 entry: &SessionEntry,
468 request: &native::ExecuteRequest,
469 ) -> Result<String, Status> {
470 use native::execute_request::Command;
471 match &request.command {
472 Some(Command::Sql(sql)) => Ok(sql.clone()),
473 Some(Command::PreparedStatementId(id)) => {
474 let (_, binding) = entry
475 .prepared_binding_by_id(*id)
476 .ok_or_else(|| Status::not_found("prepared statement not found"))?;
477 if !prepared::CatalogState::capture(&self.db).is_compatible(&binding) {
478 return Err(structured_status(
479 Code::FailedPrecondition,
480 ErrorCategory::SchemaVersionMismatch,
481 "prepared statement schema changed",
482 ));
483 }
484 let values = request
485 .parameters
486 .iter()
487 .map(|encoded| {
488 bincode::deserialize::<mongreldb_protocol::request::ParameterValue>(encoded)
489 .map_err(|_| Status::invalid_argument("invalid prepared parameter"))
490 })
491 .collect::<Result<Vec<_>, _>>()?;
492 let literals = values
493 .iter()
494 .map(parameter_literal)
495 .collect::<Result<Vec<_>, _>>()?;
496 bind_numbered_parameters(&binding.sql, &literals)
497 }
498 Some(Command::AdminCommand(command)) => String::from_utf8(command.clone())
499 .map_err(|_| Status::invalid_argument("admin command must be UTF-8 SQL")),
500 None => Err(Status::invalid_argument("execute command is required")),
501 }
502 }
503}
504
505#[tonic::async_trait]
506impl native::auth_service_server::AuthService for NativeRuntime {
507 async fn authenticate(
508 &self,
509 request: Request<native::AuthenticateRequest>,
510 ) -> Result<Response<native::AuthenticateResponse>, Status> {
511 let request = request.into_inner();
512 validate_native_context(request.context.as_ref())?;
513 let (principal, identity, response_identity) = match request.credential {
514 Some(native::authenticate_request::Credential::Password(password)) => {
515 let username = password.username;
516 let password = zeroize::Zeroizing::new(password.password);
517 let principal = self
518 .db
519 .authenticate_principal(&username, password.as_str())
520 .map_err(core_status)?
521 .ok_or_else(|| Status::unauthenticated("invalid credentials"))?;
522 let identity = AuthenticatedIdentity::CatalogUser {
523 username: principal.username.clone(),
524 user_id: principal.user_id,
525 created_version: principal.created_epoch,
526 };
527 let response_identity = native_identity(Some(&principal));
528 (Some(principal), identity, response_identity)
529 }
530 Some(native::authenticate_request::Credential::MysqlCachingSha2(credential)) => {
531 if credential.nonce.len() != 20
532 || (!credential.proof.is_empty() && credential.proof.len() != 32)
533 {
534 return Err(Status::invalid_argument(
535 "invalid caching_sha2_password proof",
536 ));
537 }
538 let principal = self
539 .db
540 .authenticate_mysql_caching_sha2_principal(
541 &credential.username,
542 &credential.nonce,
543 &credential.proof,
544 )
545 .ok_or_else(|| Status::unauthenticated("invalid credentials"))?;
546 let identity = AuthenticatedIdentity::CatalogUser {
547 username: principal.username.clone(),
548 user_id: principal.user_id,
549 created_version: principal.created_epoch,
550 };
551 let response_identity = native_identity(Some(&principal));
552 (Some(principal), identity, response_identity)
553 }
554 None if !self.db.require_auth_enabled() => (
555 None,
556 AuthenticatedIdentity::Credentialless,
557 native_identity(None),
558 ),
559 Some(native::authenticate_request::Credential::ServiceToken(credential)) => {
560 let auth = self.external_auth.as_ref().ok_or_else(|| {
561 Status::unauthenticated("service-token authentication is not configured")
562 })?;
563 let secret = zeroize::Zeroizing::new(credential.secret);
564 let token = auth
565 .service_tokens
566 .read()
567 .map_err(|_| Status::internal("service-token registry poisoned"))?
568 .authenticate(&credential.token_id, secret.as_str(), now_unix_seconds())
569 .cloned()
570 .ok_or_else(|| Status::unauthenticated("invalid credentials"))?;
571 let principal = self
572 .db
573 .resolve_principal(&token.principal)
574 .ok_or_else(|| Status::unauthenticated("invalid credentials"))?;
575 let identity = AuthenticatedIdentity::ExternalPrincipal {
576 provider: "service_token".into(),
577 subject: token.token_id.clone(),
578 username: principal.username.clone(),
579 user_id: principal.user_id,
580 created_version: principal.created_epoch,
581 scopes: token.scopes.clone(),
582 };
583 let response_identity =
584 native_external_identity(&principal, token.token_id, token.scopes);
585 (Some(principal), identity, response_identity)
586 }
587 Some(native::authenticate_request::Credential::Oidc(credential)) => {
588 let oidc = self
589 .external_auth
590 .as_ref()
591 .and_then(|auth| auth.oidc.clone())
592 .ok_or_else(|| {
593 Status::unauthenticated("OIDC authentication is not configured")
594 })?;
595 let token = credential.compact_jws;
596 let verified = tokio::task::spawn_blocking(move || {
597 oidc.cache
598 .verify(&token, &oidc.validation, now_unix_seconds())
599 })
600 .await
601 .map_err(|_| Status::internal("OIDC verifier task failed"))?
602 .map_err(|_| Status::unauthenticated("invalid credentials"))?;
603 let principal = self
604 .db
605 .resolve_principal(&verified.principal)
606 .ok_or_else(|| Status::unauthenticated("invalid credentials"))?;
607 let identity = AuthenticatedIdentity::ExternalPrincipal {
608 provider: "oidc".into(),
609 subject: verified.principal.clone(),
610 username: principal.username.clone(),
611 user_id: principal.user_id,
612 created_version: principal.created_epoch,
613 scopes: verified.scopes.clone(),
614 };
615 let response_identity =
616 native_external_identity(&principal, verified.principal, verified.scopes);
617 (Some(principal), identity, response_identity)
618 }
619 None => return Err(Status::unauthenticated("credentials are required")),
620 };
621 let auth_token = self.issue_auth_grant(principal.clone(), identity)?;
622 Ok(Response::new(native::AuthenticateResponse {
623 identity: Some(response_identity),
624 auth_token,
625 }))
626 }
627
628 async fn begin_scram(
629 &self,
630 request: Request<native::BeginScramRequest>,
631 ) -> Result<Response<native::BeginScramResponse>, Status> {
632 let request = request.into_inner();
633 validate_native_context(request.context.as_ref())?;
634 if !request
635 .client_first_bare
636 .starts_with(&format!("n={},", request.username))
637 {
638 return Err(Status::invalid_argument(
639 "SCRAM username does not match client-first message",
640 ));
641 }
642 let (verifier, fake_user) = match self.db.user_scram_verifier(&request.username) {
643 Some(verifier) => (verifier, false),
644 None => {
645 let salt = mongreldb_types::ids::QueryId::new_random();
646 (
647 mongreldb_core::ScramVerifier::from_password(
648 "invalid-user-password",
649 salt.as_bytes(),
650 mongreldb_core::security_hardening::SCRAM_SHA_256_MIN_ITERATIONS,
651 )
652 .map_err(|error| Status::internal(error.to_string()))?,
653 true,
654 )
655 }
656 };
657 let server_nonce = mongreldb_types::ids::QueryId::new_random().to_hex();
658 let session = mongreldb_core::ScramServerSession::begin(
659 verifier,
660 request.client_first_bare,
661 &request.client_nonce,
662 &server_nonce,
663 mongreldb_core::ScramChannelBindingPolicy::Disabled,
664 Vec::new(),
665 )
666 .map_err(|error| Status::invalid_argument(error.to_string()))?;
667 let server_first = session.server_first_message().to_owned();
668 let exchange_id = mongreldb_types::ids::QueryId::new_random()
669 .as_bytes()
670 .to_vec();
671 let now = now_unix_micros();
672 let mut exchanges = self
673 .scram_exchanges
674 .lock()
675 .map_err(|_| Status::internal("SCRAM exchange store poisoned"))?;
676 exchanges.retain(|_, exchange| exchange.expires_unix_micros > now);
677 if exchanges.len() >= MAX_AUTH_TOKENS {
678 return Err(Status::resource_exhausted("too many SCRAM exchanges"));
679 }
680 exchanges.insert(
681 exchange_id.clone(),
682 PendingScram {
683 username: request.username,
684 session,
685 fake_user,
686 expires_unix_micros: now.saturating_add(AUTH_TOKEN_TTL.as_micros() as u64),
687 },
688 );
689 Ok(Response::new(native::BeginScramResponse {
690 exchange_id,
691 server_first,
692 }))
693 }
694
695 async fn finish_scram(
696 &self,
697 request: Request<native::FinishScramRequest>,
698 ) -> Result<Response<native::FinishScramResponse>, Status> {
699 let request = request.into_inner();
700 validate_native_context(request.context.as_ref())?;
701 let exchange = self
702 .scram_exchanges
703 .lock()
704 .map_err(|_| Status::internal("SCRAM exchange store poisoned"))?
705 .remove(&request.exchange_id)
706 .ok_or_else(|| Status::unauthenticated("invalid SCRAM exchange"))?;
707 if exchange.expires_unix_micros <= now_unix_micros() {
708 return Err(Status::unauthenticated("invalid SCRAM exchange"));
709 }
710 let server_final = exchange
711 .session
712 .finish(&request.client_final_without_proof, &request.client_proof)
713 .map_err(|_| Status::unauthenticated("invalid credentials"))?;
714 if exchange.fake_user {
715 return Err(Status::unauthenticated("invalid credentials"));
716 }
717 let principal = self
718 .db
719 .resolve_principal(&exchange.username)
720 .ok_or_else(|| Status::unauthenticated("invalid credentials"))?;
721 let identity = AuthenticatedIdentity::CatalogUser {
722 username: principal.username.clone(),
723 user_id: principal.user_id,
724 created_version: principal.created_epoch,
725 };
726 let auth_token = self.issue_auth_grant(Some(principal.clone()), identity)?;
727 Ok(Response::new(native::FinishScramResponse {
728 server_final,
729 authentication: Some(native::AuthenticateResponse {
730 identity: Some(native_identity(Some(&principal))),
731 auth_token,
732 }),
733 }))
734 }
735}
736
737#[tonic::async_trait]
738impl native::session_service_server::SessionService for NativeRuntime {
739 async fn open_session(
740 &self,
741 request: Request<native::OpenSessionRequest>,
742 ) -> Result<Response<native::OpenSessionResponse>, Status> {
743 let request = request.into_inner();
744 validate_native_context(request.context.as_ref())?;
745 if request.database_id != self.sessions.database_id().as_bytes() {
746 return Err(Status::not_found("database not found"));
747 }
748 let grant = self.consume_auth_grant(&request.auth_token)?;
749 let owner = grant.principal.as_ref().map_or_else(
750 || "anonymous".into(),
751 |principal| principal.username.clone(),
752 );
753 let session = MongrelSession::open_with_external_modules_as(
754 Arc::clone(&self.db),
755 std::iter::empty(),
756 grant.principal,
757 )
758 .map_err(query_status)?
759 .with_query_registry(Arc::clone(&self.query_registry));
760 let token = self
761 .sessions
762 .create_with_identity(session, owner, grant.identity)
763 .ok_or_else(|| Status::resource_exhausted("session limit reached"))?;
764 Ok(Response::new(native::OpenSessionResponse {
765 session_id: hex_id(&token)?,
766 }))
767 }
768
769 async fn close_session(
770 &self,
771 request: Request<native::CloseSessionRequest>,
772 ) -> Result<Response<native::Empty>, Status> {
773 let request = request.into_inner();
774 validate_native_context(request.context.as_ref())?;
775 let token = id_hex(&request.session_id, "session id")?;
776 if !self.sessions.close_by_token(&token) {
777 return Err(Status::not_found("session not found"));
778 }
779 self.query_registry
780 .cancel_session(&token, mongreldb_core::CancellationReason::SessionClosed);
781 Ok(Response::new(native::Empty {}))
782 }
783}
784
785#[tonic::async_trait]
786impl native::query_service_server::QueryService for NativeRuntime {
787 type ExecuteStreamStream =
788 Pin<Box<dyn Stream<Item = Result<native::ArrowFrame, Status>> + Send + 'static>>;
789
790 async fn prepare(
791 &self,
792 request: Request<native::PrepareRequest>,
793 ) -> Result<Response<native::PrepareResponse>, Status> {
794 let request = request.into_inner();
795 validate_native_context(request.context.as_ref())?;
796 let (_, entry) = self.session(&request.session_id, "query")?;
797 let _guard = entry.lock.lock().await;
798 let statement_id = entry.allocate_statement_id();
799 let name = format!("native_{}", statement_id.get());
800 entry
801 .session()
802 .run(&format!("PREPARE {name} AS {}", request.sql))
803 .await
804 .map_err(query_status)?;
805 let catalog = prepared::CatalogState::capture(&self.db);
806 entry.insert_prepared_binding(
807 name,
808 prepared::build_binding(statement_id, request.sql, Vec::new(), &catalog),
809 );
810 Ok(Response::new(native::PrepareResponse {
811 statement_id: statement_id.get(),
812 schema_version: catalog.catalog_version.get(),
813 }))
814 }
815
816 async fn execute(
817 &self,
818 request: Request<native::ExecuteRequest>,
819 ) -> Result<Response<native::ExecuteResponse>, Status> {
820 let request = request.into_inner();
821 validate_native_context(request.context.as_ref())?;
822 let (session_id, entry) = self.session(&request.session_id, "query")?;
823 let id = query_id(&request.query_id)?;
824 let sql = self.resolve_sql(&entry, &request).await?;
825 let idempotency = self
826 .begin_idempotency(&request, &entry, &session_id, &sql)
827 .await?;
828 if let NativeIdempotency::Replay(receipt) = idempotency {
829 return Ok(Response::new(replayed_response(&request.query_id, receipt)));
830 }
831 let execution = match idempotency {
832 NativeIdempotency::Execute(execution) => Some(execution),
833 NativeIdempotency::None => None,
834 NativeIdempotency::Replay(_) => unreachable!(),
835 };
836 let (query_id, batches) = match self
837 .execute_sql(entry, session_id, id, sql, request.context.as_ref())
838 .await
839 {
840 Ok(result) => result,
841 Err(error) => {
842 if let Some(execution) = execution {
843 self.finish_idempotency(execution, id);
844 }
845 return Err(error);
846 }
847 };
848 let frames = batches
849 .batches()
850 .iter()
851 .enumerate()
852 .map(|(sequence, batch)| encode_batch(batch, sequence as u64, false))
853 .chain(std::iter::once(Ok(native::ArrowFrame {
854 ipc: Vec::new(),
855 sequence: batches.batches().len() as u64,
856 end_of_stream: true,
857 })))
858 .collect::<Result<Vec<_>, Status>>();
859 let frames = match frames {
860 Ok(frames) => {
861 batches.complete().map_err(query_status)?;
862 frames
863 }
864 Err(error) => {
865 batches.fail_serialization();
866 if let Some(execution) = execution {
867 self.finish_idempotency(execution, id);
868 }
869 return Err(error);
870 }
871 };
872 let status = self.query_registry.status(id);
873 if let Some(execution) = execution {
874 self.finish_idempotency(execution, id);
875 }
876 let durable = status
877 .as_ref()
878 .map(|status| proto_durable(&status.durable_outcome, status.serialization_outcome))
879 .unwrap_or_default();
880 Ok(Response::new(native::ExecuteResponse {
881 query_id,
882 rows_affected: 0,
883 frames,
884 idempotency_replayed: false,
885 committed: durable.committed,
886 commit_epoch: durable.last_commit_epoch.unwrap_or(0),
887 original_query_id: id.as_bytes().to_vec(),
888 durable: Some(durable),
889 }))
890 }
891
892 async fn execute_stream(
893 &self,
894 request: Request<native::ExecuteRequest>,
895 ) -> Result<Response<Self::ExecuteStreamStream>, Status> {
896 let request = request.into_inner();
897 validate_native_context(request.context.as_ref())?;
898 if request
899 .context
900 .as_ref()
901 .is_some_and(|context| !context.idempotency_key.is_empty())
902 {
903 return Err(Status::invalid_argument(
904 "idempotent writes require buffered Execute",
905 ));
906 }
907 let (session_id, entry) = self.session(&request.session_id, "query")?;
908 let id = query_id(&request.query_id)?;
909 let sql = self.resolve_sql(&entry, &request).await?;
910 let stream = entry
911 .session()
912 .run_stream_with_options(
913 &sql,
914 SqlQueryOptions {
915 query_id: Some(id),
916 timeout: request_timeout(request.context.as_ref())?,
917 owner: Some(entry.owner.clone()),
918 session_id: Some(session_id),
919 parent_control: None,
920 },
921 )
922 .await
923 .map_err(query_status)?;
924 let output = futures::stream::try_unfold(
925 (stream, 0_u64, false),
926 |(mut stream, sequence, done)| async move {
927 if done {
928 return Ok(None);
929 }
930 match stream.next().await {
931 Some(Ok(batch)) => {
932 let frame = encode_batch(&batch, sequence, false)?;
933 Ok(Some((frame, (stream, sequence + 1, false))))
934 }
935 Some(Err(error)) => Err(stream_status(error)),
936 None => Ok(Some((
937 native::ArrowFrame {
938 ipc: Vec::new(),
939 sequence,
940 end_of_stream: true,
941 },
942 (stream, sequence + 1, true),
943 ))),
944 }
945 },
946 );
947 Ok(Response::new(Box::pin(output)))
948 }
949
950 async fn cancel_query(
951 &self,
952 request: Request<native::CancelQueryRequest>,
953 ) -> Result<Response<native::CancelQueryResponse>, Status> {
954 let request = request.into_inner();
955 validate_native_context(request.context.as_ref())?;
956 let (session_id, _) = self.session(&request.session_id, "query")?;
957 let query_id = query_id(&request.query_id)?;
958 let status = self
959 .query_registry
960 .status(query_id)
961 .ok_or_else(|| Status::not_found("query not found"))?;
962 if status.session_id.as_deref() != Some(&session_id) {
963 return Err(Status::not_found("query not found"));
964 }
965 let durable = proto_durable(&status.durable_outcome, status.serialization_outcome);
966 let outcome = match self.query_registry.cancel(query_id) {
967 CancelOutcome::Accepted => native::CancelOutcome::Accepted,
968 CancelOutcome::AlreadyCancelling => native::CancelOutcome::AlreadyCancelling,
969 CancelOutcome::TooLate => native::CancelOutcome::TooLate,
970 CancelOutcome::AlreadyFinished => native::CancelOutcome::AlreadyFinished,
971 CancelOutcome::NotFound => native::CancelOutcome::NotFound,
972 };
973 let durable = self
975 .query_registry
976 .status(query_id)
977 .map(|status| proto_durable(&status.durable_outcome, status.serialization_outcome))
978 .unwrap_or(durable);
979 Ok(Response::new(native::CancelQueryResponse {
980 outcome: outcome as i32,
981 durable: Some(durable),
982 }))
983 }
984
985 async fn get_query_status(
986 &self,
987 request: Request<native::GetQueryStatusRequest>,
988 ) -> Result<Response<native::QueryStatusResponse>, Status> {
989 let request = request.into_inner();
990 validate_native_context(request.context.as_ref())?;
991 let (session_id, _) = self.session(&request.session_id, "query")?;
992 let id = query_id(&request.query_id)?;
993 let status = self
994 .query_registry
995 .status(id)
996 .ok_or_else(|| Status::not_found("query not found"))?;
997 if status.session_id.as_deref() != Some(&session_id) {
998 return Err(Status::not_found("query not found"));
999 }
1000 Ok(Response::new(native::QueryStatusResponse {
1001 query_id: id.as_bytes().to_vec(),
1002 phase: native_phase(status.phase) as i32,
1003 error: status.terminal_error.map(|error| native::ErrorDetail {
1004 category_code: 0,
1005 category: format!("{:?}", error.category),
1006 message: error.code,
1007 retryable: false,
1008 metadata: HashMap::new(),
1009 }),
1010 durable: Some(proto_durable(
1011 &status.durable_outcome,
1012 status.serialization_outcome,
1013 )),
1014 }))
1015 }
1016}
1017
1018fn proto_durable(
1019 outcome: &mongreldb_query::DurableOutcome,
1020 serialization: SerializationOutcome,
1021) -> native::DurableOutcome {
1022 let last_commit_hlc = outcome
1023 .commit_ts
1024 .map(|ts| {
1025 let mut bytes = Vec::with_capacity(16);
1026 bytes.extend_from_slice(&ts.physical_micros.to_be_bytes());
1027 bytes.extend_from_slice(&ts.logical.to_be_bytes());
1028 bytes.extend_from_slice(&ts.node_tiebreaker.to_be_bytes());
1029 bytes
1030 })
1031 .unwrap_or_default();
1032 native::DurableOutcome {
1033 committed: outcome.committed,
1034 committed_statements: outcome.committed_statements as u64,
1035 last_commit_hlc,
1036 first_commit_statement_index: outcome
1037 .first_commit_statement_index
1038 .map(|index| index as u64),
1039 last_commit_statement_index: outcome
1040 .last_commit_statement_index
1041 .map(|index| index as u64),
1042 completed_statements: outcome.committed_statements as u64,
1043 current_statement_index: outcome
1044 .last_commit_statement_index
1045 .map(|index| index as u64)
1046 .unwrap_or(0),
1047 terminal_state: if outcome.committed {
1048 "committed".into()
1049 } else {
1050 String::new()
1051 },
1052 serialization_state: serialization_state_name(serialization).into(),
1053 last_commit_epoch: outcome.last_commit_epoch,
1054 }
1055}
1056
1057fn serialization_state_name(outcome: SerializationOutcome) -> &'static str {
1058 match outcome {
1059 SerializationOutcome::NotStarted => "not_started",
1060 SerializationOutcome::InProgress => "in_progress",
1061 SerializationOutcome::Succeeded => "succeeded",
1062 SerializationOutcome::Failed => "failed",
1063 }
1064}
1065
1066fn receipt_durable(receipt: &crate::sql_idempotency::SqlDurableReceipt) -> native::DurableOutcome {
1067 native::DurableOutcome {
1068 committed: receipt.outcome.committed,
1069 committed_statements: receipt.outcome.committed_statements as u64,
1070 last_commit_hlc: Vec::new(),
1071 first_commit_statement_index: receipt
1072 .outcome
1073 .first_commit_statement_index
1074 .map(|index| index as u64),
1075 last_commit_statement_index: receipt
1076 .outcome
1077 .last_commit_statement_index
1078 .map(|index| index as u64),
1079 completed_statements: receipt.outcome.committed_statements as u64,
1080 current_statement_index: 0,
1081 terminal_state: if receipt.outcome.committed {
1082 "committed".into()
1083 } else {
1084 String::new()
1085 },
1086 serialization_state: "succeeded".into(),
1088 last_commit_epoch: receipt.outcome.last_commit_epoch,
1089 }
1090}
1091
1092fn replayed_response(
1093 current_query_id: &[u8],
1094 receipt: crate::sql_idempotency::SqlDurableReceipt,
1095) -> native::ExecuteResponse {
1096 let durable = receipt_durable(&receipt);
1097 native::ExecuteResponse {
1098 query_id: current_query_id.to_vec(),
1099 rows_affected: 0,
1100 frames: Vec::new(),
1101 idempotency_replayed: true,
1102 committed: durable.committed,
1103 commit_epoch: durable.last_commit_epoch.unwrap_or(0),
1104 original_query_id: hex_id(&receipt.original_query_id).unwrap_or_default(),
1105 durable: Some(durable),
1106 }
1107}
1108
1109#[tonic::async_trait]
1110impl native::transaction_service_server::TransactionService for NativeRuntime {
1111 async fn begin(
1112 &self,
1113 request: Request<native::BeginTransactionRequest>,
1114 ) -> Result<Response<native::BeginTransactionResponse>, Status> {
1115 let request = request.into_inner();
1116 validate_native_context(request.context.as_ref())?;
1117 let (_, entry) = self.session(&request.session_id, "transaction")?;
1118 let sql = match native::IsolationLevel::try_from(request.isolation) {
1119 Ok(native::IsolationLevel::ReadCommitted) => {
1120 "BEGIN; SET TRANSACTION ISOLATION LEVEL READ COMMITTED"
1121 }
1122 Ok(native::IsolationLevel::Serializable) => {
1123 "BEGIN; SET TRANSACTION ISOLATION LEVEL SERIALIZABLE"
1124 }
1125 _ => "BEGIN",
1126 };
1127 entry.session().run(sql).await.map_err(query_status)?;
1128 Ok(Response::new(native::BeginTransactionResponse {
1129 transaction_id: TransactionId::new_random().as_bytes().to_vec(),
1130 }))
1131 }
1132
1133 async fn commit(
1134 &self,
1135 request: Request<native::TransactionRequest>,
1136 ) -> Result<Response<native::Empty>, Status> {
1137 transaction_sql(self, request.into_inner(), "COMMIT").await
1138 }
1139
1140 async fn rollback(
1141 &self,
1142 request: Request<native::TransactionRequest>,
1143 ) -> Result<Response<native::Empty>, Status> {
1144 transaction_sql(self, request.into_inner(), "ROLLBACK").await
1145 }
1146}
1147
1148#[tonic::async_trait]
1149impl native::catalog_service_server::CatalogService for NativeRuntime {
1150 async fn get_schema(
1151 &self,
1152 request: Request<native::GetSchemaRequest>,
1153 ) -> Result<Response<native::GetSchemaResponse>, Status> {
1154 let request = request.into_inner();
1155 validate_native_context(request.context.as_ref())?;
1156 if request.database_id != self.sessions.database_id().as_bytes() {
1157 return Err(Status::not_found("database not found"));
1158 }
1159 let (_, entry) = self.session(&request.session_id, "catalog:read")?;
1160 let principal = self.session_principal(&entry)?;
1161 self.db
1162 .require_for(
1163 principal.as_ref(),
1164 &mongreldb_core::Permission::Select {
1165 table: request.table.clone(),
1166 },
1167 )
1168 .map_err(core_status)?;
1169 let table = self.db.table(&request.table).map_err(core_status)?;
1170 let schema = table.lock().schema().clone();
1171 let schema_json = serde_json::to_vec(&schema)
1172 .map_err(|error| Status::internal(format!("schema encode failed: {error}")))?;
1173 Ok(Response::new(native::GetSchemaResponse {
1174 table: request.table,
1175 schema_version: schema.schema_id,
1176 columns: schema
1177 .columns
1178 .into_iter()
1179 .map(|column| native::ColumnSchema {
1180 name: column.name,
1181 data_type: format!("{:?}", column.ty),
1182 nullable: column.flags.contains(mongreldb_core::ColumnFlags::NULLABLE),
1183 })
1184 .collect(),
1185 schema_json,
1186 }))
1187 }
1188
1189 async fn create_table(
1190 &self,
1191 request: Request<native::CreateTableRequest>,
1192 ) -> Result<Response<native::CreateTableResponse>, Status> {
1193 let request = request.into_inner();
1194 validate_native_context(request.context.as_ref())?;
1195 let (_, entry) = self.session(&request.session_id, "catalog:write")?;
1196 let _guard = entry.lock.lock().await;
1197 let principal = self.session_principal(&entry)?;
1198 self.db
1199 .require_for(principal.as_ref(), &mongreldb_core::Permission::Ddl)
1200 .map_err(core_status)?;
1201 let mut schema = if request.schema_json.is_empty() {
1202 let columns = request
1203 .columns
1204 .into_iter()
1205 .map(native_create_column)
1206 .collect::<Result<Vec<_>, _>>()?;
1207 let uniques = request
1208 .uniques
1209 .into_iter()
1210 .map(|constraint| {
1211 Ok(mongreldb_core::constraint::UniqueConstraint {
1212 id: native_u16(constraint.id, "unique constraint id")?,
1213 name: constraint.name,
1214 columns: native_u16s(constraint.columns, "unique constraint column")?,
1215 })
1216 })
1217 .collect::<Result<Vec<_>, Status>>()?;
1218 let foreign_keys = request
1219 .foreign_keys
1220 .into_iter()
1221 .map(native_foreign_key)
1222 .collect::<Result<Vec<_>, _>>()?;
1223 mongreldb_core::Schema {
1224 schema_id: request.schema_id,
1225 columns,
1226 indexes: Vec::new(),
1227 colocation: Vec::new(),
1228 constraints: mongreldb_core::constraint::TableConstraints {
1229 uniques,
1230 foreign_keys,
1231 checks: Vec::new(),
1232 },
1233 clustered: false,
1234 }
1235 } else {
1236 let schema: mongreldb_core::Schema = serde_json::from_slice(&request.schema_json)
1237 .map_err(|error| {
1238 Status::invalid_argument(format!("invalid schema_json: {error}"))
1239 })?;
1240 if schema.schema_id != request.schema_id {
1241 return Err(Status::invalid_argument(
1242 "schema_json schema_id does not match request schema_id",
1243 ));
1244 }
1245 schema
1246 };
1247 if let Ok(existing) = self.db.table(&request.table) {
1248 let existing = existing.lock().schema().clone();
1249 schema.schema_id = existing.schema_id;
1250 if serde_json::to_vec(&schema).map_err(|_| Status::internal("schema encode failed"))?
1251 != serde_json::to_vec(&existing)
1252 .map_err(|_| Status::internal("schema encode failed"))?
1253 {
1254 return Err(Status::already_exists(
1255 "table exists with a different schema",
1256 ));
1257 }
1258 entry
1259 .session()
1260 .refresh_database_table(&request.table)
1261 .map_err(query_status)?;
1262 return Ok(Response::new(native::CreateTableResponse {
1263 table_id: self.db.table_id(&request.table).map_err(core_status)?,
1264 schema_version: self.db.catalog_version(),
1265 }));
1266 }
1267 let table_id = self
1268 .db
1269 .create_table(&request.table, schema)
1270 .map_err(core_status)?;
1271 entry
1272 .session()
1273 .refresh_database_table(&request.table)
1274 .map_err(query_status)?;
1275 Ok(Response::new(native::CreateTableResponse {
1276 table_id,
1277 schema_version: self.db.catalog_version(),
1278 }))
1279 }
1280}
1281
1282fn native_create_column(column: native::CreateColumn) -> Result<mongreldb_core::ColumnDef, Status> {
1283 let ty = match native::ColumnType::try_from(column.data_type)
1284 .map_err(|_| Status::invalid_argument("unknown native column type"))?
1285 {
1286 native::ColumnType::Bool => mongreldb_core::TypeId::Bool,
1287 native::ColumnType::Int8 => mongreldb_core::TypeId::Int8,
1288 native::ColumnType::Int16 => mongreldb_core::TypeId::Int16,
1289 native::ColumnType::Int32 => mongreldb_core::TypeId::Int32,
1290 native::ColumnType::Int64 => mongreldb_core::TypeId::Int64,
1291 native::ColumnType::Uint8 => mongreldb_core::TypeId::UInt8,
1292 native::ColumnType::Uint16 => mongreldb_core::TypeId::UInt16,
1293 native::ColumnType::Uint32 => mongreldb_core::TypeId::UInt32,
1294 native::ColumnType::Uint64 => mongreldb_core::TypeId::UInt64,
1295 native::ColumnType::Float32 => mongreldb_core::TypeId::Float32,
1296 native::ColumnType::Float64 => mongreldb_core::TypeId::Float64,
1297 native::ColumnType::TimestampNanos => mongreldb_core::TypeId::TimestampNanos,
1298 native::ColumnType::Date32 => mongreldb_core::TypeId::Date32,
1299 native::ColumnType::Date64 => mongreldb_core::TypeId::Date64,
1300 native::ColumnType::Time64 => mongreldb_core::TypeId::Time64,
1301 native::ColumnType::Bytes => mongreldb_core::TypeId::Bytes,
1302 native::ColumnType::Json => mongreldb_core::TypeId::Json,
1303 native::ColumnType::Decimal128 => mongreldb_core::TypeId::Decimal128 {
1304 precision: u8::try_from(column.decimal_precision)
1305 .map_err(|_| Status::invalid_argument("decimal precision exceeds u8"))?,
1306 scale: i8::try_from(column.decimal_scale)
1307 .map_err(|_| Status::invalid_argument("decimal scale exceeds i8"))?,
1308 },
1309 native::ColumnType::Unspecified => {
1310 return Err(Status::invalid_argument("native column type is required"))
1311 }
1312 };
1313 let mut flags = mongreldb_core::ColumnFlags::empty();
1314 if column.nullable {
1315 flags = flags.with(mongreldb_core::ColumnFlags::NULLABLE);
1316 }
1317 if column.primary_key {
1318 flags = flags.with(mongreldb_core::ColumnFlags::PRIMARY_KEY);
1319 }
1320 if column.auto_increment {
1321 flags = flags.with(mongreldb_core::ColumnFlags::AUTO_INCREMENT);
1322 }
1323 Ok(mongreldb_core::ColumnDef {
1324 id: native_u16(column.id, "column id")?,
1325 name: column.name,
1326 ty,
1327 flags,
1328 default_value: None,
1329 embedding_source: None,
1330 })
1331}
1332
1333fn native_foreign_key(
1334 foreign_key: native::ForeignKey,
1335) -> Result<mongreldb_core::constraint::ForeignKey, Status> {
1336 let action = |value| -> Result<mongreldb_core::constraint::FkAction, Status> {
1337 match native::ForeignKeyAction::try_from(value)
1338 .map_err(|_| Status::invalid_argument("unknown foreign-key action"))?
1339 {
1340 native::ForeignKeyAction::Unspecified | native::ForeignKeyAction::Restrict => {
1341 Ok(mongreldb_core::constraint::FkAction::Restrict)
1342 }
1343 native::ForeignKeyAction::Cascade => Ok(mongreldb_core::constraint::FkAction::Cascade),
1344 native::ForeignKeyAction::SetNull => Ok(mongreldb_core::constraint::FkAction::SetNull),
1345 }
1346 };
1347 Ok(mongreldb_core::constraint::ForeignKey {
1348 id: native_u16(foreign_key.id, "foreign-key id")?,
1349 name: foreign_key.name,
1350 columns: native_u16s(foreign_key.columns, "foreign-key column")?,
1351 ref_table: foreign_key.referenced_table,
1352 ref_columns: native_u16s(foreign_key.referenced_columns, "referenced column")?,
1353 on_delete: action(foreign_key.on_delete)?,
1354 on_update: action(foreign_key.on_update)?,
1355 })
1356}
1357
1358fn native_u16(value: u32, field: &str) -> Result<u16, Status> {
1359 u16::try_from(value).map_err(|_| Status::invalid_argument(format!("{field} exceeds u16")))
1360}
1361
1362fn native_u16s(values: Vec<u32>, field: &str) -> Result<Vec<u16>, Status> {
1363 values
1364 .into_iter()
1365 .map(|value| native_u16(value, field))
1366 .collect()
1367}
1368
1369#[tonic::async_trait]
1370impl native::admin_service_server::AdminService for NativeRuntime {
1371 async fn execute_admin(
1372 &self,
1373 request: Request<native::ExecuteAdminRequest>,
1374 ) -> Result<Response<native::Empty>, Status> {
1375 let request = request.into_inner();
1376 validate_native_context(request.context.as_ref())?;
1377 let (_, entry) = self.session(&request.session_id, "admin")?;
1378 let sql = String::from_utf8(request.command)
1379 .map_err(|_| Status::invalid_argument("admin command must be UTF-8 SQL"))?;
1380 entry.session().run(&sql).await.map_err(query_status)?;
1381 Ok(Response::new(native::Empty {}))
1382 }
1383}
1384
1385#[tonic::async_trait]
1386impl native::health_service_server::HealthService for NativeRuntime {
1387 async fn status(
1388 &self,
1389 request: Request<native::HealthRequest>,
1390 ) -> Result<Response<native::HealthResponse>, Status> {
1391 validate_native_context(request.get_ref().context.as_ref())?;
1392 Ok(Response::new(native::HealthResponse {
1393 serving: self.db.lifecycle_state() == mongreldb_core::LifecycleState::Open,
1394 detail: "ready".into(),
1395 }))
1396 }
1397}
1398
1399async fn transaction_sql(
1400 runtime: &NativeRuntime,
1401 request: native::TransactionRequest,
1402 sql: &str,
1403) -> Result<Response<native::Empty>, Status> {
1404 validate_native_context(request.context.as_ref())?;
1405 let (_, entry) = runtime.session(&request.session_id, "transaction")?;
1406 entry.session().run(sql).await.map_err(query_status)?;
1407 Ok(Response::new(native::Empty {}))
1408}
1409
1410fn query_id(bytes: &[u8]) -> Result<QueryId, Status> {
1411 QueryId::from_str(&id_hex(bytes, "query id")?)
1412 .map_err(|_| Status::invalid_argument("query id must be 16 bytes"))
1413}
1414
1415fn id_hex(bytes: &[u8], label: &str) -> Result<String, Status> {
1416 if bytes.len() != 16 {
1417 return Err(Status::invalid_argument(format!(
1418 "{label} must be 16 bytes"
1419 )));
1420 }
1421 Ok(bytes.iter().map(|byte| format!("{byte:02x}")).collect())
1422}
1423
1424fn hex_id(text: &str) -> Result<Vec<u8>, Status> {
1425 if text.len() != 32 {
1426 return Err(Status::internal("invalid server session id"));
1427 }
1428 (0..16)
1429 .map(|index| {
1430 u8::from_str_radix(&text[index * 2..index * 2 + 2], 16)
1431 .map_err(|_| Status::internal("invalid server session id"))
1432 })
1433 .collect()
1434}
1435
1436fn request_timeout(context: Option<&native::RequestContext>) -> Result<Option<Duration>, Status> {
1437 let deadline = context.map_or(0, |context| context.deadline_unix_micros);
1438 if deadline == 0 {
1439 return Ok(None);
1440 }
1441 let remaining = deadline.saturating_sub(now_unix_micros());
1442 if remaining == 0 {
1443 return Err(structured_status(
1444 Code::DeadlineExceeded,
1445 ErrorCategory::DeadlineExceeded,
1446 "request deadline exceeded",
1447 ));
1448 }
1449 Ok(Some(Duration::from_micros(remaining)))
1450}
1451
1452fn now_unix_micros() -> u64 {
1453 SystemTime::now()
1454 .duration_since(UNIX_EPOCH)
1455 .unwrap_or_default()
1456 .as_micros()
1457 .min(u128::from(u64::MAX)) as u64
1458}
1459
1460fn now_unix_seconds() -> u64 {
1461 SystemTime::now()
1462 .duration_since(UNIX_EPOCH)
1463 .unwrap_or_default()
1464 .as_secs()
1465}
1466
1467fn native_identity(principal: Option<&Principal>) -> native::AuthenticatedIdentity {
1468 match principal {
1469 Some(principal) => native::AuthenticatedIdentity {
1470 principal_id: principal.user_id,
1471 principal_name: principal.username.clone(),
1472 roles: principal.roles.clone(),
1473 scopes: principal
1474 .permissions
1475 .iter()
1476 .map(|permission| format!("{permission:?}"))
1477 .collect(),
1478 },
1479 None => native::AuthenticatedIdentity {
1480 principal_id: 0,
1481 principal_name: "anonymous".into(),
1482 roles: Vec::new(),
1483 scopes: Vec::new(),
1484 },
1485 }
1486}
1487
1488fn native_external_identity(
1489 principal: &Principal,
1490 label: String,
1491 scopes: Vec<String>,
1492) -> native::AuthenticatedIdentity {
1493 native::AuthenticatedIdentity {
1494 principal_id: principal.user_id,
1495 principal_name: label,
1496 roles: principal.roles.clone(),
1497 scopes,
1498 }
1499}
1500
1501fn parameter_literal(
1502 value: &mongreldb_protocol::request::ParameterValue,
1503) -> Result<String, Status> {
1504 use mongreldb_protocol::request::ParameterValue;
1505 match value {
1506 ParameterValue::Null => Ok("NULL".into()),
1507 ParameterValue::Bool(value) => Ok(if *value { "TRUE" } else { "FALSE" }.into()),
1508 ParameterValue::Integer(value) => Ok(value.to_string()),
1509 ParameterValue::Float(value) if value.is_finite() => Ok(value.to_string()),
1510 ParameterValue::Float(_) => Err(Status::invalid_argument("non-finite float parameter")),
1511 ParameterValue::Text(value) => Ok(format!("'{}'", value.replace('\'', "''"))),
1512 ParameterValue::Bytes(value) => Ok(format!(
1513 "X'{}'",
1514 value
1515 .iter()
1516 .map(|byte| format!("{byte:02x}"))
1517 .collect::<String>()
1518 )),
1519 }
1520}
1521
1522fn bind_numbered_parameters(sql: &str, literals: &[String]) -> Result<String, Status> {
1523 #[derive(Clone, Copy)]
1524 enum State {
1525 Normal,
1526 Quote(char),
1527 LineComment,
1528 BlockComment,
1529 }
1530
1531 let mut output = String::with_capacity(sql.len());
1532 let mut used = vec![false; literals.len()];
1533 let mut state = State::Normal;
1534 let mut chars = sql.chars().peekable();
1535 while let Some(character) = chars.next() {
1536 output.push(character);
1537 match state {
1538 State::Normal if matches!(character, '\'' | '"' | '`') => {
1539 state = State::Quote(character);
1540 }
1541 State::Normal if character == '-' && chars.peek() == Some(&'-') => {
1542 output.push(chars.next().expect("peeked"));
1543 state = State::LineComment;
1544 }
1545 State::Normal if character == '/' && chars.peek() == Some(&'*') => {
1546 output.push(chars.next().expect("peeked"));
1547 state = State::BlockComment;
1548 }
1549 State::Normal if character == '$' && chars.peek().is_some_and(char::is_ascii_digit) => {
1550 output.pop();
1551 let mut digits = String::new();
1552 while chars.peek().is_some_and(char::is_ascii_digit) {
1553 digits.push(chars.next().expect("peeked"));
1554 }
1555 let index = digits
1556 .parse::<usize>()
1557 .ok()
1558 .and_then(|index| index.checked_sub(1))
1559 .filter(|index| *index < literals.len())
1560 .ok_or_else(|| {
1561 Status::invalid_argument("prepared parameter is out of range")
1562 })?;
1563 used[index] = true;
1564 output.push_str(&literals[index]);
1565 }
1566 State::Quote(_) if character == '\\' => {
1567 if let Some(escaped) = chars.next() {
1568 output.push(escaped);
1569 }
1570 }
1571 State::Quote(end) if character == end && chars.peek() == Some(&end) => {
1572 output.push(chars.next().expect("peeked"));
1573 }
1574 State::Quote(end) if character == end => state = State::Normal,
1575 State::LineComment if character == '\n' => state = State::Normal,
1576 State::BlockComment if character == '*' && chars.peek() == Some(&'/') => {
1577 output.push(chars.next().expect("peeked"));
1578 state = State::Normal;
1579 }
1580 _ => {}
1581 }
1582 }
1583 if used.iter().any(|used| !used) {
1584 return Err(Status::invalid_argument(
1585 "prepared parameter count does not match SQL placeholders",
1586 ));
1587 }
1588 Ok(output)
1589}
1590
1591fn encode_batch(
1592 batch: &arrow::record_batch::RecordBatch,
1593 sequence: u64,
1594 end_of_stream: bool,
1595) -> Result<native::ArrowFrame, Status> {
1596 let mut ipc = Vec::new();
1597 {
1598 let mut writer = StreamWriter::try_new(&mut ipc, &batch.schema())
1599 .map_err(|error| Status::internal(error.to_string()))?;
1600 writer
1601 .write(batch)
1602 .and_then(|_| writer.finish())
1603 .map_err(|error| Status::internal(error.to_string()))?;
1604 }
1605 Ok(native::ArrowFrame {
1606 ipc,
1607 sequence,
1608 end_of_stream,
1609 })
1610}
1611
1612fn core_status(error: MongrelError) -> Status {
1613 let category = error.category();
1614 let code = match category {
1615 ErrorCategory::Unauthenticated => Code::Unauthenticated,
1616 ErrorCategory::PermissionDenied => Code::PermissionDenied,
1617 ErrorCategory::DeadlineExceeded => Code::DeadlineExceeded,
1618 ErrorCategory::ResourceExhausted => Code::ResourceExhausted,
1619 ErrorCategory::StaleMetadata
1620 | ErrorCategory::SchemaVersionMismatch
1621 | ErrorCategory::ClusterVersionMismatch => Code::FailedPrecondition,
1622 ErrorCategory::TransactionConflict
1623 | ErrorCategory::SerializationFailure
1624 | ErrorCategory::Deadlock => Code::Aborted,
1625 _ => Code::Internal,
1626 };
1627 structured_status(code, category, &error.to_string())
1628}
1629
1630fn query_status(error: MongrelQueryError) -> Status {
1631 query_status_ref(&error)
1632}
1633
1634fn query_status_ref(error: &MongrelQueryError) -> Status {
1635 match error {
1636 MongrelQueryError::Core(error) => {
1637 let category = error.category();
1638 let code = match category {
1639 ErrorCategory::Unauthenticated => Code::Unauthenticated,
1640 ErrorCategory::PermissionDenied => Code::PermissionDenied,
1641 ErrorCategory::DeadlineExceeded => Code::DeadlineExceeded,
1642 ErrorCategory::ResourceExhausted => Code::ResourceExhausted,
1643 ErrorCategory::StaleMetadata
1644 | ErrorCategory::SchemaVersionMismatch
1645 | ErrorCategory::ClusterVersionMismatch => Code::FailedPrecondition,
1646 ErrorCategory::TransactionConflict
1647 | ErrorCategory::SerializationFailure
1648 | ErrorCategory::Deadlock => Code::Aborted,
1649 _ => Code::Internal,
1650 };
1651 structured_status(code, category, &error.to_string())
1652 }
1653 MongrelQueryError::DeadlineExceeded { .. } => structured_status(
1654 Code::DeadlineExceeded,
1655 ErrorCategory::DeadlineExceeded,
1656 &error.to_string(),
1657 ),
1658 MongrelQueryError::QueryCancelled { .. } => structured_status(
1659 Code::Cancelled,
1660 ErrorCategory::Cancelled,
1661 &error.to_string(),
1662 ),
1663 MongrelQueryError::QueryRegistryFull | MongrelQueryError::ResultLimitExceeded { .. } => {
1664 structured_status(
1665 Code::ResourceExhausted,
1666 ErrorCategory::ResourceExhausted,
1667 &error.to_string(),
1668 )
1669 }
1670 MongrelQueryError::TransactionAborted => structured_status(
1671 Code::Aborted,
1672 ErrorCategory::TransactionAborted,
1673 &error.to_string(),
1674 ),
1675 MongrelQueryError::OutcomeUnknown { .. } => structured_status(
1676 Code::Unknown,
1677 ErrorCategory::CommitOutcomeUnknown,
1678 &error.to_string(),
1679 ),
1680 error => Status::internal(error.to_string()),
1681 }
1682}
1683
1684fn stream_status(error: impl std::error::Error + 'static) -> Status {
1685 let message = error.to_string();
1686 let mut source = Some(&error as &(dyn std::error::Error + 'static));
1687 while let Some(current) = source {
1688 if let Some(error) = current.downcast_ref::<MongrelQueryError>() {
1689 return query_status_ref(error);
1690 }
1691 source = current.source();
1692 }
1693 Status::internal(message)
1694}
1695
1696fn structured_status(code: Code, category: ErrorCategory, message: &str) -> Status {
1697 let detail = native::ErrorDetail {
1698 category_code: category.code(),
1699 category: category.to_string(),
1700 message: message.into(),
1701 retryable: category.retry_class() != RetryClass::Never,
1702 metadata: HashMap::new(),
1703 };
1704 Status::with_details(code, message, detail.encode_to_vec().into())
1705}
1706
1707fn native_phase(phase: SqlQueryPhase) -> native::QueryPhase {
1708 match phase {
1709 SqlQueryPhase::Queued => native::QueryPhase::Queued,
1710 SqlQueryPhase::Planning => native::QueryPhase::Planning,
1711 SqlQueryPhase::Executing
1712 | SqlQueryPhase::Streaming
1713 | SqlQueryPhase::CommitCritical
1714 | SqlQueryPhase::Cancelling => native::QueryPhase::Executing,
1715 SqlQueryPhase::Serializing => native::QueryPhase::Serializing,
1716 SqlQueryPhase::Completed => native::QueryPhase::Completed,
1717 SqlQueryPhase::Failed => native::QueryPhase::Failed,
1718 SqlQueryPhase::Cancelled => native::QueryPhase::Cancelled,
1719 }
1720}