use std::collections::HashMap;
use std::collections::HashSet;
use anyhow::Result;
use super::types::{FilterOperator, FilterValue, FilterLogical, SortDirection, SortSpec};
use super::condition::FilterCondition;
use super::query::QueryFilter;
use super::validation::{is_valid_field, sanitize_field_name};
const BUILTIN_PG_TYPES: &[&str] = &[
"uuid", "numeric", "decimal", "integer", "int", "int4", "int8",
"bigint", "smallint", "int2", "real", "float", "float4", "float8",
"double precision", "boolean", "bool", "text", "varchar", "char",
"timestamp", "timestamptz", "date", "time", "timetz",
"interval", "jsonb", "json", "bytea", "inet", "cidr", "macaddr",
];
pub(crate) fn is_custom_enum_type(col_type: &str) -> bool {
!BUILTIN_PG_TYPES.contains(&col_type.to_lowercase().as_str())
}
const AUDIT_METADATA_FIELDS: &[&str] = &["created_at", "updated_at", "deleted_at"];
pub(crate) fn audit_metadata_sql_expr(field: &str) -> Option<String> {
if AUDIT_METADATA_FIELDS.contains(&field) {
Some(format!("(metadata->>'{}')::timestamptz", field))
} else {
None
}
}
pub(crate) fn to_snake_case(s: &str) -> String {
let mut result = String::with_capacity(s.len() + 4);
for (i, ch) in s.chars().enumerate() {
if ch.is_uppercase() && i > 0 {
result.push('_');
}
result.push(ch.to_lowercase().next().unwrap_or(ch));
}
result
}
fn normalize_enum_value(value: String) -> String {
if value.contains(',') {
value.split(',')
.map(|v| to_snake_case(v.trim()))
.collect::<Vec<_>>()
.join(",")
} else {
to_snake_case(&value)
}
}
pub fn parse_filters(
params: &HashMap<String, String>,
column_types: &HashMap<String, String>,
allowed_fields: Option<&HashSet<String>>,
) -> Result<QueryFilter> {
let mut filter = QueryFilter::new();
let mut or_conditions: Vec<FilterCondition> = Vec::new();
for (key, value) in params {
if let Some(bracket_pos) = key.find('[') {
if key.ends_with(']') {
let field = &key[..bracket_pos];
let sanitized_field = sanitize_field_name(field)?;
if let Some(allowed) = allowed_fields {
if !is_valid_field(&sanitized_field, allowed) {
continue; }
}
let operator_str = &key[bracket_pos + 1..key.len() - 1];
if field.eq_ignore_ascii_case("orderby") || field.eq_ignore_ascii_case("sort") {
let sort_field = match sanitize_field_name(operator_str) {
Ok(f) => f,
Err(_) => continue,
};
if let Some(allowed) = allowed_fields {
if !is_valid_field(&sort_field, allowed) {
continue;
}
}
let v = value.trim();
let (f, dir) = if let Some(stripped) = v.strip_prefix('-') {
(stripped, crate::filter::types::SortDirection::Desc)
} else if v.eq_ignore_ascii_case("desc") {
(sort_field.as_str(), crate::filter::types::SortDirection::Desc)
} else {
(sort_field.as_str(), crate::filter::types::SortDirection::Asc)
};
let _ = f;
filter.add_sort(SortSpec::new(sort_field, dir));
continue;
}
match operator_str.to_ascii_lowercase().as_str() {
"orderby" => {
let sort_field = audit_metadata_sql_expr(&sanitized_field)
.unwrap_or_else(|| sanitized_field.clone());
filter.add_sort(SortSpec::new(
sort_field,
SortDirection::from_str(value)
));
continue;
}
"or" | "orwhere" => {
let or_value = match column_types.get(&sanitized_field) {
Some(col_type) if is_custom_enum_type(col_type) => {
normalize_enum_value(value.clone())
}
_ => value.clone(),
};
let condition = FilterCondition::new(
sanitized_field.clone(),
FilterOperator::Equal,
FilterValue::from_string(or_value, false)
).with_logical(FilterLogical::Or);
let condition = match column_types.get(&sanitized_field) {
Some(col_type) => condition.with_column_type(col_type.clone()),
None => condition,
};
or_conditions.push(condition);
continue;
}
_ => {
if let Some(op) = FilterOperator::from_str(operator_str) {
let filter_value = if let Some(col_type) = column_types.get(&sanitized_field) {
if is_custom_enum_type(col_type) {
normalize_enum_value(value.clone())
} else {
value.clone()
}
} else {
value.clone()
};
let condition_field = audit_metadata_sql_expr(&sanitized_field)
.unwrap_or_else(|| sanitized_field.clone());
let is_list = matches!(
op,
FilterOperator::In
| FilterOperator::NotIn
| FilterOperator::Between
| FilterOperator::NotBetween
);
let condition = FilterCondition::new(
condition_field,
op.clone(),
FilterValue::from_string(filter_value, is_list)
);
let condition = if audit_metadata_sql_expr(&sanitized_field).is_some() {
condition.with_column_type("timestamptz".to_string())
} else if let Some(col_type) = column_types.get(&sanitized_field) {
condition.with_column_type(col_type.clone())
} else {
condition
};
filter.add_condition(condition);
}
}
}
}
} else {
match key.to_ascii_lowercase().as_str() {
"orderby" | "sort" => {
if value.contains(',') {
for part in value.split(',') {
let part = part.trim();
let direction = if part.starts_with('-') {
SortDirection::Desc
} else {
SortDirection::Asc
};
let field = part.trim_start_matches('-');
if let Ok(sanitized) = sanitize_field_name(field) {
let sort_field = audit_metadata_sql_expr(&sanitized)
.unwrap_or_else(|| sanitized.clone());
if let Some(allowed) = allowed_fields {
if is_valid_field(&sanitized, allowed) {
filter.add_sort(SortSpec::new(sort_field, direction));
}
} else {
filter.add_sort(SortSpec::new(sort_field, direction));
}
}
}
} else {
let (bare, suffix_dir) = match value.split_once(':') {
Some((f, d)) => (
f,
if d.trim().eq_ignore_ascii_case("desc") {
Some(SortDirection::Desc)
} else {
Some(SortDirection::Asc)
},
),
None => (value.as_str(), None),
};
let direction = match (bare.starts_with('-'), suffix_dir) {
(true, _) => SortDirection::Desc,
(false, Some(d)) => d,
(false, None) => SortDirection::Asc,
};
let field = bare.trim_start_matches('-');
if let Ok(sanitized) = sanitize_field_name(field) {
let sort_field = audit_metadata_sql_expr(&sanitized)
.unwrap_or_else(|| sanitized.clone());
if let Some(allowed) = allowed_fields {
if is_valid_field(&sanitized, allowed) {
filter.add_sort(SortSpec::new(sort_field, direction));
}
} else {
filter.add_sort(SortSpec::new(sort_field, direction));
}
}
}
}
"search" => {
filter.search = Some(value.clone());
}
"searchfields" => {
filter.search_fields = value.split(',').filter_map(|s| {
let trimmed = s.trim();
sanitize_field_name(trimmed).ok().filter(|sanitized| {
if let Some(allowed) = allowed_fields {
is_valid_field(sanitized, allowed)
} else {
true
}
})
}).collect();
}
"limit" => {
if let Ok(l) = value.parse::<u32>() {
filter.limit = Some(l);
}
}
"offset" => {
if let Ok(o) = value.parse::<u32>() {
filter.offset = Some(o);
}
}
"page" => {
if let Ok(p) = value.parse::<u32>() {
filter.page = Some(p);
}
}
"pagesize" | "perpage" | "per_page" => {
if let Ok(p) = value.parse::<u32>() {
filter.limit = Some(p);
}
}
"after" => {
filter.cursor_after = Some(value.clone());
}
"before" => {
filter.cursor_before = Some(value.clone());
}
"estimate" | "estimate_total" => {
filter.estimate_total = matches!(value.as_str(), "1" | "true" | "yes");
}
"__base_condition" => {
filter.base_conditions.push(value.clone());
}
_ => {
if !matches!(key.as_str(), "fields" | "include" | "with") {
let sanitized_field = match sanitize_field_name(key) {
Ok(f) => f,
Err(_) => continue, };
if let Some(allowed) = allowed_fields {
if !is_valid_field(&sanitized_field, allowed) {
continue; }
}
let filter_value = if let Some(col_type) = column_types.get(&sanitized_field) {
if is_custom_enum_type(col_type) {
normalize_enum_value(value.clone())
} else {
value.clone()
}
} else {
value.clone()
};
let condition_field = audit_metadata_sql_expr(&sanitized_field)
.unwrap_or_else(|| sanitized_field.clone());
let condition = FilterCondition::new(
condition_field,
FilterOperator::Equal,
FilterValue::from_string(filter_value, value.contains(','))
);
let condition = if audit_metadata_sql_expr(&sanitized_field).is_some() {
condition.with_column_type("timestamptz".to_string())
} else if let Some(col_type) = column_types.get(&sanitized_field) {
condition.with_column_type(col_type.clone())
} else {
condition
};
filter.add_condition(condition);
}
}
}
}
}
for condition in or_conditions {
filter.add_condition(condition);
}
Ok(filter)
}
#[cfg(test)]
mod orderby_wire_tests {
use super::*;
use std::collections::HashMap;
#[test]
fn orderby_bracket_key_populates_sorts() {
let mut params = HashMap::new();
params.insert("orderby[employeeNumber]".to_string(), "desc".to_string());
let f = parse_filters(¶ms, &HashMap::new(), None).unwrap();
assert_eq!(f.sorts.len(), 1, "sorts: {:?}", f.sorts);
assert_eq!(f.sorts[0].field, "employee_number");
assert!(matches!(f.sorts[0].direction, crate::filter::types::SortDirection::Desc));
}
#[test]
fn orderby_plain_key_populates_sorts() {
let mut params = HashMap::new();
params.insert("orderby".to_string(), "employeeNumber:desc".to_string());
let f = parse_filters(¶ms, &HashMap::new(), None).unwrap();
assert!(!f.sorts.is_empty(), "sorts: {:?}", f.sorts);
}
}