use crate::lang::Language;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct QueryBindingRule {
pub lang: Language,
pub construct: String,
pub sql_arg: usize,
}
#[derive(Debug, Clone)]
pub struct BindingRules {
rules: Vec<QueryBindingRule>,
}
impl BindingRules {
pub fn empty() -> Self {
BindingRules { rules: Vec::new() }
}
pub fn with_defaults() -> Self {
let mut rules = BindingRules::empty();
for (lang, construct, sql_arg) in DEFAULT_RULES {
rules.register(QueryBindingRule {
lang: *lang,
construct: (*construct).to_string(),
sql_arg: *sql_arg,
});
}
rules
}
pub fn register(&mut self, rule: QueryBindingRule) {
self.rules.push(rule);
}
pub fn for_language(&self, lang: Language) -> impl Iterator<Item = &QueryBindingRule> {
self.rules.iter().filter(move |r| r.lang == lang)
}
pub fn is_empty(&self) -> bool {
self.rules.is_empty()
}
pub fn len(&self) -> usize {
self.rules.len()
}
}
const DEFAULT_RULES: &[(Language, &str, usize)] = &[
(Language::Rust, "sqlx::query", 0),
(Language::Rust, "sqlx::query_as", 0),
(Language::Rust, "sqlx::query_scalar", 0),
(Language::Rust, "diesel::sql_query", 0),
(Language::TypeScript, "knex.raw", 0),
(Language::JavaScript, "knex.raw", 0),
(Language::Python, "execute", 0),
(Language::Python, "text", 0),
];
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn with_defaults_has_rust_sqlx_query_rule() {
let rules = BindingRules::with_defaults();
let found = rules
.for_language(Language::Rust)
.find(|r| r.construct == "sqlx::query");
let rule = found.expect("sqlx::query rule must be present");
assert_eq!(rule.sql_arg, 0);
}
#[test]
fn with_defaults_has_python_execute_rule() {
let rules = BindingRules::with_defaults();
let found = rules
.for_language(Language::Python)
.find(|r| r.construct == "execute");
assert!(found.is_some());
}
#[test]
fn for_language_rust_yields_only_rust_rules() {
let rules = BindingRules::with_defaults();
let rust_rules: Vec<_> = rules.for_language(Language::Rust).collect();
assert!(rust_rules.iter().all(|r| r.lang == Language::Rust));
let constructs: Vec<&str> = rust_rules.iter().map(|r| r.construct.as_str()).collect();
for expected in [
"sqlx::query",
"sqlx::query_as",
"sqlx::query_scalar",
"diesel::sql_query",
] {
assert!(
constructs.contains(&expected),
"missing default rust rule: {expected}"
);
}
}
#[test]
fn for_language_with_no_rules_is_empty() {
let rules = BindingRules::with_defaults();
assert_eq!(rules.for_language(Language::Go).count(), 0);
}
#[test]
fn register_on_empty_registry_surfaces_new_rule() {
let mut rules = BindingRules::empty();
rules.register(QueryBindingRule {
lang: Language::Ruby,
construct: "ActiveRecord::Base.connection.execute".to_string(),
sql_arg: 0,
});
let found: Vec<_> = rules.for_language(Language::Ruby).collect();
assert_eq!(found.len(), 1);
assert_eq!(found[0].construct, "ActiveRecord::Base.connection.execute");
}
#[test]
fn empty_registry_yields_nothing_for_any_language() {
let rules = BindingRules::empty();
assert!(rules.is_empty());
assert_eq!(rules.len(), 0);
assert_eq!(rules.for_language(Language::Rust).count(), 0);
}
}