use crate::prelude::*;
use crate::sql::SQL;
use crate::traits::SQLParam;
use percent_encoding::{AsciiSet, CONTROLS, utf8_percent_encode};
const COMPONENT: &AsciiSet = &CONTROLS
.add(b' ')
.add(b'"')
.add(b'#')
.add(b'$')
.add(b'%')
.add(b'&')
.add(b'+')
.add(b',')
.add(b'/')
.add(b':')
.add(b';')
.add(b'<')
.add(b'=')
.add(b'>')
.add(b'?')
.add(b'@')
.add(b'[')
.add(b'\\')
.add(b']')
.add(b'^')
.add(b'`')
.add(b'{')
.add(b'|')
.add(b'}');
pub fn comment<'a, V: SQLParam>(text: impl AsRef<str>) -> SQL<'a, V> {
let text = text.as_ref();
if text.is_empty() {
return SQL::empty();
}
let sanitized = sanitize_string_input(text);
let mut out = String::with_capacity(sanitized.len() + 4);
out.push_str("/*");
out.push_str(&sanitized);
out.push_str("*/");
SQL::raw(out)
}
pub fn comment_tags<'a, V, I, K, Val>(pairs: I) -> SQL<'a, V>
where
V: SQLParam,
I: IntoIterator<Item = (K, Val)>,
K: AsRef<str>,
Val: AsRef<str>,
{
let mut parts: Vec<String> = Vec::new();
for (k, v) in pairs {
let v = v.as_ref();
if v.is_empty() {
continue;
}
let ek = sanitize_object_element(k.as_ref());
let ev = sanitize_object_element(v);
let mut entry = String::with_capacity(ek.len() + ev.len() + 3);
entry.push_str(&ek);
entry.push_str("='");
entry.push_str(&ev);
entry.push('\'');
parts.push(entry);
}
if parts.is_empty() {
return SQL::empty();
}
parts.sort();
let total: usize = parts.iter().map(String::len).sum::<usize>() + parts.len() + 3;
let mut out = String::with_capacity(total);
out.push_str("/*");
for (i, p) in parts.iter().enumerate() {
if i > 0 {
out.push(',');
}
out.push_str(p);
}
out.push_str("*/");
SQL::raw(out)
}
#[inline]
fn sanitize_string_input(input: &str) -> String {
input.replace("/*", "/ *").replace("*/", "* /")
}
#[inline]
fn sanitize_object_element(s: &str) -> String {
let encoded = utf8_percent_encode(s, COMPONENT).to_string();
if encoded.as_bytes().contains(&b'\'') {
encoded.replace('\'', "\\'")
} else {
encoded
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dialect::{Dialect, SQLiteDialect};
#[derive(Clone, Debug)]
struct TestValue;
impl SQLParam for TestValue {
const DIALECT: Dialect = Dialect::SQLite;
type DialectMarker = SQLiteDialect;
}
fn render(s: SQL<'_, TestValue>) -> String {
s.sql()
}
#[test]
fn empty_string_yields_empty_sql() {
let s: SQL<'_, TestValue> = comment("");
assert_eq!(render(s), "");
}
#[test]
fn plain_string_is_wrapped() {
let s: SQL<'_, TestValue> = comment("hello world");
assert_eq!(render(s), "/*hello world*/");
}
#[test]
fn string_input_sanitises_comment_terminators() {
let s: SQL<'_, TestValue> = comment("/* nested */ end");
assert_eq!(render(s), "/*/ * nested * / end*/");
}
#[test]
fn tags_are_sorted_and_url_encoded() {
let s: SQL<'_, TestValue> = comment_tags([("route", "/users/:id"), ("action", "update")]);
assert_eq!(render(s), "/*action='update',route='%2Fusers%2F%3Aid'*/");
}
#[test]
fn empty_values_are_skipped() {
let s: SQL<'_, TestValue> = comment_tags([("a", ""), ("b", "ok")]);
assert_eq!(render(s), "/*b='ok'*/");
}
#[test]
fn all_empty_yields_empty_sql() {
let s: SQL<'_, TestValue> = comment_tags([("a", ""), ("b", "")]);
assert_eq!(render(s), "");
}
#[test]
fn quote_in_value_is_escaped_after_url_encoding() {
let s: SQL<'_, TestValue> = comment_tags([("k", "it's")]);
assert_eq!(render(s), r"/*k='it\'s'*/");
}
#[test]
fn multibyte_utf8_is_percent_encoded_per_byte() {
let s: SQL<'_, TestValue> = comment_tags([("name", "café")]);
assert_eq!(render(s), "/*name='caf%C3%A9'*/");
}
#[test]
fn unreserved_set_is_preserved() {
let s: SQL<'_, TestValue> = comment_tags([("k", "abcXYZ012-_.!~*()")]);
assert_eq!(render(s), "/*k='abcXYZ012-_.!~*()'*/");
}
}