use axum::http::{HeaderMap, HeaderValue};
use bytes::Bytes;
use fraiseql_core::security::SecurityContext;
use futures::{StreamExt as _, stream};
use super::{
super::{
export_config::ExportConfig,
handler::{ResolvedGetQuery, RestError, RestHandler, set_request_id},
},
guard_formula_injection,
helpers::determine_columns,
};
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 async fn handle_csv_get(
handler: &RestHandler<'_>,
export_config: &ExportConfig,
relative_path: &str,
query_pairs: &[(&str, &str)],
headers: &HeaderMap,
security_context: Option<&SecurityContext>,
) -> Result<CsvResponse, RestError> {
let resolved = handler.resolve_streaming_get_query(
relative_path,
query_pairs,
headers,
security_context,
)?;
let ResolvedGetQuery {
query_name,
query_match,
variables,
params,
..
} = 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 select_columns = super::helpers::export_columns(&query_match);
let rows = super::helpers::export_rows(
handler.executor(),
query_match,
variables,
security_context.cloned(),
params.requested_pagination.export_total(),
)
.await?;
let csv_stream = stream::unfold(
CsvStreamState {
chunks: rows.ready_chunks(usize::try_from(batch_size).unwrap_or(usize::MAX)),
delimiter: ascii_delimiter(export_config.csv_delimiter),
include_bom: export_config.csv_include_bom,
select_columns,
columns: None,
header_emitted: false,
finished: false,
},
|mut state| async move {
if state.finished {
return None;
}
let bytes = serialize_next_csv_chunk(&mut state).await?;
Some((Ok(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 {
chunks: futures::stream::ReadyChunks<fraiseql_core::runtime::JsonRowStream>,
delimiter: u8,
include_bom: bool,
select_columns: Option<Vec<String>>,
columns: Option<Vec<String>>,
header_emitted: bool,
finished: bool,
}
async fn serialize_next_csv_chunk(state: &mut CsvStreamState) -> Option<Bytes> {
let Some(chunk) = state.chunks.next().await else {
if !state.header_emitted {
if let Some(cols) = state.select_columns.clone() {
state.columns = Some(cols.clone());
state.finished = true;
state.header_emitted = true;
return Some(
write_csv_payload(&cols, &[], state.delimiter, state.include_bom, true)
.unwrap_or_else(|err| err),
);
}
}
return None;
};
let mut rows = Vec::with_capacity(chunk.len());
let mut failure = None;
for row in chunk {
match row {
Ok(row) => rows.push(row),
Err(e) => {
failure = Some(e.to_string());
break;
},
}
}
if state.columns.is_none() && !rows.is_empty() {
state.columns = Some(determine_columns(state.select_columns.as_deref(), &rows));
}
let mut bytes = if let Some(columns) = state.columns.clone() {
let emit_header = !state.header_emitted;
state.header_emitted = true;
match write_csv_payload(
&columns,
&rows,
state.delimiter,
state.include_bom && emit_header,
emit_header,
) {
Ok(b) => b.to_vec(),
Err(err) => {
state.finished = true;
return Some(err);
},
}
} else {
Vec::new()
};
if let Some(message) = failure {
bytes.extend_from_slice(&error_csv_line(&message));
state.finished = true;
}
Some(Bytes::from(bytes))
}
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 fn ascii_delimiter(c: char) -> u8 {
if c.is_ascii() && c.len_utf8() == 1 {
c as u8
} else {
b','
}
}
#[cfg(test)]
mod tests;