use anyhow::{anyhow, Result}; use axum::{ extract::{Extension, Path, Query, State}, http::{Method, StatusCode}, Json, }; use serde_json::Value; use sqlx::postgres::PgRow; use sqlx::Column; use sqlx::Row; use sqlx::TypeInfo; use std::collections::HashMap; use crate::{ auth::Claims, models::blacklist::BlacklistEntry, state::{AppState, CacheEntry}, }; /// Returns (sql, ordered_param_values, cache_key). /// body_cols: (col_name, typed_value) pairs from request body. /// filter_cols: (col_name, string_value) pairs from query params. fn coerce_id(id_val: &str) -> Value { if let Ok(n) = id_val.parse::() { Value::Number(n.into()) } else { Value::String(id_val.to_string()) } } pub fn build_query( method: &str, table: &str, id: Option<&str>, body_cols: &[(String, Value)], filter_cols: &[(String, String)], ) -> Result<(String, Vec, String)> { if !crate::routes::is_valid_identifier(table) { return Err(anyhow!("invalid identifier: {}", table)); } for (col, _) in body_cols.iter() { if !crate::routes::is_valid_identifier(col) { return Err(anyhow!("invalid identifier: {}", col)); } } for (col, _) in filter_cols.iter() { if !crate::routes::is_valid_identifier(col) { return Err(anyhow!("invalid identifier: {}", col)); } } let mut sorted_body = body_cols.to_vec(); sorted_body.sort_by(|a, b| a.0.cmp(&b.0)); let mut sorted_filters = filter_cols.to_vec(); sorted_filters.sort_by(|a, b| a.0.cmp(&b.0)); match method.to_uppercase().as_str() { "GET" => { if let Some(id_val) = id { let sql = format!("SELECT * FROM {} WHERE id = $1", table); let key = format!("GET:{}:~id", table); Ok((sql, vec![coerce_id(id_val)], key)) } else if sorted_filters.is_empty() { let sql = format!("SELECT * FROM {}", table); let key = format!("GET:{}:", table); Ok((sql, vec![], key)) } else { let col_names: Vec = sorted_filters.iter().map(|(c, _)| c.clone()).collect(); let where_clause: Vec = col_names .iter() .enumerate() .map(|(i, c)| format!("{} = ${}", c, i + 1)) .collect(); let sql = format!( "SELECT * FROM {} WHERE {}", table, where_clause.join(" AND ") ); let params: Vec = sorted_filters .iter() .map(|(_, v)| Value::String(v.clone())) .collect(); let key = format!("GET:{}:{}", table, col_names.join(",")); Ok((sql, params, key)) } } "POST" => { if sorted_body.is_empty() { return Err(anyhow!("POST requires a body with at least one field")); } let cols: Vec = sorted_body.iter().map(|(c, _)| c.clone()).collect(); let placeholders: Vec = (1..=cols.len()).map(|i| format!("${}", i)).collect(); let sql = format!( "INSERT INTO {} ({}) VALUES ({}) RETURNING *", table, cols.join(", "), placeholders.join(", ") ); let params: Vec = sorted_body.iter().map(|(_, v)| v.clone()).collect(); let key = format!("POST:{}:{}", table, cols.join(",")); Ok((sql, params, key)) } "PUT" => { let id_val = id.ok_or_else(|| anyhow!("PUT requires an id"))?; if sorted_body.is_empty() { return Err(anyhow!("PUT requires a body with at least one field")); } let cols: Vec = sorted_body.iter().map(|(c, _)| c.clone()).collect(); let set_clause: Vec = cols .iter() .enumerate() .map(|(i, c)| format!("{} = ${}", c, i + 1)) .collect(); let id_placeholder = cols.len() + 1; let sql = format!( "UPDATE {} SET {} WHERE id = ${} RETURNING *", table, set_clause.join(", "), id_placeholder ); let mut params: Vec = sorted_body.iter().map(|(_, v)| v.clone()).collect(); params.push(coerce_id(id_val)); let key = format!("PUT:{}:{}:by_id", table, cols.join(",")); Ok((sql, params, key)) } "DELETE" => { let id_val = id.ok_or_else(|| anyhow!("DELETE requires an id"))?; let sql = format!("DELETE FROM {} WHERE id = $1", table); let key = format!("DELETE:{}:by_id", table); Ok((sql, vec![coerce_id(id_val)], key)) } m => Err(anyhow!("unsupported method: {}", m)), } } pub fn pg_row_to_json(row: PgRow) -> Value { let columns = row.columns(); let mut map = serde_json::Map::new(); for col in columns { let name = col.name().to_string(); let type_name = col.type_info().name(); let val = match type_name { "INT2" => row .try_get::(col.ordinal()) .map(|v| Value::Number(i64::from(v).into())) .unwrap_or(Value::Null), "INT4" | "SERIAL" => row .try_get::(col.ordinal()) .map(|v| Value::Number(i64::from(v).into())) .unwrap_or(Value::Null), "INT8" => row .try_get::(col.ordinal()) .map(|v| Value::Number(v.into())) .unwrap_or(Value::Null), "FLOAT4" | "FLOAT8" => row .try_get::(col.ordinal()) .ok() .and_then(serde_json::Number::from_f64) .map(Value::Number) .unwrap_or(Value::Null), "BOOL" => row .try_get::(col.ordinal()) .map(Value::Bool) .unwrap_or(Value::Null), "UUID" => row .try_get::(col.ordinal()) .map(|v| Value::String(v.to_string())) .unwrap_or(Value::Null), "TIMESTAMPTZ" | "TIMESTAMP" => row .try_get::, _>(col.ordinal()) .map(|v| Value::String(v.to_rfc3339())) .unwrap_or(Value::Null), _ => row .try_get::(col.ordinal()) .map(Value::String) .unwrap_or(Value::Null), }; map.insert(name, val); } Value::Object(map) } async fn reload_blacklist(state: &AppState) -> Result<(), StatusCode> { let entries = sqlx::query_as::<_, BlacklistEntry>( "SELECT id, pattern, method, reason, active, bypass_mask, created_at FROM blacklist ORDER BY id", ) .fetch_all(&state.pool) .await .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; state.blacklist_cache.load(entries).await; Ok(()) } async fn reload_cors(state: &AppState) -> Result<(), StatusCode> { let origins: Vec = sqlx::query_scalar::<_, String>("SELECT origin FROM cors_origins ORDER BY id") .fetch_all(&state.pool) .await .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; state.cors_cache.load(origins).await; Ok(()) } fn strip_password_hash(v: Value) -> Value { match v { Value::Object(mut m) => { m.remove("password_hash"); Value::Object(m) } Value::Array(arr) => Value::Array( arr.into_iter() .map(|item| match item { Value::Object(mut m) => { m.remove("password_hash"); Value::Object(m) } other => other, }) .collect(), ), other => other, } } pub async fn handle_crud( State(state): State, method: Method, Extension(claims): Extension, Path(params): Path>, Query(query_params): Query>, body: Option>>, ) -> Result, StatusCode> { let table = params.get("table").ok_or(StatusCode::BAD_REQUEST)?.clone(); let id = params.get("id").map(|s| s.as_str()); let method_str = method.as_str(); // Enforce permission bits before doing any work. let required_bit = match method_str.to_uppercase().as_str() { "GET" => crate::auth::permissions::READ, "POST" | "PUT" => crate::auth::permissions::WRITE, "DELETE" => crate::auth::permissions::DELETE, _ => return Err(StatusCode::METHOD_NOT_ALLOWED), }; if !claims.has_permission(required_bit) { return Err(StatusCode::FORBIDDEN); } // Collect body as typed Values directly — no sentinel needed. let mut body_cols: Vec<(String, Value)> = body .map(|Json(b)| b.into_iter().collect()) .unwrap_or_default(); // Hash the password field for the users table before building the query. if table == "users" && matches!(method_str.to_uppercase().as_str(), "POST" | "PUT") { if let Some(pos) = body_cols.iter().position(|(k, _)| k == "password") { let (_, val) = body_cols.remove(pos); if let Value::String(plaintext) = val { let hash = bcrypt::hash(&plaintext, bcrypt::DEFAULT_COST) .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; body_cols.push(("password_hash".to_string(), Value::String(hash))); } } } let filter_cols: Vec<(String, String)> = query_params.into_iter().collect(); let (sql, params_vals, cache_key) = build_query(method_str, &table, id, &body_cols, &filter_cols) .map_err(|_| StatusCode::BAD_REQUEST)?; let sql = if let Some(entry) = state.query_cache.get(&cache_key) { entry.sql.clone() } else { let entry = CacheEntry::new(sql.clone()); state .query_cache .insert(cache_key, entry.clone(), state.config.cache_max_capacity); entry.sql }; let mut q = sqlx::query(&sql); for val in ¶ms_vals { match val { Value::Null => q = q.bind(Option::::None), Value::Bool(b) => q = q.bind(*b), Value::Number(n) => { if let Some(i) = n.as_i64() { q = q.bind(i); } else if let Some(f) = n.as_f64() { q = q.bind(f); } else { q = q.bind(n.to_string()); } } Value::String(s) => q = q.bind(s.as_str()), other => q = q.bind(other.to_string()), } } let response = match method_str.to_uppercase().as_str() { "GET" => { let rows = q.fetch_all(&state.pool).await.map_err(|e| { tracing::error!("GET {}: {}", table, e); StatusCode::INTERNAL_SERVER_ERROR })?; let mut v = Value::Array(rows.into_iter().map(pg_row_to_json).collect()); if table == "users" { v = strip_password_hash(v); } v } "POST" | "PUT" => { let row = q.fetch_one(&state.pool).await.map_err(|e| { if matches!(e, sqlx::Error::RowNotFound) { StatusCode::NOT_FOUND } else { tracing::error!("{} {}: {}", method_str, table, e); StatusCode::INTERNAL_SERVER_ERROR } })?; let mut v = pg_row_to_json(row); if table == "users" { v = strip_password_hash(v); } v } "DELETE" => { let rows_affected = q .execute(&state.pool) .await .map_err(|e| { tracing::error!("DELETE {}: {}", table, e); StatusCode::INTERNAL_SERVER_ERROR })? .rows_affected(); if rows_affected == 0 { return Err(StatusCode::NOT_FOUND); } serde_json::json!({ "deleted": true }) } _ => return Err(StatusCode::METHOD_NOT_ALLOWED), }; // Reload in-memory caches after mutations to their backing tables. if table == "blacklist" && method_str != "GET" { reload_blacklist(&state).await?; } if table == "cors_origins" && method_str != "GET" { reload_cors(&state).await?; } Ok(Json(response)) } #[cfg(test)] mod tests { use super::*; #[test] fn test_build_select_all() { let (sql, params, key) = build_query("GET", "orders", None, &[], &[]).unwrap(); assert_eq!(sql, "SELECT * FROM orders"); assert!(params.is_empty()); assert_eq!(key, "GET:orders:"); } #[test] fn test_build_select_by_id() { let (sql, params, key) = build_query("GET", "orders", Some("42"), &[], &[]).unwrap(); assert_eq!(sql, "SELECT * FROM orders WHERE id = $1"); assert_eq!(params, vec![Value::Number(42.into())]); assert_eq!(key, "GET:orders:~id"); } #[test] fn test_build_insert() { let cols = vec![ ("email".into(), Value::String("a@b.com".into())), ("name".into(), Value::String("Alice".into())), ]; let (sql, params, key) = build_query("POST", "users", None, &cols, &[]).unwrap(); assert_eq!( sql, "INSERT INTO users (email, name) VALUES ($1, $2) RETURNING *" ); assert_eq!( params, vec![ Value::String("a@b.com".into()), Value::String("Alice".into()) ] ); assert_eq!(key, "POST:users:email,name"); } #[test] fn test_build_update() { let cols = vec![("name".into(), Value::String("Bob".into()))]; let (sql, params, key) = build_query("PUT", "users", Some("7"), &cols, &[]).unwrap(); assert_eq!(sql, "UPDATE users SET name = $1 WHERE id = $2 RETURNING *"); assert_eq!( params, vec![Value::String("Bob".into()), Value::Number(7.into())] ); assert_eq!(key, "PUT:users:name:by_id"); } #[test] fn test_build_delete() { let (sql, params, key) = build_query("DELETE", "users", Some("3"), &[], &[]).unwrap(); assert_eq!(sql, "DELETE FROM users WHERE id = $1"); assert_eq!(params, vec![Value::Number(3.into())]); assert_eq!(key, "DELETE:users:by_id"); } #[test] fn test_build_select_with_filters() { let filters = vec![ ("status".into(), "active".into()), ("role".into(), "admin".into()), ]; let (sql, params, key) = build_query("GET", "users", None, &[], &filters).unwrap(); assert_eq!(sql, "SELECT * FROM users WHERE role = $1 AND status = $2"); assert_eq!( params, vec![ Value::String("admin".into()), Value::String("active".into()) ] ); assert_eq!(key, "GET:users:role,status"); } #[test] fn test_build_null_body_value() { let cols = vec![("note".into(), Value::Null)]; let (sql, params, key) = build_query("POST", "items", None, &cols, &[]).unwrap(); assert_eq!(sql, "INSERT INTO items (note) VALUES ($1) RETURNING *"); assert_eq!(params, vec![Value::Null]); assert_eq!(key, "POST:items:note"); } #[test] fn test_rejects_invalid_table_name() { let result = build_query("GET", "users; DROP TABLE users--", None, &[], &[]); assert!(result.is_err()); } }