use std::{
any::TypeId,
collections::{HashMap, HashSet},
fs,
path::{Path, PathBuf},
sync::Arc,
};
use serde::Deserialize;
const TABLE_TOKEN: &str = "{{table}}";
pub trait StaticLogicalTable: Send + 'static {
const TABLE: &'static str;
const COLUMNS: &'static [&'static str];
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum DatabaseNameMappingError {
DirectoryUnreadable,
InvalidFile,
InvalidLogicalName,
InvalidPhysicalName,
DuplicateLogicalTable,
DuplicatePhysicalTable,
DuplicateLogicalColumn,
DuplicatePhysicalColumn,
MissingTable,
ExtraTable,
MissingColumn,
ExtraColumn,
DuplicateOperation,
InvalidTemplate,
}
impl DatabaseNameMappingError {
pub const fn code(self) -> &'static str {
match self {
Self::DirectoryUnreadable => "db.name_mapping.directory_unreadable",
Self::InvalidFile => "db.name_mapping.invalid_file",
Self::InvalidLogicalName => "db.name_mapping.invalid_logical_name",
Self::InvalidPhysicalName => "db.name_mapping.invalid_physical_name",
Self::DuplicateLogicalTable => "db.name_mapping.duplicate_logical_table",
Self::DuplicatePhysicalTable => "db.name_mapping.duplicate_physical_table",
Self::DuplicateLogicalColumn => "db.name_mapping.duplicate_logical_column",
Self::DuplicatePhysicalColumn => "db.name_mapping.duplicate_physical_column",
Self::MissingTable => "db.name_mapping.missing_table",
Self::ExtraTable => "db.name_mapping.extra_table",
Self::MissingColumn => "db.name_mapping.missing_column",
Self::ExtraColumn => "db.name_mapping.extra_column",
Self::DuplicateOperation => "db.name_mapping.duplicate_operation",
Self::InvalidTemplate => "db.name_mapping.invalid_template",
}
}
}
#[derive(Clone)]
pub(crate) struct MappingStartupConfig {
directory: PathBuf,
tables: Vec<TableDeclaration>,
queries: Vec<OperationDeclaration>,
writes: Vec<OperationDeclaration>,
}
impl MappingStartupConfig {
pub(crate) fn new(directory: PathBuf) -> Self {
Self {
directory,
tables: Vec::new(),
queries: Vec::new(),
writes: Vec::new(),
}
}
pub(crate) fn set_directory(&mut self, directory: PathBuf) {
self.directory = directory;
}
pub(crate) fn register_table<T: StaticLogicalTable>(&mut self) {
self.tables.push(TableDeclaration {
type_id: TypeId::of::<T>(),
logical: T::TABLE,
columns: T::COLUMNS,
});
}
pub(crate) fn register_query<O: 'static>(
&mut self,
operation: &'static str,
table: &'static str,
columns: &'static [&'static str],
template: &'static str,
) {
self.queries.push(OperationDeclaration {
type_id: TypeId::of::<O>(),
operation,
table,
columns,
template,
});
}
pub(crate) fn register_write<O: 'static>(
&mut self,
operation: &'static str,
table: &'static str,
columns: &'static [&'static str],
template: &'static str,
) {
self.writes.push(OperationDeclaration {
type_id: TypeId::of::<O>(),
operation,
table,
columns,
template,
});
}
}
#[derive(Clone)]
struct TableDeclaration {
#[allow(dead_code)]
type_id: TypeId,
logical: &'static str,
columns: &'static [&'static str],
}
#[derive(Clone)]
struct OperationDeclaration {
type_id: TypeId,
operation: &'static str,
table: &'static str,
columns: &'static [&'static str],
template: &'static str,
}
#[derive(Default)]
pub(crate) struct PhysicalOperationPlans {
enabled: bool,
queries: HashMap<TypeId, Arc<str>>,
writes: HashMap<TypeId, Arc<str>>,
}
pub(crate) enum OperationSql {
Legacy,
Mapped(Arc<str>),
Missing,
}
impl PhysicalOperationPlans {
pub(crate) fn disabled() -> Self {
Self::default()
}
pub(crate) fn query<O: 'static>(&self) -> OperationSql {
if self.enabled {
self.queries
.get(&TypeId::of::<O>())
.cloned()
.map_or(OperationSql::Missing, OperationSql::Mapped)
} else {
OperationSql::Legacy
}
}
pub(crate) fn enabled(&self) -> bool {
self.enabled
}
pub(crate) fn write<O: 'static>(&self) -> OperationSql {
if self.enabled {
self.writes
.get(&TypeId::of::<O>())
.cloned()
.map_or(OperationSql::Missing, OperationSql::Mapped)
} else {
OperationSql::Legacy
}
}
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct MappingFile {
table: NamePair,
columns: Vec<NamePair>,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct NamePair {
from: String,
to: String,
}
struct FrozenTable {
physical: String,
columns: HashMap<String, String>,
}
pub(crate) fn freeze(
config: Option<MappingStartupConfig>,
) -> Result<PhysicalOperationPlans, DatabaseNameMappingError> {
let Some(config) = config else {
return Ok(PhysicalOperationPlans::disabled());
};
let mappings = load_directory(&config.directory)?;
validate_complete(&mappings, &config.tables)?;
Ok(PhysicalOperationPlans {
enabled: true,
queries: compile_operations(&mappings, &config.queries)?,
writes: compile_operations(&mappings, &config.writes)?,
})
}
fn load_directory(
directory: &Path,
) -> Result<HashMap<String, FrozenTable>, DatabaseNameMappingError> {
let entries =
fs::read_dir(directory).map_err(|_| DatabaseNameMappingError::DirectoryUnreadable)?;
let mut logical_tables = HashMap::new();
let mut physical_tables = HashSet::new();
for entry in entries {
let entry = entry.map_err(|_| DatabaseNameMappingError::DirectoryUnreadable)?;
let path = entry.path();
if !path.is_file() || path.extension().and_then(|value| value.to_str()) != Some("json") {
return Err(DatabaseNameMappingError::InvalidFile);
}
let bytes = fs::read(path).map_err(|_| DatabaseNameMappingError::InvalidFile)?;
let mapping: MappingFile =
serde_json::from_slice(&bytes).map_err(|_| DatabaseNameMappingError::InvalidFile)?;
validate_logical(&mapping.table.from)?;
validate_physical(&mapping.table.to)?;
if !physical_tables.insert(mapping.table.to.clone()) {
return Err(DatabaseNameMappingError::DuplicatePhysicalTable);
}
let mut columns = HashMap::new();
let mut physical_columns = HashSet::new();
for column in mapping.columns {
validate_logical(&column.from)?;
validate_physical(&column.to)?;
if columns.insert(column.from, column.to.clone()).is_some() {
return Err(DatabaseNameMappingError::DuplicateLogicalColumn);
}
if !physical_columns.insert(column.to) {
return Err(DatabaseNameMappingError::DuplicatePhysicalColumn);
}
}
let table = FrozenTable {
physical: mapping.table.to,
columns,
};
if logical_tables.insert(mapping.table.from, table).is_some() {
return Err(DatabaseNameMappingError::DuplicateLogicalTable);
}
}
Ok(logical_tables)
}
fn validate_complete(
mappings: &HashMap<String, FrozenTable>,
declarations: &[TableDeclaration],
) -> Result<(), DatabaseNameMappingError> {
let mut declared = HashSet::new();
for declaration in declarations {
validate_logical(declaration.logical)?;
if !declared.insert(declaration.logical) {
return Err(DatabaseNameMappingError::DuplicateLogicalTable);
}
let mapping = mappings
.get(declaration.logical)
.ok_or(DatabaseNameMappingError::MissingTable)?;
let expected: HashSet<_> = declaration.columns.iter().copied().collect();
if expected.len() != declaration.columns.len() {
return Err(DatabaseNameMappingError::DuplicateLogicalColumn);
}
if expected
.iter()
.any(|column| !mapping.columns.contains_key(*column))
{
return Err(DatabaseNameMappingError::MissingColumn);
}
if mapping
.columns
.keys()
.any(|column| !expected.contains(column.as_str()))
{
return Err(DatabaseNameMappingError::ExtraColumn);
}
}
if mappings
.keys()
.any(|table| !declared.contains(table.as_str()))
{
return Err(DatabaseNameMappingError::ExtraTable);
}
Ok(())
}
fn compile_operations(
mappings: &HashMap<String, FrozenTable>,
operations: &[OperationDeclaration],
) -> Result<HashMap<TypeId, Arc<str>>, DatabaseNameMappingError> {
let mut output = HashMap::new();
for operation in operations {
if operation.operation.is_empty() {
return Err(DatabaseNameMappingError::InvalidTemplate);
}
let table = mappings
.get(operation.table)
.ok_or(DatabaseNameMappingError::MissingTable)?;
let mut sql = operation
.template
.replace(TABLE_TOKEN, "e_identifier(&table.physical));
if sql == operation.template {
return Err(DatabaseNameMappingError::InvalidTemplate);
}
for logical in operation.columns {
let physical = table
.columns
.get(*logical)
.ok_or(DatabaseNameMappingError::MissingColumn)?;
let token = format!("{{{{column:{logical}}}}}");
let replaced = sql.replace(&token, "e_identifier(physical));
if replaced == sql {
return Err(DatabaseNameMappingError::InvalidTemplate);
}
sql = replaced;
}
if sql.contains("{{") || sql.contains("}}") {
return Err(DatabaseNameMappingError::InvalidTemplate);
}
if output.insert(operation.type_id, Arc::from(sql)).is_some() {
return Err(DatabaseNameMappingError::DuplicateOperation);
}
}
Ok(output)
}
fn validate_logical(value: &str) -> Result<(), DatabaseNameMappingError> {
if value.trim().is_empty() {
Err(DatabaseNameMappingError::InvalidLogicalName)
} else {
Ok(())
}
}
fn validate_physical(value: &str) -> Result<(), DatabaseNameMappingError> {
let mut bytes = value.bytes();
let Some(first) = bytes.next() else {
return Err(DatabaseNameMappingError::InvalidPhysicalName);
};
if !(first.is_ascii_alphabetic() || first == b'_')
|| !bytes.all(|byte| byte.is_ascii_alphanumeric() || byte == b'_')
{
return Err(DatabaseNameMappingError::InvalidPhysicalName);
}
Ok(())
}
fn quote_identifier(value: &str) -> String {
format!("`{value}`")
}
#[cfg(test)]
mod tests {
use std::{
fs,
path::PathBuf,
sync::atomic::{AtomicU64, Ordering},
};
use super::*;
static NEXT: AtomicU64 = AtomicU64::new(0);
struct Orders;
impl StaticLogicalTable for Orders {
const TABLE: &'static str = "订单";
const COLUMNS: &'static [&'static str] = &["订单号", "金额"];
}
struct FindOrder;
struct UpdateOrder;
fn directory() -> PathBuf {
let path = std::env::temp_dir().join(format!(
"saddle-db-name-mapping-{}-{}",
std::process::id(),
NEXT.fetch_add(1, Ordering::Relaxed)
));
fs::create_dir(&path).unwrap();
path
}
fn write_mapping(directory: &Path, value: &str) {
fs::write(directory.join("orders.json"), value).unwrap();
}
fn config(directory: PathBuf) -> MappingStartupConfig {
let mut config = MappingStartupConfig::new(directory);
config.register_table::<Orders>();
config.register_query::<FindOrder>(
"order.find",
"订单",
&["订单号", "金额"],
"SELECT {{column:金额}} FROM {{table}} WHERE {{column:订单号}} = ?",
);
config.register_write::<UpdateOrder>(
"order.update",
"订单",
&["金额", "订单号"],
"UPDATE {{table}} SET {{column:金额}} = ? WHERE {{column:订单号}} = ?",
);
config
}
#[test]
fn complete_mapping_freezes_query_and_write_plans() {
let directory = directory();
write_mapping(
&directory,
r#"{"table":{"from":"订单","to":"t_order_v2"},"columns":[{"from":"订单号","to":"c_order_id"},{"from":"金额","to":"c_amount_v2"}]}"#,
);
let plans = freeze(Some(config(directory.clone()))).unwrap();
match plans.query::<FindOrder>() {
OperationSql::Mapped(sql) => assert_eq!(
sql.as_ref(),
"SELECT `c_amount_v2` FROM `t_order_v2` WHERE `c_order_id` = ?"
),
_ => panic!("mapped query plan missing"),
}
match plans.write::<UpdateOrder>() {
OperationSql::Mapped(sql) => assert_eq!(
sql.as_ref(),
"UPDATE `t_order_v2` SET `c_amount_v2` = ? WHERE `c_order_id` = ?"
),
_ => panic!("mapped write plan missing"),
}
fs::remove_dir_all(directory).unwrap();
}
#[test]
fn incomplete_extra_duplicate_and_unsafe_mappings_fail_closed() {
let cases = [
(
r#"{"table":{"from":"订单","to":"t_order"},"columns":[{"from":"订单号","to":"c_order_id"}]}"#,
DatabaseNameMappingError::MissingColumn,
),
(
r#"{"table":{"from":"订单","to":"t_order"},"columns":[{"from":"订单号","to":"c_order_id"},{"from":"金额","to":"c_amount"},{"from":"额外","to":"c_extra"}]}"#,
DatabaseNameMappingError::ExtraColumn,
),
(
r#"{"table":{"from":"订单","to":"t_order"},"columns":[{"from":"订单号","to":"same"},{"from":"金额","to":"same"}]}"#,
DatabaseNameMappingError::DuplicatePhysicalColumn,
),
(
r#"{"table":{"from":"订单","to":"t_order;drop"},"columns":[{"from":"订单号","to":"c_order_id"},{"from":"金额","to":"c_amount"}]}"#,
DatabaseNameMappingError::InvalidPhysicalName,
),
];
for (value, expected) in cases {
let directory = directory();
write_mapping(&directory, value);
assert_eq!(
freeze(Some(config(directory.clone()))).err(),
Some(expected)
);
fs::remove_dir_all(directory).unwrap();
}
}
#[test]
fn unregistered_operation_has_no_raw_sql_fallback_when_mapping_is_enabled() {
let directory = directory();
write_mapping(
&directory,
r#"{"table":{"from":"订单","to":"t_order"},"columns":[{"from":"订单号","to":"c_order_id"},{"from":"金额","to":"c_amount"}]}"#,
);
let plans = freeze(Some(config(directory.clone()))).unwrap();
struct Unregistered;
assert!(matches!(
plans.query::<Unregistered>(),
OperationSql::Missing
));
assert!(matches!(
plans.write::<Unregistered>(),
OperationSql::Missing
));
fs::remove_dir_all(directory).unwrap();
}
}