phase 3 complete

This commit is contained in:
2026-08-17 14:08:20 +01:00
parent f639b19c13
commit d52bcf2d64
11 changed files with 1260 additions and 221 deletions
+866
View File
@@ -0,0 +1,866 @@
use crate::db;
use crate::types::{
Config, LandingDto, MonthlyAggregate, PriceTrend, SpeciesDto, SummaryDto, TopSpecies,
};
use axum::{
Json, Router,
body::Body,
extract::{Query, State},
http::{StatusCode, header},
response::{IntoResponse, Response},
routing::get,
};
use duckdb::params;
use rust_embed::Embed;
use serde::Deserialize;
use std::sync::Arc;
use tokio::sync::Mutex;
use tower_http::cors::CorsLayer;
use tower_http::trace::TraceLayer;
#[derive(Clone)]
pub struct AppState {
pub conn: Arc<Mutex<duckdb::Connection>>,
#[allow(dead_code)]
pub config: Arc<Config>,
}
#[derive(Embed)]
#[folder = "static/"]
struct StaticAssets;
#[derive(Debug, Clone, Deserialize)]
pub struct LandingsQuery {
pub month: Option<String>,
pub species: Option<String>,
pub gear: Option<String>,
pub zone: Option<String>,
pub measure: Option<String>,
#[serde(default = "default_limit")]
pub limit: u32,
}
fn default_limit() -> u32 {
10000
}
#[derive(Debug, Clone, Deserialize)]
pub struct SummaryQuery {
pub species: Option<String>,
}
#[derive(Debug)]
pub struct ApiError {
pub status: StatusCode,
pub message: String,
}
impl IntoResponse for ApiError {
fn into_response(self) -> Response {
let body = serde_json::json!({ "error": self.message });
(self.status, Json(body)).into_response()
}
}
impl From<duckdb::Error> for ApiError {
fn from(e: duckdb::Error) -> Self {
tracing::error!("DuckDB error: {e}");
ApiError {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: format!("Database error: {e}"),
}
}
}
impl From<db::DbError> for ApiError {
fn from(e: db::DbError) -> Self {
tracing::error!("DB error: {e}");
ApiError {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: format!("Database error: {e}"),
}
}
}
type ApiResult<T> = std::result::Result<T, ApiError>;
pub fn build_router(state: AppState) -> Router {
Router::new()
.route("/healthz", get(healthz))
.route("/api/species", get(get_species))
.route("/api/landings", get(get_landings))
.route("/api/summary", get(get_summary))
.route("/api/export.parquet", get(export_parquet))
.fallback(static_handler)
.layer(CorsLayer::permissive())
.layer(TraceLayer::new_for_http())
.with_state(state)
}
async fn healthz(State(_state): State<AppState>) -> impl IntoResponse {
StatusCode::OK
}
async fn get_species(State(state): State<AppState>) -> ApiResult<Json<Vec<SpeciesDto>>> {
let conn = state.conn.clone();
let rows = tokio::task::spawn_blocking(move || -> db::Result<Vec<SpeciesDto>> {
let conn = conn.blocking_lock();
let mut stmt = conn.prepare("SELECT code, label FROM species ORDER BY code")?;
let rows = stmt.query_map([], |row| {
Ok(SpeciesDto {
code: row.get(0)?,
label: row.get(1)?,
})
})?;
Ok(rows.collect::<std::result::Result<Vec<_>, duckdb::Error>>()?)
})
.await
.map_err(|e| ApiError {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: format!("Task join error: {e}"),
})??;
Ok(Json(rows))
}
async fn get_landings(
State(state): State<AppState>,
Query(params): Query<LandingsQuery>,
) -> ApiResult<Json<Vec<LandingDto>>> {
let conn = state.conn.clone();
let rows = tokio::task::spawn_blocking(move || -> db::Result<Vec<LandingDto>> {
let conn = conn.blocking_lock();
let mut sql = String::from(
"SELECT month, species_code, species_label, gear_code, zone_code, \
processing_code, preservation_code, shipsize_code, measure_code, value \
FROM landings WHERE 1=1",
);
let mut args: Vec<Box<dyn duckdb::ToSql>> = Vec::new();
let mut idx = 1;
if let Some(ref month) = params.month {
sql.push_str(&format!(" AND month = ${idx}"));
args.push(Box::new(month.clone()));
idx += 1;
}
if let Some(ref species) = params.species {
sql.push_str(&format!(" AND species_code = ${idx}"));
args.push(Box::new(species.clone()));
idx += 1;
}
if let Some(ref gear) = params.gear {
sql.push_str(&format!(" AND gear_code = ${idx}"));
args.push(Box::new(gear.clone()));
idx += 1;
}
if let Some(ref zone) = params.zone {
sql.push_str(&format!(" AND zone_code = ${idx}"));
args.push(Box::new(zone.clone()));
idx += 1;
}
if let Some(ref measure) = params.measure {
sql.push_str(&format!(" AND measure_code = ${idx}"));
args.push(Box::new(measure.clone()));
idx += 1;
}
sql.push_str(&format!(
" ORDER BY month, species_code, measure_code LIMIT ${idx}"
));
args.push(Box::new(params.limit as i64));
let arg_refs: Vec<&dyn duckdb::ToSql> = args.iter().map(|b| b.as_ref()).collect();
let mut stmt = conn.prepare(&sql)?;
let rows = stmt.query_map(arg_refs.as_slice(), |row| {
Ok(LandingDto {
month: row.get(0)?,
species_code: row.get(1)?,
species_label: row.get(2)?,
gear_code: row.get(3)?,
zone_code: row.get(4)?,
processing_code: row.get(5)?,
preservation_code: row.get(6)?,
shipsize_code: row.get(7)?,
measure_code: row.get(8)?,
value: row.get(9)?,
})
})?;
Ok(rows.collect::<std::result::Result<Vec<_>, duckdb::Error>>()?)
})
.await
.map_err(|e| ApiError {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: format!("Task join error: {e}"),
})??;
Ok(Json(rows))
}
async fn get_summary(
State(state): State<AppState>,
Query(params): Query<SummaryQuery>,
) -> ApiResult<Json<SummaryDto>> {
let conn = state.conn.clone();
let species_filter = params.species.clone();
let result = tokio::task::spawn_blocking(move || -> db::Result<SummaryDto> {
let conn = conn.blocking_lock();
let where_clause = if species_filter.is_some() {
" WHERE species_code = ?".to_string()
} else {
String::new()
};
let monthly_sql = format!(
"SELECT month, \
SUM(CASE WHEN measure_code = 'MASS' THEN value END) AS total_mass, \
SUM(CASE WHEN measure_code = 'VALUE' THEN value END) AS total_value \
FROM landings{where_clause} \
GROUP BY month ORDER BY month"
);
let monthly: Vec<MonthlyAggregate> = if let Some(ref sp) = species_filter {
let mut stmt = conn.prepare(&monthly_sql)?;
let rows = stmt.query_map(params![sp], |row| {
Ok(MonthlyAggregate {
month: row.get(0)?,
total_mass: row.get(1)?,
total_value: row.get(2)?,
})
})?;
Ok::<Vec<MonthlyAggregate>, duckdb::Error>(
rows.collect::<std::result::Result<Vec<_>, _>>()?,
)?
} else {
let mut stmt = conn.prepare(&monthly_sql)?;
let rows = stmt.query_map([], |row| {
Ok(MonthlyAggregate {
month: row.get(0)?,
total_mass: row.get(1)?,
total_value: row.get(2)?,
})
})?;
Ok::<Vec<MonthlyAggregate>, duckdb::Error>(
rows.collect::<std::result::Result<Vec<_>, _>>()?,
)?
};
let top_sql = format!(
"SELECT species_code, \
COALESCE(MAX(species_label), species_code) AS species_label, \
SUM(CASE WHEN measure_code = 'VALUE' THEN value END) AS total_value, \
SUM(CASE WHEN measure_code = 'MASS' THEN value END) AS total_mass \
FROM landings{where_clause} \
GROUP BY species_code, species_label \
ORDER BY total_value DESC NULLS LAST \
LIMIT 10"
);
let top_species: Vec<TopSpecies> = if let Some(ref sp) = species_filter {
let mut stmt = conn.prepare(&top_sql)?;
let rows = stmt.query_map(params![sp], |row| {
Ok(TopSpecies {
species_code: row.get(0)?,
species_label: row.get(1)?,
total_value: row.get(2)?,
total_mass: row.get(3)?,
})
})?;
Ok::<Vec<TopSpecies>, duckdb::Error>(rows.collect::<std::result::Result<Vec<_>, _>>()?)?
} else {
let mut stmt = conn.prepare(&top_sql)?;
let rows = stmt.query_map([], |row| {
Ok(TopSpecies {
species_code: row.get(0)?,
species_label: row.get(1)?,
total_value: row.get(2)?,
total_mass: row.get(3)?,
})
})?;
Ok::<Vec<TopSpecies>, duckdb::Error>(rows.collect::<std::result::Result<Vec<_>, _>>()?)?
};
let price_sql = format!(
"WITH monthly_mass AS ( \
SELECT month, SUM(value) AS mass FROM landings \
WHERE measure_code = 'MASS'{where_clause} \
GROUP BY month \
), \
monthly_value AS ( \
SELECT month, SUM(value) AS value FROM landings \
WHERE measure_code = 'VALUE'{where_clause} \
GROUP BY month \
) \
SELECT m.month, \
CASE WHEN m.mass IS NOT NULL AND m.mass > 0 \
THEN v.value / m.mass END AS price_per_kg \
FROM monthly_mass m \
JOIN monthly_value v ON m.month = v.month \
ORDER BY m.month"
);
let price_trend: Vec<PriceTrend> = if let Some(ref sp) = species_filter {
let mut stmt = conn.prepare(&price_sql)?;
let rows = stmt.query_map(params![sp, sp], |row| {
Ok(PriceTrend {
month: row.get(0)?,
price_per_kg: row.get(1)?,
})
})?;
Ok::<Vec<PriceTrend>, duckdb::Error>(rows.collect::<std::result::Result<Vec<_>, _>>()?)?
} else {
let mut stmt = conn.prepare(&price_sql)?;
let rows = stmt.query_map([], |row| {
Ok(PriceTrend {
month: row.get(0)?,
price_per_kg: row.get(1)?,
})
})?;
Ok::<Vec<PriceTrend>, duckdb::Error>(rows.collect::<std::result::Result<Vec<_>, _>>()?)?
};
Ok(SummaryDto {
monthly,
top_species,
price_trend,
})
})
.await
.map_err(|e| ApiError {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: format!("Task join error: {e}"),
})??;
Ok(Json(result))
}
async fn export_parquet(State(state): State<AppState>) -> ApiResult<Response> {
let conn = state.conn.clone();
let (file_path, file) =
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}"),
})?;
let (_kept_file, path) = tmp.keep().map_err(|e| ApiError {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: format!("Failed to persist temp file: {e}"),
})?;
let path_str = path
.to_str()
.ok_or_else(|| ApiError {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: "Temp path contains invalid UTF-8".to_string(),
})?
.to_string();
db::export_parquet(&conn, &path_str)?;
let file = std::fs::File::open(&path_str).map_err(|e| ApiError {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: format!("Failed to open exported file: {e}"),
})?;
Ok((path_str, file))
})
.await
.map_err(|e| ApiError {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: format!("Task join error: {e}"),
})??;
let tokio_file = tokio::fs::File::from(file);
let stream = tokio_util::io::ReaderStream::new(tokio_file);
let body = Body::from_stream(stream);
let cleanup_path = file_path.clone();
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_secs(300)).await;
if let Err(e) = tokio::fs::remove_file(&cleanup_path).await {
tracing::warn!("Failed to clean up parquet file {}: {e}", cleanup_path);
}
});
Ok(Response::builder()
.status(StatusCode::OK)
.header(header::CONTENT_TYPE, "application/octet-stream")
.header(
header::CONTENT_DISPOSITION,
"attachment; filename=\"landings.parquet\"",
)
.body(body)
.unwrap())
}
async fn static_handler(uri: axum::http::Uri) -> Response {
let path = uri.path().trim_start_matches('/');
let asset = StaticAssets::get(path).or_else(|| StaticAssets::get("index.html"));
match asset {
Some(file) => {
let mime = mime_guess::from_path(path)
.first_or_octet_stream()
.as_ref()
.to_string();
Response::builder()
.status(StatusCode::OK)
.header(header::CONTENT_TYPE, mime)
.body(Body::from(file.data.into_owned()))
.unwrap()
}
None => Response::builder()
.status(StatusCode::NOT_FOUND)
.body(Body::from("Not found"))
.unwrap(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::Landing;
use duckdb::Connection;
use std::collections::HashMap;
fn test_state() -> AppState {
let conn = Connection::open_in_memory().unwrap();
db::init_schema(&conn).unwrap();
let rows = vec![
Landing {
month: "2024M01".to_string(),
species_code: "COD".to_string(),
species_label: "Toskur".to_string(),
gear_code: "TOTAL".to_string(),
zone_code: "TOTAL".to_string(),
processing_code: "TOTAL".to_string(),
preservation_code: "TOTAL".to_string(),
shipsize_code: "TOTAL".to_string(),
measure_code: "MASS".to_string(),
value: Some(1000.0),
},
Landing {
month: "2024M01".to_string(),
species_code: "COD".to_string(),
species_label: "Toskur".to_string(),
gear_code: "TOTAL".to_string(),
zone_code: "TOTAL".to_string(),
processing_code: "TOTAL".to_string(),
preservation_code: "TOTAL".to_string(),
shipsize_code: "TOTAL".to_string(),
measure_code: "VALUE".to_string(),
value: Some(5000.0),
},
Landing {
month: "2024M01".to_string(),
species_code: "HER".to_string(),
species_label: "Sild".to_string(),
gear_code: "TOTAL".to_string(),
zone_code: "TOTAL".to_string(),
processing_code: "TOTAL".to_string(),
preservation_code: "TOTAL".to_string(),
shipsize_code: "TOTAL".to_string(),
measure_code: "MASS".to_string(),
value: Some(2000.0),
},
Landing {
month: "2024M01".to_string(),
species_code: "HER".to_string(),
species_label: "Sild".to_string(),
gear_code: "TOTAL".to_string(),
zone_code: "TOTAL".to_string(),
processing_code: "TOTAL".to_string(),
preservation_code: "TOTAL".to_string(),
shipsize_code: "TOTAL".to_string(),
measure_code: "VALUE".to_string(),
value: Some(3000.0),
},
Landing {
month: "2024M02".to_string(),
species_code: "COD".to_string(),
species_label: "Toskur".to_string(),
gear_code: "TOTAL".to_string(),
zone_code: "TOTAL".to_string(),
processing_code: "TOTAL".to_string(),
preservation_code: "TOTAL".to_string(),
shipsize_code: "TOTAL".to_string(),
measure_code: "MASS".to_string(),
value: Some(1500.0),
},
Landing {
month: "2024M02".to_string(),
species_code: "COD".to_string(),
species_label: "Toskur".to_string(),
gear_code: "TOTAL".to_string(),
zone_code: "TOTAL".to_string(),
processing_code: "TOTAL".to_string(),
preservation_code: "TOTAL".to_string(),
shipsize_code: "TOTAL".to_string(),
measure_code: "VALUE".to_string(),
value: Some(7500.0),
},
];
db::update_lookups(
&conn,
&HashMap::from([(
"Species (ASFIS2022)".to_string(),
HashMap::from([
("COD".to_string(), "Toskur".to_string()),
("HER".to_string(), "Sild".to_string()),
]),
)]),
)
.unwrap();
db::upsert_landings(&conn, &rows).unwrap();
AppState {
conn: Arc::new(Mutex::new(conn)),
config: Arc::new(Config::default()),
}
}
#[tokio::test]
async fn test_healthz_returns_ok() {
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();
});
let resp = reqwest::get(format!("http://{addr}/healthz"))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_get_species_returns_all() {
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();
});
let resp = reqwest::get(format!("http://{addr}/api/species"))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: Vec<SpeciesDto> = resp.json().await.unwrap();
assert!(body.iter().any(|s| s.code == "COD" && s.label == "Toskur"));
assert!(body.iter().any(|s| s.code == "HER" && s.label == "Sild"));
}
#[tokio::test]
async fn test_get_landings_no_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();
});
let resp = reqwest::get(format!("http://{addr}/api/landings"))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: Vec<LandingDto> = resp.json().await.unwrap();
assert!(!body.is_empty(), "Should have landing rows");
assert_eq!(body.len(), 6);
}
#[tokio::test]
async fn test_get_landings_filter_by_species() {
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();
});
let resp = reqwest::get(format!("http://{addr}/api/landings?species=COD"))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: Vec<LandingDto> = resp.json().await.unwrap();
assert!(body.iter().all(|l| l.species_code == "COD"));
assert_eq!(body.len(), 4);
}
#[tokio::test]
async fn test_get_landings_filter_by_measure() {
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();
});
let resp = reqwest::get(format!("http://{addr}/api/landings?measure=MASS"))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: Vec<LandingDto> = resp.json().await.unwrap();
assert!(body.iter().all(|l| l.measure_code == "MASS"));
assert_eq!(body.len(), 3);
}
#[tokio::test]
async fn test_get_summary_monthly_aggregates() {
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();
});
let resp = reqwest::get(format!("http://{addr}/api/summary"))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: SummaryDto = resp.json().await.unwrap();
assert_eq!(body.monthly.len(), 2);
let jan = body.monthly.iter().find(|m| m.month == "2024M01").unwrap();
assert!((jan.total_mass.unwrap() - 3000.0).abs() < f64::EPSILON);
assert!((jan.total_value.unwrap() - 8000.0).abs() < f64::EPSILON);
let feb = body.monthly.iter().find(|m| m.month == "2024M02").unwrap();
assert!((feb.total_mass.unwrap() - 1500.0).abs() < f64::EPSILON);
assert!((feb.total_value.unwrap() - 7500.0).abs() < f64::EPSILON);
}
#[tokio::test]
async fn test_get_summary_top_species() {
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();
});
let resp = reqwest::get(format!("http://{addr}/api/summary"))
.await
.unwrap();
let body: SummaryDto = resp.json().await.unwrap();
assert!(!body.top_species.is_empty());
assert_eq!(body.top_species[0].species_code, "COD");
assert!((body.top_species[0].total_value.unwrap() - 12500.0).abs() < f64::EPSILON);
}
#[tokio::test]
async fn test_get_summary_price_trend() {
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();
});
let resp = reqwest::get(format!("http://{addr}/api/summary"))
.await
.unwrap();
let body: SummaryDto = resp.json().await.unwrap();
assert_eq!(body.price_trend.len(), 2);
let jan = body
.price_trend
.iter()
.find(|p| p.month == "2024M01")
.unwrap();
let expected_jan = 8000.0 / 3000.0;
assert!((jan.price_per_kg.unwrap() - expected_jan).abs() < 0.01);
let feb = body
.price_trend
.iter()
.find(|p| p.month == "2024M02")
.unwrap();
assert!((feb.price_per_kg.unwrap() - 5.0).abs() < f64::EPSILON);
}
#[tokio::test]
async fn test_get_landings_faroese_label_preserved() {
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();
});
let resp = reqwest::get(format!("http://{addr}/api/landings?species=COD"))
.await
.unwrap();
let body: Vec<LandingDto> = resp.json().await.unwrap();
let cod = &body[0];
assert_eq!(cod.species_label, "Toskur");
assert!(!cod.species_label.contains('\u{FFFD}'));
}
#[tokio::test]
async fn test_get_landings_empty_result() {
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();
});
let resp = reqwest::get(format!("http://{addr}/api/landings?species=NONEXISTENT"))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: Vec<LandingDto> = resp.json().await.unwrap();
assert!(body.is_empty());
}
#[tokio::test]
async fn test_get_landings_limit_applied() {
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();
});
let resp = reqwest::get(format!("http://{addr}/api/landings?limit=2"))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: Vec<LandingDto> = resp.json().await.unwrap();
assert_eq!(body.len(), 2);
}
#[tokio::test]
async fn test_export_parquet() {
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();
});
let resp = reqwest::get(format!("http://{addr}/api/export.parquet"))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(
resp.headers().get(header::CONTENT_DISPOSITION).unwrap(),
"attachment; filename=\"landings.parquet\""
);
let bytes = resp.bytes().await.unwrap();
assert!(bytes.len() > 4);
assert_eq!(&bytes[..4], b"PAR1");
}
#[tokio::test]
async fn test_concurrent_export_requests() {
let state = test_state();
let app = build_router(state.clone());
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();
});
let resp1 = reqwest::get(format!("http://{addr}/api/export.parquet"))
.await
.unwrap();
let resp2 = reqwest::get(format!("http://{addr}/api/export.parquet"))
.await
.unwrap();
assert_eq!(resp1.status(), StatusCode::OK);
assert_eq!(resp2.status(), StatusCode::OK);
let bytes1 = resp1.bytes().await.unwrap();
let bytes2 = resp2.bytes().await.unwrap();
assert_eq!(&bytes1[..4], b"PAR1");
assert_eq!(&bytes2[..4], b"PAR1");
assert_ne!(bytes1.len(), 0);
assert_ne!(bytes2.len(), 0);
}
#[tokio::test]
async fn test_sql_injection_safe() {
let state = test_state();
let state_clone = state.clone();
let app = build_router(state.clone());
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();
});
let malicious_params = "?species='; DROP TABLE landings; --";
let url = format!("http://{addr}/api/landings{malicious_params}");
let resp = reqwest::get(&url).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let count: i64 = tokio::task::spawn_blocking(move || {
let conn = state_clone.conn.blocking_lock();
conn.query_row("SELECT COUNT(*) FROM landings", [], |row| row.get(0))
.unwrap()
})
.await
.unwrap();
assert_eq!(count, 6);
}
}