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