use std::fmt::{Display, Write};
use crate::{SqlFormat, ToSql};
pub struct EscapeWriter<W> {
escape_char: char,
preserve_newlines: bool,
writer: W,
}
impl<'a> EscapeWriter<&'a mut String> {
fn escape<D: Display + ?Sized>(into: &'a mut String, escape: char, display: &D) {
Self {
escape_char: escape,
preserve_newlines: false,
writer: into,
}
.write(display)
}
fn escape_pretty<D: Display + ?Sized>(into: &'a mut String, escape: char, display: &D) {
Self {
escape_char: escape,
preserve_newlines: true,
writer: into,
}
.write(display)
}
fn write<D: Display + ?Sized>(&mut self, display: &D) {
let _ = self.write_fmt(format_args!("{display}"));
}
}
impl<W: Write> Write for EscapeWriter<W> {
fn write_str(&mut self, s: &str) -> std::fmt::Result {
for c in s.chars() {
self.write_char(c)?;
}
Ok(())
}
fn write_char(&mut self, c: char) -> std::fmt::Result {
match c {
'\0' => {
self.writer.write_str("\\0")?;
}
'\r' => {
self.writer.write_str("\\r")?;
}
'\t' => {
self.writer.write_str("\\t")?;
}
'\n' if self.preserve_newlines => {
self.writer.write_char('\n')?;
}
'\n' => {
self.writer.write_str("\\n")?;
}
'\x08' => {
self.writer.write_str("\\u{8}")?;
}
'\x0C' => {
self.writer.write_str("\\f")?;
}
'\\' => {
self.writer.write_str("\\\\")?;
}
x if x == self.escape_char => {
self.writer.write_char('\\')?;
self.writer.write_char(x)?;
}
_ => self.writer.write_char(c)?,
}
Ok(())
}
}
pub struct QuoteStr<'a>(pub &'a str);
impl ToSql for QuoteStr<'_> {
fn fmt_sql(&self, f: &mut String, fmt: SqlFormat) {
let s = self.0;
let quote = if s.contains('\'') {
'\"'
} else {
'\''
};
f.push(quote);
if fmt.is_pretty() {
EscapeWriter::escape_pretty(f, quote, self.0);
} else {
EscapeWriter::escape(f, quote, self.0);
}
f.push(quote);
}
}
pub struct EscapeSqonIdent<'a>(pub &'a str);
impl ToSql for EscapeSqonIdent<'_> {
fn fmt_sql(&self, f: &mut String, _fmt: SqlFormat) {
let s = self.0;
if s.is_empty()
|| s.starts_with(|x: char| x.is_ascii_digit())
|| s.contains(|x: char| !x.is_ascii_alphanumeric() && x != '_')
{
f.push('`');
EscapeWriter::escape(f, '`', self.0);
f.push('`');
} else {
f.push_str(s)
}
}
}
pub struct EscapeObjectKey<'a>(pub &'a str);
impl ToSql for EscapeObjectKey<'_> {
fn fmt_sql(&self, f: &mut String, _fmt: SqlFormat) {
let s = self.0;
if s.is_empty()
|| s.starts_with(|x: char| x.is_ascii_digit())
|| s.contains(|x: char| !x.is_ascii_alphanumeric() && x != '_')
{
f.push('\"');
EscapeWriter::escape(f, '"', self.0);
f.push('\"');
} else {
f.push_str(s)
}
}
}
pub struct EscapeRecordKey<'a>(pub &'a str);
impl ToSql for EscapeRecordKey<'_> {
fn fmt_sql(&self, f: &mut String, _fmt: SqlFormat) {
let s = self.0;
if s.is_empty()
|| s.contains(|x: char| !x.is_ascii_alphanumeric() && x != '_')
|| !s.contains(|x: char| !x.is_ascii_digit() && x != '_')
{
f.push('`');
EscapeWriter::escape(f, '`', self.0);
f.push('`');
} else {
f.push_str(s)
}
}
}
pub fn decode_backtick_ident(s: &str) -> anyhow::Result<String> {
let Some(interior) = s.strip_prefix('`').and_then(|rest| rest.strip_suffix('`')) else {
return Ok(s.to_string());
};
parse_common::unescape_cow(interior)
.map(|unescaped| unescaped.into_owned())
.map_err(|e| anyhow::anyhow!("{} in `{interior}`", e.message))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn decode_backtick_ident_plain() {
assert_eq!(decode_backtick_ident("tobie").unwrap(), "tobie");
assert_eq!(decode_backtick_ident("needs escaping").unwrap(), "needs escaping");
}
#[test]
fn decode_backtick_ident_roundtrips_escape_writer() {
let mut escaped = String::new();
EscapeRecordKey("needs escaping").fmt_sql(&mut escaped, SqlFormat::SingleLine);
assert_eq!(escaped, "`needs escaping`");
assert_eq!(decode_backtick_ident(&escaped).unwrap(), "needs escaping");
let mut escaped = String::new();
EscapeRecordKey("he said `hi`").fmt_sql(&mut escaped, SqlFormat::SingleLine);
assert_eq!(decode_backtick_ident(&escaped).unwrap(), "he said `hi`");
let mut escaped = String::new();
EscapeRecordKey("tab\there\nand \0 null").fmt_sql(&mut escaped, SqlFormat::SingleLine);
assert_eq!(decode_backtick_ident(&escaped).unwrap(), "tab\there\nand \0 null");
}
#[test]
fn decode_backtick_ident_accepts_every_lexer_escape() {
assert_eq!(decode_backtick_ident(r"`a\bc`").unwrap(), "a\x08c");
assert_eq!(decode_backtick_ident(r"`a\'c`").unwrap(), "a'c");
assert_eq!(decode_backtick_ident(r#"`a\"c`"#).unwrap(), "a\"c");
assert_eq!(decode_backtick_ident(r"`a\⟩c`").unwrap(), "a⟩c");
assert_eq!(decode_backtick_ident("`\\u0021`").unwrap(), "!");
assert_eq!(decode_backtick_ident(r"`\u{1F600}`").unwrap(), "😀");
}
#[test]
fn decode_backtick_ident_rejects_invalid_escape() {
assert!(decode_backtick_ident(r"`a\xc`").is_err());
assert!(decode_backtick_ident(r"`a\`").is_err());
}
}