use std::collections::HashSet;
use super::parser::to_snake_case;
use anyhow::Result;
pub trait FilterableEntity {
fn filterable_fields() -> HashSet<String>;
fn sortable_fields() -> HashSet<String> {
Self::filterable_fields()
}
}
pub fn is_valid_field(field: &str, allowed_fields: &HashSet<String>) -> bool {
allowed_fields.contains(field)
}
pub fn sanitize_field_name(field: &str) -> Result<String> {
if field.is_empty() {
return Err(anyhow::anyhow!("Field name cannot be empty"));
}
if !field.chars().all(|c| c.is_alphanumeric() || c == '_') {
return Err(anyhow::anyhow!("Invalid field name: '{}'", field));
}
let lower = field.to_lowercase();
if lower.contains("--") || lower.contains("/*") || lower.contains(";")
|| lower.contains("drop ") || lower.contains("delete ") || lower.contains("truncate ")
|| lower.contains("update ") || lower.contains("insert ") || lower.contains("exec ")
|| lower.contains("execute ") || lower.contains("script>") {
return Err(anyhow::anyhow!("Potentially dangerous field name: '{}'", field));
}
Ok(to_snake_case(field))
}
#[cfg(test)]
mod casing_tests {
use super::sanitize_field_name;
#[test]
fn camel_case_fields_fold_to_snake_columns() {
assert_eq!(sanitize_field_name("employeeId").unwrap(), "employee_id");
assert_eq!(sanitize_field_name("date").unwrap(), "date");
assert_eq!(sanitize_field_name("employee_id").unwrap(), "employee_id");
assert_eq!(sanitize_field_name("billableCost").unwrap(), "billable_cost");
}
#[test]
fn dangerous_names_still_refused() {
assert!(sanitize_field_name("drop table").is_err());
assert!(sanitize_field_name("").is_err());
}
}