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
13pub 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 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, "e_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, "e_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}