800 lines
25 KiB
Rust
800 lines
25 KiB
Rust
use serde::{Deserialize, Serialize};
|
|
|
|
use crate::models::connection::DatabaseType;
|
|
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
#[serde(rename_all = "camelCase")]
|
|
pub struct SqlFileRequest {
|
|
pub execution_id: String,
|
|
pub connection_id: String,
|
|
pub database: String,
|
|
pub file_path: String,
|
|
pub continue_on_error: bool,
|
|
}
|
|
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
#[serde(rename_all = "camelCase")]
|
|
pub struct SqlFilePreview {
|
|
pub file_name: String,
|
|
pub file_path: String,
|
|
pub size_bytes: u64,
|
|
pub preview: String,
|
|
}
|
|
|
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
|
#[serde(rename_all = "camelCase")]
|
|
pub enum SqlFileStatus {
|
|
Started,
|
|
Running,
|
|
StatementDone,
|
|
StatementFailed,
|
|
Done,
|
|
Error,
|
|
Cancelled,
|
|
}
|
|
|
|
#[derive(Debug, Clone, PartialEq, Eq)]
|
|
pub enum SqlFileStatementAction {
|
|
Execute(String),
|
|
Skip,
|
|
}
|
|
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
#[serde(rename_all = "camelCase")]
|
|
pub struct SqlFileProgress {
|
|
pub execution_id: String,
|
|
pub status: SqlFileStatus,
|
|
pub statement_index: usize,
|
|
pub success_count: usize,
|
|
pub failure_count: usize,
|
|
pub affected_rows: u64,
|
|
pub elapsed_ms: u128,
|
|
pub statement_summary: String,
|
|
pub error: Option<String>,
|
|
}
|
|
|
|
#[derive(Default)]
|
|
pub struct SqlStatementSplitter {
|
|
buffer: String,
|
|
in_single_quote: bool,
|
|
in_double_quote: bool,
|
|
in_backtick: bool,
|
|
in_line_comment: bool,
|
|
in_block_comment: bool,
|
|
dollar_quote_tag: Option<String>,
|
|
previous: Option<char>,
|
|
}
|
|
|
|
impl SqlStatementSplitter {
|
|
pub fn push_chunk(&mut self, chunk: &str) -> Vec<String> {
|
|
let mut statements = Vec::new();
|
|
let chars = chunk.chars().collect::<Vec<_>>();
|
|
let mut i = 0;
|
|
|
|
while i < chars.len() {
|
|
if let Some(tag) = &self.dollar_quote_tag {
|
|
let tag_chars = tag.chars().collect::<Vec<_>>();
|
|
if starts_with_chars(&chars, i, &tag_chars) {
|
|
for tag_ch in &tag_chars {
|
|
self.buffer.push(*tag_ch);
|
|
self.previous = Some(*tag_ch);
|
|
}
|
|
i += tag_chars.len();
|
|
self.dollar_quote_tag = None;
|
|
continue;
|
|
}
|
|
|
|
let ch = chars[i];
|
|
self.buffer.push(ch);
|
|
self.previous = Some(ch);
|
|
i += 1;
|
|
continue;
|
|
}
|
|
|
|
let ch = chars[i];
|
|
let next = chars.get(i + 1).copied();
|
|
|
|
if self.in_line_comment {
|
|
self.buffer.push(ch);
|
|
if ch == '\n' {
|
|
self.in_line_comment = false;
|
|
}
|
|
self.previous = Some(ch);
|
|
i += 1;
|
|
continue;
|
|
}
|
|
|
|
if self.in_block_comment {
|
|
self.buffer.push(ch);
|
|
if self.previous == Some('*') && ch == '/' {
|
|
self.in_block_comment = false;
|
|
}
|
|
self.previous = Some(ch);
|
|
i += 1;
|
|
continue;
|
|
}
|
|
|
|
if !self.in_single_quote && !self.in_double_quote && !self.in_backtick {
|
|
if self.previous == Some('-') && ch == '-' {
|
|
self.in_line_comment = true;
|
|
self.buffer.push(ch);
|
|
self.previous = Some(ch);
|
|
i += 1;
|
|
continue;
|
|
}
|
|
if self.previous == Some('/') && ch == '*' {
|
|
self.in_block_comment = true;
|
|
self.buffer.push(ch);
|
|
self.previous = Some(ch);
|
|
i += 1;
|
|
continue;
|
|
}
|
|
if ch == '-' && next == Some('-') {
|
|
self.in_line_comment = true;
|
|
self.buffer.push(ch);
|
|
self.previous = Some(ch);
|
|
i += 1;
|
|
continue;
|
|
}
|
|
if ch == '/' && next == Some('*') {
|
|
self.in_block_comment = true;
|
|
self.buffer.push(ch);
|
|
self.previous = Some(ch);
|
|
i += 1;
|
|
continue;
|
|
}
|
|
if let Some(tag) = dollar_quote_tag_at(&chars, i) {
|
|
for tag_ch in tag.chars() {
|
|
self.buffer.push(tag_ch);
|
|
self.previous = Some(tag_ch);
|
|
}
|
|
i += tag.chars().count();
|
|
self.dollar_quote_tag = Some(tag);
|
|
continue;
|
|
}
|
|
}
|
|
|
|
match ch {
|
|
'\'' if !self.in_double_quote && !self.in_backtick && self.previous != Some('\\') => {
|
|
self.in_single_quote = !self.in_single_quote;
|
|
self.buffer.push(ch);
|
|
}
|
|
'"' if !self.in_single_quote && !self.in_backtick && self.previous != Some('\\') => {
|
|
self.in_double_quote = !self.in_double_quote;
|
|
self.buffer.push(ch);
|
|
}
|
|
'`' if !self.in_single_quote && !self.in_double_quote => {
|
|
self.in_backtick = !self.in_backtick;
|
|
self.buffer.push(ch);
|
|
}
|
|
';' if !self.in_single_quote && !self.in_double_quote && !self.in_backtick => {
|
|
self.push_current_statement(&mut statements);
|
|
}
|
|
_ => self.buffer.push(ch),
|
|
}
|
|
|
|
self.previous = Some(ch);
|
|
i += 1;
|
|
}
|
|
|
|
statements
|
|
}
|
|
|
|
pub fn finish(mut self) -> Vec<String> {
|
|
let mut statements = Vec::new();
|
|
self.push_current_statement(&mut statements);
|
|
statements
|
|
}
|
|
|
|
fn push_current_statement(&mut self, statements: &mut Vec<String>) {
|
|
let statement = self.buffer.trim();
|
|
if has_executable_sql(statement) {
|
|
statements.push(statement.to_string());
|
|
}
|
|
self.buffer.clear();
|
|
self.previous = None;
|
|
}
|
|
}
|
|
|
|
pub fn split_sql_statements(sql: &str) -> Vec<String> {
|
|
let mut splitter = SqlStatementSplitter::default();
|
|
let mut statements = splitter.push_chunk(sql);
|
|
statements.extend(splitter.finish());
|
|
statements
|
|
}
|
|
|
|
pub fn split_sql_batches(sql: &str) -> Vec<String> {
|
|
let mut batches = Vec::new();
|
|
let mut current_start = 0;
|
|
let lines: Vec<&str> = sql.split('\n').collect();
|
|
let mut offset = 0;
|
|
|
|
for line in &lines {
|
|
let line_start = offset;
|
|
let line_end = offset + line.len();
|
|
offset = line_end + 1; // +1 for the '\n'
|
|
|
|
let trimmed = line.trim();
|
|
if trimmed.eq_ignore_ascii_case("go")
|
|
|| trimmed.to_ascii_lowercase().starts_with("go ") && trimmed[2..].trim().is_empty()
|
|
{
|
|
let batch = sql[current_start..line_start].trim();
|
|
if has_executable_sql(batch) {
|
|
batches.push(batch.to_string());
|
|
}
|
|
current_start = line_end.min(sql.len());
|
|
if current_start < sql.len() && sql.as_bytes()[current_start] == b'\n' {
|
|
current_start += 1;
|
|
}
|
|
}
|
|
}
|
|
|
|
let trailing = sql[current_start..].trim();
|
|
if has_executable_sql(trailing) {
|
|
batches.push(trailing.to_string());
|
|
}
|
|
|
|
if batches.is_empty() {
|
|
let trimmed = sql.trim();
|
|
if !trimmed.is_empty() {
|
|
batches.push(trimmed.to_string());
|
|
}
|
|
}
|
|
|
|
batches
|
|
}
|
|
|
|
pub fn statement_summary(statement: &str) -> String {
|
|
const MAX_LEN: usize = 120;
|
|
|
|
let collapsed = statement.split_whitespace().collect::<Vec<_>>().join(" ");
|
|
if collapsed.chars().count() <= MAX_LEN {
|
|
return collapsed;
|
|
}
|
|
|
|
collapsed.chars().take(MAX_LEN).collect()
|
|
}
|
|
|
|
pub fn prepare_sql_file_statement(
|
|
statement: &str,
|
|
db_type: &DatabaseType,
|
|
driver_profile: Option<&str>,
|
|
) -> SqlFileStatementAction {
|
|
let statement = statement.trim();
|
|
let is_mysql_compatible_target = is_mysql_compatible_import_target(db_type, driver_profile);
|
|
if is_mysql_compatible_target && is_mysql_lock_table_statement(statement) {
|
|
return SqlFileStatementAction::Skip;
|
|
}
|
|
|
|
let Some(body) = mysql_executable_comment_body(statement) else {
|
|
if is_mysql_compatible_target && is_mysql_session_restore_statement(statement) {
|
|
return SqlFileStatementAction::Skip;
|
|
}
|
|
return SqlFileStatementAction::Execute(statement.to_string());
|
|
};
|
|
|
|
if !is_mysql_compatible_target {
|
|
return SqlFileStatementAction::Skip;
|
|
}
|
|
|
|
let body = body.trim();
|
|
if body.is_empty() || is_mysql_key_toggle_statement(body) || is_mysql_session_restore_statement(body) {
|
|
return SqlFileStatementAction::Skip;
|
|
}
|
|
|
|
SqlFileStatementAction::Execute(body.to_string())
|
|
}
|
|
|
|
pub fn starts_with_executable_sql_keyword(sql: &str, keywords: &[&str]) -> bool {
|
|
let Some(token) = first_executable_sql_token(sql) else {
|
|
return false;
|
|
};
|
|
keywords.iter().any(|keyword| token.eq_ignore_ascii_case(keyword))
|
|
}
|
|
|
|
fn is_mysql_compatible_import_target(db_type: &DatabaseType, driver_profile: Option<&str>) -> bool {
|
|
matches!(db_type, DatabaseType::Mysql | DatabaseType::Doris | DatabaseType::StarRocks | DatabaseType::Goldendb)
|
|
|| driver_profile.map(|profile| profile.to_ascii_lowercase()).is_some_and(|profile| {
|
|
matches!(
|
|
profile.as_str(),
|
|
"mariadb" | "tidb" | "oceanbase" | "custom_mysql" | "doris" | "starrocks" | "selectdb" | "goldendb"
|
|
)
|
|
})
|
|
}
|
|
|
|
fn mysql_executable_comment_body(statement: &str) -> Option<&str> {
|
|
let bytes = statement.as_bytes();
|
|
let start = leading_mysql_executable_comment_start(statement)?;
|
|
let body_start = if bytes.get(start + 2) == Some(&b'!') { start + 3 } else { start + 4 };
|
|
let mut body_start = body_start;
|
|
while body_start < bytes.len() && (bytes[body_start].is_ascii_digit() || bytes[body_start].is_ascii_whitespace()) {
|
|
body_start += 1;
|
|
}
|
|
|
|
let close = find_block_comment_close(bytes, body_start)?;
|
|
if has_executable_sql(&statement[close + 2..]) {
|
|
return None;
|
|
}
|
|
|
|
Some(&statement[body_start..close])
|
|
}
|
|
|
|
fn leading_mysql_executable_comment_start(statement: &str) -> Option<usize> {
|
|
let bytes = statement.as_bytes();
|
|
let mut i = 0;
|
|
|
|
while i < bytes.len() {
|
|
while i < bytes.len() && bytes[i].is_ascii_whitespace() {
|
|
i += 1;
|
|
}
|
|
|
|
if i + 1 < bytes.len() && bytes[i] == b'-' && bytes[i + 1] == b'-' {
|
|
i += 2;
|
|
while i < bytes.len() && bytes[i] != b'\n' {
|
|
i += 1;
|
|
}
|
|
continue;
|
|
}
|
|
|
|
if i + 1 < bytes.len() && bytes[i] == b'/' && bytes[i + 1] == b'*' {
|
|
if i + 2 < bytes.len() && (bytes[i + 2] == b'!' || (i + 3 < bytes.len() && &bytes[i + 2..i + 4] == b"M!")) {
|
|
return Some(i);
|
|
}
|
|
|
|
let close = find_block_comment_close(bytes, i + 2)?;
|
|
i = close + 2;
|
|
continue;
|
|
}
|
|
|
|
return None;
|
|
}
|
|
|
|
None
|
|
}
|
|
|
|
fn find_block_comment_close(bytes: &[u8], mut start: usize) -> Option<usize> {
|
|
while start + 1 < bytes.len() {
|
|
if bytes[start] == b'*' && bytes[start + 1] == b'/' {
|
|
return Some(start);
|
|
}
|
|
start += 1;
|
|
}
|
|
None
|
|
}
|
|
|
|
fn is_mysql_key_toggle_statement(statement: &str) -> bool {
|
|
let upper = statement.split_whitespace().collect::<Vec<_>>().join(" ").to_ascii_uppercase();
|
|
upper.starts_with("ALTER TABLE ") && (upper.ends_with(" ENABLE KEYS") || upper.ends_with(" DISABLE KEYS"))
|
|
}
|
|
|
|
fn is_mysql_lock_table_statement(statement: &str) -> bool {
|
|
let executable = leading_executable_sql(statement);
|
|
let upper = executable.split_whitespace().collect::<Vec<_>>().join(" ").to_ascii_uppercase();
|
|
upper == "UNLOCK TABLES" || (upper.starts_with("LOCK TABLES ") && upper.ends_with(" WRITE"))
|
|
}
|
|
|
|
fn is_mysql_session_restore_statement(statement: &str) -> bool {
|
|
let executable = leading_executable_sql(statement);
|
|
let upper = executable.split_whitespace().collect::<Vec<_>>().join(" ").to_ascii_uppercase();
|
|
if !upper.starts_with("SET ") {
|
|
return false;
|
|
}
|
|
|
|
let assignment = upper.trim_start_matches("SET ").trim();
|
|
if assignment.starts_with('@') {
|
|
return false;
|
|
}
|
|
|
|
assignment.contains("= @OLD_")
|
|
|| assignment.contains("=@OLD_")
|
|
|| assignment.contains("= @SAVED_")
|
|
|| assignment.contains("=@SAVED_")
|
|
}
|
|
|
|
fn leading_executable_sql(sql: &str) -> &str {
|
|
let bytes = sql.as_bytes();
|
|
let mut i = 0;
|
|
|
|
while i < bytes.len() {
|
|
while i < bytes.len() && bytes[i].is_ascii_whitespace() {
|
|
i += 1;
|
|
}
|
|
|
|
if i + 1 < bytes.len() && bytes[i] == b'-' && bytes[i + 1] == b'-' {
|
|
i += 2;
|
|
while i < bytes.len() && bytes[i] != b'\n' {
|
|
i += 1;
|
|
}
|
|
continue;
|
|
}
|
|
|
|
if i + 1 < bytes.len() && bytes[i] == b'/' && bytes[i + 1] == b'*' {
|
|
if i + 2 < bytes.len() && (bytes[i + 2] == b'!' || (i + 3 < bytes.len() && &bytes[i + 2..i + 4] == b"M!")) {
|
|
break;
|
|
}
|
|
|
|
let Some(close) = find_block_comment_close(bytes, i + 2) else {
|
|
return &sql[sql.len()..];
|
|
};
|
|
i = close + 2;
|
|
continue;
|
|
}
|
|
|
|
break;
|
|
}
|
|
|
|
&sql[i..]
|
|
}
|
|
|
|
fn first_executable_sql_token(sql: &str) -> Option<&str> {
|
|
let bytes = sql.as_bytes();
|
|
let mut i = 0;
|
|
|
|
while i < bytes.len() {
|
|
while i < bytes.len() && bytes[i].is_ascii_whitespace() {
|
|
i += 1;
|
|
}
|
|
|
|
if i + 1 < bytes.len() && bytes[i] == b'-' && bytes[i + 1] == b'-' {
|
|
i += 2;
|
|
while i < bytes.len() && bytes[i] != b'\n' {
|
|
i += 1;
|
|
}
|
|
continue;
|
|
}
|
|
|
|
if i + 1 < bytes.len() && bytes[i] == b'/' && bytes[i + 1] == b'*' {
|
|
if i + 2 < bytes.len() && (bytes[i + 2] == b'!' || (i + 3 < bytes.len() && &bytes[i + 2..i + 4] == b"M!")) {
|
|
i += if bytes[i + 2] == b'!' { 3 } else { 4 };
|
|
while i < bytes.len() && (bytes[i].is_ascii_digit() || bytes[i].is_ascii_whitespace()) {
|
|
i += 1;
|
|
}
|
|
break;
|
|
}
|
|
|
|
i += 2;
|
|
while i + 1 < bytes.len() && !(bytes[i] == b'*' && bytes[i + 1] == b'/') {
|
|
i += 1;
|
|
}
|
|
i = (i + 2).min(bytes.len());
|
|
continue;
|
|
}
|
|
|
|
break;
|
|
}
|
|
|
|
let start = i;
|
|
while i < bytes.len() && (bytes[i].is_ascii_alphabetic() || bytes[i] == b'_') {
|
|
i += 1;
|
|
}
|
|
|
|
(i > start).then_some(&sql[start..i])
|
|
}
|
|
|
|
fn starts_with_chars(chars: &[char], start: usize, needle: &[char]) -> bool {
|
|
start + needle.len() <= chars.len() && chars[start..start + needle.len()] == *needle
|
|
}
|
|
|
|
fn dollar_quote_tag_at(chars: &[char], start: usize) -> Option<String> {
|
|
if chars.get(start) != Some(&'$') {
|
|
return None;
|
|
}
|
|
|
|
match chars.get(start + 1) {
|
|
Some('$') => return Some("$$".to_string()),
|
|
Some(ch) if ch.is_ascii_alphabetic() || *ch == '_' => {}
|
|
_ => return None,
|
|
}
|
|
|
|
let mut end = start + 2;
|
|
while let Some(ch) = chars.get(end) {
|
|
if *ch == '$' {
|
|
return Some(chars[start..=end].iter().collect());
|
|
}
|
|
if !ch.is_ascii_alphanumeric() && *ch != '_' {
|
|
return None;
|
|
}
|
|
end += 1;
|
|
}
|
|
|
|
None
|
|
}
|
|
|
|
fn has_executable_sql(statement: &str) -> bool {
|
|
let chars = statement.chars().collect::<Vec<_>>();
|
|
let mut in_line_comment = false;
|
|
let mut in_block_comment = false;
|
|
let mut previous = None;
|
|
let mut i = 0;
|
|
|
|
while i < chars.len() {
|
|
let ch = chars[i];
|
|
let next = chars.get(i + 1).copied();
|
|
|
|
if in_line_comment {
|
|
if ch == '\n' {
|
|
in_line_comment = false;
|
|
}
|
|
previous = Some(ch);
|
|
i += 1;
|
|
continue;
|
|
}
|
|
|
|
if in_block_comment {
|
|
if previous == Some('*') && ch == '/' {
|
|
in_block_comment = false;
|
|
}
|
|
previous = Some(ch);
|
|
i += 1;
|
|
continue;
|
|
}
|
|
|
|
if ch == '-' && next == Some('-') {
|
|
in_line_comment = true;
|
|
previous = Some(ch);
|
|
i += 1;
|
|
continue;
|
|
}
|
|
|
|
if ch == '/' && next == Some('*') {
|
|
if is_mysql_executable_comment_start(&chars, i) {
|
|
return true;
|
|
}
|
|
in_block_comment = true;
|
|
previous = Some(ch);
|
|
i += 1;
|
|
continue;
|
|
}
|
|
|
|
if !ch.is_whitespace() {
|
|
return true;
|
|
}
|
|
|
|
previous = Some(ch);
|
|
i += 1;
|
|
}
|
|
|
|
false
|
|
}
|
|
|
|
fn is_mysql_executable_comment_start(chars: &[char], start: usize) -> bool {
|
|
chars.get(start) == Some(&'/')
|
|
&& chars.get(start + 1) == Some(&'*')
|
|
&& (chars.get(start + 2) == Some(&'!')
|
|
|| (chars.get(start + 2) == Some(&'M') && chars.get(start + 3) == Some(&'!')))
|
|
}
|
|
|
|
#[cfg(test)]
|
|
fn split_sql_script(sql: &str) -> Result<Vec<String>, String> {
|
|
Ok(split_sql_statements(sql))
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use crate::models::connection::DatabaseType;
|
|
|
|
use super::{
|
|
prepare_sql_file_statement, split_sql_script, starts_with_executable_sql_keyword, SqlFileStatementAction,
|
|
SqlStatementSplitter,
|
|
};
|
|
|
|
#[test]
|
|
fn splits_semicolon_delimited_statements() {
|
|
assert_eq!(
|
|
split_sql_script("CREATE TABLE a(id int); INSERT INTO a VALUES (1);").unwrap(),
|
|
vec!["CREATE TABLE a(id int)", "INSERT INTO a VALUES (1)"]
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn ignores_semicolons_inside_quotes_and_comments() {
|
|
let sql = "\
|
|
INSERT INTO logs VALUES ('a;b', \"c;d\", `weird;name`);\n\
|
|
-- comment ; ignored\n\
|
|
/* block ; ignored */\n\
|
|
SELECT 1;";
|
|
assert_eq!(
|
|
split_sql_script(sql).unwrap(),
|
|
vec![
|
|
"INSERT INTO logs VALUES ('a;b', \"c;d\", `weird;name`)",
|
|
"-- comment ; ignored\n/* block ; ignored */\nSELECT 1",
|
|
]
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn emits_trailing_statement_without_semicolon() {
|
|
assert_eq!(
|
|
split_sql_script("CREATE TABLE a(id int);\nINSERT INTO a VALUES (1)").unwrap(),
|
|
vec!["CREATE TABLE a(id int)", "INSERT INTO a VALUES (1)"]
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn line_comment_openers_can_span_chunks() {
|
|
let mut splitter = SqlStatementSplitter::default();
|
|
|
|
assert_eq!(splitter.push_chunk("SELECT 1; -"), vec!["SELECT 1"]);
|
|
assert_eq!(splitter.push_chunk("- comment ; ignored\nSELECT 2;"), vec!["-- comment ; ignored\nSELECT 2"]);
|
|
assert_eq!(splitter.finish(), Vec::<String>::new());
|
|
}
|
|
|
|
#[test]
|
|
fn block_comment_openers_can_span_chunks() {
|
|
let mut splitter = SqlStatementSplitter::default();
|
|
|
|
assert_eq!(splitter.push_chunk("SELECT 1; /"), vec!["SELECT 1"]);
|
|
assert_eq!(splitter.push_chunk("* comment ; ignored */\nSELECT 2;"), vec!["/* comment ; ignored */\nSELECT 2"]);
|
|
assert_eq!(splitter.finish(), Vec::<String>::new());
|
|
}
|
|
|
|
#[test]
|
|
fn skips_comment_only_tail_after_statement() {
|
|
assert_eq!(
|
|
split_sql_script("CREATE TABLE a(id int); -- done\n/* no more sql */").unwrap(),
|
|
vec!["CREATE TABLE a(id int)"]
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn keeps_postgres_dollar_quoted_function_body_together() {
|
|
let sql = "\
|
|
CREATE FUNCTION bump_counter()\n\
|
|
RETURNS trigger AS $$\n\
|
|
BEGIN\n\
|
|
PERFORM 1;\n\
|
|
RETURN NEW;\n\
|
|
END;\n\
|
|
$$ LANGUAGE plpgsql;\n\
|
|
SELECT 1;";
|
|
|
|
assert_eq!(
|
|
split_sql_script(sql).unwrap(),
|
|
vec![
|
|
"CREATE FUNCTION bump_counter()\nRETURNS trigger AS $$\nBEGIN\nPERFORM 1;\nRETURN NEW;\nEND;\n$$ LANGUAGE plpgsql",
|
|
"SELECT 1",
|
|
]
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn keeps_mysql_executable_comments_as_statements() {
|
|
assert_eq!(
|
|
split_sql_script("/*!40101 SET @OLD_CHARACTER_SET_CLIENT=@@CHARACTER_SET_CLIENT */;\nSELECT 1;",).unwrap(),
|
|
vec!["/*!40101 SET @OLD_CHARACTER_SET_CLIENT=@@CHARACTER_SET_CLIENT */", "SELECT 1",]
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn detects_result_set_keyword_after_comments() {
|
|
assert!(starts_with_executable_sql_keyword("-- comment\nselect * from users;", &["SELECT"]));
|
|
assert!(starts_with_executable_sql_keyword(
|
|
"/* comment */\nWITH rows AS (SELECT 1) SELECT * FROM rows;",
|
|
&["WITH"]
|
|
));
|
|
assert!(!starts_with_executable_sql_keyword("-- comment only\n", &["SELECT"]));
|
|
}
|
|
|
|
#[test]
|
|
fn detects_mysql_executable_comment_keyword() {
|
|
assert!(starts_with_executable_sql_keyword("/*!40101 SELECT 1 */", &["SELECT"]));
|
|
assert!(starts_with_executable_sql_keyword("/*M! SELECT 1 */", &["SELECT"]));
|
|
}
|
|
|
|
#[test]
|
|
fn prepares_mysql_executable_comments_for_mysql_compatible_imports() {
|
|
assert_eq!(
|
|
prepare_sql_file_statement("/*!40101 SET NAMES utf8mb4 */", &DatabaseType::Mysql, None),
|
|
SqlFileStatementAction::Execute("SET NAMES utf8mb4".to_string())
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn skips_mysql_key_toggle_comments_for_mysql_compatible_imports() {
|
|
assert_eq!(
|
|
prepare_sql_file_statement(" /*!40000 ALTER TABLE `dd_admin` ENABLE KEYS */", &DatabaseType::Mysql, None),
|
|
SqlFileStatementAction::Skip
|
|
);
|
|
assert_eq!(
|
|
prepare_sql_file_statement("/*!40000 ALTER TABLE `dd_admin` DISABLE KEYS */", &DatabaseType::Mysql, None),
|
|
SqlFileStatementAction::Skip
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn skips_mysql_lock_table_statements_for_mysql_compatible_imports() {
|
|
assert_eq!(
|
|
prepare_sql_file_statement("LOCK TABLES `dd_geo_json` WRITE", &DatabaseType::Mysql, None),
|
|
SqlFileStatementAction::Skip
|
|
);
|
|
assert_eq!(
|
|
prepare_sql_file_statement("UNLOCK TABLES", &DatabaseType::Mysql, None),
|
|
SqlFileStatementAction::Skip
|
|
);
|
|
assert_eq!(
|
|
prepare_sql_file_statement(
|
|
"-- Dumping data for table `dd_geo_json`\nLOCK TABLES `dd_geo_json` WRITE",
|
|
&DatabaseType::Mysql,
|
|
None
|
|
),
|
|
SqlFileStatementAction::Skip
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn skips_mysql_session_restore_statements_for_mysql_compatible_imports() {
|
|
assert_eq!(
|
|
prepare_sql_file_statement(
|
|
"/*!40101 SET character_set_client = @saved_cs_client */",
|
|
&DatabaseType::Mysql,
|
|
None
|
|
),
|
|
SqlFileStatementAction::Skip
|
|
);
|
|
assert_eq!(
|
|
prepare_sql_file_statement("/*!40103 SET TIME_ZONE=@OLD_TIME_ZONE */", &DatabaseType::Mysql, None),
|
|
SqlFileStatementAction::Skip
|
|
);
|
|
assert_eq!(
|
|
prepare_sql_file_statement("SET FOREIGN_KEY_CHECKS=@OLD_FOREIGN_KEY_CHECKS", &DatabaseType::Mysql, None),
|
|
SqlFileStatementAction::Skip
|
|
);
|
|
assert_eq!(
|
|
prepare_sql_file_statement(
|
|
"/*!40101 SET @saved_cs_client = @@character_set_client */",
|
|
&DatabaseType::Mysql,
|
|
None
|
|
),
|
|
SqlFileStatementAction::Execute("SET @saved_cs_client = @@character_set_client".to_string())
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn skips_mysql_executable_comments_for_non_mysql_imports() {
|
|
assert_eq!(
|
|
prepare_sql_file_statement(
|
|
"/*!40101 SET character_set_client = @saved_cs_client */",
|
|
&DatabaseType::Postgres,
|
|
None
|
|
),
|
|
SqlFileStatementAction::Skip
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn split_batches_by_go() {
|
|
assert_eq!(super::split_sql_batches("SELECT 1\nGO\nSELECT 2"), vec!["SELECT 1", "SELECT 2"]);
|
|
}
|
|
|
|
#[test]
|
|
fn split_batches_go_case_insensitive() {
|
|
assert_eq!(
|
|
super::split_sql_batches("SELECT 1\ngo\nSELECT 2\nGo\nSELECT 3"),
|
|
vec!["SELECT 1", "SELECT 2", "SELECT 3"]
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn split_batches_go_with_surrounding_whitespace() {
|
|
assert_eq!(super::split_sql_batches("SELECT 1\n GO \nSELECT 2"), vec!["SELECT 1", "SELECT 2"]);
|
|
}
|
|
|
|
#[test]
|
|
fn split_batches_no_go_returns_whole() {
|
|
assert_eq!(
|
|
super::split_sql_batches("DECLARE @x INT = 1;\nSELECT @x;"),
|
|
vec!["DECLARE @x INT = 1;\nSELECT @x;"]
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn split_batches_skips_empty_batches() {
|
|
assert_eq!(super::split_sql_batches("SELECT 1\nGO\n\nGO\nSELECT 2"), vec!["SELECT 1", "SELECT 2"]);
|
|
}
|
|
|
|
#[test]
|
|
fn split_batches_trailing_go() {
|
|
assert_eq!(super::split_sql_batches("SELECT 1\nGO"), vec!["SELECT 1"]);
|
|
}
|
|
}
|