#![allow(dead_code)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CopyFormat {
Text,
Csv,
Binary,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct CopyStatement {
pub table: String,
pub columns: Vec<String>,
pub to_stdout: bool,
pub format: CopyFormat,
}
pub(crate) fn parse_copy(sql: &str) -> Option<CopyStatement> {
let s = sql.trim().trim_end_matches(';').trim();
let mut rest = strip_kw(s, "COPY")?;
rest = rest.trim_start();
let (table, mut rest) = take_ident(rest)?;
rest = rest.trim_start();
let mut columns = Vec::new();
if let Some(after_paren) = rest.strip_prefix('(') {
let close = after_paren.find(')')?;
let cols = &after_paren[..close];
for c in cols.split(',') {
let c = c.trim().trim_matches('"').trim();
if c.is_empty() {
return None;
}
columns.push(c.to_string());
}
rest = after_paren[close + 1..].trim_start();
}
let to_stdout = if let Some(r) = strip_kw(rest, "FROM") {
rest = strip_kw(r.trim_start(), "STDIN")?;
false
} else if let Some(r) = strip_kw(rest, "TO") {
rest = strip_kw(r.trim_start(), "STDOUT")?;
true
} else {
return None;
};
let format = parse_format(rest.trim_start());
Some(CopyStatement {
table,
columns,
to_stdout,
format,
})
}
fn parse_format(mut rest: &str) -> CopyFormat {
if rest.is_empty() {
return CopyFormat::Text;
}
if let Some(r) = strip_kw(rest, "WITH") {
rest = r.trim_start();
}
let body = rest.trim();
let scan = body.trim_start_matches('(').trim_end_matches(')');
let lower = scan.to_ascii_lowercase();
if lower.contains("binary") {
CopyFormat::Binary
} else if lower.contains("csv") {
CopyFormat::Csv
} else {
CopyFormat::Text
}
}
fn strip_kw<'a>(s: &'a str, kw: &str) -> Option<&'a str> {
let s = s.trim_start();
if s.len() >= kw.len() && s[..kw.len()].eq_ignore_ascii_case(kw) {
let after = &s[kw.len()..];
if after.is_empty() || after.starts_with(|c: char| c.is_whitespace() || c == '(') {
return Some(after);
}
}
None
}
fn take_ident(s: &str) -> Option<(String, &str)> {
let s = s.trim_start();
if let Some(after_q) = s.strip_prefix('"') {
let end = after_q.find('"')?;
return Some((after_q[..end].to_string(), &after_q[end + 1..]));
}
let end = s
.find(|c: char| c.is_whitespace() || c == '(')
.unwrap_or(s.len());
if end == 0 {
return None;
}
Some((s[..end].to_string(), &s[end..]))
}
pub(crate) fn decode_text_field(raw: &str) -> Option<String> {
if raw == "\\N" {
return None;
}
let mut out = String::with_capacity(raw.len());
let mut chars = raw.chars();
while let Some(c) = chars.next() {
if c == '\\' {
match chars.next() {
Some('t') => out.push('\t'),
Some('n') => out.push('\n'),
Some('r') => out.push('\r'),
Some('b') => out.push('\u{8}'),
Some('f') => out.push('\u{c}'),
Some('v') => out.push('\u{b}'),
Some('\\') => out.push('\\'),
Some(other) => out.push(other), None => out.push('\\'),
}
} else {
out.push(c);
}
}
Some(out)
}
pub(crate) fn parse_text_rows(data: &[u8]) -> Vec<Vec<Option<String>>> {
let text = String::from_utf8_lossy(data);
let mut rows = Vec::new();
for line in text.split('\n') {
let line = line.strip_suffix('\r').unwrap_or(line);
if line == "\\." {
break;
}
if line.is_empty() {
continue;
}
rows.push(line.split('\t').map(decode_text_field).collect());
}
rows
}
fn quote_ident(id: &str) -> String {
format!("\"{}\"", id.replace('"', "\"\""))
}
fn sql_value(v: &Option<String>) -> String {
match v {
None => "NULL".to_string(),
Some(s) => format!("'{}'", s.replace('\'', "''")),
}
}
pub(crate) fn build_insert_sql(
table: &str,
columns: &[String],
rows: &[Vec<Option<String>>],
) -> Option<String> {
if rows.is_empty() {
return None;
}
let cols_clause = if columns.is_empty() {
String::new()
} else {
let q: Vec<String> = columns.iter().map(|c| quote_ident(c)).collect();
format!(" ({})", q.join(", "))
};
let values: Vec<String> = rows
.iter()
.map(|r| {
let vs: Vec<String> = r.iter().map(sql_value).collect();
format!("({})", vs.join(", "))
})
.collect();
Some(format!(
"INSERT INTO {}{} VALUES {}",
quote_ident(table),
cols_clause,
values.join(", ")
))
}
pub(crate) fn encode_text_field(v: Option<&[u8]>) -> String {
match v {
None => "\\N".to_string(),
Some(bytes) => {
let s = String::from_utf8_lossy(bytes);
let mut out = String::with_capacity(s.len() + 2);
for c in s.chars() {
match c {
'\\' => out.push_str("\\\\"),
'\t' => out.push_str("\\t"),
'\n' => out.push_str("\\n"),
'\r' => out.push_str("\\r"),
other => out.push(other),
}
}
out
}
}
}
pub(crate) fn encode_text_row(fields: &[Option<Vec<u8>>]) -> Vec<u8> {
let parts: Vec<String> = fields.iter().map(|f| encode_text_field(f.as_deref())).collect();
let mut line = parts.join("\t");
line.push('\n');
line.into_bytes()
}
fn finish_csv_field(field: &str, was_quoted: bool) -> Option<String> {
if !was_quoted && field.is_empty() {
None
} else {
Some(field.to_string())
}
}
pub(crate) fn parse_csv_rows(data: &[u8]) -> Vec<Vec<Option<String>>> {
let text = String::from_utf8_lossy(data);
let mut rows: Vec<Vec<Option<String>>> = Vec::new();
let mut row: Vec<Option<String>> = Vec::new();
let mut field = String::new();
let mut in_quotes = false;
let mut was_quoted = false;
let mut field_started = false;
let mut chars = text.chars().peekable();
while let Some(c) = chars.next() {
if in_quotes {
if c == '"' {
if chars.peek() == Some(&'"') {
chars.next();
field.push('"');
} else {
in_quotes = false;
}
} else {
field.push(c);
}
continue;
}
match c {
'"' if !field_started => {
in_quotes = true;
was_quoted = true;
field_started = true;
}
',' => {
row.push(finish_csv_field(&field, was_quoted));
field.clear();
was_quoted = false;
field_started = false;
}
'\r' => {} '\n' => {
row.push(finish_csv_field(&field, was_quoted));
rows.push(std::mem::take(&mut row));
field.clear();
was_quoted = false;
field_started = false;
}
other => {
field.push(other);
field_started = true;
}
}
}
if field_started || was_quoted || !row.is_empty() {
row.push(finish_csv_field(&field, was_quoted));
rows.push(row);
}
rows
}
pub(crate) fn encode_csv_field(v: Option<&[u8]>) -> String {
match v {
None => String::new(),
Some(bytes) => {
let s = String::from_utf8_lossy(bytes);
if s.contains(',') || s.contains('"') || s.contains('\n') || s.contains('\r') {
format!("\"{}\"", s.replace('"', "\"\""))
} else {
s.into_owned()
}
}
}
}
pub(crate) fn encode_csv_row(fields: &[Option<Vec<u8>>]) -> Vec<u8> {
let parts: Vec<String> = fields.iter().map(|f| encode_csv_field(f.as_deref())).collect();
let mut line = parts.join(",");
line.push('\n');
line.into_bytes()
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
#[test]
fn from_stdin_basic() {
let c = parse_copy("COPY users FROM STDIN").unwrap();
assert_eq!(c.table, "users");
assert!(!c.to_stdout);
assert!(c.columns.is_empty());
assert_eq!(c.format, CopyFormat::Text);
}
#[test]
fn from_stdin_cols_csv() {
let c = parse_copy("COPY users (id, name) FROM STDIN WITH (FORMAT csv)").unwrap();
assert_eq!(c.table, "users");
assert_eq!(c.columns, vec!["id", "name"]);
assert!(!c.to_stdout);
assert_eq!(c.format, CopyFormat::Csv);
}
#[test]
fn to_stdout_binary_and_legacy() {
let c = parse_copy("COPY t TO STDOUT (FORMAT binary)").unwrap();
assert!(c.to_stdout);
assert_eq!(c.format, CopyFormat::Binary);
let l = parse_copy("COPY t FROM STDIN WITH BINARY").unwrap();
assert_eq!(l.format, CopyFormat::Binary);
}
#[test]
fn non_copy_and_file_copy_return_none() {
assert!(parse_copy("SELECT 1").is_none());
assert!(parse_copy("COPY t FROM '/tmp/x.csv'").is_none());
assert!(parse_copy("COPYISH t FROM STDIN").is_none());
}
#[test]
fn decode_field_null_and_escapes() {
assert_eq!(decode_text_field("\\N"), None);
assert_eq!(decode_text_field("plain"), Some("plain".to_string()));
assert_eq!(decode_text_field("a\\tb\\nc"), Some("a\tb\nc".to_string()));
assert_eq!(decode_text_field("a\\\\b"), Some("a\\b".to_string()));
assert_eq!(decode_text_field(""), Some(String::new()));
}
#[test]
fn parse_rows_with_null_and_terminator() {
let data = b"1\thello\n2\t\\N\n\\.\n3\tignored";
let rows = parse_text_rows(data);
assert_eq!(rows.len(), 2); assert_eq!(rows[0], vec![Some("1".to_string()), Some("hello".to_string())]);
assert_eq!(rows[1], vec![Some("2".to_string()), None]);
}
#[test]
fn insert_sql_is_injection_safe() {
let rows = vec![vec![
Some("1".to_string()),
Some("x'); DROP TABLE users;--".to_string()),
]];
let sql = build_insert_sql("users", &["id".to_string(), "name".to_string()], &rows).unwrap();
assert!(sql.contains("'x''); DROP TABLE users;--'"));
assert!(sql.starts_with("INSERT INTO \"users\" (\"id\", \"name\") VALUES ("));
let nrows = vec![vec![None, Some("a".to_string())]];
let s2 = build_insert_sql("t", &[], &nrows).unwrap();
assert!(s2.contains("VALUES (NULL, 'a')"));
assert_eq!(super::quote_ident("we\"ird"), "\"we\"\"ird\"");
}
#[test]
fn encode_field_null_and_escapes() {
assert_eq!(encode_text_field(None), "\\N");
assert_eq!(encode_text_field(Some(b"plain")), "plain");
assert_eq!(encode_text_field(Some(b"a\tb\nc")), "a\\tb\\nc");
assert_eq!(encode_text_field(Some(b"a\\b")), "a\\\\b");
assert_eq!(encode_text_field(Some(b"")), ""); }
#[test]
fn encode_decode_roundtrip() {
let fields: Vec<Option<Vec<u8>>> =
vec![Some(b"1".to_vec()), None, Some(b"has\ttab\nand nl".to_vec())];
let line = encode_text_row(&fields);
assert_eq!(line.last(), Some(&b'\n'));
let rows = parse_text_rows(&line);
assert_eq!(rows.len(), 1);
assert_eq!(
rows[0],
vec![Some("1".to_string()), None, Some("has\ttab\nand nl".to_string())]
);
}
#[test]
fn csv_parse_quoting_and_null() {
let data = b"1,,\"\"\n2,\"a,b\",\"she said \"\"hi\"\"\"\n3,\"line1\nline2\",x\n";
let rows = parse_csv_rows(data);
assert_eq!(rows.len(), 3);
assert_eq!(rows[0], vec![Some("1".into()), None, Some("".into())]);
assert_eq!(
rows[1],
vec![Some("2".into()), Some("a,b".into()), Some("she said \"hi\"".into())]
);
assert_eq!(rows[2], vec![Some("3".into()), Some("line1\nline2".into()), Some("x".into())]);
}
#[test]
fn csv_encode_quotes_when_needed() {
assert_eq!(encode_csv_field(None), ""); assert_eq!(encode_csv_field(Some(b"plain")), "plain");
assert_eq!(encode_csv_field(Some(b"a,b")), "\"a,b\"");
assert_eq!(encode_csv_field(Some(b"she \"q\"")), "\"she \"\"q\"\"\"");
assert_eq!(encode_csv_field(Some(b"l1\nl2")), "\"l1\nl2\"");
}
#[test]
fn csv_encode_decode_roundtrip() {
let fields: Vec<Option<Vec<u8>>> =
vec![Some(b"1".to_vec()), None, Some(b"a,b\"c\nd".to_vec())];
let line = encode_csv_row(&fields);
let rows = parse_csv_rows(&line);
assert_eq!(rows.len(), 1);
assert_eq!(
rows[0],
vec![Some("1".to_string()), None, Some("a,b\"c\nd".to_string())]
);
}
}