Skip to main content

saddle_db/
name_mapping.rs

1use std::{
2    any::TypeId,
3    collections::{HashMap, HashSet},
4    fs,
5    path::{Path, PathBuf},
6    sync::Arc,
7};
8
9use serde::Deserialize;
10
11const TABLE_TOKEN: &str = "{{table}}";
12
13/// A generated logical table definition used to prove mapping completeness.
14pub trait StaticLogicalTable: Send + 'static {
15    const TABLE: &'static str;
16    const COLUMNS: &'static [&'static str];
17}
18
19#[derive(Clone, Copy, Debug, Eq, PartialEq)]
20pub enum DatabaseNameMappingError {
21    DirectoryUnreadable,
22    InvalidFile,
23    InvalidLogicalName,
24    InvalidPhysicalName,
25    DuplicateLogicalTable,
26    DuplicatePhysicalTable,
27    DuplicateLogicalColumn,
28    DuplicatePhysicalColumn,
29    MissingTable,
30    ExtraTable,
31    MissingColumn,
32    ExtraColumn,
33    DuplicateOperation,
34    InvalidTemplate,
35}
36
37impl DatabaseNameMappingError {
38    /// Stable machine-readable startup error code. Paths and mapping contents
39    /// are deliberately excluded from the public error boundary.
40    pub const fn code(self) -> &'static str {
41        match self {
42            Self::DirectoryUnreadable => "db.name_mapping.directory_unreadable",
43            Self::InvalidFile => "db.name_mapping.invalid_file",
44            Self::InvalidLogicalName => "db.name_mapping.invalid_logical_name",
45            Self::InvalidPhysicalName => "db.name_mapping.invalid_physical_name",
46            Self::DuplicateLogicalTable => "db.name_mapping.duplicate_logical_table",
47            Self::DuplicatePhysicalTable => "db.name_mapping.duplicate_physical_table",
48            Self::DuplicateLogicalColumn => "db.name_mapping.duplicate_logical_column",
49            Self::DuplicatePhysicalColumn => "db.name_mapping.duplicate_physical_column",
50            Self::MissingTable => "db.name_mapping.missing_table",
51            Self::ExtraTable => "db.name_mapping.extra_table",
52            Self::MissingColumn => "db.name_mapping.missing_column",
53            Self::ExtraColumn => "db.name_mapping.extra_column",
54            Self::DuplicateOperation => "db.name_mapping.duplicate_operation",
55            Self::InvalidTemplate => "db.name_mapping.invalid_template",
56        }
57    }
58}
59
60#[derive(Clone)]
61pub(crate) struct MappingStartupConfig {
62    directory: PathBuf,
63    tables: Vec<TableDeclaration>,
64    queries: Vec<OperationDeclaration>,
65    writes: Vec<OperationDeclaration>,
66}
67
68impl MappingStartupConfig {
69    pub(crate) fn new(directory: PathBuf) -> Self {
70        Self {
71            directory,
72            tables: Vec::new(),
73            queries: Vec::new(),
74            writes: Vec::new(),
75        }
76    }
77
78    pub(crate) fn set_directory(&mut self, directory: PathBuf) {
79        self.directory = directory;
80    }
81
82    pub(crate) fn register_table<T: StaticLogicalTable>(&mut self) {
83        self.tables.push(TableDeclaration {
84            type_id: TypeId::of::<T>(),
85            logical: T::TABLE,
86            columns: T::COLUMNS,
87        });
88    }
89
90    pub(crate) fn register_query<O: 'static>(
91        &mut self,
92        operation: &'static str,
93        table: &'static str,
94        columns: &'static [&'static str],
95        template: &'static str,
96    ) {
97        self.queries.push(OperationDeclaration {
98            type_id: TypeId::of::<O>(),
99            operation,
100            table,
101            columns,
102            template,
103        });
104    }
105
106    pub(crate) fn register_write<O: 'static>(
107        &mut self,
108        operation: &'static str,
109        table: &'static str,
110        columns: &'static [&'static str],
111        template: &'static str,
112    ) {
113        self.writes.push(OperationDeclaration {
114            type_id: TypeId::of::<O>(),
115            operation,
116            table,
117            columns,
118            template,
119        });
120    }
121}
122
123#[derive(Clone)]
124struct TableDeclaration {
125    #[allow(dead_code)]
126    type_id: TypeId,
127    logical: &'static str,
128    columns: &'static [&'static str],
129}
130
131#[derive(Clone)]
132struct OperationDeclaration {
133    type_id: TypeId,
134    operation: &'static str,
135    table: &'static str,
136    columns: &'static [&'static str],
137    template: &'static str,
138}
139
140#[derive(Default)]
141pub(crate) struct PhysicalOperationPlans {
142    enabled: bool,
143    queries: HashMap<TypeId, Arc<str>>,
144    writes: HashMap<TypeId, Arc<str>>,
145}
146
147pub(crate) enum OperationSql {
148    Legacy,
149    Mapped(Arc<str>),
150    Missing,
151}
152
153impl PhysicalOperationPlans {
154    pub(crate) fn disabled() -> Self {
155        Self::default()
156    }
157
158    pub(crate) fn query<O: 'static>(&self) -> OperationSql {
159        if self.enabled {
160            self.queries
161                .get(&TypeId::of::<O>())
162                .cloned()
163                .map_or(OperationSql::Missing, OperationSql::Mapped)
164        } else {
165            OperationSql::Legacy
166        }
167    }
168
169    pub(crate) fn enabled(&self) -> bool {
170        self.enabled
171    }
172
173    pub(crate) fn write<O: 'static>(&self) -> OperationSql {
174        if self.enabled {
175            self.writes
176                .get(&TypeId::of::<O>())
177                .cloned()
178                .map_or(OperationSql::Missing, OperationSql::Mapped)
179        } else {
180            OperationSql::Legacy
181        }
182    }
183}
184
185#[derive(Deserialize)]
186#[serde(deny_unknown_fields)]
187struct MappingFile {
188    table: NamePair,
189    columns: Vec<NamePair>,
190}
191
192#[derive(Deserialize)]
193#[serde(deny_unknown_fields)]
194struct NamePair {
195    from: String,
196    to: String,
197}
198
199struct FrozenTable {
200    physical: String,
201    columns: HashMap<String, String>,
202}
203
204pub(crate) fn freeze(
205    config: Option<MappingStartupConfig>,
206) -> Result<PhysicalOperationPlans, DatabaseNameMappingError> {
207    let Some(config) = config else {
208        return Ok(PhysicalOperationPlans::disabled());
209    };
210    let mappings = load_directory(&config.directory)?;
211    validate_complete(&mappings, &config.tables)?;
212    Ok(PhysicalOperationPlans {
213        enabled: true,
214        queries: compile_operations(&mappings, &config.queries)?,
215        writes: compile_operations(&mappings, &config.writes)?,
216    })
217}
218
219fn load_directory(
220    directory: &Path,
221) -> Result<HashMap<String, FrozenTable>, DatabaseNameMappingError> {
222    let entries =
223        fs::read_dir(directory).map_err(|_| DatabaseNameMappingError::DirectoryUnreadable)?;
224    let mut logical_tables = HashMap::new();
225    let mut physical_tables = HashSet::new();
226    for entry in entries {
227        let entry = entry.map_err(|_| DatabaseNameMappingError::DirectoryUnreadable)?;
228        let path = entry.path();
229        if !path.is_file() || path.extension().and_then(|value| value.to_str()) != Some("json") {
230            return Err(DatabaseNameMappingError::InvalidFile);
231        }
232        let bytes = fs::read(path).map_err(|_| DatabaseNameMappingError::InvalidFile)?;
233        let mapping: MappingFile =
234            serde_json::from_slice(&bytes).map_err(|_| DatabaseNameMappingError::InvalidFile)?;
235        validate_logical(&mapping.table.from)?;
236        validate_physical(&mapping.table.to)?;
237        if !physical_tables.insert(mapping.table.to.clone()) {
238            return Err(DatabaseNameMappingError::DuplicatePhysicalTable);
239        }
240        let mut columns = HashMap::new();
241        let mut physical_columns = HashSet::new();
242        for column in mapping.columns {
243            validate_logical(&column.from)?;
244            validate_physical(&column.to)?;
245            if columns.insert(column.from, column.to.clone()).is_some() {
246                return Err(DatabaseNameMappingError::DuplicateLogicalColumn);
247            }
248            if !physical_columns.insert(column.to) {
249                return Err(DatabaseNameMappingError::DuplicatePhysicalColumn);
250            }
251        }
252        let table = FrozenTable {
253            physical: mapping.table.to,
254            columns,
255        };
256        if logical_tables.insert(mapping.table.from, table).is_some() {
257            return Err(DatabaseNameMappingError::DuplicateLogicalTable);
258        }
259    }
260    Ok(logical_tables)
261}
262
263fn validate_complete(
264    mappings: &HashMap<String, FrozenTable>,
265    declarations: &[TableDeclaration],
266) -> Result<(), DatabaseNameMappingError> {
267    let mut declared = HashSet::new();
268    for declaration in declarations {
269        validate_logical(declaration.logical)?;
270        if !declared.insert(declaration.logical) {
271            return Err(DatabaseNameMappingError::DuplicateLogicalTable);
272        }
273        let mapping = mappings
274            .get(declaration.logical)
275            .ok_or(DatabaseNameMappingError::MissingTable)?;
276        let expected: HashSet<_> = declaration.columns.iter().copied().collect();
277        if expected.len() != declaration.columns.len() {
278            return Err(DatabaseNameMappingError::DuplicateLogicalColumn);
279        }
280        if expected
281            .iter()
282            .any(|column| !mapping.columns.contains_key(*column))
283        {
284            return Err(DatabaseNameMappingError::MissingColumn);
285        }
286        if mapping
287            .columns
288            .keys()
289            .any(|column| !expected.contains(column.as_str()))
290        {
291            return Err(DatabaseNameMappingError::ExtraColumn);
292        }
293    }
294    if mappings
295        .keys()
296        .any(|table| !declared.contains(table.as_str()))
297    {
298        return Err(DatabaseNameMappingError::ExtraTable);
299    }
300    Ok(())
301}
302
303fn compile_operations(
304    mappings: &HashMap<String, FrozenTable>,
305    operations: &[OperationDeclaration],
306) -> Result<HashMap<TypeId, Arc<str>>, DatabaseNameMappingError> {
307    let mut output = HashMap::new();
308    for operation in operations {
309        if operation.operation.is_empty() {
310            return Err(DatabaseNameMappingError::InvalidTemplate);
311        }
312        let table = mappings
313            .get(operation.table)
314            .ok_or(DatabaseNameMappingError::MissingTable)?;
315        let mut sql = operation
316            .template
317            .replace(TABLE_TOKEN, &quote_identifier(&table.physical));
318        if sql == operation.template {
319            return Err(DatabaseNameMappingError::InvalidTemplate);
320        }
321        for logical in operation.columns {
322            let physical = table
323                .columns
324                .get(*logical)
325                .ok_or(DatabaseNameMappingError::MissingColumn)?;
326            let token = format!("{{{{column:{logical}}}}}");
327            let replaced = sql.replace(&token, &quote_identifier(physical));
328            if replaced == sql {
329                return Err(DatabaseNameMappingError::InvalidTemplate);
330            }
331            sql = replaced;
332        }
333        if sql.contains("{{") || sql.contains("}}") {
334            return Err(DatabaseNameMappingError::InvalidTemplate);
335        }
336        if output.insert(operation.type_id, Arc::from(sql)).is_some() {
337            return Err(DatabaseNameMappingError::DuplicateOperation);
338        }
339    }
340    Ok(output)
341}
342
343fn validate_logical(value: &str) -> Result<(), DatabaseNameMappingError> {
344    if value.trim().is_empty() {
345        Err(DatabaseNameMappingError::InvalidLogicalName)
346    } else {
347        Ok(())
348    }
349}
350
351fn validate_physical(value: &str) -> Result<(), DatabaseNameMappingError> {
352    let mut bytes = value.bytes();
353    let Some(first) = bytes.next() else {
354        return Err(DatabaseNameMappingError::InvalidPhysicalName);
355    };
356    if !(first.is_ascii_alphabetic() || first == b'_')
357        || !bytes.all(|byte| byte.is_ascii_alphanumeric() || byte == b'_')
358    {
359        return Err(DatabaseNameMappingError::InvalidPhysicalName);
360    }
361    Ok(())
362}
363
364fn quote_identifier(value: &str) -> String {
365    format!("`{value}`")
366}
367
368#[cfg(test)]
369mod tests {
370    use std::{
371        fs,
372        path::PathBuf,
373        sync::atomic::{AtomicU64, Ordering},
374    };
375
376    use super::*;
377
378    static NEXT: AtomicU64 = AtomicU64::new(0);
379
380    struct Orders;
381
382    impl StaticLogicalTable for Orders {
383        const TABLE: &'static str = "订单";
384        const COLUMNS: &'static [&'static str] = &["订单号", "金额"];
385    }
386
387    struct FindOrder;
388    struct UpdateOrder;
389
390    fn directory() -> PathBuf {
391        let path = std::env::temp_dir().join(format!(
392            "saddle-db-name-mapping-{}-{}",
393            std::process::id(),
394            NEXT.fetch_add(1, Ordering::Relaxed)
395        ));
396        fs::create_dir(&path).unwrap();
397        path
398    }
399
400    fn write_mapping(directory: &Path, value: &str) {
401        fs::write(directory.join("orders.json"), value).unwrap();
402    }
403
404    fn config(directory: PathBuf) -> MappingStartupConfig {
405        let mut config = MappingStartupConfig::new(directory);
406        config.register_table::<Orders>();
407        config.register_query::<FindOrder>(
408            "order.find",
409            "订单",
410            &["订单号", "金额"],
411            "SELECT {{column:金额}} FROM {{table}} WHERE {{column:订单号}} = ?",
412        );
413        config.register_write::<UpdateOrder>(
414            "order.update",
415            "订单",
416            &["金额", "订单号"],
417            "UPDATE {{table}} SET {{column:金额}} = ? WHERE {{column:订单号}} = ?",
418        );
419        config
420    }
421
422    #[test]
423    fn complete_mapping_freezes_query_and_write_plans() {
424        let directory = directory();
425        write_mapping(
426            &directory,
427            r#"{"table":{"from":"订单","to":"t_order_v2"},"columns":[{"from":"订单号","to":"c_order_id"},{"from":"金额","to":"c_amount_v2"}]}"#,
428        );
429        let plans = freeze(Some(config(directory.clone()))).unwrap();
430        match plans.query::<FindOrder>() {
431            OperationSql::Mapped(sql) => assert_eq!(
432                sql.as_ref(),
433                "SELECT `c_amount_v2` FROM `t_order_v2` WHERE `c_order_id` = ?"
434            ),
435            _ => panic!("mapped query plan missing"),
436        }
437        match plans.write::<UpdateOrder>() {
438            OperationSql::Mapped(sql) => assert_eq!(
439                sql.as_ref(),
440                "UPDATE `t_order_v2` SET `c_amount_v2` = ? WHERE `c_order_id` = ?"
441            ),
442            _ => panic!("mapped write plan missing"),
443        }
444        fs::remove_dir_all(directory).unwrap();
445    }
446
447    #[test]
448    fn incomplete_extra_duplicate_and_unsafe_mappings_fail_closed() {
449        let cases = [
450            (
451                r#"{"table":{"from":"订单","to":"t_order"},"columns":[{"from":"订单号","to":"c_order_id"}]}"#,
452                DatabaseNameMappingError::MissingColumn,
453            ),
454            (
455                r#"{"table":{"from":"订单","to":"t_order"},"columns":[{"from":"订单号","to":"c_order_id"},{"from":"金额","to":"c_amount"},{"from":"额外","to":"c_extra"}]}"#,
456                DatabaseNameMappingError::ExtraColumn,
457            ),
458            (
459                r#"{"table":{"from":"订单","to":"t_order"},"columns":[{"from":"订单号","to":"same"},{"from":"金额","to":"same"}]}"#,
460                DatabaseNameMappingError::DuplicatePhysicalColumn,
461            ),
462            (
463                r#"{"table":{"from":"订单","to":"t_order;drop"},"columns":[{"from":"订单号","to":"c_order_id"},{"from":"金额","to":"c_amount"}]}"#,
464                DatabaseNameMappingError::InvalidPhysicalName,
465            ),
466        ];
467        for (value, expected) in cases {
468            let directory = directory();
469            write_mapping(&directory, value);
470            assert_eq!(
471                freeze(Some(config(directory.clone()))).err(),
472                Some(expected)
473            );
474            fs::remove_dir_all(directory).unwrap();
475        }
476    }
477
478    #[test]
479    fn unregistered_operation_has_no_raw_sql_fallback_when_mapping_is_enabled() {
480        let directory = directory();
481        write_mapping(
482            &directory,
483            r#"{"table":{"from":"订单","to":"t_order"},"columns":[{"from":"订单号","to":"c_order_id"},{"from":"金额","to":"c_amount"}]}"#,
484        );
485        let plans = freeze(Some(config(directory.clone()))).unwrap();
486        struct Unregistered;
487        assert!(matches!(
488            plans.query::<Unregistered>(),
489            OperationSql::Missing
490        ));
491        assert!(matches!(
492            plans.write::<Unregistered>(),
493            OperationSql::Missing
494        ));
495        fs::remove_dir_all(directory).unwrap();
496    }
497}