use crate::query::ast::*;
use crate::query::error::{ParseError, QueryError};
use crate::query::parser::parse_eql;
use crate::types::Atom;
use std::collections::HashMap;
#[derive(Debug, Clone)]
pub struct PreparedQuery {
pub source: String,
pub ast: QueryAst,
pub params: Vec<String>,
}
#[derive(Debug, Clone)]
pub struct ParamSlot {
pub name: String,
pub position: usize,
}
pub fn prepare(query: &str) -> Result<PreparedQuery, ParseError> {
let params = extract_params(query);
let mut substituted = query.to_string();
for param in ¶ms {
substituted = substituted.replace(&format!("${}", param), "0");
}
let ast = parse_eql(&substituted)?;
Ok(PreparedQuery {
source: query.to_string(),
ast,
params,
})
}
pub fn bind_params(
prepared: &PreparedQuery,
params: &HashMap<String, Atom>,
) -> Result<QueryAst, QueryError> {
let mut query = prepared.source.clone();
for param_name in &prepared.params {
let value = params
.get(param_name)
.ok_or_else(|| QueryError::Internal(format!("missing parameter: ${}", param_name)))?;
let value_str = atom_to_query_literal(value);
query = query.replace(&format!("${}", param_name), &value_str);
}
parse_eql(&query).map_err(|e| QueryError::Parse {
message: e.message,
line: e.line,
column: e.column,
})
}
fn extract_params(query: &str) -> Vec<String> {
let mut params = Vec::new();
let mut chars = query.chars().peekable();
while let Some(c) = chars.next() {
if c == '$' {
let mut name = String::new();
while let Some(&nc) = chars.peek() {
if nc.is_alphanumeric() || nc == '_' {
name.push(nc);
chars.next();
} else {
break;
}
}
if !name.is_empty() && !params.contains(&name) {
params.push(name);
}
}
}
params
}
fn atom_to_query_literal(atom: &Atom) -> String {
match atom {
Atom::Float(f) => format!("{}", f),
Atom::Int(i) => format!("{}", i),
Atom::Text(s) => format!("'{}'", s.replace('\'', "''")),
Atom::Null => "NULL".to_string(),
Atom::Bytes(_) => "NULL".to_string(), Atom::Vector(_, _) => "NULL".to_string(),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_prepare_extracts_params() {
let prepared =
prepare("SELECT * FROM 'sensor/*' WHERE value > $threshold LIMIT $count").unwrap();
assert_eq!(prepared.params, vec!["threshold", "count"]);
}
#[test]
fn test_bind_params_substitutes() {
let prepared = prepare("SELECT * FROM 'k' WHERE value > $min").unwrap();
let mut params = HashMap::new();
params.insert("min".to_string(), Atom::Float(42.0));
let ast = bind_params(&prepared, ¶ms).unwrap();
match ast {
QueryAst::Select(q) => {
assert!(q.where_clause.is_some());
}
_ => panic!("expected Select"),
}
}
#[test]
fn test_bind_missing_param_errors() {
let prepared = prepare("SELECT * FROM 'k' WHERE value > $x").unwrap();
let params = HashMap::new();
let result = bind_params(&prepared, ¶ms);
assert!(result.is_err());
}
#[test]
fn test_no_params() {
let prepared = prepare("SELECT * FROM 'k'").unwrap();
assert!(prepared.params.is_empty());
let ast = bind_params(&prepared, &HashMap::new()).unwrap();
match ast {
QueryAst::Select(_) => {}
_ => panic!("expected Select"),
}
}
#[test]
fn test_text_param_escaping() {
let prepared = prepare("SELECT * FROM 'k' WHERE key = $name").unwrap();
let mut params = HashMap::new();
params.insert("name".to_string(), Atom::Text("hello_world".to_string()));
let ast = bind_params(&prepared, ¶ms).unwrap();
match ast {
QueryAst::Select(q) => assert!(q.where_clause.is_some()),
_ => panic!("expected Select"),
}
}
}