use crate::table_dao::model::{CrudType, table};
use sqlx::{Column, Executor, Row, SqliteConnection, SqlitePool, Statement, TypeInfo};
use std::collections::HashMap;
use std::sync::{LazyLock, Mutex};
static TYPE_OVERRIDES: LazyLock<Mutex<HashMap<String, String>>> =
LazyLock::new(|| Mutex::new(HashMap::new()));
pub(crate) fn set_type_overrides(overrides: &HashMap<String, String>) {
*TYPE_OVERRIDES.lock().unwrap() = overrides.clone();
}
pub async fn all_entities(conn: &SqlitePool) -> Vec<table::Entity> {
sqlx::query(
"SELECT name,sql FROM sqlite_master WHERE type = 'table' and name NOT LIKE 'sqlite_%'",
)
.fetch_all(conn)
.await
.unwrap()
.iter()
.map(|it| {
let sql: String = it.get("sql");
let table_info = read_table_info_from_sql(&sql);
table_info
})
.collect()
}
pub async fn get_column_from_sql(conn: &SqlitePool, sql: &str) -> Vec<table::Column> {
let Ok(describe) = conn.describe(sql).await else {
panic!("获取列数据失败:{}", sql);
};
describe
.columns()
.iter()
.enumerate()
.map(|(i, it)| {
let is_nullable = describe.nullable(i).unwrap_or(true);
let name = it.name();
if name.contains("!:") {
let name_arr: Vec<&str> = name.split("!:").collect();
table::Column {
name: name_arr[0].to_string(),
data_type: name_arr[1].to_string(),
is_nullable: false,
..Default::default()
}
} else if name.contains(":") {
let name_arr: Vec<&str> = name.split(":").collect();
table::Column {
name: name_arr[0].to_string(),
data_type: name_arr[1].to_string(),
is_nullable: true,
..Default::default()
}
} else {
table::Column {
name: name.to_string(),
data_type: it.type_info().name().to_string(),
is_nullable,
..Default::default()
}
}
})
.collect()
}
pub fn read_table_info_from_sql(sql: &str) -> table::Entity {
let mut sql = sql.replace("\r\n", "\n");
sql = sql.replace("\r", "\n");
let (name, comment) = find_table_name(&sql);
let open_kh_index = sql.find('(').unwrap(); let close_kh_index = sql.rfind(')').unwrap(); let columns_sql = sql[open_kh_index + 1..close_kh_index].to_string();
let columns = columns_sql
.split('\n')
.filter_map(|line| read_column_info_from_sql_line(line))
.collect::<Vec<_>>();
table::Entity {
name,
columns,
comment,
}
}
fn find_table_name(sql: &str) -> (String, String) {
let mut sql = sql.to_string();
while sql.contains(" ") {
sql = sql.replace(" ", " ");
}
sql = sql.replace("\r\n", "\n");
sql = sql.replace("\r", "\n");
let create_table_index = sql.to_uppercase().find("CREATE TABLE").unwrap();
sql = sql[create_table_index + 12..].trim().to_string();
if sql.starts_with("IF NOT EXISTS") {
sql = sql[13..].trim().to_string();
}
let open_kh_index = sql.find('(').unwrap(); let block_sql = sql[..open_kh_index].to_string();
let comment_index = block_sql.find("--");
let mut comment = String::new();
let name = if let Some(comment_index) = comment_index {
comment = block_sql[comment_index..].trim().to_string();
block_sql[..comment_index].trim().to_string()
} else {
block_sql.trim().to_string()
};
(name, comment)
}
fn read_column_info_from_sql_line(line: &str) -> Option<table::Column> {
let line = line.to_string();
let comment_index = line.find("--");
let (line, comment) = if let Some(index) = comment_index {
(
line[..index].trim().to_string(),
line[index + 2..].trim().to_string(),
)
} else {
(line.trim().to_string(), "".to_string())
};
let mut line = line;
while line.contains(" ") {
line = line.replace(" ", " ");
}
let next_space_index = line.find(' ').unwrap_or(line.len());
let column_name = &line[..next_space_index];
if column_name.is_empty() {
return None;
}
if !column_name
.replace("_", "")
.chars()
.all(|c| c.is_ascii_alphanumeric())
{
return None;
}
if column_name.to_uppercase() == "UNIQUE"
|| column_name.to_uppercase() == "PRIMARY"
|| column_name.to_uppercase() == "FOREIGN"
{
return None;
}
let line = line[next_space_index..].trim().to_string();
let next_space_index = line
.find(' ')
.unwrap_or(line.len())
.min(line.find(',').unwrap_or(line.len()));
let data_type = &line[..next_space_index];
let line = line[next_space_index..].trim().to_string().to_uppercase();
let is_primary_key = line.contains("PRIMARY KEY"); let is_nullable = !line.contains("NOT NULL"); let is_auto_increment = line.contains("AUTOINCREMENT") || line.contains("AUTO_INCREMENT");
Some(table::Column {
name: column_name.to_string(),
nick: "".to_string(),
data_type: data_type.to_string(),
is_primary_key,
default_value: None,
is_nullable,
is_auto_increment,
comment,
})
}
#[derive(Debug, Clone, Copy)]
enum ParseSqlArgsState {
Normal,
SingleQuote,
DoubleQuote,
}
pub fn parse_sql_args(sql: &str) -> (String, Vec<String>) {
let chars: Vec<char> = sql.chars().collect();
let mut state = ParseSqlArgsState::Normal;
let mut args = Vec::new();
let mut output = String::new();
let mut i = 0;
while i < chars.len() {
match state {
ParseSqlArgsState::Normal => match chars[i] {
'\'' => {
output.push(chars[i]);
state = ParseSqlArgsState::SingleQuote;
}
'"' => {
output.push(chars[i]);
state = ParseSqlArgsState::DoubleQuote;
}
':' => {
if (i > 0 && chars[i - 1] == ':')
|| (i + 1 < chars.len() && chars[i + 1] == ':')
{
output.push(chars[i]);
} else {
let start = i + 1;
if start < chars.len()
&& (chars[start].is_ascii_alphabetic() || chars[start] == '_')
{
let mut end = start + 1;
while end < chars.len()
&& (chars[end].is_ascii_alphanumeric() || chars[end] == '_')
{
end += 1;
}
args.push(chars[start..end].iter().collect());
output.push('?');
i = end - 1;
} else {
output.push(chars[i]);
}
}
}
_ => {
output.push(chars[i]);
}
},
ParseSqlArgsState::SingleQuote => {
output.push(chars[i]);
if chars[i] == '\'' {
if i + 1 < chars.len() && chars[i + 1] == '\'' {
i += 1;
output.push(chars[i]);
} else {
state = ParseSqlArgsState::Normal;
}
}
}
ParseSqlArgsState::DoubleQuote => {
output.push(chars[i]);
if chars[i] == '"' {
if i + 1 < chars.len() && chars[i + 1] == '"' {
i += 1;
output.push(chars[i]);
} else {
state = ParseSqlArgsState::Normal;
}
}
}
}
i += 1;
}
(output, args)
}
pub fn db_type_to_rust(db_type: &str, name: &str) -> String {
if let Some(rust_type) = TYPE_OVERRIDES.lock().unwrap().get(name) {
return rust_type.clone();
}
match db_type.to_uppercase().as_str() {
"INTEGER" | "INT" => "i64".to_string(),
"BIGINT" => "i64".to_string(),
"INT8" => "i64".to_string(),
"INT16" => "i64".to_string(),
"INT32" => "i64".to_string(),
"INT64" => "i64".to_string(),
"VARCHAR" | "TEXT" => "String".to_string(),
"BOOLEAN" => "bool".to_string(),
"FLOAT" => "f32".to_string(),
"DOUBLE" => "f64".to_string(),
"DATETIME" => "chrono::NaiveDateTime".to_string(),
_ => {
if db_type.to_uppercase().starts_with("VARCHAR")
|| db_type.to_uppercase().starts_with("CHAR")
{
"String".to_string()
} else if name.to_lowercase().ends_with("date") {
"chrono::NaiveDateTime".to_string()
} else {
db_type.to_string()
}
} }
}