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 Ok(Response::new(native::GetSchemaResponse {
1060 table: request.table,
1061 schema_version: schema.schema_id,
1062 columns: schema
1063 .columns
1064 .into_iter()
1065 .map(|column| native::ColumnSchema {
1066 name: column.name,
1067 data_type: format!("{:?}", column.ty),
1068 nullable: column.flags.contains(mongreldb_core::ColumnFlags::NULLABLE),
1069 })
1070 .collect(),
1071 }))
1072 }
1073
1074 async fn create_table(
1075 &self,
1076 request: Request<native::CreateTableRequest>,
1077 ) -> Result<Response<native::CreateTableResponse>, Status> {
1078 let request = request.into_inner();
1079 validate_native_context(request.context.as_ref())?;
1080 let (_, entry) = self.session(&request.session_id, "catalog:write")?;
1081 let _guard = entry.lock.lock().await;
1082 let principal = self.session_principal(&entry)?;
1083 self.db
1084 .require_for(principal.as_ref(), &mongreldb_core::Permission::Ddl)
1085 .map_err(core_status)?;
1086 let columns = request
1087 .columns
1088 .into_iter()
1089 .map(native_create_column)
1090 .collect::<Result<Vec<_>, _>>()?;
1091 let uniques = request
1092 .uniques
1093 .into_iter()
1094 .map(|constraint| {
1095 Ok(mongreldb_core::constraint::UniqueConstraint {
1096 id: native_u16(constraint.id, "unique constraint id")?,
1097 name: constraint.name,
1098 columns: native_u16s(constraint.columns, "unique constraint column")?,
1099 })
1100 })
1101 .collect::<Result<Vec<_>, Status>>()?;
1102 let foreign_keys = request
1103 .foreign_keys
1104 .into_iter()
1105 .map(native_foreign_key)
1106 .collect::<Result<Vec<_>, _>>()?;
1107 let mut schema = mongreldb_core::Schema {
1108 schema_id: request.schema_id,
1109 columns,
1110 indexes: Vec::new(),
1111 colocation: Vec::new(),
1112 constraints: mongreldb_core::constraint::TableConstraints {
1113 uniques,
1114 foreign_keys,
1115 checks: Vec::new(),
1116 },
1117 clustered: false,
1118 };
1119 if let Ok(existing) = self.db.table(&request.table) {
1120 let existing = existing.lock().schema().clone();
1121 schema.schema_id = existing.schema_id;
1122 if serde_json::to_vec(&schema).map_err(|_| Status::internal("schema encode failed"))?
1123 != serde_json::to_vec(&existing)
1124 .map_err(|_| Status::internal("schema encode failed"))?
1125 {
1126 return Err(Status::already_exists(
1127 "table exists with a different schema",
1128 ));
1129 }
1130 entry
1131 .session()
1132 .refresh_database_table(&request.table)
1133 .map_err(query_status)?;
1134 return Ok(Response::new(native::CreateTableResponse {
1135 table_id: self.db.table_id(&request.table).map_err(core_status)?,
1136 schema_version: self.db.catalog_version(),
1137 }));
1138 }
1139 let table_id = self
1140 .db
1141 .create_table(&request.table, schema)
1142 .map_err(core_status)?;
1143 entry
1144 .session()
1145 .refresh_database_table(&request.table)
1146 .map_err(query_status)?;
1147 Ok(Response::new(native::CreateTableResponse {
1148 table_id,
1149 schema_version: self.db.catalog_version(),
1150 }))
1151 }
1152}
1153
1154fn native_create_column(column: native::CreateColumn) -> Result<mongreldb_core::ColumnDef, Status> {
1155 let ty = match native::ColumnType::try_from(column.data_type)
1156 .map_err(|_| Status::invalid_argument("unknown native column type"))?
1157 {
1158 native::ColumnType::Bool => mongreldb_core::TypeId::Bool,
1159 native::ColumnType::Int8 => mongreldb_core::TypeId::Int8,
1160 native::ColumnType::Int16 => mongreldb_core::TypeId::Int16,
1161 native::ColumnType::Int32 => mongreldb_core::TypeId::Int32,
1162 native::ColumnType::Int64 => mongreldb_core::TypeId::Int64,
1163 native::ColumnType::Uint8 => mongreldb_core::TypeId::UInt8,
1164 native::ColumnType::Uint16 => mongreldb_core::TypeId::UInt16,
1165 native::ColumnType::Uint32 => mongreldb_core::TypeId::UInt32,
1166 native::ColumnType::Uint64 => mongreldb_core::TypeId::UInt64,
1167 native::ColumnType::Float32 => mongreldb_core::TypeId::Float32,
1168 native::ColumnType::Float64 => mongreldb_core::TypeId::Float64,
1169 native::ColumnType::TimestampNanos => mongreldb_core::TypeId::TimestampNanos,
1170 native::ColumnType::Date32 => mongreldb_core::TypeId::Date32,
1171 native::ColumnType::Date64 => mongreldb_core::TypeId::Date64,
1172 native::ColumnType::Time64 => mongreldb_core::TypeId::Time64,
1173 native::ColumnType::Bytes => mongreldb_core::TypeId::Bytes,
1174 native::ColumnType::Json => mongreldb_core::TypeId::Json,
1175 native::ColumnType::Decimal128 => mongreldb_core::TypeId::Decimal128 {
1176 precision: u8::try_from(column.decimal_precision)
1177 .map_err(|_| Status::invalid_argument("decimal precision exceeds u8"))?,
1178 scale: i8::try_from(column.decimal_scale)
1179 .map_err(|_| Status::invalid_argument("decimal scale exceeds i8"))?,
1180 },
1181 native::ColumnType::Unspecified => {
1182 return Err(Status::invalid_argument("native column type is required"))
1183 }
1184 };
1185 let mut flags = mongreldb_core::ColumnFlags::empty();
1186 if column.nullable {
1187 flags = flags.with(mongreldb_core::ColumnFlags::NULLABLE);
1188 }
1189 if column.primary_key {
1190 flags = flags.with(mongreldb_core::ColumnFlags::PRIMARY_KEY);
1191 }
1192 if column.auto_increment {
1193 flags = flags.with(mongreldb_core::ColumnFlags::AUTO_INCREMENT);
1194 }
1195 Ok(mongreldb_core::ColumnDef {
1196 id: native_u16(column.id, "column id")?,
1197 name: column.name,
1198 ty,
1199 flags,
1200 default_value: None,
1201 embedding_source: None,
1202 })
1203}
1204
1205fn native_foreign_key(
1206 foreign_key: native::ForeignKey,
1207) -> Result<mongreldb_core::constraint::ForeignKey, Status> {
1208 let action = |value| -> Result<mongreldb_core::constraint::FkAction, Status> {
1209 match native::ForeignKeyAction::try_from(value)
1210 .map_err(|_| Status::invalid_argument("unknown foreign-key action"))?
1211 {
1212 native::ForeignKeyAction::Unspecified | native::ForeignKeyAction::Restrict => {
1213 Ok(mongreldb_core::constraint::FkAction::Restrict)
1214 }
1215 native::ForeignKeyAction::Cascade => Ok(mongreldb_core::constraint::FkAction::Cascade),
1216 native::ForeignKeyAction::SetNull => Ok(mongreldb_core::constraint::FkAction::SetNull),
1217 }
1218 };
1219 Ok(mongreldb_core::constraint::ForeignKey {
1220 id: native_u16(foreign_key.id, "foreign-key id")?,
1221 name: foreign_key.name,
1222 columns: native_u16s(foreign_key.columns, "foreign-key column")?,
1223 ref_table: foreign_key.referenced_table,
1224 ref_columns: native_u16s(foreign_key.referenced_columns, "referenced column")?,
1225 on_delete: action(foreign_key.on_delete)?,
1226 on_update: action(foreign_key.on_update)?,
1227 })
1228}
1229
1230fn native_u16(value: u32, field: &str) -> Result<u16, Status> {
1231 u16::try_from(value).map_err(|_| Status::invalid_argument(format!("{field} exceeds u16")))
1232}
1233
1234fn native_u16s(values: Vec<u32>, field: &str) -> Result<Vec<u16>, Status> {
1235 values
1236 .into_iter()
1237 .map(|value| native_u16(value, field))
1238 .collect()
1239}
1240
1241#[tonic::async_trait]
1242impl native::admin_service_server::AdminService for NativeRuntime {
1243 async fn execute_admin(
1244 &self,
1245 request: Request<native::ExecuteAdminRequest>,
1246 ) -> Result<Response<native::Empty>, Status> {
1247 let request = request.into_inner();
1248 validate_native_context(request.context.as_ref())?;
1249 let (_, entry) = self.session(&request.session_id, "admin")?;
1250 let sql = String::from_utf8(request.command)
1251 .map_err(|_| Status::invalid_argument("admin command must be UTF-8 SQL"))?;
1252 entry.session().run(&sql).await.map_err(query_status)?;
1253 Ok(Response::new(native::Empty {}))
1254 }
1255}
1256
1257#[tonic::async_trait]
1258impl native::health_service_server::HealthService for NativeRuntime {
1259 async fn status(
1260 &self,
1261 request: Request<native::HealthRequest>,
1262 ) -> Result<Response<native::HealthResponse>, Status> {
1263 validate_native_context(request.get_ref().context.as_ref())?;
1264 Ok(Response::new(native::HealthResponse {
1265 serving: self.db.lifecycle_state() == mongreldb_core::LifecycleState::Open,
1266 detail: "ready".into(),
1267 }))
1268 }
1269}
1270
1271async fn transaction_sql(
1272 runtime: &NativeRuntime,
1273 request: native::TransactionRequest,
1274 sql: &str,
1275) -> Result<Response<native::Empty>, Status> {
1276 validate_native_context(request.context.as_ref())?;
1277 let (_, entry) = runtime.session(&request.session_id, "transaction")?;
1278 entry.session().run(sql).await.map_err(query_status)?;
1279 Ok(Response::new(native::Empty {}))
1280}
1281
1282fn query_id(bytes: &[u8]) -> Result<QueryId, Status> {
1283 QueryId::from_str(&id_hex(bytes, "query id")?)
1284 .map_err(|_| Status::invalid_argument("query id must be 16 bytes"))
1285}
1286
1287fn id_hex(bytes: &[u8], label: &str) -> Result<String, Status> {
1288 if bytes.len() != 16 {
1289 return Err(Status::invalid_argument(format!(
1290 "{label} must be 16 bytes"
1291 )));
1292 }
1293 Ok(bytes.iter().map(|byte| format!("{byte:02x}")).collect())
1294}
1295
1296fn hex_id(text: &str) -> Result<Vec<u8>, Status> {
1297 if text.len() != 32 {
1298 return Err(Status::internal("invalid server session id"));
1299 }
1300 (0..16)
1301 .map(|index| {
1302 u8::from_str_radix(&text[index * 2..index * 2 + 2], 16)
1303 .map_err(|_| Status::internal("invalid server session id"))
1304 })
1305 .collect()
1306}
1307
1308fn request_timeout(context: Option<&native::RequestContext>) -> Result<Option<Duration>, Status> {
1309 let deadline = context.map_or(0, |context| context.deadline_unix_micros);
1310 if deadline == 0 {
1311 return Ok(None);
1312 }
1313 let remaining = deadline.saturating_sub(now_unix_micros());
1314 if remaining == 0 {
1315 return Err(structured_status(
1316 Code::DeadlineExceeded,
1317 ErrorCategory::DeadlineExceeded,
1318 "request deadline exceeded",
1319 ));
1320 }
1321 Ok(Some(Duration::from_micros(remaining)))
1322}
1323
1324fn now_unix_micros() -> u64 {
1325 SystemTime::now()
1326 .duration_since(UNIX_EPOCH)
1327 .unwrap_or_default()
1328 .as_micros()
1329 .min(u128::from(u64::MAX)) as u64
1330}
1331
1332fn now_unix_seconds() -> u64 {
1333 SystemTime::now()
1334 .duration_since(UNIX_EPOCH)
1335 .unwrap_or_default()
1336 .as_secs()
1337}
1338
1339fn native_identity(principal: Option<&Principal>) -> native::AuthenticatedIdentity {
1340 match principal {
1341 Some(principal) => native::AuthenticatedIdentity {
1342 principal_id: principal.user_id,
1343 principal_name: principal.username.clone(),
1344 roles: principal.roles.clone(),
1345 scopes: principal
1346 .permissions
1347 .iter()
1348 .map(|permission| format!("{permission:?}"))
1349 .collect(),
1350 },
1351 None => native::AuthenticatedIdentity {
1352 principal_id: 0,
1353 principal_name: "anonymous".into(),
1354 roles: Vec::new(),
1355 scopes: Vec::new(),
1356 },
1357 }
1358}
1359
1360fn native_external_identity(
1361 principal: &Principal,
1362 label: String,
1363 scopes: Vec<String>,
1364) -> native::AuthenticatedIdentity {
1365 native::AuthenticatedIdentity {
1366 principal_id: principal.user_id,
1367 principal_name: label,
1368 roles: principal.roles.clone(),
1369 scopes,
1370 }
1371}
1372
1373fn parameter_literal(
1374 value: &mongreldb_protocol::request::ParameterValue,
1375) -> Result<String, Status> {
1376 use mongreldb_protocol::request::ParameterValue;
1377 match value {
1378 ParameterValue::Null => Ok("NULL".into()),
1379 ParameterValue::Bool(value) => Ok(if *value { "TRUE" } else { "FALSE" }.into()),
1380 ParameterValue::Integer(value) => Ok(value.to_string()),
1381 ParameterValue::Float(value) if value.is_finite() => Ok(value.to_string()),
1382 ParameterValue::Float(_) => Err(Status::invalid_argument("non-finite float parameter")),
1383 ParameterValue::Text(value) => Ok(format!("'{}'", value.replace('\'', "''"))),
1384 ParameterValue::Bytes(value) => Ok(format!(
1385 "X'{}'",
1386 value
1387 .iter()
1388 .map(|byte| format!("{byte:02x}"))
1389 .collect::<String>()
1390 )),
1391 }
1392}
1393
1394fn bind_numbered_parameters(sql: &str, literals: &[String]) -> Result<String, Status> {
1395 #[derive(Clone, Copy)]
1396 enum State {
1397 Normal,
1398 Quote(char),
1399 LineComment,
1400 BlockComment,
1401 }
1402
1403 let mut output = String::with_capacity(sql.len());
1404 let mut used = vec![false; literals.len()];
1405 let mut state = State::Normal;
1406 let mut chars = sql.chars().peekable();
1407 while let Some(character) = chars.next() {
1408 output.push(character);
1409 match state {
1410 State::Normal if matches!(character, '\'' | '"' | '`') => {
1411 state = State::Quote(character);
1412 }
1413 State::Normal if character == '-' && chars.peek() == Some(&'-') => {
1414 output.push(chars.next().expect("peeked"));
1415 state = State::LineComment;
1416 }
1417 State::Normal if character == '/' && chars.peek() == Some(&'*') => {
1418 output.push(chars.next().expect("peeked"));
1419 state = State::BlockComment;
1420 }
1421 State::Normal if character == '$' && chars.peek().is_some_and(char::is_ascii_digit) => {
1422 output.pop();
1423 let mut digits = String::new();
1424 while chars.peek().is_some_and(char::is_ascii_digit) {
1425 digits.push(chars.next().expect("peeked"));
1426 }
1427 let index = digits
1428 .parse::<usize>()
1429 .ok()
1430 .and_then(|index| index.checked_sub(1))
1431 .filter(|index| *index < literals.len())
1432 .ok_or_else(|| {
1433 Status::invalid_argument("prepared parameter is out of range")
1434 })?;
1435 used[index] = true;
1436 output.push_str(&literals[index]);
1437 }
1438 State::Quote(_) if character == '\\' => {
1439 if let Some(escaped) = chars.next() {
1440 output.push(escaped);
1441 }
1442 }
1443 State::Quote(end) if character == end && chars.peek() == Some(&end) => {
1444 output.push(chars.next().expect("peeked"));
1445 }
1446 State::Quote(end) if character == end => state = State::Normal,
1447 State::LineComment if character == '\n' => state = State::Normal,
1448 State::BlockComment if character == '*' && chars.peek() == Some(&'/') => {
1449 output.push(chars.next().expect("peeked"));
1450 state = State::Normal;
1451 }
1452 _ => {}
1453 }
1454 }
1455 if used.iter().any(|used| !used) {
1456 return Err(Status::invalid_argument(
1457 "prepared parameter count does not match SQL placeholders",
1458 ));
1459 }
1460 Ok(output)
1461}
1462
1463fn encode_batch(
1464 batch: &arrow::record_batch::RecordBatch,
1465 sequence: u64,
1466 end_of_stream: bool,
1467) -> Result<native::ArrowFrame, Status> {
1468 let mut ipc = Vec::new();
1469 {
1470 let mut writer = StreamWriter::try_new(&mut ipc, &batch.schema())
1471 .map_err(|error| Status::internal(error.to_string()))?;
1472 writer
1473 .write(batch)
1474 .and_then(|_| writer.finish())
1475 .map_err(|error| Status::internal(error.to_string()))?;
1476 }
1477 Ok(native::ArrowFrame {
1478 ipc,
1479 sequence,
1480 end_of_stream,
1481 })
1482}
1483
1484fn core_status(error: MongrelError) -> Status {
1485 let category = error.category();
1486 let code = match category {
1487 ErrorCategory::Unauthenticated => Code::Unauthenticated,
1488 ErrorCategory::PermissionDenied => Code::PermissionDenied,
1489 ErrorCategory::DeadlineExceeded => Code::DeadlineExceeded,
1490 ErrorCategory::ResourceExhausted => Code::ResourceExhausted,
1491 ErrorCategory::StaleMetadata
1492 | ErrorCategory::SchemaVersionMismatch
1493 | ErrorCategory::ClusterVersionMismatch => Code::FailedPrecondition,
1494 ErrorCategory::TransactionConflict
1495 | ErrorCategory::SerializationFailure
1496 | ErrorCategory::Deadlock => Code::Aborted,
1497 _ => Code::Internal,
1498 };
1499 structured_status(code, category, &error.to_string())
1500}
1501
1502fn query_status(error: MongrelQueryError) -> Status {
1503 query_status_ref(&error)
1504}
1505
1506fn query_status_ref(error: &MongrelQueryError) -> Status {
1507 match error {
1508 MongrelQueryError::Core(error) => {
1509 let category = error.category();
1510 let code = match category {
1511 ErrorCategory::Unauthenticated => Code::Unauthenticated,
1512 ErrorCategory::PermissionDenied => Code::PermissionDenied,
1513 ErrorCategory::DeadlineExceeded => Code::DeadlineExceeded,
1514 ErrorCategory::ResourceExhausted => Code::ResourceExhausted,
1515 ErrorCategory::StaleMetadata
1516 | ErrorCategory::SchemaVersionMismatch
1517 | ErrorCategory::ClusterVersionMismatch => Code::FailedPrecondition,
1518 ErrorCategory::TransactionConflict
1519 | ErrorCategory::SerializationFailure
1520 | ErrorCategory::Deadlock => Code::Aborted,
1521 _ => Code::Internal,
1522 };
1523 structured_status(code, category, &error.to_string())
1524 }
1525 MongrelQueryError::DeadlineExceeded { .. } => structured_status(
1526 Code::DeadlineExceeded,
1527 ErrorCategory::DeadlineExceeded,
1528 &error.to_string(),
1529 ),
1530 MongrelQueryError::QueryCancelled { .. } => structured_status(
1531 Code::Cancelled,
1532 ErrorCategory::Cancelled,
1533 &error.to_string(),
1534 ),
1535 MongrelQueryError::QueryRegistryFull | MongrelQueryError::ResultLimitExceeded { .. } => {
1536 structured_status(
1537 Code::ResourceExhausted,
1538 ErrorCategory::ResourceExhausted,
1539 &error.to_string(),
1540 )
1541 }
1542 MongrelQueryError::TransactionAborted => structured_status(
1543 Code::Aborted,
1544 ErrorCategory::TransactionAborted,
1545 &error.to_string(),
1546 ),
1547 MongrelQueryError::OutcomeUnknown { .. } => structured_status(
1548 Code::Unknown,
1549 ErrorCategory::CommitOutcomeUnknown,
1550 &error.to_string(),
1551 ),
1552 error => Status::internal(error.to_string()),
1553 }
1554}
1555
1556fn stream_status(error: impl std::error::Error + 'static) -> Status {
1557 let message = error.to_string();
1558 let mut source = Some(&error as &(dyn std::error::Error + 'static));
1559 while let Some(current) = source {
1560 if let Some(error) = current.downcast_ref::<MongrelQueryError>() {
1561 return query_status_ref(error);
1562 }
1563 source = current.source();
1564 }
1565 Status::internal(message)
1566}
1567
1568fn structured_status(code: Code, category: ErrorCategory, message: &str) -> Status {
1569 let detail = native::ErrorDetail {
1570 category_code: category.code(),
1571 category: category.to_string(),
1572 message: message.into(),
1573 retryable: category.retry_class() != RetryClass::Never,
1574 metadata: HashMap::new(),
1575 };
1576 Status::with_details(code, message, detail.encode_to_vec().into())
1577}
1578
1579fn native_phase(phase: SqlQueryPhase) -> native::QueryPhase {
1580 match phase {
1581 SqlQueryPhase::Queued => native::QueryPhase::Queued,
1582 SqlQueryPhase::Planning => native::QueryPhase::Planning,
1583 SqlQueryPhase::Executing
1584 | SqlQueryPhase::Streaming
1585 | SqlQueryPhase::CommitCritical
1586 | SqlQueryPhase::Cancelling => native::QueryPhase::Executing,
1587 SqlQueryPhase::Serializing => native::QueryPhase::Serializing,
1588 SqlQueryPhase::Completed => native::QueryPhase::Completed,
1589 SqlQueryPhase::Failed => native::QueryPhase::Failed,
1590 SqlQueryPhase::Cancelled => native::QueryPhase::Cancelled,
1591 }
1592}