449 lines
16 KiB
Rust
449 lines
16 KiB
Rust
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::<i64>() {
|
|
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<Value>, 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<String> =
|
|
sorted_filters.iter().map(|(c, _)| c.clone()).collect();
|
|
let where_clause: Vec<String> = 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<Value> = 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<String> = sorted_body.iter().map(|(c, _)| c.clone()).collect();
|
|
let placeholders: Vec<String> = (1..=cols.len()).map(|i| format!("${}", i)).collect();
|
|
let sql = format!(
|
|
"INSERT INTO {} ({}) VALUES ({}) RETURNING *",
|
|
table,
|
|
cols.join(", "),
|
|
placeholders.join(", ")
|
|
);
|
|
let params: Vec<Value> = 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<String> = sorted_body.iter().map(|(c, _)| c.clone()).collect();
|
|
let set_clause: Vec<String> = 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<Value> = 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::<i16, _>(col.ordinal())
|
|
.map(|v| Value::Number(i64::from(v).into()))
|
|
.unwrap_or(Value::Null),
|
|
"INT4" | "SERIAL" => row
|
|
.try_get::<i32, _>(col.ordinal())
|
|
.map(|v| Value::Number(i64::from(v).into()))
|
|
.unwrap_or(Value::Null),
|
|
"INT8" => row
|
|
.try_get::<i64, _>(col.ordinal())
|
|
.map(|v| Value::Number(v.into()))
|
|
.unwrap_or(Value::Null),
|
|
"FLOAT4" | "FLOAT8" => row
|
|
.try_get::<f64, _>(col.ordinal())
|
|
.ok()
|
|
.and_then(serde_json::Number::from_f64)
|
|
.map(Value::Number)
|
|
.unwrap_or(Value::Null),
|
|
"BOOL" => row
|
|
.try_get::<bool, _>(col.ordinal())
|
|
.map(Value::Bool)
|
|
.unwrap_or(Value::Null),
|
|
"UUID" => row
|
|
.try_get::<uuid::Uuid, _>(col.ordinal())
|
|
.map(|v| Value::String(v.to_string()))
|
|
.unwrap_or(Value::Null),
|
|
"TIMESTAMPTZ" | "TIMESTAMP" => row
|
|
.try_get::<chrono::DateTime<chrono::Utc>, _>(col.ordinal())
|
|
.map(|v| Value::String(v.to_rfc3339()))
|
|
.unwrap_or(Value::Null),
|
|
_ => row
|
|
.try_get::<String, _>(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<String> =
|
|
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<AppState>,
|
|
method: Method,
|
|
Extension(claims): Extension<Claims>,
|
|
Path(params): Path<HashMap<String, String>>,
|
|
Query(query_params): Query<HashMap<String, String>>,
|
|
body: Option<Json<HashMap<String, Value>>>,
|
|
) -> Result<Json<Value>, 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::<String>::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());
|
|
}
|
|
}
|