use sqlx::query::{Query, QueryScalar};
use sqlx::{Database, Sqlite};
pub use crate::string::highlighted::TextHighlighter;
pub trait FtsQueryExt {
#[must_use]
fn bind_highlight(self, highlighter: TextHighlighter) -> Self;
#[must_use]
fn bind_highlightable(self, highlighter: TextHighlighter, highlightable: &str) -> Self;
#[must_use]
fn bind_match_query(self, input: &str) -> Option<Self>
where
Self: Sized;
}
impl FtsQueryExt for Query<'_, Sqlite, <Sqlite as Database>::Arguments> {
fn bind_highlight(self, highlighter: TextHighlighter) -> Self {
let [open, close] = highlighter.markers();
self.bind(open.to_string()).bind(close.to_string())
}
fn bind_highlightable(self, highlighter: TextHighlighter, highlightable: &str) -> Self {
self.bind(highlighter.sanitize(highlightable).into_owned())
}
fn bind_match_query(self, input: &str) -> Option<Self> {
match_expression(input).map(|expr| self.bind(expr))
}
}
impl<O> FtsQueryExt for QueryScalar<'_, Sqlite, O, <Sqlite as Database>::Arguments> {
fn bind_highlight(self, highlighter: TextHighlighter) -> Self {
let [open, close] = highlighter.markers();
self.bind(open.to_string()).bind(close.to_string())
}
fn bind_highlightable(self, highlighter: TextHighlighter, highlightable: &str) -> Self {
self.bind(highlighter.sanitize(highlightable).into_owned())
}
fn bind_match_query(self, input: &str) -> Option<Self> {
match_expression(input).map(|expr| self.bind(expr))
}
}
#[must_use]
pub fn match_expression(input: &str) -> Option<String> {
input
.split_whitespace()
.map(|term| format!("\"{}\"", term.replace('"', "\"\"")))
.reduce(|expr, term| format!("{expr} {term}"))
}
#[cfg(test)]
mod tests {
use std::borrow::Cow;
use std::ops::Range;
use pretty_assertions::assert_eq;
use rstest::rstest;
use sqlx::Row;
use super::*;
async fn highlight_roundtrip(h: TextHighlighter, body: &str, term: &str) -> String {
let sqlite = crate::db::sqlite::Sqlite::builder_in_memory().open().await.unwrap();
let mut conn = sqlite.pool().acquire().await.unwrap();
crate::db::query::<Sqlite>("create virtual table docs using fts5(body)")
.execute(&mut *conn)
.await
.unwrap();
crate::db::query::<Sqlite>("insert into docs (body) values (?)")
.bind_highlightable(h, body)
.execute(&mut *conn)
.await
.unwrap();
let row = crate::db::query::<Sqlite>(
"select highlight(docs, 0, ?, ?) as body from docs where docs match ?",
)
.bind_highlight(h)
.bind(term)
.fetch_one(&mut *conn)
.await
.unwrap();
row.try_get("body").unwrap()
}
#[rstest]
#[case::start_of_body("error here", "error", vec![3..8], vec!["error"])]
#[case::middle_of_body("the build failed now", "failed", vec![13..19], vec!["failed"])]
#[case::multibyte_prefix("café error", "error", vec![9..14], vec!["error"])]
#[case::multibyte_match("the café", "café", vec![7..12], vec!["café"])]
#[case::two_matches_case_preserved("Error and error", "error", vec![3..8, 19..24], vec!["Error", "error"])]
#[case::stray_source_marker("a\u{E000}b error", "error", vec![6..11], vec!["error"])]
#[tokio::test]
async fn highlight_round_trips_through_fts5(
#[case] body: &str,
#[case] term: &str,
#[case] expected: Vec<Range<usize>>,
#[case] hits: Vec<&str>,
) {
let h = TextHighlighter::default();
let marked = highlight_roundtrip(h, body, term).await;
assert_eq!(h.sanitize(&marked), h.sanitize(body));
let hl = h.as_highlighted(marked.as_str());
let ranges: Vec<Range<usize>> = hl.ranges().collect();
assert_eq!(ranges.as_slice(), expected.as_slice());
let raw: &str = hl.as_ref();
let got: Vec<&str> = ranges.iter().map(|r| &raw[r.clone()]).collect();
assert_eq!(got.as_slice(), hits.as_slice());
}
#[tokio::test]
async fn highlight_of_an_unmatched_column_yields_no_ranges() {
let h = TextHighlighter::default();
let sqlite = crate::db::sqlite::Sqlite::builder_in_memory().open().await.unwrap();
let mut conn = sqlite.pool().acquire().await.unwrap();
crate::db::query::<Sqlite>("create virtual table docs using fts5(title, body)")
.execute(&mut *conn)
.await
.unwrap();
crate::db::query::<Sqlite>("insert into docs (title, body) values (?, ?)")
.bind_highlightable(h, "alpha")
.bind_highlightable(h, "beta gamma")
.execute(&mut *conn)
.await
.unwrap();
let row = crate::db::query::<Sqlite>(
"select highlight(docs, 1, ?, ?) as body from docs where docs match ?",
)
.bind_highlight(h)
.bind("alpha")
.fetch_one(&mut *conn)
.await
.unwrap();
let marked: String = row.try_get("body").unwrap();
assert_eq!(h.as_highlighted(marked.as_str()).ranges().count(), 0);
assert!(matches!(h.sanitize(&marked), Cow::Borrowed(_)));
}
#[rstest]
#[case::single_term("error", Some("\"error\""))]
#[case::terms_are_anded("build failed", Some("\"build\" \"failed\""))]
#[case::embedded_quote_is_doubled("a\"b", Some("\"a\"\"b\""))]
#[case::surrounding_whitespace_trimmed(" spaced ", Some("\"spaced\""))]
#[case::operator_char_kept_literal("refused:", Some("\"refused:\""))]
#[case::empty_is_none("", None)]
#[case::whitespace_only_is_none(" \t\n ", None)]
fn match_expression_escapes_free_form_input(
#[case] input: &str,
#[case] expected: Option<&str>,
) {
assert_eq!(match_expression(input).as_deref(), expected);
}
#[tokio::test]
async fn bind_match_query_neutralizes_fts5_operators() {
let sqlite = crate::db::sqlite::Sqlite::builder_in_memory().open().await.unwrap();
let mut conn = sqlite.pool().acquire().await.unwrap();
crate::db::query::<Sqlite>("create virtual table docs using fts5(body)")
.execute(&mut *conn)
.await
.unwrap();
crate::db::query::<Sqlite>("insert into docs(body) values ('connection refused: oops')")
.execute(&mut *conn)
.await
.unwrap();
let row = crate::db::query::<Sqlite>("select count(*) as n from docs where docs match ?")
.bind_match_query("refused:")
.expect("non-blank query")
.fetch_one(&mut *conn)
.await
.unwrap();
assert_eq!(row.try_get::<i64, _>("n").unwrap(), 1);
}
}