phase 4 first review and QA

This commit is contained in:
2026-08-17 17:32:24 +01:00
parent 1a6840d16c
commit b072e2c3dc
3 changed files with 145 additions and 16 deletions
+74 -9
View File
@@ -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();