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 let 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 receipt.attach_commit_receipt(commit_receipt);
460 execution.commit(receipt);
461 } else if super::can_abort_idempotency_intent(&status) {
462 execution.abort();
463 }
464 }
465
466 async fn resolve_sql(
467 &self,
468 entry: &SessionEntry,
469 request: &native::ExecuteRequest,
470 ) -> Result<String, Status> {
471 use native::execute_request::Command;
472 match &request.command {
473 Some(Command::Sql(sql)) => Ok(sql.clone()),
474 Some(Command::PreparedStatementId(id)) => {
475 let (_, binding) = entry
476 .prepared_binding_by_id(*id)
477 .ok_or_else(|| Status::not_found("prepared statement not found"))?;
478 if !prepared::CatalogState::capture(&self.db).is_compatible(&binding) {
479 return Err(structured_status(
480 Code::FailedPrecondition,
481 ErrorCategory::SchemaVersionMismatch,
482 "prepared statement schema changed",
483 ));
484 }
485 let values = request
486 .parameters
487 .iter()
488 .map(|encoded| {
489 bincode::deserialize::<mongreldb_protocol::request::ParameterValue>(encoded)
490 .map_err(|_| Status::invalid_argument("invalid prepared parameter"))
491 })
492 .collect::<Result<Vec<_>, _>>()?;
493 let literals = values
494 .iter()
495 .map(parameter_literal)
496 .collect::<Result<Vec<_>, _>>()?;
497 bind_numbered_parameters(&binding.sql, &literals)
498 }
499 Some(Command::AdminCommand(command)) => String::from_utf8(command.clone())
500 .map_err(|_| Status::invalid_argument("admin command must be UTF-8 SQL")),
501 None => Err(Status::invalid_argument("execute command is required")),
502 }
503 }
504}
505
506#[tonic::async_trait]
507impl native::auth_service_server::AuthService for NativeRuntime {
508 async fn authenticate(
509 &self,
510 request: Request<native::AuthenticateRequest>,
511 ) -> Result<Response<native::AuthenticateResponse>, Status> {
512 let request = request.into_inner();
513 validate_native_context(request.context.as_ref())?;
514 let (principal, identity, response_identity) = match request.credential {
515 Some(native::authenticate_request::Credential::Password(password)) => {
516 let username = password.username;
517 let password = zeroize::Zeroizing::new(password.password);
518 let principal = self
519 .db
520 .authenticate_principal(&username, password.as_str())
521 .map_err(core_status)?
522 .ok_or_else(|| Status::unauthenticated("invalid credentials"))?;
523 let identity = AuthenticatedIdentity::CatalogUser {
524 username: principal.username.clone(),
525 user_id: principal.user_id,
526 created_version: principal.created_epoch,
527 };
528 let response_identity = native_identity(Some(&principal));
529 (Some(principal), identity, response_identity)
530 }
531 Some(native::authenticate_request::Credential::MysqlCachingSha2(credential)) => {
532 if credential.nonce.len() != 20
533 || (!credential.proof.is_empty() && credential.proof.len() != 32)
534 {
535 return Err(Status::invalid_argument(
536 "invalid caching_sha2_password proof",
537 ));
538 }
539 let principal = self
540 .db
541 .authenticate_mysql_caching_sha2_principal(
542 &credential.username,
543 &credential.nonce,
544 &credential.proof,
545 )
546 .ok_or_else(|| Status::unauthenticated("invalid credentials"))?;
547 let identity = AuthenticatedIdentity::CatalogUser {
548 username: principal.username.clone(),
549 user_id: principal.user_id,
550 created_version: principal.created_epoch,
551 };
552 let response_identity = native_identity(Some(&principal));
553 (Some(principal), identity, response_identity)
554 }
555 None if !self.db.require_auth_enabled() => (
556 None,
557 AuthenticatedIdentity::Credentialless,
558 native_identity(None),
559 ),
560 Some(native::authenticate_request::Credential::ServiceToken(credential)) => {
561 let auth = self.external_auth.as_ref().ok_or_else(|| {
562 Status::unauthenticated("service-token authentication is not configured")
563 })?;
564 let secret = zeroize::Zeroizing::new(credential.secret);
565 let token = auth
566 .service_tokens
567 .read()
568 .map_err(|_| Status::internal("service-token registry poisoned"))?
569 .authenticate(&credential.token_id, secret.as_str(), now_unix_seconds())
570 .cloned()
571 .ok_or_else(|| Status::unauthenticated("invalid credentials"))?;
572 let principal = self
573 .db
574 .resolve_principal(&token.principal)
575 .ok_or_else(|| Status::unauthenticated("invalid credentials"))?;
576 let identity = AuthenticatedIdentity::ExternalPrincipal {
577 provider: "service_token".into(),
578 subject: token.token_id.clone(),
579 username: principal.username.clone(),
580 user_id: principal.user_id,
581 created_version: principal.created_epoch,
582 scopes: token.scopes.clone(),
583 };
584 let response_identity =
585 native_external_identity(&principal, token.token_id, token.scopes);
586 (Some(principal), identity, response_identity)
587 }
588 Some(native::authenticate_request::Credential::Oidc(credential)) => {
589 let oidc = self
590 .external_auth
591 .as_ref()
592 .and_then(|auth| auth.oidc.clone())
593 .ok_or_else(|| {
594 Status::unauthenticated("OIDC authentication is not configured")
595 })?;
596 let token = credential.compact_jws;
597 let verified = tokio::task::spawn_blocking(move || {
598 oidc.cache
599 .verify(&token, &oidc.validation, now_unix_seconds())
600 })
601 .await
602 .map_err(|_| Status::internal("OIDC verifier task failed"))?
603 .map_err(|_| Status::unauthenticated("invalid credentials"))?;
604 let principal = self
605 .db
606 .resolve_principal(&verified.principal)
607 .ok_or_else(|| Status::unauthenticated("invalid credentials"))?;
608 let identity = AuthenticatedIdentity::ExternalPrincipal {
609 provider: "oidc".into(),
610 subject: verified.principal.clone(),
611 username: principal.username.clone(),
612 user_id: principal.user_id,
613 created_version: principal.created_epoch,
614 scopes: verified.scopes.clone(),
615 };
616 let response_identity =
617 native_external_identity(&principal, verified.principal, verified.scopes);
618 (Some(principal), identity, response_identity)
619 }
620 None => return Err(Status::unauthenticated("credentials are required")),
621 };
622 let auth_token = self.issue_auth_grant(principal.clone(), identity)?;
623 Ok(Response::new(native::AuthenticateResponse {
624 identity: Some(response_identity),
625 auth_token,
626 }))
627 }
628
629 async fn begin_scram(
630 &self,
631 request: Request<native::BeginScramRequest>,
632 ) -> Result<Response<native::BeginScramResponse>, Status> {
633 let request = request.into_inner();
634 validate_native_context(request.context.as_ref())?;
635 if !request
636 .client_first_bare
637 .starts_with(&format!("n={},", request.username))
638 {
639 return Err(Status::invalid_argument(
640 "SCRAM username does not match client-first message",
641 ));
642 }
643 let (verifier, fake_user) = match self.db.user_scram_verifier(&request.username) {
644 Some(verifier) => (verifier, false),
645 None => {
646 let salt = mongreldb_types::ids::QueryId::new_random();
647 (
648 mongreldb_core::ScramVerifier::from_password(
649 "invalid-user-password",
650 salt.as_bytes(),
651 mongreldb_core::security_hardening::SCRAM_SHA_256_MIN_ITERATIONS,
652 )
653 .map_err(|error| Status::internal(error.to_string()))?,
654 true,
655 )
656 }
657 };
658 let server_nonce = mongreldb_types::ids::QueryId::new_random().to_hex();
659 let session = mongreldb_core::ScramServerSession::begin(
660 verifier,
661 request.client_first_bare,
662 &request.client_nonce,
663 &server_nonce,
664 mongreldb_core::ScramChannelBindingPolicy::Disabled,
665 Vec::new(),
666 )
667 .map_err(|error| Status::invalid_argument(error.to_string()))?;
668 let server_first = session.server_first_message().to_owned();
669 let exchange_id = mongreldb_types::ids::QueryId::new_random()
670 .as_bytes()
671 .to_vec();
672 let now = now_unix_micros();
673 let mut exchanges = self
674 .scram_exchanges
675 .lock()
676 .map_err(|_| Status::internal("SCRAM exchange store poisoned"))?;
677 exchanges.retain(|_, exchange| exchange.expires_unix_micros > now);
678 if exchanges.len() >= MAX_AUTH_TOKENS {
679 return Err(Status::resource_exhausted("too many SCRAM exchanges"));
680 }
681 exchanges.insert(
682 exchange_id.clone(),
683 PendingScram {
684 username: request.username,
685 session,
686 fake_user,
687 expires_unix_micros: now.saturating_add(AUTH_TOKEN_TTL.as_micros() as u64),
688 },
689 );
690 Ok(Response::new(native::BeginScramResponse {
691 exchange_id,
692 server_first,
693 }))
694 }
695
696 async fn finish_scram(
697 &self,
698 request: Request<native::FinishScramRequest>,
699 ) -> Result<Response<native::FinishScramResponse>, Status> {
700 let request = request.into_inner();
701 validate_native_context(request.context.as_ref())?;
702 let exchange = self
703 .scram_exchanges
704 .lock()
705 .map_err(|_| Status::internal("SCRAM exchange store poisoned"))?
706 .remove(&request.exchange_id)
707 .ok_or_else(|| Status::unauthenticated("invalid SCRAM exchange"))?;
708 if exchange.expires_unix_micros <= now_unix_micros() {
709 return Err(Status::unauthenticated("invalid SCRAM exchange"));
710 }
711 let server_final = exchange
712 .session
713 .finish(&request.client_final_without_proof, &request.client_proof)
714 .map_err(|_| Status::unauthenticated("invalid credentials"))?;
715 if exchange.fake_user {
716 return Err(Status::unauthenticated("invalid credentials"));
717 }
718 let principal = self
719 .db
720 .resolve_principal(&exchange.username)
721 .ok_or_else(|| Status::unauthenticated("invalid credentials"))?;
722 let identity = AuthenticatedIdentity::CatalogUser {
723 username: principal.username.clone(),
724 user_id: principal.user_id,
725 created_version: principal.created_epoch,
726 };
727 let auth_token = self.issue_auth_grant(Some(principal.clone()), identity)?;
728 Ok(Response::new(native::FinishScramResponse {
729 server_final,
730 authentication: Some(native::AuthenticateResponse {
731 identity: Some(native_identity(Some(&principal))),
732 auth_token,
733 }),
734 }))
735 }
736}
737
738#[tonic::async_trait]
739impl native::session_service_server::SessionService for NativeRuntime {
740 async fn open_session(
741 &self,
742 request: Request<native::OpenSessionRequest>,
743 ) -> Result<Response<native::OpenSessionResponse>, Status> {
744 let request = request.into_inner();
745 validate_native_context(request.context.as_ref())?;
746 if request.database_id != self.sessions.database_id().as_bytes() {
747 return Err(Status::not_found("database not found"));
748 }
749 let grant = self.consume_auth_grant(&request.auth_token)?;
750 let owner = grant.principal.as_ref().map_or_else(
751 || "anonymous".into(),
752 |principal| principal.username.clone(),
753 );
754 let session = MongrelSession::open_with_external_modules_as(
755 Arc::clone(&self.db),
756 std::iter::empty(),
757 grant.principal,
758 )
759 .map_err(query_status)?
760 .with_query_registry(Arc::clone(&self.query_registry));
761 let token = self
762 .sessions
763 .create_with_identity(session, owner, grant.identity)
764 .ok_or_else(|| Status::resource_exhausted("session limit reached"))?;
765 Ok(Response::new(native::OpenSessionResponse {
766 session_id: hex_id(&token)?,
767 }))
768 }
769
770 async fn close_session(
771 &self,
772 request: Request<native::CloseSessionRequest>,
773 ) -> Result<Response<native::Empty>, Status> {
774 let request = request.into_inner();
775 validate_native_context(request.context.as_ref())?;
776 let token = id_hex(&request.session_id, "session id")?;
777 if !self.sessions.close_by_token(&token) {
778 return Err(Status::not_found("session not found"));
779 }
780 self.query_registry
781 .cancel_session(&token, mongreldb_core::CancellationReason::SessionClosed);
782 Ok(Response::new(native::Empty {}))
783 }
784}
785
786#[tonic::async_trait]
787impl native::query_service_server::QueryService for NativeRuntime {
788 type ExecuteStreamStream =
789 Pin<Box<dyn Stream<Item = Result<native::ArrowFrame, Status>> + Send + 'static>>;
790
791 async fn prepare(
792 &self,
793 request: Request<native::PrepareRequest>,
794 ) -> Result<Response<native::PrepareResponse>, Status> {
795 let request = request.into_inner();
796 validate_native_context(request.context.as_ref())?;
797 let (_, entry) = self.session(&request.session_id, "query")?;
798 let _guard = entry.lock.lock().await;
799 let statement_id = entry.allocate_statement_id();
800 let name = format!("native_{}", statement_id.get());
801 entry
802 .session()
803 .run(&format!("PREPARE {name} AS {}", request.sql))
804 .await
805 .map_err(query_status)?;
806 let catalog = prepared::CatalogState::capture(&self.db);
807 entry.insert_prepared_binding(
808 name,
809 prepared::build_binding(statement_id, request.sql, Vec::new(), &catalog),
810 );
811 Ok(Response::new(native::PrepareResponse {
812 statement_id: statement_id.get(),
813 schema_version: catalog.catalog_version.get(),
814 }))
815 }
816
817 async fn execute(
818 &self,
819 request: Request<native::ExecuteRequest>,
820 ) -> Result<Response<native::ExecuteResponse>, Status> {
821 let request = request.into_inner();
822 validate_native_context(request.context.as_ref())?;
823 let (session_id, entry) = self.session(&request.session_id, "query")?;
824 let id = query_id(&request.query_id)?;
825 let sql = self.resolve_sql(&entry, &request).await?;
826 let idempotency = self
827 .begin_idempotency(&request, &entry, &session_id, &sql)
828 .await?;
829 if let NativeIdempotency::Replay(receipt) = idempotency {
830 return Ok(Response::new(replayed_response(&request.query_id, receipt)));
831 }
832 let execution = match idempotency {
833 NativeIdempotency::Execute(execution) => Some(execution),
834 NativeIdempotency::None => None,
835 NativeIdempotency::Replay(_) => unreachable!(),
836 };
837 let (query_id, batches) = match self
838 .execute_sql(entry, session_id, id, sql, request.context.as_ref())
839 .await
840 {
841 Ok(result) => result,
842 Err(error) => {
843 if let Some(execution) = execution {
844 self.finish_idempotency(execution, id);
845 }
846 return Err(error);
847 }
848 };
849 let frames = batches
850 .batches()
851 .iter()
852 .enumerate()
853 .map(|(sequence, batch)| encode_batch(batch, sequence as u64, false))
854 .chain(std::iter::once(Ok(native::ArrowFrame {
855 ipc: Vec::new(),
856 sequence: batches.batches().len() as u64,
857 end_of_stream: true,
858 })))
859 .collect::<Result<Vec<_>, Status>>();
860 let frames = match frames {
861 Ok(frames) => {
862 batches.complete().map_err(query_status)?;
863 frames
864 }
865 Err(error) => {
866 batches.fail_serialization();
867 if let Some(execution) = execution {
868 self.finish_idempotency(execution, id);
869 }
870 return Err(error);
871 }
872 };
873 let status = self.query_registry.status(id);
874 if let Some(execution) = execution {
875 self.finish_idempotency(execution, id);
876 }
877 let durable = status
878 .as_ref()
879 .map(|status| proto_durable(&status.durable_outcome, status.serialization_outcome))
880 .unwrap_or_default();
881 Ok(Response::new(native::ExecuteResponse {
882 query_id,
883 rows_affected: 0,
884 frames,
885 idempotency_replayed: false,
886 committed: durable.committed,
887 commit_epoch: durable.last_commit_epoch.unwrap_or(0),
888 original_query_id: id.as_bytes().to_vec(),
889 durable: Some(durable),
890 }))
891 }
892
893 async fn execute_stream(
894 &self,
895 request: Request<native::ExecuteRequest>,
896 ) -> Result<Response<Self::ExecuteStreamStream>, Status> {
897 let request = request.into_inner();
898 validate_native_context(request.context.as_ref())?;
899 if request
900 .context
901 .as_ref()
902 .is_some_and(|context| !context.idempotency_key.is_empty())
903 {
904 return Err(Status::invalid_argument(
905 "idempotent writes require buffered Execute",
906 ));
907 }
908 let (session_id, entry) = self.session(&request.session_id, "query")?;
909 let id = query_id(&request.query_id)?;
910 let sql = self.resolve_sql(&entry, &request).await?;
911 let stream = entry
912 .session()
913 .run_stream_with_options(
914 &sql,
915 SqlQueryOptions {
916 query_id: Some(id),
917 timeout: request_timeout(request.context.as_ref())?,
918 owner: Some(entry.owner.clone()),
919 session_id: Some(session_id),
920 parent_control: None,
921 },
922 )
923 .await
924 .map_err(query_status)?;
925 let output = futures::stream::try_unfold(
926 (stream, 0_u64, false),
927 |(mut stream, sequence, done)| async move {
928 if done {
929 return Ok(None);
930 }
931 match stream.next().await {
932 Some(Ok(batch)) => {
933 let frame = encode_batch(&batch, sequence, false)?;
934 Ok(Some((frame, (stream, sequence + 1, false))))
935 }
936 Some(Err(error)) => Err(stream_status(error)),
937 None => Ok(Some((
938 native::ArrowFrame {
939 ipc: Vec::new(),
940 sequence,
941 end_of_stream: true,
942 },
943 (stream, sequence + 1, true),
944 ))),
945 }
946 },
947 );
948 Ok(Response::new(Box::pin(output)))
949 }
950
951 async fn cancel_query(
952 &self,
953 request: Request<native::CancelQueryRequest>,
954 ) -> Result<Response<native::CancelQueryResponse>, Status> {
955 let request = request.into_inner();
956 validate_native_context(request.context.as_ref())?;
957 let (session_id, _) = self.session(&request.session_id, "query")?;
958 let query_id = query_id(&request.query_id)?;
959 let status = self
960 .query_registry
961 .status(query_id)
962 .ok_or_else(|| Status::not_found("query not found"))?;
963 if status.session_id.as_deref() != Some(&session_id) {
964 return Err(Status::not_found("query not found"));
965 }
966 let durable = proto_durable(&status.durable_outcome, status.serialization_outcome);
967 let outcome = match self.query_registry.cancel(query_id) {
968 CancelOutcome::Accepted => native::CancelOutcome::Accepted,
969 CancelOutcome::AlreadyCancelling => native::CancelOutcome::AlreadyCancelling,
970 CancelOutcome::TooLate => native::CancelOutcome::TooLate,
971 CancelOutcome::AlreadyFinished => native::CancelOutcome::AlreadyFinished,
972 CancelOutcome::NotFound => native::CancelOutcome::NotFound,
973 };
974 let durable = self
976 .query_registry
977 .status(query_id)
978 .map(|status| proto_durable(&status.durable_outcome, status.serialization_outcome))
979 .unwrap_or(durable);
980 Ok(Response::new(native::CancelQueryResponse {
981 outcome: outcome as i32,
982 durable: Some(durable),
983 }))
984 }
985
986 async fn get_query_status(
987 &self,
988 request: Request<native::GetQueryStatusRequest>,
989 ) -> Result<Response<native::QueryStatusResponse>, Status> {
990 let request = request.into_inner();
991 validate_native_context(request.context.as_ref())?;
992 let (session_id, _) = self.session(&request.session_id, "query")?;
993 let id = query_id(&request.query_id)?;
994 let status = self
995 .query_registry
996 .status(id)
997 .ok_or_else(|| Status::not_found("query not found"))?;
998 if status.session_id.as_deref() != Some(&session_id) {
999 return Err(Status::not_found("query not found"));
1000 }
1001 Ok(Response::new(native::QueryStatusResponse {
1002 query_id: id.as_bytes().to_vec(),
1003 phase: native_phase(status.phase) as i32,
1004 error: status.terminal_error.map(|error| native::ErrorDetail {
1005 category_code: 0,
1006 category: format!("{:?}", error.category),
1007 message: error.code,
1008 retryable: false,
1009 metadata: HashMap::new(),
1010 }),
1011 durable: Some(proto_durable(
1012 &status.durable_outcome,
1013 status.serialization_outcome,
1014 )),
1015 }))
1016 }
1017}
1018
1019fn proto_durable(
1020 outcome: &mongreldb_query::DurableOutcome,
1021 serialization: SerializationOutcome,
1022) -> native::DurableOutcome {
1023 let last_commit_hlc = outcome
1024 .commit_ts
1025 .map(|ts| {
1026 let mut bytes = Vec::with_capacity(16);
1027 bytes.extend_from_slice(&ts.physical_micros.to_be_bytes());
1028 bytes.extend_from_slice(&ts.logical.to_be_bytes());
1029 bytes.extend_from_slice(&ts.node_tiebreaker.to_be_bytes());
1030 bytes
1031 })
1032 .unwrap_or_default();
1033 native::DurableOutcome {
1034 committed: outcome.committed,
1035 committed_statements: outcome.committed_statements as u64,
1036 last_commit_hlc,
1037 first_commit_statement_index: outcome
1038 .first_commit_statement_index
1039 .map(|index| index as u64),
1040 last_commit_statement_index: outcome
1041 .last_commit_statement_index
1042 .map(|index| index as u64),
1043 completed_statements: outcome.committed_statements as u64,
1044 current_statement_index: outcome
1045 .last_commit_statement_index
1046 .map(|index| index as u64)
1047 .unwrap_or(0),
1048 terminal_state: if outcome.committed {
1049 "committed".into()
1050 } else {
1051 String::new()
1052 },
1053 serialization_state: serialization_state_name(serialization).into(),
1054 last_commit_epoch: outcome.last_commit_epoch,
1055 }
1056}
1057
1058fn serialization_state_name(outcome: SerializationOutcome) -> &'static str {
1059 match outcome {
1060 SerializationOutcome::NotStarted => "not_started",
1061 SerializationOutcome::InProgress => "in_progress",
1062 SerializationOutcome::Succeeded => "succeeded",
1063 SerializationOutcome::Failed => "failed",
1064 }
1065}
1066
1067fn receipt_durable(receipt: &crate::sql_idempotency::SqlDurableReceipt) -> native::DurableOutcome {
1068 native::DurableOutcome {
1069 committed: receipt.outcome.committed,
1070 committed_statements: receipt.outcome.committed_statements as u64,
1071 last_commit_hlc: Vec::new(),
1072 first_commit_statement_index: receipt
1073 .outcome
1074 .first_commit_statement_index
1075 .map(|index| index as u64),
1076 last_commit_statement_index: receipt
1077 .outcome
1078 .last_commit_statement_index
1079 .map(|index| index as u64),
1080 completed_statements: receipt.outcome.committed_statements as u64,
1081 current_statement_index: 0,
1082 terminal_state: if receipt.outcome.committed {
1083 "committed".into()
1084 } else {
1085 String::new()
1086 },
1087 serialization_state: "succeeded".into(),
1089 last_commit_epoch: receipt.outcome.last_commit_epoch,
1090 }
1091}
1092
1093fn replayed_response(
1094 current_query_id: &[u8],
1095 receipt: crate::sql_idempotency::SqlDurableReceipt,
1096) -> native::ExecuteResponse {
1097 let durable = receipt_durable(&receipt);
1098 native::ExecuteResponse {
1099 query_id: current_query_id.to_vec(),
1100 rows_affected: 0,
1101 frames: Vec::new(),
1102 idempotency_replayed: true,
1103 committed: durable.committed,
1104 commit_epoch: durable.last_commit_epoch.unwrap_or(0),
1105 original_query_id: hex_id(&receipt.original_query_id).unwrap_or_default(),
1106 durable: Some(durable),
1107 }
1108}
1109
1110#[tonic::async_trait]
1111impl native::transaction_service_server::TransactionService for NativeRuntime {
1112 async fn begin(
1113 &self,
1114 request: Request<native::BeginTransactionRequest>,
1115 ) -> Result<Response<native::BeginTransactionResponse>, Status> {
1116 let request = request.into_inner();
1117 validate_native_context(request.context.as_ref())?;
1118 let (_, entry) = self.session(&request.session_id, "transaction")?;
1119 let sql = match native::IsolationLevel::try_from(request.isolation) {
1120 Ok(native::IsolationLevel::ReadCommitted) => {
1121 "BEGIN; SET TRANSACTION ISOLATION LEVEL READ COMMITTED"
1122 }
1123 Ok(native::IsolationLevel::Serializable) => {
1124 "BEGIN; SET TRANSACTION ISOLATION LEVEL SERIALIZABLE"
1125 }
1126 _ => "BEGIN",
1127 };
1128 entry.session().run(sql).await.map_err(query_status)?;
1129 Ok(Response::new(native::BeginTransactionResponse {
1130 transaction_id: TransactionId::new_random().as_bytes().to_vec(),
1131 }))
1132 }
1133
1134 async fn commit(
1135 &self,
1136 request: Request<native::TransactionRequest>,
1137 ) -> Result<Response<native::Empty>, Status> {
1138 transaction_sql(self, request.into_inner(), "COMMIT").await
1139 }
1140
1141 async fn rollback(
1142 &self,
1143 request: Request<native::TransactionRequest>,
1144 ) -> Result<Response<native::Empty>, Status> {
1145 transaction_sql(self, request.into_inner(), "ROLLBACK").await
1146 }
1147}
1148
1149#[tonic::async_trait]
1150impl native::catalog_service_server::CatalogService for NativeRuntime {
1151 async fn get_schema(
1152 &self,
1153 request: Request<native::GetSchemaRequest>,
1154 ) -> Result<Response<native::GetSchemaResponse>, Status> {
1155 let request = request.into_inner();
1156 validate_native_context(request.context.as_ref())?;
1157 if request.database_id != self.sessions.database_id().as_bytes() {
1158 return Err(Status::not_found("database not found"));
1159 }
1160 let (_, entry) = self.session(&request.session_id, "catalog:read")?;
1161 let principal = self.session_principal(&entry)?;
1162 self.db
1163 .require_for(
1164 principal.as_ref(),
1165 &mongreldb_core::Permission::Select {
1166 table: request.table.clone(),
1167 },
1168 )
1169 .map_err(core_status)?;
1170 let table = self.db.table(&request.table).map_err(core_status)?;
1171 let schema = table.lock().schema().clone();
1172 let schema_json = serde_json::to_vec(&schema)
1173 .map_err(|error| Status::internal(format!("schema encode failed: {error}")))?;
1174 Ok(Response::new(native::GetSchemaResponse {
1175 table: request.table,
1176 schema_version: schema.schema_id,
1177 columns: schema
1178 .columns
1179 .into_iter()
1180 .map(|column| native::ColumnSchema {
1181 name: column.name,
1182 data_type: format!("{:?}", column.ty),
1183 nullable: column.flags.contains(mongreldb_core::ColumnFlags::NULLABLE),
1184 })
1185 .collect(),
1186 schema_json,
1187 }))
1188 }
1189
1190 async fn create_table(
1191 &self,
1192 request: Request<native::CreateTableRequest>,
1193 ) -> Result<Response<native::CreateTableResponse>, Status> {
1194 let request = request.into_inner();
1195 validate_native_context(request.context.as_ref())?;
1196 let (_, entry) = self.session(&request.session_id, "catalog:write")?;
1197 let _guard = entry.lock.lock().await;
1198 let principal = self.session_principal(&entry)?;
1199 self.db
1200 .require_for(principal.as_ref(), &mongreldb_core::Permission::Ddl)
1201 .map_err(core_status)?;
1202 let mut schema = if request.schema_json.is_empty() {
1203 let columns = request
1204 .columns
1205 .into_iter()
1206 .map(native_create_column)
1207 .collect::<Result<Vec<_>, _>>()?;
1208 let uniques = request
1209 .uniques
1210 .into_iter()
1211 .map(|constraint| {
1212 Ok(mongreldb_core::constraint::UniqueConstraint {
1213 id: native_u16(constraint.id, "unique constraint id")?,
1214 name: constraint.name,
1215 columns: native_u16s(constraint.columns, "unique constraint column")?,
1216 })
1217 })
1218 .collect::<Result<Vec<_>, Status>>()?;
1219 let foreign_keys = request
1220 .foreign_keys
1221 .into_iter()
1222 .map(native_foreign_key)
1223 .collect::<Result<Vec<_>, _>>()?;
1224 mongreldb_core::Schema {
1225 schema_id: request.schema_id,
1226 columns,
1227 indexes: Vec::new(),
1228 colocation: Vec::new(),
1229 constraints: mongreldb_core::constraint::TableConstraints {
1230 uniques,
1231 foreign_keys,
1232 checks: Vec::new(),
1233 },
1234 clustered: false,
1235 }
1236 } else {
1237 let schema: mongreldb_core::Schema = serde_json::from_slice(&request.schema_json)
1238 .map_err(|error| {
1239 Status::invalid_argument(format!("invalid schema_json: {error}"))
1240 })?;
1241 if schema.schema_id != request.schema_id {
1242 return Err(Status::invalid_argument(
1243 "schema_json schema_id does not match request schema_id",
1244 ));
1245 }
1246 schema
1247 };
1248 if let Ok(existing) = self.db.table(&request.table) {
1249 let existing = existing.lock().schema().clone();
1250 schema.schema_id = existing.schema_id;
1251 if serde_json::to_vec(&schema).map_err(|_| Status::internal("schema encode failed"))?
1252 != serde_json::to_vec(&existing)
1253 .map_err(|_| Status::internal("schema encode failed"))?
1254 {
1255 return Err(Status::already_exists(
1256 "table exists with a different schema",
1257 ));
1258 }
1259 entry
1260 .session()
1261 .refresh_database_table(&request.table)
1262 .map_err(query_status)?;
1263 return Ok(Response::new(native::CreateTableResponse {
1264 table_id: self.db.table_id(&request.table).map_err(core_status)?,
1265 schema_version: self.db.catalog_version(),
1266 }));
1267 }
1268 let table_id = self
1269 .db
1270 .create_table(&request.table, schema)
1271 .map_err(core_status)?;
1272 entry
1273 .session()
1274 .refresh_database_table(&request.table)
1275 .map_err(query_status)?;
1276 Ok(Response::new(native::CreateTableResponse {
1277 table_id,
1278 schema_version: self.db.catalog_version(),
1279 }))
1280 }
1281}
1282
1283fn native_create_column(column: native::CreateColumn) -> Result<mongreldb_core::ColumnDef, Status> {
1284 let ty = match native::ColumnType::try_from(column.data_type)
1285 .map_err(|_| Status::invalid_argument("unknown native column type"))?
1286 {
1287 native::ColumnType::Bool => mongreldb_core::TypeId::Bool,
1288 native::ColumnType::Int8 => mongreldb_core::TypeId::Int8,
1289 native::ColumnType::Int16 => mongreldb_core::TypeId::Int16,
1290 native::ColumnType::Int32 => mongreldb_core::TypeId::Int32,
1291 native::ColumnType::Int64 => mongreldb_core::TypeId::Int64,
1292 native::ColumnType::Uint8 => mongreldb_core::TypeId::UInt8,
1293 native::ColumnType::Uint16 => mongreldb_core::TypeId::UInt16,
1294 native::ColumnType::Uint32 => mongreldb_core::TypeId::UInt32,
1295 native::ColumnType::Uint64 => mongreldb_core::TypeId::UInt64,
1296 native::ColumnType::Float32 => mongreldb_core::TypeId::Float32,
1297 native::ColumnType::Float64 => mongreldb_core::TypeId::Float64,
1298 native::ColumnType::TimestampNanos => mongreldb_core::TypeId::TimestampNanos,
1299 native::ColumnType::Date32 => mongreldb_core::TypeId::Date32,
1300 native::ColumnType::Date64 => mongreldb_core::TypeId::Date64,
1301 native::ColumnType::Time64 => mongreldb_core::TypeId::Time64,
1302 native::ColumnType::Bytes => mongreldb_core::TypeId::Bytes,
1303 native::ColumnType::Json => mongreldb_core::TypeId::Json,
1304 native::ColumnType::Decimal128 => mongreldb_core::TypeId::Decimal128 {
1305 precision: u8::try_from(column.decimal_precision)
1306 .map_err(|_| Status::invalid_argument("decimal precision exceeds u8"))?,
1307 scale: i8::try_from(column.decimal_scale)
1308 .map_err(|_| Status::invalid_argument("decimal scale exceeds i8"))?,
1309 },
1310 native::ColumnType::Unspecified => {
1311 return Err(Status::invalid_argument("native column type is required"))
1312 }
1313 };
1314 let mut flags = mongreldb_core::ColumnFlags::empty();
1315 if column.nullable {
1316 flags = flags.with(mongreldb_core::ColumnFlags::NULLABLE);
1317 }
1318 if column.primary_key {
1319 flags = flags.with(mongreldb_core::ColumnFlags::PRIMARY_KEY);
1320 }
1321 if column.auto_increment {
1322 flags = flags.with(mongreldb_core::ColumnFlags::AUTO_INCREMENT);
1323 }
1324 Ok(mongreldb_core::ColumnDef {
1325 id: native_u16(column.id, "column id")?,
1326 name: column.name,
1327 ty,
1328 flags,
1329 default_value: None,
1330 embedding_source: None,
1331 })
1332}
1333
1334fn native_foreign_key(
1335 foreign_key: native::ForeignKey,
1336) -> Result<mongreldb_core::constraint::ForeignKey, Status> {
1337 let action = |value| -> Result<mongreldb_core::constraint::FkAction, Status> {
1338 match native::ForeignKeyAction::try_from(value)
1339 .map_err(|_| Status::invalid_argument("unknown foreign-key action"))?
1340 {
1341 native::ForeignKeyAction::Unspecified | native::ForeignKeyAction::Restrict => {
1342 Ok(mongreldb_core::constraint::FkAction::Restrict)
1343 }
1344 native::ForeignKeyAction::Cascade => Ok(mongreldb_core::constraint::FkAction::Cascade),
1345 native::ForeignKeyAction::SetNull => Ok(mongreldb_core::constraint::FkAction::SetNull),
1346 }
1347 };
1348 Ok(mongreldb_core::constraint::ForeignKey {
1349 id: native_u16(foreign_key.id, "foreign-key id")?,
1350 name: foreign_key.name,
1351 columns: native_u16s(foreign_key.columns, "foreign-key column")?,
1352 ref_table: foreign_key.referenced_table,
1353 ref_columns: native_u16s(foreign_key.referenced_columns, "referenced column")?,
1354 on_delete: action(foreign_key.on_delete)?,
1355 on_update: action(foreign_key.on_update)?,
1356 })
1357}
1358
1359fn native_u16(value: u32, field: &str) -> Result<u16, Status> {
1360 u16::try_from(value).map_err(|_| Status::invalid_argument(format!("{field} exceeds u16")))
1361}
1362
1363fn native_u16s(values: Vec<u32>, field: &str) -> Result<Vec<u16>, Status> {
1364 values
1365 .into_iter()
1366 .map(|value| native_u16(value, field))
1367 .collect()
1368}
1369
1370#[tonic::async_trait]
1371impl native::admin_service_server::AdminService for NativeRuntime {
1372 async fn execute_admin(
1373 &self,
1374 request: Request<native::ExecuteAdminRequest>,
1375 ) -> Result<Response<native::Empty>, Status> {
1376 let request = request.into_inner();
1377 validate_native_context(request.context.as_ref())?;
1378 let (_, entry) = self.session(&request.session_id, "admin")?;
1379 let sql = String::from_utf8(request.command)
1380 .map_err(|_| Status::invalid_argument("admin command must be UTF-8 SQL"))?;
1381 entry.session().run(&sql).await.map_err(query_status)?;
1382 Ok(Response::new(native::Empty {}))
1383 }
1384}
1385
1386#[tonic::async_trait]
1387impl native::health_service_server::HealthService for NativeRuntime {
1388 async fn status(
1389 &self,
1390 request: Request<native::HealthRequest>,
1391 ) -> Result<Response<native::HealthResponse>, Status> {
1392 validate_native_context(request.get_ref().context.as_ref())?;
1393 Ok(Response::new(native::HealthResponse {
1394 serving: self.db.lifecycle_state() == mongreldb_core::LifecycleState::Open,
1395 detail: "ready".into(),
1396 }))
1397 }
1398}
1399
1400async fn transaction_sql(
1401 runtime: &NativeRuntime,
1402 request: native::TransactionRequest,
1403 sql: &str,
1404) -> Result<Response<native::Empty>, Status> {
1405 validate_native_context(request.context.as_ref())?;
1406 let (_, entry) = runtime.session(&request.session_id, "transaction")?;
1407 entry.session().run(sql).await.map_err(query_status)?;
1408 Ok(Response::new(native::Empty {}))
1409}
1410
1411fn query_id(bytes: &[u8]) -> Result<QueryId, Status> {
1412 QueryId::from_str(&id_hex(bytes, "query id")?)
1413 .map_err(|_| Status::invalid_argument("query id must be 16 bytes"))
1414}
1415
1416fn id_hex(bytes: &[u8], label: &str) -> Result<String, Status> {
1417 if bytes.len() != 16 {
1418 return Err(Status::invalid_argument(format!(
1419 "{label} must be 16 bytes"
1420 )));
1421 }
1422 Ok(bytes.iter().map(|byte| format!("{byte:02x}")).collect())
1423}
1424
1425fn hex_id(text: &str) -> Result<Vec<u8>, Status> {
1426 if text.len() != 32 {
1427 return Err(Status::internal("invalid server session id"));
1428 }
1429 (0..16)
1430 .map(|index| {
1431 u8::from_str_radix(&text[index * 2..index * 2 + 2], 16)
1432 .map_err(|_| Status::internal("invalid server session id"))
1433 })
1434 .collect()
1435}
1436
1437fn request_timeout(context: Option<&native::RequestContext>) -> Result<Option<Duration>, Status> {
1438 let deadline = context.map_or(0, |context| context.deadline_unix_micros);
1439 if deadline == 0 {
1440 return Ok(None);
1441 }
1442 let remaining = deadline.saturating_sub(now_unix_micros());
1443 if remaining == 0 {
1444 return Err(structured_status(
1445 Code::DeadlineExceeded,
1446 ErrorCategory::DeadlineExceeded,
1447 "request deadline exceeded",
1448 ));
1449 }
1450 Ok(Some(Duration::from_micros(remaining)))
1451}
1452
1453fn now_unix_micros() -> u64 {
1454 SystemTime::now()
1455 .duration_since(UNIX_EPOCH)
1456 .unwrap_or_default()
1457 .as_micros()
1458 .min(u128::from(u64::MAX)) as u64
1459}
1460
1461fn now_unix_seconds() -> u64 {
1462 SystemTime::now()
1463 .duration_since(UNIX_EPOCH)
1464 .unwrap_or_default()
1465 .as_secs()
1466}
1467
1468fn native_identity(principal: Option<&Principal>) -> native::AuthenticatedIdentity {
1469 match principal {
1470 Some(principal) => native::AuthenticatedIdentity {
1471 principal_id: principal.user_id,
1472 principal_name: principal.username.clone(),
1473 roles: principal.roles.clone(),
1474 scopes: principal
1475 .permissions
1476 .iter()
1477 .map(|permission| format!("{permission:?}"))
1478 .collect(),
1479 },
1480 None => native::AuthenticatedIdentity {
1481 principal_id: 0,
1482 principal_name: "anonymous".into(),
1483 roles: Vec::new(),
1484 scopes: Vec::new(),
1485 },
1486 }
1487}
1488
1489fn native_external_identity(
1490 principal: &Principal,
1491 label: String,
1492 scopes: Vec<String>,
1493) -> native::AuthenticatedIdentity {
1494 native::AuthenticatedIdentity {
1495 principal_id: principal.user_id,
1496 principal_name: label,
1497 roles: principal.roles.clone(),
1498 scopes,
1499 }
1500}
1501
1502fn parameter_literal(
1503 value: &mongreldb_protocol::request::ParameterValue,
1504) -> Result<String, Status> {
1505 use mongreldb_protocol::request::ParameterValue;
1506 match value {
1507 ParameterValue::Null => Ok("NULL".into()),
1508 ParameterValue::Bool(value) => Ok(if *value { "TRUE" } else { "FALSE" }.into()),
1509 ParameterValue::Integer(value) => Ok(value.to_string()),
1510 ParameterValue::Float(value) if value.is_finite() => Ok(value.to_string()),
1511 ParameterValue::Float(_) => Err(Status::invalid_argument("non-finite float parameter")),
1512 ParameterValue::Text(value) => Ok(format!("'{}'", value.replace('\'', "''"))),
1513 ParameterValue::Bytes(value) => Ok(format!(
1514 "X'{}'",
1515 value
1516 .iter()
1517 .map(|byte| format!("{byte:02x}"))
1518 .collect::<String>()
1519 )),
1520 }
1521}
1522
1523fn bind_numbered_parameters(sql: &str, literals: &[String]) -> Result<String, Status> {
1524 #[derive(Clone, Copy)]
1525 enum State {
1526 Normal,
1527 Quote(char),
1528 LineComment,
1529 BlockComment,
1530 }
1531
1532 let mut output = String::with_capacity(sql.len());
1533 let mut used = vec![false; literals.len()];
1534 let mut state = State::Normal;
1535 let mut chars = sql.chars().peekable();
1536 while let Some(character) = chars.next() {
1537 output.push(character);
1538 match state {
1539 State::Normal if matches!(character, '\'' | '"' | '`') => {
1540 state = State::Quote(character);
1541 }
1542 State::Normal if character == '-' && chars.peek() == Some(&'-') => {
1543 output.push(chars.next().expect("peeked"));
1544 state = State::LineComment;
1545 }
1546 State::Normal if character == '/' && chars.peek() == Some(&'*') => {
1547 output.push(chars.next().expect("peeked"));
1548 state = State::BlockComment;
1549 }
1550 State::Normal if character == '$' && chars.peek().is_some_and(char::is_ascii_digit) => {
1551 output.pop();
1552 let mut digits = String::new();
1553 while chars.peek().is_some_and(char::is_ascii_digit) {
1554 digits.push(chars.next().expect("peeked"));
1555 }
1556 let index = digits
1557 .parse::<usize>()
1558 .ok()
1559 .and_then(|index| index.checked_sub(1))
1560 .filter(|index| *index < literals.len())
1561 .ok_or_else(|| {
1562 Status::invalid_argument("prepared parameter is out of range")
1563 })?;
1564 used[index] = true;
1565 output.push_str(&literals[index]);
1566 }
1567 State::Quote(_) if character == '\\' => {
1568 if let Some(escaped) = chars.next() {
1569 output.push(escaped);
1570 }
1571 }
1572 State::Quote(end) if character == end && chars.peek() == Some(&end) => {
1573 output.push(chars.next().expect("peeked"));
1574 }
1575 State::Quote(end) if character == end => state = State::Normal,
1576 State::LineComment if character == '\n' => state = State::Normal,
1577 State::BlockComment if character == '*' && chars.peek() == Some(&'/') => {
1578 output.push(chars.next().expect("peeked"));
1579 state = State::Normal;
1580 }
1581 _ => {}
1582 }
1583 }
1584 if used.iter().any(|used| !used) {
1585 return Err(Status::invalid_argument(
1586 "prepared parameter count does not match SQL placeholders",
1587 ));
1588 }
1589 Ok(output)
1590}
1591
1592fn encode_batch(
1593 batch: &arrow::record_batch::RecordBatch,
1594 sequence: u64,
1595 end_of_stream: bool,
1596) -> Result<native::ArrowFrame, Status> {
1597 let mut ipc = Vec::new();
1598 {
1599 let mut writer = StreamWriter::try_new(&mut ipc, &batch.schema())
1600 .map_err(|error| Status::internal(error.to_string()))?;
1601 writer
1602 .write(batch)
1603 .and_then(|_| writer.finish())
1604 .map_err(|error| Status::internal(error.to_string()))?;
1605 }
1606 Ok(native::ArrowFrame {
1607 ipc,
1608 sequence,
1609 end_of_stream,
1610 })
1611}
1612
1613fn core_status(error: MongrelError) -> Status {
1614 let category = error.category();
1615 let code = match category {
1616 ErrorCategory::Unauthenticated => Code::Unauthenticated,
1617 ErrorCategory::PermissionDenied => Code::PermissionDenied,
1618 ErrorCategory::DeadlineExceeded => Code::DeadlineExceeded,
1619 ErrorCategory::ResourceExhausted => Code::ResourceExhausted,
1620 ErrorCategory::StaleMetadata
1621 | ErrorCategory::SchemaVersionMismatch
1622 | ErrorCategory::ClusterVersionMismatch => Code::FailedPrecondition,
1623 ErrorCategory::TransactionConflict
1624 | ErrorCategory::SerializationFailure
1625 | ErrorCategory::Deadlock => Code::Aborted,
1626 _ => Code::Internal,
1627 };
1628 structured_status(code, category, &error.to_string())
1629}
1630
1631fn query_status(error: MongrelQueryError) -> Status {
1632 query_status_ref(&error)
1633}
1634
1635fn query_status_ref(error: &MongrelQueryError) -> Status {
1636 match error {
1637 MongrelQueryError::Core(error) => {
1638 let category = error.category();
1639 let code = match category {
1640 ErrorCategory::Unauthenticated => Code::Unauthenticated,
1641 ErrorCategory::PermissionDenied => Code::PermissionDenied,
1642 ErrorCategory::DeadlineExceeded => Code::DeadlineExceeded,
1643 ErrorCategory::ResourceExhausted => Code::ResourceExhausted,
1644 ErrorCategory::StaleMetadata
1645 | ErrorCategory::SchemaVersionMismatch
1646 | ErrorCategory::ClusterVersionMismatch => Code::FailedPrecondition,
1647 ErrorCategory::TransactionConflict
1648 | ErrorCategory::SerializationFailure
1649 | ErrorCategory::Deadlock => Code::Aborted,
1650 _ => Code::Internal,
1651 };
1652 structured_status(code, category, &error.to_string())
1653 }
1654 MongrelQueryError::DeadlineExceeded { .. } => structured_status(
1655 Code::DeadlineExceeded,
1656 ErrorCategory::DeadlineExceeded,
1657 &error.to_string(),
1658 ),
1659 MongrelQueryError::QueryCancelled { .. } => structured_status(
1660 Code::Cancelled,
1661 ErrorCategory::Cancelled,
1662 &error.to_string(),
1663 ),
1664 MongrelQueryError::QueryRegistryFull | MongrelQueryError::ResultLimitExceeded { .. } => {
1665 structured_status(
1666 Code::ResourceExhausted,
1667 ErrorCategory::ResourceExhausted,
1668 &error.to_string(),
1669 )
1670 }
1671 MongrelQueryError::TransactionAborted => structured_status(
1672 Code::Aborted,
1673 ErrorCategory::TransactionAborted,
1674 &error.to_string(),
1675 ),
1676 MongrelQueryError::OutcomeUnknown { .. } => structured_status(
1677 Code::Unknown,
1678 ErrorCategory::CommitOutcomeUnknown,
1679 &error.to_string(),
1680 ),
1681 error => Status::internal(error.to_string()),
1682 }
1683}
1684
1685fn stream_status(error: impl std::error::Error + 'static) -> Status {
1686 let message = error.to_string();
1687 let mut source = Some(&error as &(dyn std::error::Error + 'static));
1688 while let Some(current) = source {
1689 if let Some(error) = current.downcast_ref::<MongrelQueryError>() {
1690 return query_status_ref(error);
1691 }
1692 source = current.source();
1693 }
1694 Status::internal(message)
1695}
1696
1697fn structured_status(code: Code, category: ErrorCategory, message: &str) -> Status {
1698 let detail = native::ErrorDetail {
1699 category_code: category.code(),
1700 category: category.to_string(),
1701 message: message.into(),
1702 retryable: category.retry_class() != RetryClass::Never,
1703 metadata: HashMap::new(),
1704 };
1705 Status::with_details(code, message, detail.encode_to_vec().into())
1706}
1707
1708fn native_phase(phase: SqlQueryPhase) -> native::QueryPhase {
1709 match phase {
1710 SqlQueryPhase::Queued => native::QueryPhase::Queued,
1711 SqlQueryPhase::Planning => native::QueryPhase::Planning,
1712 SqlQueryPhase::Executing
1713 | SqlQueryPhase::Streaming
1714 | SqlQueryPhase::CommitCritical
1715 | SqlQueryPhase::Cancelling => native::QueryPhase::Executing,
1716 SqlQueryPhase::Serializing => native::QueryPhase::Serializing,
1717 SqlQueryPhase::Completed => native::QueryPhase::Completed,
1718 SqlQueryPhase::Failed => native::QueryPhase::Failed,
1719 SqlQueryPhase::Cancelled => native::QueryPhase::Cancelled,
1720 }
1721}