use crate::Result;
use crate::error::DbError;
use crate::tpl::cache;
use dashmap::DashMap;
use glob::glob;
use quick_xml::events::{BytesStart, Event};
use quick_xml::reader::Reader;
use std::fs;
use std::path::Path;
use std::sync::{Arc, OnceLock};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum StatementType {
Select,
Insert,
Update,
Delete,
Sql,
}
impl StatementType {
fn from_str(s: &str) -> Option<Self> {
match s {
"select" => Some(StatementType::Select),
"insert" => Some(StatementType::Insert),
"update" => Some(StatementType::Update),
"delete" => Some(StatementType::Delete),
"sql" => Some(StatementType::Sql),
_ => None,
}
}
}
#[derive(Debug, Clone)]
pub struct SqlStatement {
pub r#type: StatementType,
pub database_type: Option<String>,
pub content: Option<String>,
pub return_key: bool,
}
pub type StatementStore = DashMap<String, DashMap<String, Vec<Arc<SqlStatement>>>>;
static STATEMENTS: OnceLock<StatementStore> = OnceLock::new();
pub fn load(pattern: &str) -> Result<()> {
let paths = glob(pattern)
.map_err(|e| DbError::MapperLoadError(format!("无效的 glob 模式: {} - {}", pattern, e)))?;
for entry in paths {
let path: std::path::PathBuf = entry.map_err(|e: glob::GlobError| {
DbError::MapperLoadError(format!("无法读取路径: {} - {}", pattern, e))
})?;
if path.is_file() {
load_file(&path)?;
}
}
Ok(())
}
pub fn load_assets(assets: Vec<(&str, &str)>) -> Result<()> {
for (source, content) in assets {
parse_and_register(content, source)?;
}
Ok(())
}
pub fn find_statement(full_id: &str, db_type: &str) -> Option<Arc<SqlStatement>> {
let (namespace, id) = full_id.rsplit_once('.')?;
let ns_map = STATEMENTS.get()?.get(namespace)?;
let statements = ns_map.get(id)?;
let mut fallback = None;
for stmt in statements.value().iter() {
match stmt.database_type.as_deref() {
Some(t) if t == db_type => return Some(stmt.clone()),
None => fallback = Some(stmt.clone()),
_ => {}
}
}
fallback
}
pub fn clear() {
if let Some(store) = STATEMENTS.get() {
store.clear();
}
}
fn load_file(path: &Path) -> Result<()> {
let xml_content = fs::read_to_string(path).map_err(|e| {
DbError::MapperLoadError(format!(
"读取 Mapper 文件失败: {} (cause: {})",
path.display(),
e
))
})?;
parse_and_register(&xml_content, &path.display().to_string())
}
fn parse_and_register(xml_content: &str, source: &str) -> Result<()> {
let (namespace, items) = parse_xml(xml_content, source)?;
let store = STATEMENTS.get_or_init(DashMap::new);
let ns_map = store.entry(namespace.clone()).or_default();
for mut statement in items {
if let Some(content) = &mut statement.content {
*content = content.trim().to_string();
}
if let Some(content) = &statement.content {
let full_id = format!("{}.{}", namespace, statement.id);
cache::get_ast(&full_id, content);
}
let mut statements = ns_map.entry(statement.id.clone()).or_default();
if statements
.iter()
.any(|s| s.database_type == statement.database_type)
{
return Err(DbError::MapperLoadError(format!(
"重复的 SQL ID 定义: '{}' (Database: '{:?}', Source: '{}')",
statement.id, statement.database_type, source
)));
}
statements.push(Arc::new(statement.into_sql_statement()));
}
Ok(())
}
struct ParsedItem {
r#type: StatementType,
id: String,
database_type: Option<String>,
return_key: bool,
content: Option<String>,
}
impl ParsedItem {
fn into_sql_statement(self) -> SqlStatement {
SqlStatement {
r#type: self.r#type,
database_type: self.database_type,
content: self.content,
return_key: self.return_key,
}
}
}
fn parse_xml(xml: &str, source: &str) -> Result<(String, Vec<ParsedItem>)> {
let mut reader = Reader::from_str(xml);
reader.config_mut().trim_text(true);
let mut namespace = None;
let mut items = Vec::new();
let mut buf = Vec::new();
loop {
match reader.read_event_into(&mut buf) {
Ok(Event::Start(ref e)) => {
let name = e.name();
let name_str = String::from_utf8_lossy(name.as_ref());
if name_str == "mapper" {
namespace =
get_attribute(e, "namespace").or_else(|| get_attribute(e, "Namespace"));
} else if let Some(stmt_type) = StatementType::from_str(&name_str) {
let id = get_attribute(e, "id").ok_or_else(|| {
DbError::MapperLoadError(format!("SQL 语句缺少 id 属性: {}", source))
})?;
let database_type = get_attribute(e, "databaseType");
let return_key = parse_bool(get_attribute(e, "returnKey").as_deref());
let start_pos = reader.buffer_position() as usize;
let end_pos = read_until_end_tag(&mut reader, &name_str, &mut Vec::new())?;
let tag_len = name.as_ref().len();
if end_pos < tag_len + 3 {
return Err(DbError::MapperLoadError(
"解析错误: 结束标签位置异常".to_string(),
));
}
let content_end = end_pos - (tag_len + 3);
let content = if content_end > start_pos {
let raw_content = &xml[start_pos..content_end];
quick_xml::escape::unescape(raw_content)
.map(|s| s.into_owned())
.ok()
} else {
None
};
items.push(ParsedItem {
r#type: stmt_type,
id,
database_type,
return_key,
content,
});
}
}
Ok(Event::Eof) => break,
Err(e) => {
return Err(DbError::MapperLoadError(format!(
"XML 解析错误: {} (Source: {})",
e, source
)));
}
_ => {}
}
buf.clear();
}
let namespace = namespace.ok_or_else(|| {
DbError::MapperLoadError(format!("Mapper XML 缺少 namespace 属性: {}", source))
})?;
Ok((namespace, items))
}
fn read_until_end_tag(
reader: &mut Reader<&[u8]>,
target_tag: &str,
buf: &mut Vec<u8>,
) -> Result<usize> {
let mut depth = 0;
loop {
match reader.read_event_into(buf) {
Ok(Event::Start(ref e)) => {
if e.name().as_ref() == target_tag.as_bytes() {
depth += 1;
}
}
Ok(Event::End(ref e)) => {
if e.name().as_ref() == target_tag.as_bytes() {
if depth == 0 {
return Ok(reader.buffer_position() as usize);
}
depth -= 1;
}
}
Ok(Event::Eof) => {
return Err(DbError::MapperLoadError(format!(
"未找到结束标签: </{}>",
target_tag
)));
}
Err(e) => return Err(DbError::MapperLoadError(format!("XML 解析错误: {}", e))),
_ => {}
}
buf.clear();
}
}
fn get_attribute(e: &BytesStart, key: &str) -> Option<String> {
e.attributes()
.filter_map(|a| a.ok())
.find(|a| a.key.as_ref() == key.as_bytes())
.map(|a| String::from_utf8_lossy(&a.value).into_owned())
}
fn parse_bool(s: Option<&str>) -> bool {
matches!(
s.unwrap_or("").trim().to_ascii_lowercase().as_str(),
"true" | "1" | "yes" | "on"
)
}