use std::sync::{Arc, OnceLock};
use crate::db::route_context::RouteContext;
use crate::migrate::ModelMeta;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct Alias(String);
impl Alias {
pub fn new(s: impl Into<String>) -> Self {
Alias(s.into())
}
pub fn as_str(&self) -> &str {
&self.0
}
pub fn default_alias() -> Self {
Alias("default".to_string())
}
}
impl From<&str> for Alias {
fn from(s: &str) -> Self {
Alias(s.to_string())
}
}
impl From<String> for Alias {
fn from(s: String) -> Self {
Alias(s)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Schema(String);
impl Schema {
pub fn new(s: impl Into<String>) -> Option<Self> {
let s = s.into();
let ok = (1..=63).contains(&s.len())
&& s.chars()
.next()
.is_some_and(|c| c.is_ascii_alphabetic() || c == '_')
&& s.chars().all(|c| c.is_ascii_alphanumeric() || c == '_');
ok.then_some(Schema(s))
}
pub fn as_str(&self) -> &str {
&self.0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RouteOp {
Read,
Write,
}
pub trait DatabaseRouter: Send + Sync {
fn db_for_read(&self, model: &ModelMeta, ctx: &RouteContext) -> Alias {
let _ = ctx;
default_alias_for(model)
}
fn db_for_write(&self, model: &ModelMeta, ctx: &RouteContext) -> Alias {
let _ = ctx;
default_alias_for(model)
}
fn allow_relation(&self, a: &ModelMeta, b: &ModelMeta) -> bool {
default_alias_for(a) == default_alias_for(b)
}
fn allow_migrate(&self, alias: &str, model: &ModelMeta) -> bool {
default_alias_for(model).as_str() == alias
}
fn schema_for(&self, ctx: &RouteContext) -> Option<Schema> {
let _ = ctx;
None
}
fn schema_for_table(&self, ctx: &RouteContext, table: &str) -> Option<Schema> {
let _ = table;
self.schema_for(ctx)
}
}
fn default_alias_for(model: &ModelMeta) -> Alias {
match crate::migrate::model_alias(&model.name) {
Some(a) => Alias::new(a),
None => Alias::default_alias(),
}
}
#[derive(Debug, Default)]
pub struct DefaultRouter;
impl DatabaseRouter for DefaultRouter {}
static ROUTER: OnceLock<Arc<dyn DatabaseRouter>> = OnceLock::new();
static DEFAULT: OnceLock<Arc<dyn DatabaseRouter>> = OnceLock::new();
pub(crate) fn install_router(router: Arc<dyn DatabaseRouter>) {
let _ = ROUTER.set(router);
}
pub fn install_router_from_plugin(router: Arc<dyn DatabaseRouter>) {
if ROUTER.set(router).is_err() {
tracing::warn!(
"umbral::db::router: a DatabaseRouter is already installed (via \
App::builder().router(...) or an earlier plugin); ignoring this \
plugin's router. Install exactly one router — don't combine a \
router-owning plugin with an explicit .router(...) or a second \
such plugin."
);
}
}
fn default_router_arc() -> Arc<dyn DatabaseRouter> {
DEFAULT.get_or_init(|| Arc::new(DefaultRouter)).clone()
}
pub fn router() -> Arc<dyn DatabaseRouter> {
ROUTER.get().cloned().unwrap_or_else(default_router_arc)
}
pub fn schema_qualified_table(table: &str) -> sea_query::TableRef {
use sea_query::{Alias as SqAlias, IntoTableRef};
let ctx = crate::db::route_context::current();
match router().schema_for_table(&ctx, table) {
Some(schema) => (SqAlias::new(schema.as_str()), SqAlias::new(table)).into_table_ref(),
None => SqAlias::new(table).into_table_ref(),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn schema_accepts_valid_identifiers_and_rejects_the_rest() {
assert!(Schema::new("tenant_7").is_some());
assert!(Schema::new("_private").is_some());
assert!(Schema::new("public").is_some());
assert!(Schema::new("").is_none());
assert!(Schema::new("1tenant").is_none());
assert!(Schema::new("a b").is_none());
assert!(Schema::new("drop\";--").is_none());
assert!(Schema::new("a".repeat(64)).is_none());
}
#[test]
fn alias_roundtrips() {
assert_eq!(Alias::from("replica").as_str(), "replica");
assert_eq!(Alias::default_alias().as_str(), "default");
}
#[test]
fn schema_qualified_table_is_bare_under_default_router() {
let sql = sea_query::Query::select()
.column(sea_query::Asterisk)
.from(schema_qualified_table("widget"))
.to_string(sea_query::PostgresQueryBuilder);
assert!(sql.contains("\"widget\""), "got: {sql}");
assert!(
!sql.contains(".\"widget\""),
"unexpected qualification: {sql}"
);
}
}