use super::{Dialect, Executor, Migration, Model, quote, sql};
use crate::Result;
use anyhow::anyhow;
const MAX_TERMS: usize = 16;
const MAX_TERM_CHARS: usize = 64;
pub(crate) fn terms(text: &str) -> Option<String> {
let words: Vec<String> = text
.split(|c: char| !c.is_alphanumeric())
.filter(|w| !w.is_empty())
.take(MAX_TERMS)
.map(|w| w.chars().take(MAX_TERM_CHARS).collect())
.collect();
(!words.is_empty()).then(|| words.join(" "))
}
pub(crate) fn unsearchable<M: Model>() -> Option<String> {
if M::SEARCHABLE.is_empty() {
return Some(format!(
"`{}` has no full-text search: add #[model(search = \"…\")]",
M::TABLE
));
}
if !valid_language(M::SEARCH_LANGUAGE) {
return Some(format!(
"search language `{}` isn't a plain name (letters and `_`)",
M::SEARCH_LANGUAGE
));
}
None
}
fn valid_language(language: &str) -> bool {
!language.is_empty() && language.chars().all(|c| c.is_ascii_lowercase() || c == '_')
}
fn fts_table<M: Model>() -> String {
format!("{}_search", M::TABLE)
}
pub(crate) fn bound(text: &str) -> Option<String> {
let words = terms(text)?;
let typed: String = text.chars().filter(|c| !c.is_control()).take(200).collect();
Some(format!("{words}{SEPARATOR}{}", typed.trim()))
}
const SEPARATOR: char = '\u{1f}';
fn match_expression(dialect: Dialect) -> &'static str {
match dialect {
Dialect::Sqlite => {
"(SELECT '\"' || replace(substr(v, 1, instr(v, char(31)) - 1), ' ', '\"* \"') || '\"*' \
FROM (SELECT ? AS v))"
}
Dialect::Postgres => {
"(SELECT to_tsquery(c, replace(split_part(v, chr(31), 1), ' ', ':* & ') || ':*') \
|| plainto_tsquery(c, split_part(v, chr(31), 2)) FROM (SELECT ?::text AS v, %CONFIG% AS c) AS s)"
}
}
}
fn config_literal<M: Model>() -> String {
format!("'{}'::regconfig", M::SEARCH_LANGUAGE)
}
pub(crate) fn filter_sql<M: Model>(dialect: Dialect) -> String {
match dialect {
Dialect::Sqlite => format!(
"{}.rowid IN (SELECT rowid FROM {fts} WHERE {fts} MATCH {})",
quote(M::TABLE),
match_expression(dialect),
fts = quote(&fts_table::<M>()),
),
Dialect::Postgres => format!(
"{}.\"search_vector\" @@ {}",
quote(M::TABLE),
match_expression(dialect).replace("%CONFIG%", &config_literal::<M>())
),
}
}
fn weight(index: usize) -> (char, &'static str) {
match index {
0 => ('A', "1.0"),
1 => ('B', "0.4"),
2 => ('C', "0.2"),
_ => ('D', "0.1"),
}
}
pub(crate) fn rank_sql<M: Model>(dialect: Dialect) -> String {
match dialect {
Dialect::Sqlite => {
let weights: Vec<&str> = (0..M::SEARCHABLE.len()).map(|i| weight(i).1).collect();
format!(
"COALESCE((SELECT bm25({fts}, {}) FROM {fts} WHERE {fts} MATCH {} AND {fts}.rowid = {}.rowid), 0) ASC",
weights.join(", "),
match_expression(dialect),
quote(M::TABLE),
fts = quote(&fts_table::<M>()),
)
}
Dialect::Postgres => format!(
"ts_rank({}.\"search_vector\", {}) DESC",
quote(M::TABLE),
match_expression(dialect).replace("%CONFIG%", &config_literal::<M>())
),
}
}
fn literal(text: &str) -> String {
format!("'{}'", text.replace('\'', "''"))
}
fn sqlite_up<M: Model>() -> String {
let table = quote(M::TABLE);
let fts = quote(&fts_table::<M>());
let columns: Vec<String> = M::SEARCHABLE.iter().map(|c| quote(c)).collect();
let columns = columns.join(", ");
let new: Vec<String> = M::SEARCHABLE
.iter()
.map(|c| format!("new.{}", quote(c)))
.collect();
let old: Vec<String> = M::SEARCHABLE
.iter()
.map(|c| format!("old.{}", quote(c)))
.collect();
let (new, old) = (new.join(", "), old.join(", "));
let tokenize = if M::SEARCH_LANGUAGE == "english" {
"porter unicode61 remove_diacritics 2"
} else {
"unicode61 remove_diacritics 2"
};
let trigger = |suffix: &str| quote(&format!("{}_search_{suffix}", M::TABLE));
format!(
"{drop}\n\
CREATE VIRTUAL TABLE {fts} USING fts5({columns}, content={content}, tokenize={tokenize});\n\
CREATE TRIGGER {insert} AFTER INSERT ON {table} BEGIN\n\
\x20 INSERT INTO {fts}(rowid, {columns}) VALUES (new.rowid, {new});\n\
END;\n\
CREATE TRIGGER {delete} AFTER DELETE ON {table} BEGIN\n\
\x20 INSERT INTO {fts}({fts}, rowid, {columns}) VALUES ('delete', old.rowid, {old});\n\
END;\n\
CREATE TRIGGER {update} AFTER UPDATE OF \"id\", {columns} ON {table} BEGIN\n\
\x20 INSERT INTO {fts}({fts}, rowid, {columns}) VALUES ('delete', old.rowid, {old});\n\
\x20 INSERT INTO {fts}(rowid, {columns}) VALUES (new.rowid, {new});\n\
END;\n\
INSERT INTO {fts}({fts}) VALUES ('rebuild');",
drop = sqlite_down::<M>(),
content = literal(M::TABLE),
tokenize = literal(tokenize),
insert = trigger("insert"),
delete = trigger("delete"),
update = trigger("update"),
)
}
fn sqlite_down<M: Model>() -> String {
let trigger = |suffix: &str| quote(&format!("{}_search_{suffix}", M::TABLE));
format!(
"DROP TRIGGER IF EXISTS {};\nDROP TRIGGER IF EXISTS {};\nDROP TRIGGER IF EXISTS {};\nDROP TABLE IF EXISTS {};",
trigger("insert"),
trigger("update"),
trigger("delete"),
quote(&fts_table::<M>())
)
}
fn postgres_up<M: Model>() -> String {
let vector: Vec<String> = M::SEARCHABLE
.iter()
.enumerate()
.map(|(i, c)| {
format!(
"setweight(to_tsvector({}, coalesce({}::text, '')), '{}')",
config_literal::<M>(),
quote(c),
weight(i).0
)
})
.collect();
format!(
"{drop}\n\
ALTER TABLE {table} ADD COLUMN \"search_vector\" tsvector GENERATED ALWAYS AS ({vector}) STORED;\n\
CREATE INDEX {index} ON {table} USING GIN (\"search_vector\");",
drop = postgres_down::<M>(),
table = quote(M::TABLE),
vector = vector.join(" || "),
index = quote(&format!("{}_search_index", M::TABLE)),
)
}
fn postgres_down<M: Model>() -> String {
format!(
"DROP INDEX IF EXISTS {};\nALTER TABLE {} DROP COLUMN IF EXISTS \"search_vector\";",
quote(&format!("{}_search_index", M::TABLE)),
quote(M::TABLE)
)
}
pub fn migration<M: Model>(name: &'static str) -> Migration {
if let Some(problem) = unsearchable::<M>() {
panic!("renox::db::search::migration: {problem}");
}
let leak = |sql: String| -> &'static str { Box::leak(sql.into_boxed_str()) };
Migration::new(name, "", None)
.sqlite(leak(sqlite_up::<M>()), Some(leak(sqlite_down::<M>())))
.postgres(leak(postgres_up::<M>()), Some(leak(postgres_down::<M>())))
}
pub async fn rebuild<'c, M: Model>(db: impl Executor<'c>) -> Result {
if let Some(problem) = unsearchable::<M>() {
return Err(anyhow!("{problem}").into());
}
let db = db.into_conn();
if db.dialect() == Dialect::Sqlite {
let fts = quote(&fts_table::<M>());
sql(format!("INSERT INTO {fts}({fts}) VALUES ('rebuild')"))
.execute(db)
.await?;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::terms;
#[test]
fn terms_keep_only_words() {
assert_eq!(terms(" Coffee, roast! "), Some("Coffee roast".into()));
assert_eq!(
terms("\"a\" OR b* NEAR(c) -d"),
Some("a OR b NEAR c d".into())
);
assert_eq!(
terms("it's x'; DROP TABLE posts; --"),
Some("it s x DROP TABLE posts".into())
);
assert_eq!(terms("café 東京"), Some("café 東京".into()));
assert_eq!(terms("!!! ... ---"), None);
assert_eq!(terms(""), None);
let many = "w ".repeat(40);
assert_eq!(terms(&many).unwrap().split(' ').count(), 16);
assert_eq!(terms(&"x".repeat(100)).unwrap().len(), 64);
}
}