1use std::collections::HashMap;
62
63#[derive(Debug, Clone, PartialEq)]
65pub enum DynamicSqlError {
66 ParseError(String),
68 StatementNotFound(String),
70 EvalError(String),
72 MissingParam(String),
74}
75
76impl std::fmt::Display for DynamicSqlError {
77 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
78 match self {
79 DynamicSqlError::ParseError(msg) => write!(f, "XML 解析错误: {}", msg),
80 DynamicSqlError::StatementNotFound(id) => {
81 write!(f, "找不到语句 ID: {}", id)
82 }
83 DynamicSqlError::EvalError(msg) => write!(f, "表达式求值错误: {}", msg),
84 DynamicSqlError::MissingParam(name) => write!(f, "缺少参数: {}", name),
85 }
86 }
87}
88
89impl std::error::Error for DynamicSqlError {}
90
91#[derive(Debug, Clone, Default)]
95pub struct SqlParams {
96 params: HashMap<String, ParamValue>,
97}
98
99#[derive(Debug, Clone)]
101pub enum ParamValue {
102 Null,
104 String(String),
106 Int(i64),
108 Float(f64),
110 Bool(bool),
112 Array(Vec<ParamValue>),
114}
115
116impl SqlParams {
117 pub fn new() -> Self {
119 Self::default()
120 }
121
122 pub fn set(&mut self, name: &str, value: &str) {
124 self.params
125 .insert(name.to_string(), ParamValue::String(value.to_string()));
126 }
127
128 pub fn set_int(&mut self, name: &str, value: i64) {
130 self.params.insert(name.to_string(), ParamValue::Int(value));
131 }
132
133 pub fn set_float(&mut self, name: &str, value: f64) {
135 self.params
136 .insert(name.to_string(), ParamValue::Float(value));
137 }
138
139 pub fn set_bool(&mut self, name: &str, value: bool) {
141 self.params
142 .insert(name.to_string(), ParamValue::Bool(value));
143 }
144
145 pub fn set_null(&mut self, name: &str) {
147 self.params.insert(name.to_string(), ParamValue::Null);
148 }
149
150 pub fn set_array(&mut self, name: &str, values: Vec<ParamValue>) {
152 self.params
153 .insert(name.to_string(), ParamValue::Array(values));
154 }
155
156 pub fn get(&self, name: &str) -> Option<&ParamValue> {
158 self.params.get(name)
159 }
160
161 pub fn contains(&self, name: &str) -> bool {
163 self.params.contains_key(name)
164 }
165
166 pub fn is_null(&self, name: &str) -> bool {
168 matches!(self.params.get(name), None | Some(ParamValue::Null))
169 }
170
171 pub fn is_not_null(&self, name: &str) -> bool {
173 !self.is_null(name)
174 }
175
176 pub fn names(&self) -> Vec<String> {
178 self.params.keys().cloned().collect()
179 }
180}
181
182#[derive(Debug, Clone)]
184pub struct DynamicSqlParser {
185 statements: HashMap<String, XmlNode>,
187}
188
189impl DynamicSqlParser {
190 pub fn new() -> Self {
192 Self {
193 statements: HashMap::new(),
194 }
195 }
196
197 pub fn from_xml(xml: &str) -> Result<Self, DynamicSqlError> {
199 let mut parser = Self::new();
200 parser.parse_xml(xml)?;
201 Ok(parser)
202 }
203
204 fn parse_xml(&mut self, xml: &str) -> Result<(), DynamicSqlError> {
206 let root = XmlParser::parse(xml)?;
207 for child in &root.children {
208 if let XmlNodeType::Element { name, attrs } = &child.node_type {
209 let id = attrs
210 .get("id")
211 .ok_or_else(|| DynamicSqlError::ParseError(format!("<{}> 缺少 id 属性", name)))?
212 .clone();
213 self.statements.insert(id, child.clone());
214 }
215 }
216 Ok(())
217 }
218
219 pub fn build(&self, id: &str, params: &SqlParams) -> Result<String, DynamicSqlError> {
221 let node = self
222 .statements
223 .get(id)
224 .ok_or_else(|| DynamicSqlError::StatementNotFound(id.to_string()))?;
225 let mut ctx = BuildContext::new(params);
226 self.build_node(node, &mut ctx)?;
227 Ok(self.cleanup_sql(&ctx.buffer))
228 }
229
230 pub fn build_with_binds(
232 &self,
233 id: &str,
234 params: &SqlParams,
235 ) -> Result<(String, Vec<ParamValue>), DynamicSqlError> {
236 let node = self
237 .statements
238 .get(id)
239 .ok_or_else(|| DynamicSqlError::StatementNotFound(id.to_string()))?;
240 let mut ctx = BuildContext::new(params);
241 self.build_node(node, &mut ctx)?;
242 Ok((self.cleanup_sql(&ctx.buffer), ctx.binds))
243 }
244
245 pub fn statement_ids(&self) -> Vec<String> {
247 let mut ids: Vec<String> = self.statements.keys().cloned().collect();
248 ids.sort();
249 ids
250 }
251
252 fn build_node(&self, node: &XmlNode, ctx: &mut BuildContext) -> Result<(), DynamicSqlError> {
255 match &node.node_type {
256 XmlNodeType::Text(text) => {
257 self.append_text(text, ctx)?;
258 }
259 XmlNodeType::Element { name, attrs } => {
260 match name.as_str() {
261 "select" | "insert" | "update" | "delete" => {
262 for child in &node.children {
263 self.build_node(child, ctx)?;
264 }
265 }
266 "if" => {
267 let test = attrs.get("test").ok_or_else(|| {
268 DynamicSqlError::ParseError("<if> 缺少 test 属性".into())
269 })?;
270 if eval_test(test, ctx.params)? {
271 for child in &node.children {
272 self.build_node(child, ctx)?;
273 }
274 }
275 }
276 "where" => {
277 let mut sub_ctx = BuildContext::new(ctx.params);
278 for child in &node.children {
279 self.build_node(child, &mut sub_ctx)?;
280 }
281 let content = sub_ctx.buffer.trim();
282 if !content.is_empty() {
283 let cleaned = strip_leading_and_or(content);
285 ctx.buffer.push_str(" WHERE ");
286 ctx.buffer.push_str(cleaned.trim());
287 ctx.binds.extend(sub_ctx.binds);
289 }
290 }
291 "set" => {
292 let mut sub_ctx = BuildContext::new(ctx.params);
293 for child in &node.children {
294 self.build_node(child, &mut sub_ctx)?;
295 }
296 let content = sub_ctx.buffer.trim();
297 if !content.is_empty() {
298 let cleaned = content.trim_end_matches(',').trim();
300 let normalized = normalize_set_commas(cleaned);
302 ctx.buffer.push_str(" SET ");
303 ctx.buffer.push_str(&normalized);
304 ctx.binds.extend(sub_ctx.binds);
305 }
306 }
307 "foreach" => {
308 self.build_foreach(node, attrs, ctx)?;
309 }
310 "choose" => {
311 self.build_choose(node, ctx)?;
312 }
313 "trim" => {
314 self.build_trim(node, attrs, ctx)?;
315 }
316 _ => {
317 for child in &node.children {
319 self.build_node(child, ctx)?;
320 }
321 }
322 }
323 }
324 }
325 Ok(())
326 }
327
328 fn build_foreach(
329 &self,
330 node: &XmlNode,
331 attrs: &HashMap<String, String>,
332 ctx: &mut BuildContext,
333 ) -> Result<(), DynamicSqlError> {
334 let collection = attrs
335 .get("collection")
336 .ok_or_else(|| DynamicSqlError::ParseError("<foreach> 缺少 collection 属性".into()))?;
337 let item = attrs.get("item").map(|s| s.as_str()).unwrap_or("item");
338 let separator = attrs.get("separator").map(|s| s.as_str()).unwrap_or(",");
339 let open = attrs.get("open").cloned().unwrap_or_default();
340 let close = attrs.get("close").cloned().unwrap_or_default();
341
342 let arr = match ctx.params.get(collection) {
343 Some(ParamValue::Array(arr)) => arr.clone(),
344 _ => return Ok(()),
345 };
346
347 let mut sub_params = ctx.params.clone();
349 let mut parts: Vec<String> = Vec::new();
350 for v in &arr {
351 match v {
353 ParamValue::String(s) => sub_params.set(item, s),
354 ParamValue::Int(i) => sub_params.set_int(item, *i),
355 ParamValue::Float(f) => sub_params.set_float(item, *f),
356 ParamValue::Bool(b) => sub_params.set_bool(item, *b),
357 ParamValue::Null => sub_params.set_null(item),
358 ParamValue::Array(_) => {} }
360 let mut sub_ctx = BuildContext::new(&sub_params);
361 for child in &node.children {
362 self.build_node(child, &mut sub_ctx)?;
363 }
364 parts.push(sub_ctx.buffer.trim().to_string());
365 ctx.binds.extend(sub_ctx.binds);
367 }
368 if !parts.is_empty() {
369 let joined = parts.join(separator);
370 ctx.buffer.push(' ');
371 if !open.is_empty() {
372 ctx.buffer.push_str(&open);
373 }
374 ctx.buffer.push_str(&joined);
375 if !close.is_empty() {
376 ctx.buffer.push_str(&close);
377 }
378 }
379 Ok(())
380 }
381
382 fn build_choose(&self, node: &XmlNode, ctx: &mut BuildContext) -> Result<(), DynamicSqlError> {
383 for child in &node.children {
384 if let XmlNodeType::Element { name, attrs } = &child.node_type {
385 match name.as_str() {
386 "when" => {
387 let test = attrs.get("test").ok_or_else(|| {
388 DynamicSqlError::ParseError("<when> 缺少 test 属性".into())
389 })?;
390 if eval_test(test, ctx.params)? {
391 for c in &child.children {
392 self.build_node(c, ctx)?;
393 }
394 return Ok(());
395 }
396 }
397 "otherwise" => {
398 for c in &child.children {
399 self.build_node(c, ctx)?;
400 }
401 return Ok(());
402 }
403 _ => {}
404 }
405 }
406 }
407 Ok(())
408 }
409
410 fn build_trim(
411 &self,
412 node: &XmlNode,
413 attrs: &HashMap<String, String>,
414 ctx: &mut BuildContext,
415 ) -> Result<(), DynamicSqlError> {
416 let prefix = attrs.get("prefix").cloned().unwrap_or_default();
417 let suffix = attrs.get("suffix").cloned().unwrap_or_default();
418 let prefix_overrides = attrs.get("prefixOverrides").cloned().unwrap_or_default();
419 let suffix_overrides = attrs.get("suffixOverrides").cloned().unwrap_or_default();
420
421 let mut sub_ctx = BuildContext::new(ctx.params);
422 for child in &node.children {
423 self.build_node(child, &mut sub_ctx)?;
424 }
425 let mut content = sub_ctx.buffer.trim().to_string();
426
427 if !prefix_overrides.is_empty() {
429 for ov in prefix_overrides.split('|') {
430 if content.starts_with(ov) {
431 content = content[ov.len()..].trim_start().to_string();
432 break;
433 }
434 }
435 }
436 if !suffix_overrides.is_empty() {
438 for ov in suffix_overrides.split('|') {
439 if content.ends_with(ov) {
440 content = content[..content.len() - ov.len()].trim_end().to_string();
441 break;
442 }
443 }
444 }
445
446 if !content.is_empty() {
447 ctx.buffer.push(' ');
448 if !prefix.is_empty() {
449 ctx.buffer.push_str(&prefix);
450 ctx.buffer.push(' ');
451 }
452 ctx.buffer.push_str(&content);
453 if !suffix.is_empty() {
454 ctx.buffer.push(' ');
455 ctx.buffer.push_str(&suffix);
456 }
457 ctx.binds.extend(sub_ctx.binds);
459 }
460 Ok(())
461 }
462
463 fn append_text(&self, text: &str, ctx: &mut BuildContext) -> Result<(), DynamicSqlError> {
464 let mut i = 0;
465 let bytes = text.as_bytes();
466 while i < bytes.len() {
467 if i + 1 < bytes.len() && bytes[i] == b'#' && bytes[i + 1] == b'{' {
468 let end = text[i + 2..].find('}').ok_or_else(|| {
470 DynamicSqlError::ParseError(format!("未闭合的 #{{}}: {}", &text[i..]))
471 })?;
472 let name = &text[i + 2..i + 2 + end];
473 let value = ctx
474 .params
475 .get(name)
476 .ok_or_else(|| DynamicSqlError::MissingParam(name.to_string()))?
477 .clone();
478 ctx.buffer.push('?');
479 ctx.binds.push(value);
480 i += 2 + end + 1; } else if i + 1 < bytes.len() && bytes[i] == b'$' && bytes[i + 1] == b'{' {
482 let end = text[i + 2..].find('}').ok_or_else(|| {
484 DynamicSqlError::ParseError(format!("未闭合的 ${{}}: {}", &text[i..]))
485 })?;
486 let name = &text[i + 2..i + 2 + end];
487 let value = ctx
488 .params
489 .get(name)
490 .ok_or_else(|| DynamicSqlError::MissingParam(name.to_string()))?;
491 let s = param_to_string(value);
492 ctx.buffer.push_str(&s);
493 i += 2 + end + 1;
494 } else {
495 ctx.buffer.push(bytes[i] as char);
496 i += 1;
497 }
498 }
499 Ok(())
500 }
501
502 fn cleanup_sql(&self, sql: &str) -> String {
504 let mut result = String::with_capacity(sql.len());
505 let mut prev_space = false;
506 for c in sql.chars() {
507 if c.is_whitespace() {
508 if !prev_space {
509 result.push(' ');
510 prev_space = true;
511 }
512 } else {
513 result.push(c);
514 prev_space = false;
515 }
516 }
517 result.trim().to_string()
518 }
519}
520
521impl Default for DynamicSqlParser {
522 fn default() -> Self {
523 Self::new()
524 }
525}
526
527struct BuildContext<'a> {
529 buffer: String,
530 binds: Vec<ParamValue>,
531 params: &'a SqlParams,
532}
533
534impl<'a> BuildContext<'a> {
535 fn new(params: &'a SqlParams) -> Self {
536 Self {
537 buffer: String::new(),
538 binds: Vec::new(),
539 params,
540 }
541 }
542}
543
544fn eval_test(expr: &str, params: &SqlParams) -> Result<bool, DynamicSqlError> {
555 let expr = expr.trim();
556
557 if let Some(idx) = find_keyword(expr, " or ") {
559 let left = &expr[..idx];
560 let right = &expr[idx + 4..];
561 return Ok(eval_test(left, params)? || eval_test(right, params)?);
562 }
563
564 if let Some(idx) = find_keyword(expr, " and ") {
566 let left = &expr[..idx];
567 let right = &expr[idx + 5..];
568 return Ok(eval_test(left, params)? && eval_test(right, params)?);
569 }
570
571 if let Some(stripped) = expr.strip_suffix("!= null") {
573 let name = stripped.trim();
574 return Ok(params.is_not_null(name));
575 }
576 if let Some(stripped) = expr.strip_suffix("== null") {
577 let name = stripped.trim();
578 return Ok(params.is_null(name));
579 }
580
581 if let Some(idx) = expr.find("==") {
583 let left = expr[..idx].trim();
584 let right = expr[idx + 2..].trim();
585 let actual = params.get(left);
586 let expected = right.trim_matches('\'').trim_matches('"');
587 return Ok(match actual {
588 Some(ParamValue::String(s)) => s == expected,
589 _ => false,
590 });
591 }
592 if let Some(idx) = expr.find("!=") {
593 let left = expr[..idx].trim();
594 let right = expr[idx + 2..].trim();
595 let actual = params.get(left);
596 let expected = right.trim_matches('\'').trim_matches('"');
597 return Ok(match actual {
598 Some(ParamValue::String(s)) => s != expected,
599 _ => true,
600 });
601 }
602
603 type CmpFn = fn(i64, i64) -> bool;
605 let cmps: [(&str, CmpFn); 4] = [
606 (">=", |a, b| a >= b),
607 ("<=", |a, b| a <= b),
608 (">", |a, b| a > b),
609 ("<", |a, b| a < b),
610 ];
611 for (op, cmp) in cmps {
612 if let Some(idx) = expr.find(op) {
613 let left = expr[..idx].trim();
614 let right_str = expr[idx + op.len()..].trim();
615 if let (Some(ParamValue::Int(a)), Ok(b)) = (params.get(left), right_str.parse::<i64>())
616 {
617 return Ok(cmp(*a, b));
618 }
619 return Ok(false);
620 }
621 }
622
623 Err(DynamicSqlError::EvalError(format!(
624 "无法解析表达式: {}",
625 expr
626 )))
627}
628
629fn find_keyword(expr: &str, keyword: &str) -> Option<usize> {
631 let lower = expr.to_lowercase();
632 lower.find(keyword)
633}
634
635fn strip_leading_and_or(s: &str) -> &str {
637 let trimmed = s.trim_start();
638 let lower = trimmed.to_lowercase();
639 if lower.starts_with("and ") {
640 trimmed[4..].trim_start()
641 } else if lower.starts_with("or ") {
642 trimmed[3..].trim_start()
643 } else {
644 trimmed
645 }
646}
647
648fn param_to_string(v: &ParamValue) -> String {
650 match v {
651 ParamValue::Null => "NULL".to_string(),
652 ParamValue::String(s) => escape_sql_string(s),
653 ParamValue::Int(i) => i.to_string(),
654 ParamValue::Float(f) => f.to_string(),
655 ParamValue::Bool(b) => {
656 if *b {
657 "TRUE".to_string()
658 } else {
659 "FALSE".to_string()
660 }
661 }
662 ParamValue::Array(_) => "[]".to_string(),
663 }
664}
665
666fn escape_sql_string(s: &str) -> String {
669 let mut out = String::with_capacity(s.len() + 2);
670 for ch in s.chars() {
671 match ch {
672 '\'' => out.push_str("''"),
673 '\\' => out.push_str("\\\\"),
674 '\n' => out.push_str("\\n"),
675 '\r' => out.push_str("\\r"),
676 '\0' => out.push_str("\\0"),
677 _ => out.push(ch),
678 }
679 }
680 out
681}
682
683fn normalize_set_commas(s: &str) -> String {
686 let mut out = String::with_capacity(s.len());
687 let mut chars = s.chars().peekable();
688 while let Some(ch) = chars.next() {
689 if ch == ',' {
690 out.push(',');
691 while matches!(chars.peek(), Some(c) if c.is_whitespace()) {
693 chars.next();
694 }
695 out.push(' ');
696 } else {
697 out.push(ch);
698 }
699 }
700 out
701}
702
703#[derive(Debug, Clone)]
708pub struct XmlNode {
709 pub node_type: XmlNodeType,
710 pub children: Vec<XmlNode>,
711}
712
713#[derive(Debug, Clone)]
714pub enum XmlNodeType {
715 Text(String),
716 Element {
717 name: String,
718 attrs: HashMap<String, String>,
719 },
720}
721
722struct XmlParser<'a> {
723 input: &'a str,
724 pos: usize,
725}
726
727impl<'a> XmlParser<'a> {
728 fn parse(input: &'a str) -> Result<XmlNode, DynamicSqlError> {
729 let mut parser = Self { input, pos: 0 };
730 let mut root = XmlNode {
731 node_type: XmlNodeType::Element {
732 name: "root".to_string(),
733 attrs: HashMap::new(),
734 },
735 children: Vec::new(),
736 };
737 parser.parse_children(&mut root)?;
738 Ok(root)
739 }
740
741 fn parse_children(&mut self, parent: &mut XmlNode) -> Result<(), DynamicSqlError> {
742 loop {
743 if self.pos >= self.input.len() {
744 break;
745 }
746 if self.starts_with("</") {
748 break;
749 }
750 if self.starts_with("<!--") {
752 self.skip_comment()?;
753 continue;
754 }
755 if self.starts_with("<") {
757 let element = self.parse_element()?;
758 parent.children.push(element);
759 } else {
760 let text = self.parse_text();
762 parent.children.push(XmlNode {
763 node_type: XmlNodeType::Text(text),
764 children: Vec::new(),
765 });
766 }
767 }
768 Ok(())
769 }
770
771 fn parse_element(&mut self) -> Result<XmlNode, DynamicSqlError> {
772 self.pos += 1;
774 let name = self.read_name();
776 let mut attrs = HashMap::new();
778 loop {
779 self.skip_whitespace();
780 if self.pos >= self.input.len() {
781 return Err(DynamicSqlError::ParseError("未闭合的标签".into()));
782 }
783 let c = self.input.as_bytes()[self.pos] as char;
784 if c == '>' {
785 self.pos += 1;
786 break;
787 }
788 if c == '/' {
789 if self.pos + 1 < self.input.len() && self.input.as_bytes()[self.pos + 1] == b'>' {
791 self.pos += 2;
792 return Ok(XmlNode {
793 node_type: XmlNodeType::Element { name, attrs },
794 children: Vec::new(),
795 });
796 }
797 }
798 let attr_name = self.read_name();
800 self.skip_whitespace();
801 if self.pos < self.input.len() && self.input.as_bytes()[self.pos] == b'=' {
802 self.pos += 1;
803 self.skip_whitespace();
804 let attr_value = self.read_attr_value()?;
805 attrs.insert(attr_name, attr_value);
806 }
807 }
808 let mut node = XmlNode {
810 node_type: XmlNodeType::Element {
811 name: name.clone(),
812 attrs,
813 },
814 children: Vec::new(),
815 };
816 self.parse_children(&mut node)?;
817 if self.starts_with("</") {
820 self.pos += 2;
821 let end_name = self.read_name();
822 self.skip_whitespace();
823 if self.pos < self.input.len() && self.input.as_bytes()[self.pos] == b'>' {
824 self.pos += 1;
825 }
826 if let XmlNodeType::Element { name: n, .. } = &node.node_type {
828 if n != &end_name {
829 return Err(DynamicSqlError::ParseError(format!(
830 "标签不匹配: <{}> vs </{}>",
831 n, end_name
832 )));
833 }
834 }
835 } else {
836 return Err(DynamicSqlError::ParseError(format!(
839 "未闭合的标签: <{}>(缺少 </{}>)",
840 name, name
841 )));
842 }
843 Ok(node)
844 }
845
846 fn parse_text(&mut self) -> String {
847 let start = self.pos;
848 while self.pos < self.input.len() {
849 let b = self.input.as_bytes()[self.pos];
850 if b == b'<' {
851 break;
852 }
853 self.pos += 1;
854 }
855 self.input[start..self.pos]
857 .replace(">", ">")
858 .replace("<", "<")
859 .replace("&", "&")
860 .replace(""", "\"")
861 .replace("'", "'")
862 }
863
864 fn read_name(&mut self) -> String {
865 let start = self.pos;
866 while self.pos < self.input.len() {
867 let b = self.input.as_bytes()[self.pos];
868 if b.is_ascii_alphanumeric() || b == b'_' || b == b'-' || b == b':' {
869 self.pos += 1;
870 } else {
871 break;
872 }
873 }
874 self.input[start..self.pos].to_string()
875 }
876
877 fn read_attr_value(&mut self) -> Result<String, DynamicSqlError> {
878 if self.pos >= self.input.len() {
879 return Err(DynamicSqlError::ParseError("属性值缺失".into()));
880 }
881 let quote = self.input.as_bytes()[self.pos];
882 if quote != b'"' && quote != b'\'' {
883 return Err(DynamicSqlError::ParseError(format!(
884 "属性值应以引号开头, 实际: {}",
885 quote as char
886 )));
887 }
888 self.pos += 1;
889 let start = self.pos;
890 while self.pos < self.input.len() {
891 if self.input.as_bytes()[self.pos] == quote {
892 let raw = &self.input[start..self.pos];
893 let value = raw
895 .replace(">", ">")
896 .replace("<", "<")
897 .replace("&", "&")
898 .replace(""", "\"")
899 .replace("'", "'");
900 self.pos += 1;
901 return Ok(value);
902 }
903 self.pos += 1;
904 }
905 Err(DynamicSqlError::ParseError("属性值未闭合".into()))
906 }
907
908 fn skip_whitespace(&mut self) {
909 while self.pos < self.input.len() {
910 if !self.input.as_bytes()[self.pos].is_ascii_whitespace() {
911 break;
912 }
913 self.pos += 1;
914 }
915 }
916
917 fn skip_comment(&mut self) -> Result<(), DynamicSqlError> {
918 self.pos += 4;
920 while self.pos + 2 < self.input.len() {
921 if &self.input[self.pos..self.pos + 3] == "-->" {
922 self.pos += 3;
923 return Ok(());
924 }
925 self.pos += 1;
926 }
927 Err(DynamicSqlError::ParseError("注释未闭合".into()))
928 }
929
930 fn starts_with(&self, s: &str) -> bool {
931 self.input[self.pos..].starts_with(s)
932 }
933}
934
935#[cfg(test)]
936mod tests {
937 use super::*;
938
939 #[test]
942 fn test_sql_params_set_get() {
943 let mut p = SqlParams::new();
944 p.set("name", "Alice");
945 p.set_int("age", 30);
946 p.set_bool("active", true);
947 p.set_null("deleted");
948
949 assert!(matches!(p.get("name"), Some(ParamValue::String(_))));
950 assert!(matches!(p.get("age"), Some(ParamValue::Int(30))));
951 assert!(matches!(p.get("active"), Some(ParamValue::Bool(true))));
952 assert!(matches!(p.get("deleted"), Some(ParamValue::Null)));
953 assert!(p.get("missing").is_none());
954 }
955
956 #[test]
957 fn test_sql_params_is_null() {
958 let mut p = SqlParams::new();
959 assert!(p.is_null("missing"));
960 p.set_null("x");
961 assert!(p.is_null("x"));
962 p.set("y", "val");
963 assert!(!p.is_null("y"));
964 assert!(p.is_not_null("y"));
965 }
966
967 #[test]
968 fn test_sql_params_contains() {
969 let mut p = SqlParams::new();
970 p.set("a", "1");
971 assert!(p.contains("a"));
972 assert!(!p.contains("b"));
973 }
974
975 #[test]
978 fn test_simple_select_no_params() {
979 let xml = r#"<select id="all">SELECT * FROM users</select>"#;
980 let parser = DynamicSqlParser::from_xml(xml).unwrap();
981 let params = SqlParams::new();
982 let sql = parser.build("all", ¶ms).unwrap();
983 assert_eq!(sql, "SELECT * FROM users");
984 }
985
986 #[test]
987 fn test_select_with_param_binding() {
988 let xml = r#"<select id="by_id">SELECT * FROM users WHERE id = #{id}</select>"#;
989 let parser = DynamicSqlParser::from_xml(xml).unwrap();
990 let mut params = SqlParams::new();
991 params.set_int("id", 42);
992 let sql = parser.build("by_id", ¶ms).unwrap();
993 assert_eq!(sql, "SELECT * FROM users WHERE id = ?");
994 }
995
996 #[test]
997 fn test_select_with_string_interpolation() {
998 let xml = r#"<select id="by_table">SELECT * FROM ${table}</select>"#;
999 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1000 let mut params = SqlParams::new();
1001 params.set("table", "users");
1002 let sql = parser.build("by_table", ¶ms).unwrap();
1003 assert_eq!(sql, "SELECT * FROM users");
1004 }
1005
1006 #[test]
1009 fn test_if_true() {
1010 let xml = r#"<select id="q">SELECT * FROM users WHERE 1=1 <if test="name != null">AND name = #{name}</if></select>"#;
1011 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1012 let mut params = SqlParams::new();
1013 params.set("name", "Alice");
1014 let sql = parser.build("q", ¶ms).unwrap();
1015 assert!(sql.contains("AND name = ?"));
1016 }
1017
1018 #[test]
1019 fn test_if_false() {
1020 let xml = r#"<select id="q">SELECT * FROM users WHERE 1=1 <if test="name != null">AND name = #{name}</if></select>"#;
1021 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1022 let params = SqlParams::new();
1023 let sql = parser.build("q", ¶ms).unwrap();
1024 assert!(!sql.contains("AND name"));
1025 }
1026
1027 #[test]
1028 fn test_if_null_check() {
1029 let xml = r#"<select id="q">SELECT * FROM users <if test="name == null">WHERE name IS NULL</if></select>"#;
1030 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1031
1032 let params = SqlParams::new();
1033 let sql = parser.build("q", ¶ms).unwrap();
1034 assert!(sql.contains("WHERE name IS NULL"));
1035
1036 let mut params = SqlParams::new();
1037 params.set("name", "Alice");
1038 let sql = parser.build("q", ¶ms).unwrap();
1039 assert!(!sql.contains("WHERE name IS NULL"));
1040 }
1041
1042 #[test]
1043 fn test_if_string_equals() {
1044 let xml = r#"<select id="q">SELECT * FROM users <if test="role == 'admin'">WHERE is_admin = 1</if></select>"#;
1045 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1046
1047 let mut params = SqlParams::new();
1048 params.set("role", "admin");
1049 let sql = parser.build("q", ¶ms).unwrap();
1050 assert!(sql.contains("WHERE is_admin = 1"));
1051
1052 let mut params = SqlParams::new();
1053 params.set("role", "user");
1054 let sql = parser.build("q", ¶ms).unwrap();
1055 assert!(!sql.contains("WHERE is_admin"));
1056 }
1057
1058 #[test]
1059 fn test_if_numeric_comparison() {
1060 let xml = r#"<select id="q">SELECT * FROM users <if test="age > 18">WHERE age > 18</if></select>"#;
1061 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1062
1063 let mut params = SqlParams::new();
1064 params.set_int("age", 25);
1065 let sql = parser.build("q", ¶ms).unwrap();
1066 assert!(sql.contains("WHERE age > 18"));
1067
1068 let mut params = SqlParams::new();
1069 params.set_int("age", 15);
1070 let sql = parser.build("q", ¶ms).unwrap();
1071 assert!(!sql.contains("WHERE age"));
1072 }
1073
1074 #[test]
1075 fn test_if_and_or() {
1076 let xml = r#"<select id="q">SELECT * FROM users <if test="name != null and age != null">WHERE name = #{name} AND age = #{age}</if></select>"#;
1077 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1078
1079 let mut params = SqlParams::new();
1080 params.set("name", "Alice");
1081 params.set_int("age", 30);
1082 let sql = parser.build("q", ¶ms).unwrap();
1083 assert!(sql.contains("WHERE name = ? AND age = ?"));
1084
1085 let mut params = SqlParams::new();
1086 params.set("name", "Alice");
1087 let sql = parser.build("q", ¶ms).unwrap();
1088 assert!(!sql.contains("WHERE"));
1089 }
1090
1091 #[test]
1094 fn test_where_strips_leading_and() {
1095 let xml = r#"<select id="q">SELECT * FROM users <where><if test="name != null">AND name = #{name}</if></where></select>"#;
1096 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1097 let mut params = SqlParams::new();
1098 params.set("name", "Alice");
1099 let sql = parser.build("q", ¶ms).unwrap();
1100 assert!(sql.contains("WHERE name = ?"));
1101 assert!(!sql.contains("WHERE AND"));
1102 }
1103
1104 #[test]
1105 fn test_where_empty_no_clause() {
1106 let xml = r#"<select id="q">SELECT * FROM users <where><if test="name != null">AND name = #{name}</if></where></select>"#;
1107 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1108 let params = SqlParams::new();
1109 let sql = parser.build("q", ¶ms).unwrap();
1110 assert_eq!(sql, "SELECT * FROM users");
1111 }
1112
1113 #[test]
1114 fn test_where_multiple_conditions() {
1115 let xml = r#"<select id="q">SELECT * FROM users <where>
1116 <if test="name != null">AND name = #{name}</if>
1117 <if test="age != null">AND age = #{age}</if>
1118 </where></select>"#;
1119 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1120 let mut params = SqlParams::new();
1121 params.set("name", "Alice");
1122 params.set_int("age", 30);
1123 let sql = parser.build("q", ¶ms).unwrap();
1124 assert!(sql.contains("WHERE name = ? AND age = ?"));
1125 }
1126
1127 #[test]
1130 fn test_set_strips_trailing_comma() {
1131 let xml = r#"<update id="u">UPDATE users <set><if test="name != null">name = #{name},</if><if test="age != null">age = #{age},</if></set> WHERE id = #{id}</update>"#;
1132 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1133 let mut params = SqlParams::new();
1134 params.set("name", "Alice");
1135 params.set_int("age", 30);
1136 params.set_int("id", 1);
1137 let sql = parser.build("u", ¶ms).unwrap();
1138 assert!(sql.contains("SET name = ?, age = ? WHERE id = ?"));
1140 }
1141
1142 #[test]
1143 fn test_set_single_field() {
1144 let xml = r#"<update id="u">UPDATE users <set><if test="name != null">name = #{name},</if></set> WHERE id = #{id}</update>"#;
1145 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1146 let mut params = SqlParams::new();
1147 params.set("name", "Alice");
1148 params.set_int("id", 1);
1149 let sql = parser.build("u", ¶ms).unwrap();
1150 assert!(sql.contains("SET name = ? WHERE id = ?"));
1151 }
1152
1153 #[test]
1156 fn test_foreach_basic() {
1157 let xml = r#"<select id="q">SELECT * FROM users WHERE id IN (<foreach collection="ids" item="id" separator=",">#{id}</foreach>)</select>"#;
1158 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1159 let mut params = SqlParams::new();
1160 params.set_array(
1161 "ids",
1162 vec![ParamValue::Int(1), ParamValue::Int(2), ParamValue::Int(3)],
1163 );
1164 let sql = parser.build("q", ¶ms).unwrap();
1165 assert!(sql.contains("WHERE id IN ( ?,?,?)"));
1167 }
1168
1169 #[test]
1170 fn test_foreach_empty() {
1171 let xml = r#"<select id="q">SELECT * FROM users WHERE id IN (<foreach collection="ids" item="id" separator=",">#{id}</foreach>)</select>"#;
1172 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1173 let params = SqlParams::new();
1174 let sql = parser.build("q", ¶ms).unwrap();
1175 assert!(sql.contains("WHERE id IN ( )") || sql.contains("WHERE id IN ()"));
1176 }
1177
1178 #[test]
1179 fn test_foreach_with_strings() {
1180 let xml = r#"<select id="q">SELECT * FROM users WHERE name IN (<foreach collection="names" item="n" separator=",">#{n}</foreach>)</select>"#;
1181 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1182 let mut params = SqlParams::new();
1183 params.set_array(
1184 "names",
1185 vec![
1186 ParamValue::String("Alice".into()),
1187 ParamValue::String("Bob".into()),
1188 ],
1189 );
1190 let (sql, binds) = parser.build_with_binds("q", ¶ms).unwrap();
1191 assert_eq!(binds.len(), 2);
1192 assert!(sql.contains("?"));
1193 }
1194
1195 #[test]
1198 fn test_choose_when_matches() {
1199 let xml = r#"<select id="q">SELECT * FROM users WHERE 1=1 <choose>
1200 <when test="role == 'admin'">AND is_admin = 1</when>
1201 <otherwise>AND is_admin = 0</otherwise>
1202 </choose></select>"#;
1203 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1204 let mut params = SqlParams::new();
1205 params.set("role", "admin");
1206 let sql = parser.build("q", ¶ms).unwrap();
1207 assert!(sql.contains("AND is_admin = 1"));
1208 assert!(!sql.contains("AND is_admin = 0"));
1209 }
1210
1211 #[test]
1212 fn test_choose_otherwise() {
1213 let xml = r#"<select id="q">SELECT * FROM users WHERE 1=1 <choose>
1214 <when test="role == 'admin'">AND is_admin = 1</when>
1215 <otherwise>AND is_admin = 0</otherwise>
1216 </choose></select>"#;
1217 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1218 let mut params = SqlParams::new();
1219 params.set("role", "user");
1220 let sql = parser.build("q", ¶ms).unwrap();
1221 assert!(sql.contains("AND is_admin = 0"));
1222 }
1223
1224 #[test]
1227 fn test_trim_prefix_suffix() {
1228 let xml = r#"<select id="q">SELECT * FROM users <trim prefix="WHERE" prefixOverrides="AND |OR "><if test="name != null">AND name = #{name}</if></trim></select>"#;
1229 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1230 let mut params = SqlParams::new();
1231 params.set("name", "Alice");
1232 let sql = parser.build("q", ¶ms).unwrap();
1233 assert!(sql.contains("WHERE name = ?"));
1234 }
1235
1236 #[test]
1237 fn test_trim_empty() {
1238 let xml = r#"<select id="q">SELECT * FROM users <trim prefix="WHERE" prefixOverrides="AND"><if test="name != null">AND name = #{name}</if></trim></select>"#;
1239 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1240 let params = SqlParams::new();
1241 let sql = parser.build("q", ¶ms).unwrap();
1242 assert_eq!(sql, "SELECT * FROM users");
1243 }
1244
1245 #[test]
1248 fn test_statement_not_found() {
1249 let xml = r#"<select id="a">SELECT 1</select>"#;
1250 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1251 let params = SqlParams::new();
1252 let result = parser.build("missing", ¶ms);
1253 assert!(matches!(result, Err(DynamicSqlError::StatementNotFound(_))));
1254 }
1255
1256 #[test]
1257 fn test_missing_param() {
1258 let xml = r#"<select id="q">SELECT * FROM users WHERE id = #{id}</select>"#;
1259 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1260 let params = SqlParams::new();
1261 let result = parser.build("q", ¶ms);
1262 assert!(matches!(result, Err(DynamicSqlError::MissingParam(_))));
1263 }
1264
1265 #[test]
1266 fn test_parse_error_unclosed_tag() {
1267 let xml = r#"<select id="q">SELECT 1"#;
1269 let result = DynamicSqlParser::from_xml(xml);
1270 assert!(
1271 matches!(result, Err(DynamicSqlError::ParseError(_))),
1272 "未闭合标签应返回 ParseError,实际: {:?}",
1273 result
1274 );
1275 if let Err(DynamicSqlError::ParseError(msg)) = result {
1276 assert!(
1277 msg.contains("未闭合") || msg.contains("标签") || msg.contains("EOF"),
1278 "错误信息应提及未闭合/标签/EOF: {}",
1279 msg
1280 );
1281 }
1282 }
1283
1284 #[test]
1287 fn test_multiple_statements() {
1288 let xml = r#"
1289 <select id="find_all">SELECT * FROM users</select>
1290 <select id="find_by_id">SELECT * FROM users WHERE id = #{id}</select>
1291 <insert id="insert">INSERT INTO users (name) VALUES (#{name})</insert>
1292 <update id="update">UPDATE users SET name = #{name} WHERE id = #{id}</update>
1293 <delete id="delete">DELETE FROM users WHERE id = #{id}</delete>
1294 "#;
1295 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1296 let mut ids = parser.statement_ids();
1297 ids.sort();
1298 assert_eq!(
1299 ids,
1300 vec!["delete", "find_all", "find_by_id", "insert", "update"]
1301 );
1302 }
1303
1304 #[test]
1307 fn test_xml_entities() {
1308 let xml =
1309 r#"<select id="q">SELECT * FROM users WHERE age > 18 AND age < 65</select>"#;
1310 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1311 let params = SqlParams::new();
1312 let sql = parser.build("q", ¶ms).unwrap();
1313 assert!(sql.contains("age > 18 AND age < 65"));
1314 }
1315
1316 #[test]
1319 fn test_full_dynamic_query() {
1320 let xml = r#"<select id="search">
1321 SELECT u.id, u.name, o.total
1322 FROM users u
1323 LEFT JOIN orders o ON u.id = o.user_id
1324 <where>
1325 <if test="name != null">AND u.name LIKE #{name}</if>
1326 <if test="min_age != null">AND u.age > #{min_age}</if>
1327 <if test="status != null">AND u.status = #{status}</if>
1328 </where>
1329 ORDER BY u.id
1330 </select>"#;
1331 let parser = DynamicSqlParser::from_xml(xml).unwrap();
1332 let mut params = SqlParams::new();
1333 params.set("name", "%Alice%");
1334 params.set_int("min_age", 18);
1335 let (sql, binds) = parser.build_with_binds("search", ¶ms).unwrap();
1338 assert!(sql.contains("SELECT u.id, u.name, o.total"));
1339 assert!(sql.contains("FROM users u"));
1340 assert!(sql.contains("LEFT JOIN orders o ON u.id = o.user_id"));
1341 assert!(sql.contains("WHERE u.name LIKE ? AND u.age > ?"));
1342 assert!(!sql.contains("u.status"));
1343 assert!(sql.contains("ORDER BY u.id"));
1344 assert_eq!(binds.len(), 2);
1345 }
1346}