use regex::{Regex, RegexBuilder};
use super::connection::{DatabaseObject, ObjectKind, StoredKind, StoredObject};
pub const MAX_ENTRIES: usize = 200_000;
pub const MAX_HITS: usize = 500;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum CatalogKind {
Table,
View,
Column,
Index,
Routine,
Trigger,
}
impl CatalogKind {
pub const ALL: [CatalogKind; 6] = [
CatalogKind::Table,
CatalogKind::View,
CatalogKind::Column,
CatalogKind::Index,
CatalogKind::Routine,
CatalogKind::Trigger,
];
pub fn label(self) -> &'static str {
match self {
CatalogKind::Table => "Table",
CatalogKind::View => "View",
CatalogKind::Column => "Column",
CatalogKind::Index => "Index",
CatalogKind::Routine => "Routine",
CatalogKind::Trigger => "Trigger",
}
}
pub fn keyword(self) -> &'static str {
match self {
CatalogKind::Table => "table",
CatalogKind::View => "view",
CatalogKind::Column => "column",
CatalogKind::Index => "index",
CatalogKind::Routine => "routine",
CatalogKind::Trigger => "trigger",
}
}
fn named(word: &str) -> Option<Self> {
match word.to_ascii_lowercase().as_str() {
"table" | "tables" => Some(CatalogKind::Table),
"view" | "views" => Some(CatalogKind::View),
"column" | "columns" | "col" | "cols" => Some(CatalogKind::Column),
"index" | "indexes" | "indices" => Some(CatalogKind::Index),
"routine" | "routines" | "function" | "functions" | "procedure" | "procedures" => {
Some(CatalogKind::Routine)
}
"trigger" | "triggers" => Some(CatalogKind::Trigger),
_ => None,
}
}
pub fn group(self) -> &'static str {
match self {
CatalogKind::Table | CatalogKind::View => "Tables & Views",
CatalogKind::Column => "Columns",
CatalogKind::Index => "Indexes",
CatalogKind::Routine => "Routines",
CatalogKind::Trigger => "Triggers",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CatalogEntry {
pub kind: CatalogKind,
pub object: Option<DatabaseObject>,
pub routine: Option<StoredObject>,
pub name: String,
pub detail: String,
}
impl CatalogEntry {
pub(crate) fn object(object: DatabaseObject) -> Self {
let kind = match object.kind {
ObjectKind::Table => CatalogKind::Table,
ObjectKind::View => CatalogKind::View,
};
Self {
kind,
name: object.name.clone(),
object: Some(object),
routine: None,
detail: String::new(),
}
}
pub(crate) fn routine(routine: StoredObject) -> Self {
let detail = routine.arguments.clone().unwrap_or_default();
Self {
kind: CatalogKind::Routine,
name: routine.name.clone(),
object: None,
routine: Some(routine),
detail,
}
}
pub(crate) fn member(
kind: CatalogKind,
object: DatabaseObject,
name: String,
detail: String,
) -> Self {
Self {
kind,
object: Some(object),
routine: None,
name,
detail,
}
}
pub fn owner_label(&self) -> Option<String> {
match (self.kind, &self.object, &self.routine) {
(CatalogKind::Table | CatalogKind::View, _, _) | (_, None, None) => None,
(_, Some(object), _) => Some(object.label()),
(_, None, Some(routine)) => routine.schema.clone(),
}
}
pub fn label(&self) -> String {
match &self.routine {
Some(routine) => routine.label(),
None => self.name.clone(),
}
}
pub fn kind_label(&self) -> &'static str {
match &self.routine {
Some(routine) => match routine.kind {
StoredKind::Function => "Function",
StoredKind::Procedure => "Procedure",
StoredKind::Sequence => "Sequence",
},
None => self.kind.label(),
}
}
pub fn owner(&self) -> Option<&DatabaseObject> {
self.object.as_ref()
}
fn owner_or_self(&self) -> String {
self.owner_label().unwrap_or_else(|| self.label())
}
}
#[derive(Debug, Clone, Default)]
pub struct Catalog {
pub entries: Vec<CatalogEntry>,
pub total: usize,
}
impl Catalog {
pub fn truncated(&self) -> Option<usize> {
(self.total > self.entries.len()).then_some(self.total)
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct Query {
pub text: String,
pub kinds: Vec<CatalogKind>,
pub table: Option<String>,
pub type_name: Option<String>,
}
impl Query {
pub fn parse(raw: &str) -> Self {
let mut query = Query::default();
let mut words: Vec<&str> = Vec::new();
for token in raw.split_whitespace() {
let mut consumed = false;
if let Some((prefix, value)) = token.split_once(':')
&& !value.is_empty()
{
match prefix.to_ascii_lowercase().as_str() {
"kind" => {
if let Some(kind) = CatalogKind::named(value) {
query.kinds.push(kind);
consumed = true;
}
}
"table" => {
query.table = Some(value.to_string());
consumed = true;
}
"type" => {
query.type_name = Some(value.to_string());
consumed = true;
}
_ => {}
}
}
if !consumed {
words.push(token);
}
}
query.text = words.join(" ");
query
}
#[cfg(test)]
pub(crate) fn is_empty(&self) -> bool {
self.text.is_empty()
&& self.kinds.is_empty()
&& self.table.is_none()
&& self.type_name.is_none()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Hit {
pub index: usize,
pub score: u32,
}
pub fn search(entries: &[CatalogEntry], query: &Query) -> (Vec<Hit>, usize) {
let needle = Needle::new(&query.text);
let table = query.table.as_deref().map(Needle::new);
let type_name = query.type_name.as_deref().map(Needle::new);
let mut hits = Vec::new();
for (index, entry) in entries.iter().enumerate() {
if !query.kinds.is_empty() && !query.kinds.contains(&entry.kind) {
continue;
}
if let Some(table) = &table
&& !table.is_any()
&& !table.matches(&entry.owner_or_self())
{
continue;
}
if let Some(type_name) = &type_name
&& !type_name.is_any()
&& !type_name.matches(&entry.detail)
{
continue;
}
let Some(score) = score(entry, &needle) else {
continue;
};
hits.push(Hit { index, score });
}
hits.sort_by(|a, b| {
let left = &entries[a.index];
let right = &entries[b.index];
b.score
.cmp(&a.score)
.then_with(|| rank(left.kind).cmp(&rank(right.kind)))
.then_with(|| left.owner_sort().cmp(&right.owner_sort()))
.then_with(|| left.name.cmp(&right.name))
.then_with(|| a.index.cmp(&b.index))
});
let matched = hits.len();
hits.truncate(MAX_HITS);
(hits, matched)
}
fn score(entry: &CatalogEntry, needle: &Needle) -> Option<u32> {
if needle.is_any() {
return Some(1);
}
let name = needle.name_score(&entry.name);
if name > 0 {
return Some(name * 2);
}
let owner = entry.owner_label().unwrap_or_default();
(needle.matches(&owner) || needle.matches(&entry.detail)).then_some(1)
}
fn rank(kind: CatalogKind) -> u8 {
match kind {
CatalogKind::Table => 0,
CatalogKind::View => 1,
CatalogKind::Routine => 2,
CatalogKind::Column => 3,
CatalogKind::Index => 4,
CatalogKind::Trigger => 5,
}
}
impl CatalogEntry {
fn owner_sort(&self) -> String {
self.owner_label().unwrap_or_default().to_lowercase()
}
}
enum Needle {
Any,
Pattern(Regex),
}
impl Needle {
fn new(text: &str) -> Self {
if text.is_empty() {
return Needle::Any;
}
let case_insensitive = |pattern: &str| {
RegexBuilder::new(pattern)
.case_insensitive(true)
.build()
.ok()
};
match case_insensitive(text).or_else(|| case_insensitive(®ex::escape(text))) {
Some(pattern) => Needle::Pattern(pattern),
None => Needle::Any,
}
}
fn is_any(&self) -> bool {
matches!(self, Needle::Any)
}
fn matches(&self, haystack: &str) -> bool {
match self {
Needle::Any => true,
Needle::Pattern(pattern) => pattern.is_match(haystack),
}
}
fn name_score(&self, name: &str) -> u32 {
let Needle::Pattern(pattern) = self else {
return 0;
};
let Some(found) = pattern.find(name) else {
return 0;
};
if found.start() == 0 && found.end() == name.len() {
4
} else if found.start() == 0 {
3
} else if starts_word(name, found.start()) {
2
} else {
1
}
}
}
fn starts_word(text: &str, index: usize) -> bool {
text[..index]
.chars()
.next_back()
.is_none_or(|character| !character.is_alphanumeric() && character != '_')
}
#[cfg(test)]
mod tests {
use super::*;
fn table(name: &str) -> CatalogEntry {
CatalogEntry::object(DatabaseObject {
schema: None,
name: name.into(),
kind: ObjectKind::Table,
})
}
fn column(owner: &str, name: &str, type_name: &str) -> CatalogEntry {
CatalogEntry::member(
CatalogKind::Column,
DatabaseObject {
schema: None,
name: owner.into(),
kind: ObjectKind::Table,
},
name.into(),
type_name.into(),
)
}
fn index(owner: &str, name: &str, columns: &str) -> CatalogEntry {
CatalogEntry::member(
CatalogKind::Index,
DatabaseObject {
schema: None,
name: owner.into(),
kind: ObjectKind::Table,
},
name.into(),
columns.into(),
)
}
fn fixture() -> Vec<CatalogEntry> {
vec![
table("users"),
table("orders"),
column("users", "id", "uuid"),
column("users", "email", "text"),
column("orders", "user_id", "uuid"),
index("users", "users_email_idx", "email"),
]
}
fn labels(entries: &[CatalogEntry], query: &str) -> Vec<String> {
search(entries, &Query::parse(query))
.0
.into_iter()
.map(|hit| entries[hit.index].label())
.collect()
}
#[test]
fn exact_names_outrank_prefixes_which_outrank_substrings() {
let entries = vec![table("user"), table("users"), table("superuser")];
assert_eq!(labels(&entries, "user"), ["user", "users", "superuser"]);
}
#[test]
fn a_table_outranks_a_column_on_the_same_score() {
let entries = vec![column("users", "email", "text"), table("email")];
assert_eq!(labels(&entries, "email"), ["email", "email"]);
let hits = search(&entries, &Query::parse("email")).0;
assert_eq!(entries[hits[0].index].kind, CatalogKind::Table);
}
#[test]
fn a_name_match_outranks_a_detail_match() {
let entries = vec![
column("orders", "total", "numeric"),
column("orders", "amount", "total"),
];
let hits = search(&entries, &Query::parse("total")).0;
assert_eq!(entries[hits[0].index].name, "total");
}
#[test]
fn a_member_is_found_by_its_owning_table() {
let entries = fixture();
assert!(labels(&entries, "email").contains(&"email".to_string()));
assert_eq!(labels(&entries, "type:uuid"), ["user_id", "id"]);
assert_eq!(
labels(&entries, "table:users"),
["users", "email", "id", "users_email_idx"]
);
}
#[test]
fn kind_prefixes_are_recognized_in_both_numbers() {
let entries = fixture();
assert_eq!(labels(&entries, "kind:index"), ["users_email_idx"]);
assert_eq!(labels(&entries, "kind:columns"), ["user_id", "email", "id"]);
assert_eq!(labels(&entries, "kind:table"), ["orders", "users"]);
}
#[test]
fn an_unknown_prefix_is_searched_for_as_text() {
let entries = vec![table("http://example.com")];
assert_eq!(
labels(&entries, "http://example.com"),
["http://example.com"]
);
assert!(labels(&entries, "kind:bogus").is_empty());
}
#[test]
fn the_text_is_a_case_insensitive_regex_with_a_literal_fallback() {
let entries = vec![table("Orders"), table("order_items")];
assert_eq!(labels(&entries, "^order"), ["Orders", "order_items"]);
assert!(labels(&entries, "order(").is_empty());
assert_eq!(labels(&entries, "Order_Items"), ["order_items"]);
}
#[test]
fn an_empty_query_keeps_everything() {
let entries = fixture();
assert_eq!(search(&entries, &Query::parse("")).0.len(), entries.len());
}
#[test]
fn queries_parse_prefixes_out_of_the_free_text() {
let query = Query::parse("kind:column type:uuid email");
assert_eq!(query.kinds, [CatalogKind::Column]);
assert_eq!(query.type_name.as_deref(), Some("uuid"));
assert_eq!(query.text, "email");
assert!(!query.is_empty());
assert!(Query::parse(" ").is_empty());
}
#[test]
fn the_catalog_caps_what_it_keeps_and_says_how_many_there_were() {
let entries: Vec<CatalogEntry> = (0..MAX_HITS + 10)
.map(|index| table(&format!("t{index}")))
.collect();
let (hits, matched) = search(&entries, &Query::parse(""));
assert_eq!(hits.len(), MAX_HITS);
assert_eq!(matched, MAX_HITS + 10, "the count is taken before the cap");
let exactly: Vec<CatalogEntry> = entries[..MAX_HITS].to_vec();
let (hits, matched) = search(&exactly, &Query::parse(""));
assert_eq!((hits.len(), matched), (MAX_HITS, MAX_HITS));
let catalog = Catalog {
entries: entries.clone(),
total: entries.len(),
};
assert_eq!(catalog.truncated(), None);
let capped = Catalog {
entries: entries[..MAX_HITS].to_vec(),
total: MAX_ENTRIES + 1,
};
assert_eq!(capped.truncated(), Some(MAX_ENTRIES + 1));
}
}