1use std::collections::HashMap;
19use std::sync::Arc;
20
21use common_query::Output;
22use common_recordbatch::RecordBatches;
23use common_time::timezone::system_timezone_name;
24use common_version;
25use datatypes::prelude::ConcreteDataType;
26use datatypes::schema::{ColumnSchema, Schema};
27use datatypes::vectors::StringVector;
28use once_cell::sync::Lazy;
29use regex::Regex;
30use regex::bytes::RegexSet;
31use session::SessionRef;
32use session::context::QueryContextRef;
33
34const VARIABLES_SCOPE: &str = "(GLOBAL |SESSION |LOCAL )?";
36
37static SELECT_VAR_PATTERN: Lazy<Regex> =
38 Lazy::new(|| Regex::new("(?i)^(SELECT\\s+@@(.*))").unwrap());
39static SHOW_LOWER_CASE_PATTERN: Lazy<Regex> = Lazy::new(|| {
40 Regex::new(&format!(
41 "(?i)^(SHOW {VARIABLES_SCOPE}VARIABLES LIKE 'lower_case_table_names'(.*))"
42 ))
43 .unwrap()
44});
45static SHOW_VARIABLES_LIKE_PATTERN: Lazy<Regex> = Lazy::new(|| {
46 Regex::new(&format!(
47 "(?i)^(SHOW {VARIABLES_SCOPE}VARIABLES( LIKE (.*))?)"
48 ))
49 .unwrap()
50});
51static SHOW_WARNINGS_PATTERN: Lazy<Regex> =
52 Lazy::new(|| Regex::new("(?i)^(SHOW WARNINGS)").unwrap());
53
54static SELECT_USER_OR_VAR_PATTERN: Lazy<Regex> = Lazy::new(|| {
58 Regex::new(
59 "(?i)^SELECT\\s+(?:(CURRENT_USER|SESSION_USER|SYSTEM_USER|USER)|(@[a-z0-9_$.]+))\\s*;?\\s*$",
60 )
61 .unwrap()
62});
63
64static SELECT_TIME_DIFF_FUNC_PATTERN: Lazy<Regex> =
66 Lazy::new(|| Regex::new("(?i)^(SELECT TIMEDIFF\\(NOW\\(\\), UTC_TIMESTAMP\\(\\)\\))").unwrap());
67
68static SHOW_SQL_MODE_PATTERN: Lazy<Regex> = Lazy::new(|| {
70 Regex::new(&format!(
71 "(?i)^(SHOW {VARIABLES_SCOPE}VARIABLES LIKE 'sql_mode'(.*))"
72 ))
73 .unwrap()
74});
75
76static OTHER_NOT_SUPPORTED_STMT: Lazy<RegexSet> = Lazy::new(|| {
77 RegexSet::new([
78 "(?i)^(ROLLBACK(.*))",
80 "(?i)^(COMMIT(.*))",
81 "(?i)^(START(.*))",
82 "(?i)^(BEGIN(.*))",
83
84 "(?i)^(SET NAMES(.*))",
86 "(?i)^(SET character_set_results(.*))",
87 "(?i)^(SET net_write_timeout(.*))",
88 "(?i)^(SET FOREIGN_KEY_CHECKS(.*))",
89 "(?i)^(SET AUTOCOMMIT(.*))",
90 "(?i)^(SET SQL_LOG_BIN(.*))",
91 "(?i)^(SET SESSION TRANSACTION(.*))",
92 "(?i)^(SET TRANSACTION(.*))",
93 "(?i)^(SET sql_mode(.*))",
94 "(?i)^(SET SQL_SELECT_LIMIT(.*))",
95 "(?i)^(SET PROFILING(.*))",
96
97 "(?i)^(SELECT \\$\\$)",
99
100 "(?i)^(SET SQL_QUOTE_SHOW_CREATE(.*))",
102 "(?i)^(LOCK TABLES(.*))",
103 "(?i)^(UNLOCK TABLES(.*))",
104 "(?i)^(SELECT LOGFILE_GROUP_NAME, FILE_NAME, TOTAL_EXTENTS, INITIAL_SIZE, ENGINE, EXTRA FROM INFORMATION_SCHEMA.FILES(.*))",
105
106 "(?i)^(/\\*!80003 SET(.*) \\*/)$",
108 "(?i)^(SHOW MASTER STATUS)",
109 "(?i)^(SHOW ALL SLAVES STATUS)",
110 "(?i)^(LOCK BINLOG FOR BACKUP)",
111 "(?i)^(LOCK TABLES FOR BACKUP)",
112 "(?i)^(UNLOCK BINLOG(.*))",
113 "(?i)^(/\\*!40101 SET(.*) \\*/)$",
114
115 "(?i)^(SHOW PLUGINS)",
117 "(?i)^(SHOW ENGINES)",
118 "(?i)^(SHOW @@(.*))",
119
120 "(?i)^(/\\*!40101 SET(.*) \\*/)$",
122
123 "(?i)^(/\\*!40100 SET(.*) \\*/)$",
125 "(?i)^(/\\*!40103 SET(.*) \\*/)$",
126 "(?i)^(/\\*!40111 SET(.*) \\*/)$",
127 "(?i)^(/\\*!40101 SET(.*) \\*/)$",
128 "(?i)^(/\\*!40014 SET(.*) \\*/)$",
129 "(?i)^(/\\*!40000 SET(.*) \\*/)$",
130 ]).unwrap()
131});
132
133static VAR_VALUES: Lazy<HashMap<&str, &str>> = Lazy::new(|| {
134 HashMap::from([
135 ("tx_isolation", "REPEATABLE-READ"),
136 ("session.tx_isolation", "REPEATABLE-READ"),
137 ("transaction_isolation", "REPEATABLE-READ"),
138 ("session.transaction_isolation", "REPEATABLE-READ"),
139 ("session.transaction_read_only", "0"),
140 ("max_allowed_packet", "134217728"),
141 ("interactive_timeout", "31536000"),
142 ("wait_timeout", "31536000"),
143 ("net_write_timeout", "31536000"),
144 ("version_comment", common_version::product_name()),
145 ])
146});
147
148fn select_function(name: &str, value: Option<&str>) -> RecordBatches {
153 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
154 name,
155 ConcreteDataType::string_datatype(),
156 true,
157 )]));
158 let columns = vec![Arc::new(StringVector::from(vec![value])) as _];
159 RecordBatches::try_from_columns(schema, columns)
160 .unwrap()
162}
163
164fn show_variables(name: &str, value: &str) -> RecordBatches {
169 let schema = Arc::new(Schema::new(vec![
170 ColumnSchema::new("Variable_name", ConcreteDataType::string_datatype(), true),
171 ColumnSchema::new("Value", ConcreteDataType::string_datatype(), true),
172 ]));
173 let columns = vec![
174 Arc::new(StringVector::from(vec![name])) as _,
175 Arc::new(StringVector::from(vec![value])) as _,
176 ];
177 RecordBatches::try_from_columns(schema, columns)
178 .unwrap()
180}
181
182fn select_variable(query: &str, query_context: QueryContextRef) -> Option<Output> {
183 let mut fields = vec![];
184 let mut values = vec![];
185
186 let query = query.to_lowercase();
188 let vars: Vec<&str> = query.split("@@").collect();
189 if vars.len() <= 1 {
190 return None;
191 }
192
193 for var in vars.iter().skip(1) {
195 let var = var.trim_matches(|c| c == ' ' || c == ',' || c == ';');
196 let var_as: Vec<&str> = var
197 .split(" as ")
198 .map(|x| {
199 x.trim_matches(|c| c == ' ')
200 .split_whitespace()
201 .next()
202 .unwrap_or("")
203 })
204 .collect();
205
206 let value = match var_as[0] {
208 "session.time_zone" | "time_zone" => query_context.timezone().to_string(),
209 "system_time_zone" => system_timezone_name(),
210 "max_execution_time" | "session.max_execution_time" => {
211 query_context.query_timeout_as_millis().to_string()
212 }
213 _ => VAR_VALUES
214 .get(var_as[0])
215 .map(|v| v.to_string())
216 .unwrap_or_else(|| "0".to_owned()),
217 };
218
219 values.push(Arc::new(StringVector::from(vec![value])) as _);
220 match var_as.len() {
221 1 => {
222 fields.push(ColumnSchema::new(
225 format!("@@{}", var_as[0]),
226 ConcreteDataType::string_datatype(),
227 true,
228 ));
229 }
230 2 => {
231 fields.push(ColumnSchema::new(
235 var_as[1],
236 ConcreteDataType::string_datatype(),
237 true,
238 ));
239 }
240 _ => return None,
241 }
242 }
243
244 let schema = Arc::new(Schema::new(fields));
245 let batches = RecordBatches::try_from_columns(schema, values).unwrap();
247 Some(Output::new_with_record_batches(batches))
248}
249
250fn check_select_variable(query: &str, query_context: QueryContextRef) -> Option<Output> {
251 if SELECT_VAR_PATTERN.is_match(query) {
252 select_variable(query, query_context)
253 } else {
254 None
255 }
256}
257
258fn check_select_user_or_var(query: &str, query_context: QueryContextRef) -> Option<Output> {
259 let captures = SELECT_USER_OR_VAR_PATTERN.captures(query)?;
260
261 let recordbatches = if let Some(keyword) = captures.get(1) {
262 let user = query_context.current_user();
263 select_function(keyword.as_str(), Some(user.username()))
264 } else {
265 let var = captures
268 .get(2)
269 .expect("one of the two groups always matches");
270 select_function(var.as_str(), None)
271 };
272 Some(Output::new_with_record_batches(recordbatches))
273}
274
275fn check_show_variables(query: &str) -> Option<Output> {
276 let recordbatches = if SHOW_SQL_MODE_PATTERN.is_match(query) {
277 Some(show_variables(
278 "sql_mode",
279 "ONLY_FULL_GROUP_BY STRICT_TRANS_TABLES NO_ZERO_IN_DATE NO_ZERO_DATE ERROR_FOR_DIVISION_BY_ZERO NO_ENGINE_SUBSTITUTION",
280 ))
281 } else if SHOW_LOWER_CASE_PATTERN.is_match(query) {
282 Some(show_variables("lower_case_table_names", "0"))
283 } else if SHOW_VARIABLES_LIKE_PATTERN.is_match(query) {
284 Some(show_variables("", ""))
285 } else {
286 None
287 };
288 recordbatches.map(Output::new_with_record_batches)
289}
290
291fn show_warnings(session: &SessionRef) -> RecordBatches {
293 let schema = Arc::new(Schema::new(vec![
294 ColumnSchema::new("Level", ConcreteDataType::string_datatype(), false),
295 ColumnSchema::new("Code", ConcreteDataType::uint16_datatype(), false),
296 ColumnSchema::new("Message", ConcreteDataType::string_datatype(), false),
297 ]));
298
299 let warnings = session.warnings();
300 let count = warnings.len();
301
302 let columns = if count > 0 {
303 vec![
304 Arc::new(StringVector::from(vec!["Warning"; count])) as _,
305 Arc::new(datatypes::vectors::UInt16Vector::from(vec![
306 Some(1000u16);
307 count
308 ])) as _,
309 Arc::new(StringVector::from(warnings)) as _,
310 ]
311 } else {
312 vec![
313 Arc::new(StringVector::from(Vec::<String>::new())) as _,
314 Arc::new(datatypes::vectors::UInt16Vector::from(
315 Vec::<Option<u16>>::new(),
316 )) as _,
317 Arc::new(StringVector::from(Vec::<String>::new())) as _,
318 ]
319 };
320
321 RecordBatches::try_from_columns(schema, columns).unwrap()
322}
323
324fn check_show_warnings(query: &str, session: &SessionRef) -> Option<Output> {
325 if SHOW_WARNINGS_PATTERN.is_match(query) {
326 Some(Output::new_with_record_batches(show_warnings(session)))
327 } else {
328 None
329 }
330}
331
332fn check_others(query: &str, _query_ctx: QueryContextRef) -> Option<Output> {
334 if OTHER_NOT_SUPPORTED_STMT.is_match(query.as_bytes()) {
335 return Some(Output::new_with_record_batches(RecordBatches::empty()));
336 }
337
338 let recordbatches = if SELECT_TIME_DIFF_FUNC_PATTERN.is_match(query) {
339 Some(select_function(
340 "TIMEDIFF(NOW(), UTC_TIMESTAMP())",
341 Some("00:00:00"),
342 ))
343 } else {
344 None
345 };
346 recordbatches.map(Output::new_with_record_batches)
347}
348
349fn strip_leading_comments(query: &str) -> &str {
356 let mut rest = query.trim_start();
357 loop {
358 if rest.starts_with("/*!") {
362 return rest;
363 }
364 if let Some(tail) = rest.strip_prefix("/*") {
365 let Some(end) = tail.find("*/") else {
367 return "";
368 };
369 rest = tail[end + 2..].trim_start();
370 } else if rest.starts_with('#')
371 || (rest.starts_with("--")
373 && rest[2..].chars().next().is_none_or(|c| c.is_whitespace()))
374 {
375 let Some(end) = rest.find('\n') else {
376 return "";
377 };
378 rest = rest[end + 1..].trim_start();
379 } else {
380 return rest;
381 }
382 }
383}
384
385const OTHER_LEADING_KEYWORDS: [&str; 7] = [
391 "SET", "COMMIT", "ROLLBACK", "START", "BEGIN", "LOCK", "UNLOCK",
392];
393
394fn leading_keyword(query: &str) -> &str {
397 let end = query
398 .find(|c: char| !c.is_ascii_alphabetic())
399 .unwrap_or(query.len());
400 &query[..end]
401}
402
403fn line_comment_end(bytes: &[u8], from: usize) -> usize {
405 bytes[from..]
406 .iter()
407 .position(|c| *c == b'\n')
408 .map_or(bytes.len(), |p| from + p + 1)
409}
410
411fn quoted_end(bytes: &[u8], start: usize, quote: u8) -> usize {
413 let mut i = start + 1;
414 while i < bytes.len() {
415 match bytes[i] {
416 b'\\' if quote != b'`' => i += 2,
418 c if c == quote => {
419 if bytes.get(i + 1) == Some("e) {
421 i += 2;
422 } else {
423 return i + 1;
424 }
425 }
426 _ => i += 1,
427 }
428 }
429 bytes.len()
430}
431
432fn has_trailing_statement(query: &str) -> bool {
450 let bytes = query.as_bytes();
451 let mut i = 0;
452 let mut seen_semicolon = false;
453 let mut seen_content = false;
454
455 while i < bytes.len() {
458 match bytes[i] {
459 b'/' if bytes.get(i + 1) == Some(&b'*') => {
460 if bytes.get(i + 2) == Some(&b'!') && seen_content {
461 return true;
462 }
463 i = match bytes[i + 2..].windows(2).position(|w| w == b"*/") {
464 Some(p) => i + 2 + p + 2,
465 None => bytes.len(),
467 };
468 seen_content = true;
469 }
470 b'#' => {
471 i = line_comment_end(bytes, i);
472 seen_content = true;
473 }
474 b'-' if bytes.get(i + 1) == Some(&b'-')
476 && bytes.get(i + 2).is_none_or(|c| c.is_ascii_whitespace()) =>
477 {
478 i = line_comment_end(bytes, i);
479 seen_content = true;
480 }
481 quote @ (b'\'' | b'"' | b'`') => {
482 if seen_semicolon {
483 return true;
484 }
485 seen_content = true;
486 i = quoted_end(bytes, i, quote);
487 }
488 b';' => {
489 seen_semicolon = true;
490 seen_content = true;
491 i += 1;
492 }
493 c if c.is_ascii_whitespace() => i += 1,
494 _ => {
495 if seen_semicolon {
496 return true;
497 }
498 seen_content = true;
499 i += 1;
500 }
501 }
502 }
503
504 false
505}
506
507pub(crate) fn check(
510 query: &str,
511 query_ctx: QueryContextRef,
512 session: SessionRef,
513) -> Option<Output> {
514 let query = strip_leading_comments(query);
515 let keyword = leading_keyword(query);
516
517 let absorbed = if keyword.eq_ignore_ascii_case("SELECT") {
520 check_select_variable(query, query_ctx.clone())
522 .or_else(|| check_select_user_or_var(query, query_ctx.clone()))
523 .or_else(|| check_others(query, query_ctx))
524 } else if keyword.eq_ignore_ascii_case("SHOW") {
525 check_show_variables(query)
526 .or_else(|| check_show_warnings(query, &session))
527 .or_else(|| check_others(query, query_ctx))
528 } else if query.starts_with("/*!")
529 || OTHER_LEADING_KEYWORDS
530 .iter()
531 .any(|k| k.eq_ignore_ascii_case(keyword))
532 {
533 check_others(query, query_ctx)
534 } else {
535 return None;
536 };
537
538 if absorbed.is_some() && has_trailing_statement(query) {
541 return None;
542 }
543
544 absorbed
545}
546
547#[cfg(test)]
548mod test {
549
550 use common_query::OutputData;
551 use common_time::timezone::set_default_timezone;
552 use session::Session;
553 use session::context::{Channel, QueryContext};
554
555 use super::*;
556
557 #[test]
558 fn test_check_abnormal() {
559 let session = Arc::new(Session::new(None, Channel::Mysql, Default::default(), 0));
560 let query = "🫣一点不正常的东西🫣";
561 let output = check(query, QueryContext::arc(), session.clone());
562
563 assert!(output.is_none());
564 }
565
566 #[test]
567 fn test_check() {
568 let session = Arc::new(Session::new(None, Channel::Mysql, Default::default(), 0));
569 let query = "select 1";
570 let result = check(query, QueryContext::arc(), session.clone());
571 assert!(result.is_none());
572
573 let query = "select version";
574 let output = check(query, QueryContext::arc(), session.clone());
575 assert!(output.is_none());
576
577 fn test(query: &str, expected: &str) {
578 let session = Arc::new(Session::new(None, Channel::Mysql, Default::default(), 0));
579 let output = check(query, QueryContext::arc(), session.clone());
580 match output.unwrap().data {
581 OutputData::RecordBatches(r) => {
582 assert_eq!(&r.pretty_print().unwrap(), expected)
583 }
584 _ => unreachable!(),
585 }
586 }
587
588 let query = "SELECT @@version_comment LIMIT 1";
589 let expected = "\
590+-------------------+
591| @@version_comment |
592+-------------------+
593| GreptimeDB |
594+-------------------+";
595 test(query, expected);
596
597 let query = "select @@tx_isolation, @@session.tx_isolation";
599 let expected = "\
600+-----------------+------------------------+
601| @@tx_isolation | @@session.tx_isolation |
602+-----------------+------------------------+
603| REPEATABLE-READ | REPEATABLE-READ |
604+-----------------+------------------------+";
605 test(query, expected);
606
607 set_default_timezone(Some("Asia/Shanghai")).unwrap();
609 let query = "/* mysql-connector-java-8.0.17 (Revision: 16a712ddb3f826a1933ab42b0039f7fb9eebc6ec) */SELECT @@session.auto_increment_increment AS auto_increment_increment, @@character_set_client AS character_set_client, @@character_set_connection AS character_set_connection, @@character_set_results AS character_set_results, @@character_set_server AS character_set_server, @@collation_server AS collation_server, @@collation_connection AS collation_connection, @@init_connect AS init_connect, @@interactive_timeout AS interactive_timeout, @@license AS license, @@lower_case_table_names AS lower_case_table_names, @@max_allowed_packet AS max_allowed_packet, @@net_write_timeout AS net_write_timeout, @@performance_schema AS performance_schema, @@sql_mode AS sql_mode, @@system_time_zone AS system_time_zone, @@time_zone AS time_zone, @@transaction_isolation AS transaction_isolation, @@wait_timeout AS wait_timeout;";
611 let expected = "\
612+--------------------------+----------------------+--------------------------+-----------------------+----------------------+------------------+----------------------+--------------+---------------------+---------+------------------------+--------------------+-------------------+--------------------+----------+------------------+---------------+-----------------------+--------------+
613| auto_increment_increment | character_set_client | character_set_connection | character_set_results | character_set_server | collation_server | collation_connection | init_connect | interactive_timeout | license | lower_case_table_names | max_allowed_packet | net_write_timeout | performance_schema | sql_mode | system_time_zone | time_zone | transaction_isolation | wait_timeout |
614+--------------------------+----------------------+--------------------------+-----------------------+----------------------+------------------+----------------------+--------------+---------------------+---------+------------------------+--------------------+-------------------+--------------------+----------+------------------+---------------+-----------------------+--------------+
615| 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 31536000 | 0 | 0 | 134217728 | 31536000 | 0 | 0 | Asia/Shanghai | Asia/Shanghai | REPEATABLE-READ | 31536000 |
616+--------------------------+----------------------+--------------------------+-----------------------+----------------------+------------------+----------------------+--------------+---------------------+---------+------------------------+--------------------+-------------------+--------------------+----------+------------------+---------------+-----------------------+--------------+";
617 test(query, expected);
618
619 let query = "show variables";
620 let expected = "\
621+---------------+-------+
622| Variable_name | Value |
623+---------------+-------+
624| | |
625+---------------+-------+";
626 test(query, expected);
627
628 let query = "show variables like 'lower_case_table_names'";
629 let expected = "\
630+------------------------+-------+
631| Variable_name | Value |
632+------------------------+-------+
633| lower_case_table_names | 0 |
634+------------------------+-------+";
635 test(query, expected);
636
637 let query = "SELECT TIMEDIFF(NOW(), UTC_TIMESTAMP())";
638 let expected = "\
639+----------------------------------+
640| TIMEDIFF(NOW(), UTC_TIMESTAMP()) |
641+----------------------------------+
642| 00:00:00 |
643+----------------------------------+";
644 test(query, expected);
645 }
646
647 #[test]
648 fn test_show_warnings() {
649 let session = Arc::new(Session::new(None, Channel::Mysql, Default::default(), 0));
651 let output = check("SHOW WARNINGS", QueryContext::arc(), session.clone());
652 match output.unwrap().data {
653 OutputData::RecordBatches(r) => {
654 assert_eq!(r.iter().map(|b| b.num_rows()).sum::<usize>(), 0);
655 }
656 _ => unreachable!(),
657 }
658
659 session.add_warning("Test warning message".to_string());
661 let output = check("SHOW WARNINGS", QueryContext::arc(), session.clone());
662 match output.unwrap().data {
663 OutputData::RecordBatches(r) => {
664 let expected = "\
665+---------+------+----------------------+
666| Level | Code | Message |
667+---------+------+----------------------+
668| Warning | 1000 | Test warning message |
669+---------+------+----------------------+";
670 assert_eq!(&r.pretty_print().unwrap(), expected);
671 }
672 _ => unreachable!(),
673 }
674
675 session.clear_warnings();
677 session.add_warning("First warning".to_string());
678 session.add_warning("Second warning".to_string());
679 let output = check("SHOW WARNINGS", QueryContext::arc(), session.clone());
680 match output.unwrap().data {
681 OutputData::RecordBatches(r) => {
682 let expected = "\
683+---------+------+----------------+
684| Level | Code | Message |
685+---------+------+----------------+
686| Warning | 1000 | First warning |
687| Warning | 1000 | Second warning |
688+---------+------+----------------+";
689 assert_eq!(&r.pretty_print().unwrap(), expected);
690 }
691 _ => unreachable!(),
692 }
693
694 let output = check("show warnings", QueryContext::arc(), session.clone());
696 assert!(output.is_some());
697
698 let output = check(
700 "/* ApplicationName=DBeaver */SHOW WARNINGS",
701 QueryContext::arc(),
702 session.clone(),
703 );
704 assert!(output.is_some());
705 }
706
707 #[test]
708 fn test_check_select_user_or_var() {
709 let session = Arc::new(Session::new(None, Channel::Mysql, Default::default(), 0));
710
711 fn pretty(query: &str, session: &SessionRef) -> String {
712 let output = check(query, QueryContext::arc(), session.clone())
713 .unwrap_or_else(|| panic!("{query} was not absorbed"));
714 let OutputData::RecordBatches(batches) = output.data else {
715 unreachable!()
716 };
717 batches.pretty_print().unwrap()
718 }
719
720 assert_eq!(
722 pretty("SELECT CURRENT_USER", &session),
723 "\
724+--------------+
725| CURRENT_USER |
726+--------------+
727| greptime |
728+--------------+"
729 );
730 assert_eq!(
731 pretty("select session_user;", &session),
732 "\
733+--------------+
734| session_user |
735+--------------+
736| greptime |
737+--------------+"
738 );
739
740 assert_eq!(
742 pretty("SELECT @v", &session),
743 "\
744+----+
745| @v |
746+----+
747| |
748+----+"
749 );
750
751 for query in [
754 "SELECT user FROM t",
755 "SELECT current_user, 1",
756 "SELECT @v FROM t",
757 "SELECT @v + 1",
758 "SELECT userid",
759 "SELECT 1",
760 ] {
761 assert!(
762 check(query, QueryContext::arc(), session.clone()).is_none(),
763 "{query} must not be absorbed"
764 );
765 }
766 }
767
768 #[test]
771 fn test_check_skips_multi_statement() {
772 let session = Arc::new(Session::new(None, Channel::Mysql, Default::default(), 0));
773 for query in [
774 "BEGIN; INSERT INTO t VALUES (1); COMMIT",
775 "BEGIN;\nINSERT INTO t VALUES (1)",
776 "START TRANSACTION; DELETE FROM t",
777 "COMMIT; INSERT INTO t VALUES (1)",
778 "SET NAMES utf8mb4; INSERT INTO t VALUES (1)",
779 "SELECT @@version; INSERT INTO t VALUES (1)",
780 ] {
781 assert!(
782 check(query, QueryContext::arc(), session.clone()).is_none(),
783 "{query} must not be absorbed"
784 );
785 }
786
787 for query in [
789 "BEGIN;",
790 "COMMIT; ",
791 "SET NAMES utf8mb4;\n",
792 "BEGIN; -- done",
793 "COMMIT; /* done */",
794 "BEGIN; # done",
795 "BEGIN;;",
796 "COMMIT; ; /* done */ ;",
797 "BEGIN; -- done\n",
798 ] {
799 assert!(
800 check(query, QueryContext::arc(), session.clone()).is_some(),
801 "{query} was not absorbed"
802 );
803 }
804
805 for query in [
807 "BEGIN; -- go\nINSERT INTO t VALUES (1)",
808 "COMMIT; /* go */ INSERT INTO t VALUES (1)",
809 "BEGIN;; INSERT INTO t VALUES (1)",
810 "BEGIN /* previous delimiter; -- note */; INSERT INTO t VALUES (1)",
812 "SET NAMES 'a;b'; INSERT INTO t VALUES (1)",
813 "BEGIN; /*! INSERT INTO t VALUES (1) */",
815 "BEGIN /*! INSERT INTO t VALUES (1) */",
816 "/*!40101 SET NAMES utf8mb4 */; INSERT INTO t VALUES (1)",
817 "/*!40101 SET NAMES utf8mb4 */ /*! INSERT INTO t VALUES (1) */",
818 ] {
819 assert!(
820 check(query, QueryContext::arc(), session.clone()).is_none(),
821 "{query} must not be absorbed"
822 );
823 }
824
825 for query in [
827 "BEGIN /* previous delimiter; -- note */",
828 "BEGIN -- a; b",
829 "SET NAMES 'a;b'",
830 "SET NAMES \"a;b\"",
831 "SET NAMES 'it\\'s; here'",
832 "SET NAMES 'a;b';",
833 ] {
834 assert!(
835 check(query, QueryContext::arc(), session.clone()).is_some(),
836 "{query} was not absorbed"
837 );
838 }
839 }
840
841 #[test]
842 fn test_check_skips_non_federated_keywords() {
843 let session = Arc::new(Session::new(None, Channel::Mysql, Default::default(), 0));
844 for query in [
845 "INSERT INTO t VALUES (1)",
846 "UPDATE t SET a = 1",
847 "DELETE FROM t",
848 "CREATE TABLE t (ts TIMESTAMP TIME INDEX)",
849 "WITH x AS (SELECT 1) SELECT * FROM x",
850 "TQL EVAL (0, 10, '5s') up",
851 ] {
852 assert!(
853 check(query, QueryContext::arc(), session.clone()).is_none(),
854 "{query} must not be absorbed"
855 );
856 }
857 }
858
859 #[test]
860 fn test_strip_leading_comments() {
861 assert_eq!(strip_leading_comments("SELECT 1"), "SELECT 1");
862 assert_eq!(strip_leading_comments(" \n\tSELECT 1"), "SELECT 1");
863 assert_eq!(
864 strip_leading_comments("/* ApplicationName=DataGrip 2026.2.5 */ COMMIT"),
865 "COMMIT"
866 );
867 assert_eq!(strip_leading_comments("/* a */ /* b */COMMIT"), "COMMIT");
868 assert_eq!(strip_leading_comments("-- a comment\nCOMMIT"), "COMMIT");
869 assert_eq!(strip_leading_comments("# a comment\nCOMMIT"), "COMMIT");
870 assert_eq!(strip_leading_comments("--x\nCOMMIT"), "--x\nCOMMIT");
872 assert_eq!(strip_leading_comments("/* unterminated"), "");
874 assert_eq!(strip_leading_comments("-- trailing"), "");
875 assert_eq!(
877 strip_leading_comments("/*!40101 SET NAMES utf8mb4 */"),
878 "/*!40101 SET NAMES utf8mb4 */"
879 );
880 assert_eq!(
881 strip_leading_comments("/* App */ /*!40101 SET NAMES utf8mb4 */"),
882 "/*!40101 SET NAMES utf8mb4 */"
883 );
884 assert_eq!(
886 strip_leading_comments("/* a */SELECT /* b */ 1"),
887 "SELECT /* b */ 1"
888 );
889 }
890
891 #[test]
894 fn test_check_comment_prefixed() {
895 let session = Arc::new(Session::new(None, Channel::Mysql, Default::default(), 0));
896 let prefix = "/* ApplicationName=DataGrip 2026.2.5 */ ";
897
898 for query in [
899 "SET TRANSACTION READ WRITE",
900 "SET SESSION TRANSACTION READ ONLY",
901 "SET NAMES utf8mb4",
902 "BEGIN",
903 "START TRANSACTION",
904 "COMMIT",
905 "ROLLBACK",
906 "LOCK TABLES t WRITE",
908 "UNLOCK TABLES",
909 ] {
910 let output = check(
911 &format!("{prefix}{query}"),
912 QueryContext::arc(),
913 session.clone(),
914 );
915 let OutputData::RecordBatches(batches) = output
916 .unwrap_or_else(|| panic!("{query} was not absorbed"))
917 .data
918 else {
919 unreachable!()
920 };
921 assert_eq!(batches.iter().map(|b| b.num_rows()).sum::<usize>(), 0);
922 }
923
924 for query in [
927 "/*!40101 SET NAMES utf8mb4 */",
928 "/*!40014 SET @OLD_UNIQUE_CHECKS=@@UNIQUE_CHECKS, UNIQUE_CHECKS=0 */",
929 "/*!40111 SET @OLD_SQL_NOTES=@@SQL_NOTES, SQL_NOTES=0 */",
930 "/*!80003 SET @OLD_x=1 */",
931 ] {
932 let output = check(query, QueryContext::arc(), session.clone());
933 let OutputData::RecordBatches(batches) = output
934 .unwrap_or_else(|| panic!("{query} was not absorbed"))
935 .data
936 else {
937 unreachable!()
938 };
939 assert_eq!(batches.iter().map(|b| b.num_rows()).sum::<usize>(), 0);
940 }
941
942 for query in [
945 "SELECT @@GLOBAL.event_scheduler",
946 "SHOW GLOBAL VARIABLES LIKE 'event_scheduler'",
947 ] {
948 let output = check(
949 &format!("{prefix}{query}"),
950 QueryContext::arc(),
951 session.clone(),
952 );
953 let OutputData::RecordBatches(batches) = output
954 .unwrap_or_else(|| panic!("{query} was not absorbed"))
955 .data
956 else {
957 unreachable!()
958 };
959 assert!(!batches.schema().column_schemas().is_empty(), "{query}");
960 }
961 }
962}