phase 4 first review and QA
This commit is contained in:
+74
-9
@@ -6,7 +6,7 @@ use axum::{
|
||||
Json, Router,
|
||||
body::Body,
|
||||
extract::{Query, State},
|
||||
http::{StatusCode, header},
|
||||
http::{HeaderValue, Method, StatusCode, header},
|
||||
response::{IntoResponse, Response},
|
||||
routing::get,
|
||||
};
|
||||
@@ -15,7 +15,7 @@ use rust_embed::Embed;
|
||||
use serde::Deserialize;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::Mutex;
|
||||
use tower_http::cors::CorsLayer;
|
||||
use tower_http::cors::{AllowOrigin, CorsLayer};
|
||||
use tower_http::trace::TraceLayer;
|
||||
|
||||
#[derive(Clone)]
|
||||
@@ -46,6 +46,8 @@ fn default_limit() -> u32 {
|
||||
10000
|
||||
}
|
||||
|
||||
const MAX_LIMIT: u32 = 10000;
|
||||
|
||||
#[derive(Debug, Clone, Deserialize)]
|
||||
pub struct SummaryQuery {
|
||||
pub species: Option<String>,
|
||||
@@ -87,6 +89,23 @@ impl From<db::DbError> for ApiError {
|
||||
type ApiResult<T> = std::result::Result<T, ApiError>;
|
||||
|
||||
pub fn build_router(state: AppState) -> Router {
|
||||
let cors = if state.config.allowed_origins.is_empty() {
|
||||
tracing::warn!("CORS is permissive — no allowed_origins configured");
|
||||
CorsLayer::permissive()
|
||||
} else {
|
||||
let origins: Vec<HeaderValue> = state
|
||||
.config
|
||||
.allowed_origins
|
||||
.iter()
|
||||
.filter_map(|s| s.parse().ok())
|
||||
.collect();
|
||||
|
||||
CorsLayer::new()
|
||||
.allow_methods([Method::GET])
|
||||
.allow_headers([header::CONTENT_TYPE])
|
||||
.allow_origin(AllowOrigin::list(origins))
|
||||
};
|
||||
|
||||
Router::new()
|
||||
.route("/healthz", get(healthz))
|
||||
.route("/api/species", get(get_species))
|
||||
@@ -96,7 +115,7 @@ pub fn build_router(state: AppState) -> Router {
|
||||
.route("/api/summary", get(get_summary))
|
||||
.route("/api/export.parquet", get(export_parquet))
|
||||
.fallback(static_handler)
|
||||
.layer(CorsLayer::permissive())
|
||||
.layer(cors)
|
||||
.layer(TraceLayer::new_for_http())
|
||||
.with_state(state)
|
||||
}
|
||||
@@ -132,6 +151,9 @@ async fn get_landings(
|
||||
Query(params): Query<LandingsQuery>,
|
||||
) -> ApiResult<Json<Vec<LandingDto>>> {
|
||||
let conn = state.conn.clone();
|
||||
// Apply hard cap to prevent abuse
|
||||
let capped_limit = std::cmp::min(params.limit, MAX_LIMIT);
|
||||
|
||||
let rows = tokio::task::spawn_blocking(move || -> db::Result<Vec<LandingDto>> {
|
||||
let conn = conn.blocking_lock();
|
||||
|
||||
@@ -186,7 +208,7 @@ async fn get_landings(
|
||||
sql.push_str(&format!(
|
||||
" ORDER BY month, species_code, measure_code LIMIT ${idx}"
|
||||
));
|
||||
args.push(Box::new(params.limit as i64));
|
||||
args.push(Box::new(capped_limit as i64));
|
||||
|
||||
let arg_refs: Vec<&dyn duckdb::ToSql> = args.iter().map(|b| b.as_ref()).collect();
|
||||
|
||||
@@ -363,10 +385,15 @@ async fn export_parquet(State(state): State<AppState>) -> ApiResult<Response> {
|
||||
tokio::task::spawn_blocking(move || -> ApiResult<(String, std::fs::File)> {
|
||||
let conn = conn.blocking_lock();
|
||||
|
||||
let tmp = tempfile::NamedTempFile::new().map_err(|e| ApiError {
|
||||
status: StatusCode::INTERNAL_SERVER_ERROR,
|
||||
message: format!("Failed to create temp file: {e}"),
|
||||
})?;
|
||||
// Use builder to ensure consistent naming for cleanup
|
||||
let tmp = tempfile::Builder::new()
|
||||
.prefix("hagfish-")
|
||||
.suffix(".parquet")
|
||||
.tempfile()
|
||||
.map_err(|e| ApiError {
|
||||
status: StatusCode::INTERNAL_SERVER_ERROR,
|
||||
message: format!("Failed to create temp file: {e}"),
|
||||
})?;
|
||||
|
||||
let (_kept_file, path) = tmp.keep().map_err(|e| ApiError {
|
||||
status: StatusCode::INTERNAL_SERVER_ERROR,
|
||||
@@ -604,7 +631,10 @@ mod tests {
|
||||
|
||||
AppState {
|
||||
conn: Arc::new(Mutex::new(conn)),
|
||||
config: Arc::new(Config::default()),
|
||||
config: Arc::new(Config {
|
||||
allowed_origins: vec!["https://hagfisk.poc.fló.fo".to_string()],
|
||||
..Default::default()
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -709,6 +739,41 @@ mod tests {
|
||||
assert_eq!(body.len(), 3);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_get_landings_range_filter() {
|
||||
let state = test_state();
|
||||
let app = build_router(state);
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
|
||||
tokio::spawn(async move {
|
||||
axum::serve(listener, app).await.unwrap();
|
||||
});
|
||||
|
||||
// Test single month range (same from and to)
|
||||
let resp = reqwest::get(format!(
|
||||
"http://{addr}/api/landings?month_from=2024M01&month_to=2024M01"
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
let body: Vec<LandingDto> = resp.json().await.unwrap();
|
||||
assert!(body.iter().all(|l| l.month == "2024M01"));
|
||||
assert_eq!(body.len(), 4);
|
||||
|
||||
// Test multi-month range
|
||||
let resp = reqwest::get(format!(
|
||||
"http://{addr}/api/landings?month_from=2024M01&month_to=2024M02"
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
let body: Vec<LandingDto> = resp.json().await.unwrap();
|
||||
assert!(body.len() > 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_get_summary_monthly_aggregates() {
|
||||
let state = test_state();
|
||||
|
||||
Reference in New Issue
Block a user