use std::sync::Arc;
use axum::http::{HeaderMap, HeaderValue};
use bytes::Bytes;
use fraiseql_core::{
db::traits::DatabaseAdapter,
runtime::{Executor, QueryMatch},
security::SecurityContext,
};
use futures::stream;
use super::super::{
export_config::ExportConfig,
handler::{PreferHeader, ResolvedGetQuery, RestError, RestHandler, set_request_id},
params::PaginationParams,
};
pub const CSV_CONTENT_TYPE: &str = "text/csv";
#[must_use]
pub fn accepts_csv(headers: &HeaderMap) -> bool {
headers.get("accept").and_then(|v| v.to_str().ok()).is_some_and(|accept| {
accept.split(',').any(|part| {
let media = part.split(';').next().unwrap_or(part).trim();
media.eq_ignore_ascii_case(CSV_CONTENT_TYPE)
})
})
}
pub fn validate_csv_request(
prefer: &PreferHeader,
pagination: &PaginationParams,
) -> Result<(), RestError> {
if prefer.count_exact || prefer.count_planned || prefer.count_estimated {
return Err(RestError::bad_request("count not available for streaming responses"));
}
if let PaginationParams::Offset { offset, .. } = pagination {
if *offset > 0 {
return Err(RestError::bad_request(
"pagination not available for streaming; use filters to narrow results",
));
}
}
if matches!(pagination, PaginationParams::Cursor { .. }) {
return Err(RestError::bad_request(
"pagination not available for streaming; use filters to narrow results",
));
}
Ok(())
}
pub async fn handle_csv_get<A: DatabaseAdapter + 'static>(
handler: &RestHandler<'_, A>,
export_config: &ExportConfig,
relative_path: &str,
query_pairs: &[(&str, &str)],
headers: &HeaderMap,
security_context: Option<&SecurityContext>,
) -> Result<CsvResponse, RestError> {
let resolved = handler.resolve_get_query(relative_path, query_pairs, security_context)?;
let prefer = PreferHeader::from_headers(headers);
validate_csv_request(&prefer, &resolved.params.pagination)?;
let ResolvedGetQuery {
query_name,
query_match,
variables,
..
} = resolved;
let mut response_headers = HeaderMap::new();
set_request_id(headers, &mut response_headers);
response_headers.insert("content-type", HeaderValue::from_static(CSV_CONTENT_TYPE));
let filename = sanitize_filename(&query_name);
let disposition = if filename.is_empty() {
"attachment; filename=\"export.csv\"".to_string()
} else {
format!("attachment; filename=\"{filename}.csv\"")
};
response_headers.insert(
"content-disposition",
HeaderValue::from_str(&disposition)
.unwrap_or_else(|_| HeaderValue::from_static("attachment; filename=\"export.csv\"")),
);
let batch_size = handler.config().ndjson_batch_size.max(1);
let delimiter = ascii_delimiter(export_config.csv_delimiter);
let select_columns = extract_select_columns(query_pairs);
let executor = Arc::clone(handler.executor());
let security_ctx_owned = security_context.cloned();
let csv_stream = stream::unfold(
CsvStreamState {
executor,
query_name,
query_match,
variables,
security_ctx: security_ctx_owned,
batch_size,
offset: 0,
done: false,
delimiter,
include_bom: export_config.csv_include_bom,
select_columns,
columns: None,
header_emitted: false,
},
|mut state| async move {
if state.done {
return None;
}
match fetch_and_serialize_csv_batch(&mut state).await {
Ok(Some(bytes)) => Some((Ok(bytes), state)),
Ok(None) => None,
Err(err_bytes) => {
state.done = true;
Some((Ok(err_bytes), state))
},
}
},
);
Ok(CsvResponse {
headers: response_headers,
body: CsvBody::Stream(Box::pin(csv_stream)),
})
}
fn sanitize_filename(name: &str) -> String {
name.chars()
.filter(|c| c.is_ascii_alphanumeric() || *c == '_' || *c == '-')
.collect()
}
pub struct CsvResponse {
pub headers: HeaderMap,
pub body: CsvBody,
}
#[non_exhaustive]
pub enum CsvBody {
Stream(
std::pin::Pin<
Box<dyn futures::Stream<Item = Result<Bytes, std::convert::Infallible>> + Send>,
>,
),
}
impl CsvBody {
pub fn into_body(self) -> axum::body::Body {
match self {
Self::Stream(stream) => axum::body::Body::from_stream(stream),
}
}
}
struct CsvStreamState<A: DatabaseAdapter> {
executor: Arc<Executor<A>>,
query_name: String,
query_match: QueryMatch,
variables: serde_json::Value,
security_ctx: Option<SecurityContext>,
batch_size: u64,
offset: u64,
done: bool,
delimiter: u8,
include_bom: bool,
select_columns: Option<Vec<String>>,
columns: Option<Vec<String>>,
header_emitted: bool,
}
async fn fetch_and_serialize_csv_batch<A: DatabaseAdapter>(
state: &mut CsvStreamState<A>,
) -> Result<Option<Bytes>, Bytes> {
let mut batch_vars = state.variables.clone();
if let Some(obj) = batch_vars.as_object_mut() {
obj.insert("limit".to_string(), serde_json::json!(state.batch_size));
if state.offset > 0 {
obj.insert("offset".to_string(), serde_json::json!(state.offset));
}
}
let vars_ref = if batch_vars.as_object().is_none_or(serde_json::Map::is_empty) {
None
} else {
Some(&batch_vars)
};
let result_value = match state
.executor
.execute_query_direct(&state.query_match, vars_ref, state.security_ctx.as_ref())
.await
{
Ok(r) => r,
Err(e) => {
state.done = true;
return Err(error_csv_line(&e.to_string()));
},
};
let rows = match super::helpers::extract_rows(&result_value, &state.query_name) {
Ok(r) => r,
Err(e) => {
state.done = true;
return Err(error_csv_line(&e.message));
},
};
if rows.is_empty() {
if !state.header_emitted && state.offset == 0 {
if let Some(cols) = state.select_columns.clone() {
state.columns = Some(cols);
let bytes = match serialize_batch(state, &[]) {
Ok(b) => b,
Err(err_bytes) => {
state.done = true;
return Err(err_bytes);
},
};
state.done = true;
return Ok(Some(bytes));
}
}
state.done = true;
return Ok(None);
}
if state.columns.is_none() {
state.columns = Some(determine_columns(state.select_columns.as_deref(), &rows));
}
let bytes = match serialize_batch(state, &rows) {
Ok(b) => b,
Err(err_bytes) => {
state.done = true;
return Err(err_bytes);
},
};
#[allow(clippy::cast_possible_truncation)]
let row_count = rows.len() as u64;
if row_count < state.batch_size {
state.done = true;
} else {
state.offset += state.batch_size;
}
Ok(Some(bytes))
}
fn serialize_batch<A: DatabaseAdapter>(
state: &mut CsvStreamState<A>,
rows: &[serde_json::Value],
) -> Result<Bytes, Bytes> {
let columns = state
.columns
.as_ref()
.ok_or_else(|| error_csv_line("internal error: columns not initialised"))?;
let payload = write_csv_payload(
columns,
rows,
state.delimiter,
state.include_bom && !state.header_emitted,
!state.header_emitted,
)?;
state.header_emitted = true;
Ok(payload)
}
fn write_csv_payload(
columns: &[String],
rows: &[serde_json::Value],
delimiter: u8,
emit_bom: bool,
emit_header: bool,
) -> Result<Bytes, Bytes> {
let mut buf: Vec<u8> = Vec::new();
if emit_bom {
buf.extend_from_slice("\u{FEFF}".as_bytes());
}
{
let mut wtr = csv::WriterBuilder::new().delimiter(delimiter).from_writer(&mut buf);
if emit_header {
wtr.write_record(columns.iter().map(String::as_str))
.map_err(|e| error_csv_line(&e.to_string()))?;
}
for row in rows {
let record: Vec<String> = columns
.iter()
.map(|c| value_to_csv_field(row.get(c).unwrap_or(&serde_json::Value::Null)))
.collect();
wtr.write_record(record.iter().map(String::as_str))
.map_err(|e| error_csv_line(&e.to_string()))?;
}
wtr.flush().map_err(|e| error_csv_line(&e.to_string()))?;
}
Ok(Bytes::from(buf))
}
fn error_csv_line(message: &str) -> Bytes {
let one_line: String = message.chars().map(|c| if c == '\n' { ' ' } else { c }).collect();
Bytes::from(format!("# error: {one_line}\n"))
}
fn value_to_csv_field(v: &serde_json::Value) -> String {
match v {
serde_json::Value::Null => String::new(),
serde_json::Value::Bool(b) => b.to_string(),
serde_json::Value::Number(n) => guard_formula_injection(&n.to_string()),
serde_json::Value::String(s) => guard_formula_injection(s),
other => guard_formula_injection(&serde_json::to_string(other).unwrap_or_default()),
}
}
const FORMULA_INJECTION_SENTINELS: [char; 6] = ['=', '+', '-', '@', '\t', '\r'];
pub(crate) fn guard_formula_injection(value: &str) -> String {
match value.chars().next() {
Some(c) if FORMULA_INJECTION_SENTINELS.contains(&c) => {
let mut out = String::with_capacity(value.len() + 1);
out.push('\'');
out.push_str(value);
out
},
_ => value.to_owned(),
}
}
fn determine_columns(select_columns: Option<&[String]>, rows: &[serde_json::Value]) -> Vec<String> {
if let Some(cols) = select_columns {
return cols.to_vec();
}
rows.first()
.and_then(|v| v.as_object())
.map(|m| m.keys().cloned().collect())
.unwrap_or_default()
}
fn extract_select_columns(query_pairs: &[(&str, &str)]) -> Option<Vec<String>> {
let raw = query_pairs.iter().find(|(k, _)| *k == "select").map(|(_, v)| *v)?;
let cols = parse_select_top_level(raw);
if cols.is_empty() { None } else { Some(cols) }
}
fn parse_select_top_level(select_raw: &str) -> Vec<String> {
let mut cols = Vec::new();
let mut depth = 0_usize;
let mut current = String::new();
for c in select_raw.chars() {
match c {
'(' => {
depth += 1;
current.push(c);
},
')' => {
depth = depth.saturating_sub(1);
current.push(c);
},
',' if depth == 0 => {
push_top_level(&mut cols, ¤t);
current.clear();
},
_ => current.push(c),
}
}
push_top_level(&mut cols, ¤t);
cols
}
fn push_top_level(cols: &mut Vec<String>, current: &str) {
let trimmed = current.trim();
if trimmed.is_empty() {
return;
}
let head = trimmed.split('(').next().unwrap_or("").trim();
if !head.is_empty() {
cols.push(head.to_string());
}
}
const fn ascii_delimiter(c: char) -> u8 {
if c.is_ascii() && c.len_utf8() == 1 {
c as u8
} else {
b','
}
}
#[cfg(test)]
mod tests;