use std::collections::{HashMap, HashSet};
use std::fs;
use std::path::{Path, PathBuf};
use std::sync::{Arc, Mutex, OnceLock};
use regex::Regex;
use rusqlite::{params_from_iter, Connection, OpenFlags, Result as SqliteResult, ToSql};
use serde_json::Value;
use crate::error::{Error, Result};
pub(crate) const METADATA_DB_NAME: &str = "metadata.db";
const SQLITE_PARAM_LIMIT: usize = 900;
pub(crate) const SUBSET_COLUMN: &str = "_subset_";
const SUBSET_INDEX_NAME: &str = "idx_metadata_subset";
const METADATA_SCHEMA_V1: i64 = 1;
fn is_valid_column_name(name: &str) -> bool {
lazy_static_regex().is_match(name)
}
fn lazy_static_regex() -> &'static Regex {
use std::sync::OnceLock;
static REGEX: OnceLock<Regex> = OnceLock::new();
REGEX.get_or_init(|| Regex::new(r"^[a-zA-Z_][a-zA-Z0-9_]*$").unwrap())
}
#[derive(Debug, Clone, PartialEq)]
enum Token {
Identifier(String),
Placeholder, Eq, Ne, Lt, Le, Gt, Ge, Like,
Regexp,
Between,
In,
And,
Or,
Not,
Is,
Null,
LParen,
RParen,
Comma,
Eof,
}
fn quick_safety_check(condition: &str) -> Result<()> {
let upper = condition.to_uppercase();
if condition.contains("--") || condition.contains("/*") || condition.contains("*/") {
return Err(Error::Filtering(
"SQL comments are not allowed in conditions".into(),
));
}
if condition.contains(';') {
return Err(Error::Filtering(
"Semicolons are not allowed in conditions".into(),
));
}
let dangerous_keywords = [
"SELECT", "UNION", "INSERT", "UPDATE", "DELETE", "DROP", "CREATE", "ALTER", "TRUNCATE",
"EXEC", "EXECUTE", "GRANT", "REVOKE",
];
for keyword in dangerous_keywords {
let pattern = format!(r"\b{}\b", keyword);
if Regex::new(&pattern).unwrap().is_match(&upper) {
return Err(Error::Filtering(format!(
"SQL keyword '{}' is not allowed in conditions",
keyword
)));
}
}
Ok(())
}
fn tokenize(input: &str) -> Result<Vec<Token>> {
let mut tokens = Vec::new();
let chars: Vec<char> = input.chars().collect();
let mut pos = 0;
while pos < chars.len() {
if chars[pos].is_whitespace() {
pos += 1;
continue;
}
match chars[pos] {
'?' => {
tokens.push(Token::Placeholder);
pos += 1;
continue;
}
'(' => {
tokens.push(Token::LParen);
pos += 1;
continue;
}
')' => {
tokens.push(Token::RParen);
pos += 1;
continue;
}
',' => {
tokens.push(Token::Comma);
pos += 1;
continue;
}
'=' => {
tokens.push(Token::Eq);
pos += 1;
continue;
}
_ => {}
}
if pos + 1 < chars.len() {
let two_chars: String = chars[pos..pos + 2].iter().collect();
match two_chars.as_str() {
"!=" => {
tokens.push(Token::Ne);
pos += 2;
continue;
}
"<>" => {
tokens.push(Token::Ne);
pos += 2;
continue;
}
"<=" => {
tokens.push(Token::Le);
pos += 2;
continue;
}
">=" => {
tokens.push(Token::Ge);
pos += 2;
continue;
}
_ => {}
}
}
match chars[pos] {
'<' => {
tokens.push(Token::Lt);
pos += 1;
continue;
}
'>' => {
tokens.push(Token::Gt);
pos += 1;
continue;
}
_ => {}
}
if chars[pos].is_alphabetic() || chars[pos] == '_' {
let start = pos;
while pos < chars.len() && (chars[pos].is_alphanumeric() || chars[pos] == '_') {
pos += 1;
}
let word: String = chars[start..pos].iter().collect();
let upper = word.to_uppercase();
let token = match upper.as_str() {
"AND" => Token::And,
"OR" => Token::Or,
"NOT" => Token::Not,
"IS" => Token::Is,
"NULL" => Token::Null,
"LIKE" => Token::Like,
"REGEXP" => Token::Regexp,
"BETWEEN" => Token::Between,
"IN" => Token::In,
_ => Token::Identifier(word),
};
tokens.push(token);
continue;
}
if chars[pos] == '"' {
pos += 1; let start = pos;
while pos < chars.len() && chars[pos] != '"' {
pos += 1;
}
if pos >= chars.len() {
return Err(Error::Filtering("Unterminated quoted identifier".into()));
}
let word: String = chars[start..pos].iter().collect();
tokens.push(Token::Identifier(word));
pos += 1; continue;
}
return Err(Error::Filtering(format!(
"Unexpected character '{}' in condition",
chars[pos]
)));
}
tokens.push(Token::Eof);
Ok(tokens)
}
struct ConditionValidator<'a> {
tokens: &'a [Token],
pos: usize,
valid_columns: &'a HashSet<String>,
columns_used: Vec<String>,
}
impl<'a> ConditionValidator<'a> {
fn new(tokens: &'a [Token], valid_columns: &'a HashSet<String>) -> Self {
Self {
tokens,
pos: 0,
valid_columns,
columns_used: Vec::new(),
}
}
fn current(&self) -> &Token {
self.tokens.get(self.pos).unwrap_or(&Token::Eof)
}
fn advance(&mut self) {
if self.pos < self.tokens.len() {
self.pos += 1;
}
}
fn expect(&mut self, expected: &Token) -> Result<()> {
if self.current() == expected {
self.advance();
Ok(())
} else {
Err(Error::Filtering(format!(
"Expected {:?}, found {:?}",
expected,
self.current()
)))
}
}
fn validate(&mut self) -> Result<()> {
self.parse_expr()?;
if *self.current() != Token::Eof {
return Err(Error::Filtering(format!(
"Unexpected token {:?} after expression",
self.current()
)));
}
Ok(())
}
fn parse_expr(&mut self) -> Result<()> {
self.parse_and_expr()?;
while *self.current() == Token::Or {
self.advance();
self.parse_and_expr()?;
}
Ok(())
}
fn parse_and_expr(&mut self) -> Result<()> {
self.parse_unary_expr()?;
while *self.current() == Token::And {
self.advance();
self.parse_unary_expr()?;
}
Ok(())
}
fn parse_unary_expr(&mut self) -> Result<()> {
if *self.current() == Token::Not {
self.advance();
}
self.parse_primary_expr()
}
fn parse_primary_expr(&mut self) -> Result<()> {
if *self.current() == Token::LParen {
self.advance();
self.parse_expr()?;
self.expect(&Token::RParen)?;
return Ok(());
}
let col_name = match self.current().clone() {
Token::Identifier(name) => name,
other => {
return Err(Error::Filtering(format!(
"Expected column name, found {:?}",
other
)))
}
};
let col_lower = col_name.to_lowercase();
let valid = self
.valid_columns
.iter()
.any(|c| c.to_lowercase() == col_lower);
if !valid {
return Err(Error::Filtering(format!(
"Unknown column '{}' in condition",
col_name
)));
}
self.columns_used.push(col_name);
self.advance();
match self.current() {
Token::Is => {
self.advance();
if *self.current() == Token::Not {
self.advance();
}
self.expect(&Token::Null)?;
}
Token::Not => {
self.advance();
match self.current() {
Token::Between => {
self.advance();
self.expect(&Token::Placeholder)?;
self.expect(&Token::And)?;
self.expect(&Token::Placeholder)?;
}
Token::In => {
self.advance();
self.parse_in_list()?;
}
Token::Like => {
self.advance();
self.expect(&Token::Placeholder)?;
}
Token::Regexp => {
self.advance();
self.expect(&Token::Placeholder)?;
}
_ => {
return Err(Error::Filtering(format!(
"Expected BETWEEN, IN, LIKE, or REGEXP after NOT, found {:?}",
self.current()
)));
}
}
}
Token::Between => {
self.advance();
self.expect(&Token::Placeholder)?;
self.expect(&Token::And)?;
self.expect(&Token::Placeholder)?;
}
Token::In => {
self.advance();
self.parse_in_list()?;
}
Token::Like => {
self.advance();
self.expect(&Token::Placeholder)?;
}
Token::Regexp => {
self.advance();
self.expect(&Token::Placeholder)?;
}
Token::Eq | Token::Ne | Token::Lt | Token::Le | Token::Gt | Token::Ge => {
self.advance();
self.expect(&Token::Placeholder)?;
}
other => {
return Err(Error::Filtering(format!(
"Expected operator after column name, found {:?}",
other
)));
}
}
Ok(())
}
fn parse_in_list(&mut self) -> Result<()> {
self.expect(&Token::LParen)?;
self.expect(&Token::Placeholder)?;
while *self.current() == Token::Comma {
self.advance();
self.expect(&Token::Placeholder)?;
}
self.expect(&Token::RParen)?;
Ok(())
}
}
fn get_schema_columns(conn: &Connection) -> Result<HashSet<String>> {
let mut columns = HashSet::new();
let mut stmt = conn.prepare("PRAGMA table_info(METADATA)")?;
let rows = stmt.query_map([], |row| row.get::<_, String>(1))?;
for row in rows {
columns.insert(row?);
}
Ok(columns)
}
fn is_numeric_equality(condition: &str) -> bool {
lazy_static_numeric_eq_regex().is_match(condition.trim())
}
fn lazy_static_numeric_eq_regex() -> &'static Regex {
use std::sync::OnceLock;
static REGEX: OnceLock<Regex> = OnceLock::new();
REGEX.get_or_init(|| Regex::new(r"^(\d+)\s*=\s*(\d+)$").unwrap())
}
fn validate_condition(condition: &str, valid_columns: &HashSet<String>) -> Result<()> {
if is_numeric_equality(condition) {
return Ok(());
}
quick_safety_check(condition)?;
let tokens = tokenize(condition)?;
let mut validator = ConditionValidator::new(&tokens, valid_columns);
validator.validate()?;
Ok(())
}
fn infer_sql_type(value: &Value) -> &'static str {
match value {
Value::Number(n) => {
if n.is_i64() || n.is_u64() {
"INTEGER"
} else {
"REAL"
}
}
Value::Bool(_) => "INTEGER",
Value::String(_) => "TEXT",
Value::Null => "TEXT",
Value::Array(_) | Value::Object(_) => "BLOB",
}
}
fn json_to_sql(value: &Value) -> Box<dyn ToSql> {
match value {
Value::Null => Box::new(None::<String>),
Value::Bool(b) => Box::new(if *b { 1i64 } else { 0i64 }),
Value::Number(n) => {
if let Some(i) = n.as_i64() {
Box::new(i)
} else if let Some(f) = n.as_f64() {
Box::new(f)
} else {
Box::new(n.to_string())
}
}
Value::String(s) => Box::new(s.clone()),
Value::Array(_) | Value::Object(_) => Box::new(serde_json::to_string(value).unwrap()),
}
}
pub(crate) fn get_db_path(index_path: &str) -> std::path::PathBuf {
Path::new(index_path).join(METADATA_DB_NAME)
}
static DB_READ_CONNECTIONS: OnceLock<Mutex<HashMap<PathBuf, Arc<Mutex<Connection>>>>> =
OnceLock::new();
fn open_db_read_uncached(db_path: &std::path::Path) -> Result<Connection> {
let conn = Connection::open_with_flags(
db_path,
OpenFlags::SQLITE_OPEN_READ_ONLY | OpenFlags::SQLITE_OPEN_NO_MUTEX,
)?;
conn.execute_batch(
"PRAGMA busy_timeout=5000;
PRAGMA temp_store=MEMORY;
PRAGMA query_only=ON;",
)?;
Ok(conn)
}
fn read_connection(db_path: &Path) -> Result<Arc<Mutex<Connection>>> {
let key = db_path.to_path_buf();
let connections = DB_READ_CONNECTIONS.get_or_init(|| Mutex::new(HashMap::new()));
if let Some(conn) = connections
.lock()
.expect("DB_READ_CONNECTIONS mutex poisoned while reading metadata DB cache")
.get(&key)
.cloned()
{
return Ok(conn);
}
let new_conn = Arc::new(Mutex::new(open_db_read_uncached(db_path)?));
let mut map = connections
.lock()
.expect("DB_READ_CONNECTIONS mutex poisoned while updating metadata DB cache");
Ok(map.entry(key).or_insert_with(|| new_conn).clone())
}
fn invalidate_read_connection(db_path: &Path) {
if let Some(connections) = DB_READ_CONNECTIONS.get() {
connections
.lock()
.expect("DB_READ_CONNECTIONS mutex poisoned while invalidating metadata DB cache")
.remove(db_path);
}
}
pub(crate) fn with_db_read<T>(
db_path: &std::path::Path,
f: impl FnOnce(&Connection) -> Result<T>,
) -> Result<T> {
let conn = read_connection(db_path)?;
let guard = conn
.lock()
.expect("cached metadata read connection mutex poisoned");
f(&guard)
}
pub(crate) fn open_db_write(db_path: &std::path::Path) -> Result<Connection> {
let conn = Connection::open(db_path)?;
conn.execute_batch(
"PRAGMA busy_timeout=5000;
PRAGMA journal_mode=WAL;
PRAGMA synchronous=NORMAL;
PRAGMA temp_store=MEMORY;",
)?;
Ok(conn)
}
fn validate_fixed_columns(columns: &[(&str, &str)]) -> Result<()> {
for (name, _) in columns {
if !is_valid_column_name(name) {
return Err(Error::Filtering(format!(
"Invalid column name '{}'. Column names must start with a letter or \
underscore, followed by letters, digits, or underscores.",
name
)));
}
}
Ok(())
}
fn metadata_schema_version(conn: &Connection) -> i64 {
conn.query_row("PRAGMA user_version", [], |row| row.get(0))
.unwrap_or(0)
}
fn create_subset_index(conn: &Connection) -> Result<()> {
conn.execute(
&format!(
"CREATE INDEX IF NOT EXISTS \"{}\" ON METADATA (\"{}\")",
SUBSET_INDEX_NAME, SUBSET_COLUMN
),
[],
)?;
Ok(())
}
fn ensure_fast_subset_schema(conn: &Connection) -> Result<()> {
if metadata_schema_version(conn) >= METADATA_SCHEMA_V1 {
return Ok(());
}
let has_table: i64 = conn
.query_row(
"SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name='METADATA'",
[],
|row| row.get(0),
)
.unwrap_or(0);
if has_table == 0 {
return Ok(());
}
let mut cols: Vec<(String, String)> = Vec::new(); let mut subset_is_pk = false;
{
let mut stmt = conn.prepare("PRAGMA table_info(METADATA)")?;
let rows = stmt.query_map([], |row| {
Ok((
row.get::<_, String>(1)?, row.get::<_, String>(2)?, row.get::<_, i64>(5)?, ))
})?;
for r in rows {
let (name, ty, pk) = r?;
if name == SUBSET_COLUMN && pk > 0 {
subset_is_pk = true;
}
cols.push((name, ty));
}
}
if !subset_is_pk {
create_subset_index(conn)?;
conn.execute_batch(&format!("PRAGMA user_version={}", METADATA_SCHEMA_V1))?;
return Ok(());
}
let col_defs: Vec<String> = cols
.iter()
.map(|(name, ty)| {
if name == SUBSET_COLUMN {
format!("\"{}\" INTEGER NOT NULL", SUBSET_COLUMN)
} else {
let ty = if ty.trim().is_empty() {
"TEXT"
} else {
ty.as_str()
};
format!("\"{}\" {}", name, ty)
}
})
.collect();
let all_cols = cols
.iter()
.map(|(name, _)| format!("\"{}\"", name))
.collect::<Vec<_>>()
.join(", ");
conn.execute_batch("BEGIN")?;
conn.execute("ALTER TABLE METADATA RENAME TO _METADATA_V0", [])?;
conn.execute(
&format!("CREATE TABLE METADATA ({})", col_defs.join(", ")),
[],
)?;
conn.execute(
&format!(
"INSERT INTO METADATA ({0}) SELECT {0} FROM _METADATA_V0",
all_cols
),
[],
)?;
create_subset_index(conn)?;
conn.execute("DROP TABLE _METADATA_V0", [])?;
conn.execute_batch(&format!("PRAGMA user_version={}", METADATA_SCHEMA_V1))?;
conn.execute_batch("COMMIT")?;
Ok(())
}
fn create_fixed_metadata_table(conn: &Connection, columns: &[(&str, &str)]) -> Result<()> {
let mut col_defs = vec![format!("\"{}\" INTEGER NOT NULL", SUBSET_COLUMN)];
for (name, sql_type) in columns {
col_defs.push(format!("\"{}\" {}", name, sql_type));
}
let create_sql = format!("CREATE TABLE METADATA ({})", col_defs.join(", "));
conn.execute(&create_sql, [])?;
create_subset_index(conn)?;
conn.execute_batch(&format!("PRAGMA user_version={}", METADATA_SCHEMA_V1))?;
Ok(())
}
fn insert_fixed_metadata_rows(
conn: &mut Connection,
columns: &[(&str, &str)],
metadata: &[Value],
doc_ids: &[i64],
) -> Result<usize> {
let txn = conn.transaction()?;
let mut column_names = vec![format!("\"{}\"", SUBSET_COLUMN)];
column_names.extend(columns.iter().map(|(name, _)| format!("\"{}\"", name)));
let placeholders: Vec<&str> = std::iter::repeat_n("?", columns.len() + 1).collect();
let insert_sql = format!(
"INSERT INTO METADATA ({}) VALUES ({})",
column_names.join(", "),
placeholders.join(", ")
);
{
let mut stmt = txn.prepare_cached(&insert_sql)?;
for (i, item) in metadata.iter().enumerate() {
let obj = item.as_object().ok_or_else(|| {
Error::Filtering("Expected metadata rows to be JSON objects".into())
})?;
let mut values: Vec<Box<dyn ToSql>> = vec![Box::new(doc_ids[i])];
for (column_name, _) in columns {
let value = obj.get(*column_name).unwrap_or(&Value::Null);
values.push(json_to_sql(value));
}
let params: Vec<&dyn ToSql> = values.iter().map(|value| value.as_ref()).collect();
stmt.execute(params_from_iter(params))?;
}
}
txn.commit()?;
Ok(metadata.len())
}
fn try_fixed_schema_from_first_row(
metadata: &[Value],
) -> Result<Option<Vec<(&str, &'static str)>>> {
let Some(Value::Object(first_obj)) = metadata.first() else {
return Ok(None);
};
let mut columns = Vec::with_capacity(first_obj.len());
let mut seen = HashSet::with_capacity(first_obj.len());
for (key, value) in first_obj {
if !is_valid_column_name(key) {
return Err(Error::Filtering(format!(
"Invalid column name '{}'. Column names must start with a letter or \
underscore, followed by letters, digits, or underscores.",
key
)));
}
seen.insert(key.as_str());
columns.push((key.as_str(), infer_sql_type(value)));
}
for item in &metadata[1..] {
if let Value::Object(obj) = item {
for key in obj.keys() {
if !seen.contains(key.as_str()) {
return Ok(None);
}
}
}
}
Ok(Some(columns))
}
pub fn exists(index_path: &str) -> bool {
get_db_path(index_path).exists()
}
fn create_with_fixed_columns(
index_path: &str,
columns: &[(&str, &str)],
metadata: &[Value],
doc_ids: &[i64],
) -> Result<usize> {
if metadata.len() != doc_ids.len() {
return Err(Error::Filtering(format!(
"Metadata length ({}) must match doc_ids length ({})",
metadata.len(),
doc_ids.len()
)));
}
validate_fixed_columns(columns)?;
let index_dir = Path::new(index_path);
if !index_dir.exists() {
fs::create_dir_all(index_dir)?;
}
let db_path = get_db_path(index_path);
if db_path.exists() {
invalidate_read_connection(&db_path);
fs::remove_file(&db_path)?;
}
if metadata.is_empty() {
return Ok(0);
}
let mut conn = open_db_write(&db_path)?;
create_fixed_metadata_table(&conn, columns)?;
insert_fixed_metadata_rows(&mut conn, columns, metadata, doc_ids)
}
pub fn create(index_path: &str, metadata: &[Value], doc_ids: &[i64]) -> Result<usize> {
if metadata.len() != doc_ids.len() {
return Err(Error::Filtering(format!(
"Metadata length ({}) must match doc_ids length ({})",
metadata.len(),
doc_ids.len()
)));
}
let index_dir = Path::new(index_path);
if !index_dir.exists() {
fs::create_dir_all(index_dir)?;
}
let db_path = get_db_path(index_path);
if db_path.exists() {
invalidate_read_connection(&db_path);
fs::remove_file(&db_path)?;
}
if metadata.is_empty() {
return Ok(0);
}
if let Some(columns) = try_fixed_schema_from_first_row(metadata)? {
return create_with_fixed_columns(index_path, &columns, metadata, doc_ids);
}
let mut columns: Vec<String> = Vec::new();
let mut column_types: HashMap<String, &'static str> = HashMap::new();
for item in metadata {
if let Value::Object(obj) = item {
for (key, value) in obj {
if !columns.contains(key) {
if !is_valid_column_name(key) {
return Err(Error::Filtering(format!(
"Invalid column name '{}'. Column names must start with a letter or \
underscore, followed by letters, digits, or underscores.",
key
)));
}
columns.push(key.clone());
}
if !value.is_null() && !column_types.contains_key(key) {
column_types.insert(key.clone(), infer_sql_type(value));
}
}
}
}
let mut conn = open_db_write(&db_path)?;
let mut col_defs = vec![format!("\"{}\" INTEGER NOT NULL", SUBSET_COLUMN)];
for col in &columns {
let sql_type = column_types.get(col).copied().unwrap_or("TEXT");
col_defs.push(format!("\"{}\" {}", col, sql_type));
}
let txn = conn.transaction()?;
let create_sql = format!("CREATE TABLE METADATA ({})", col_defs.join(", "));
txn.execute(&create_sql, [])?;
txn.execute(
&format!(
"CREATE INDEX IF NOT EXISTS \"{}\" ON METADATA (\"{}\")",
SUBSET_INDEX_NAME, SUBSET_COLUMN
),
[],
)?;
let placeholders: Vec<&str> = std::iter::repeat_n("?", columns.len() + 1).collect();
let insert_sql = if columns.is_empty() {
format!("INSERT INTO METADATA (\"{}\") VALUES (?)", SUBSET_COLUMN,)
} else {
let col_names: Vec<String> = columns.iter().map(|c| format!("\"{}\"", c)).collect();
format!(
"INSERT INTO METADATA (\"{}\", {}) VALUES ({})",
SUBSET_COLUMN,
col_names.join(", "),
placeholders.join(", ")
)
};
{
let mut stmt = txn.prepare(&insert_sql)?;
for (i, item) in metadata.iter().enumerate() {
let mut values: Vec<Box<dyn ToSql>> = vec![Box::new(doc_ids[i])];
if let Value::Object(obj) = item {
for col in &columns {
let value = obj.get(col).unwrap_or(&Value::Null);
values.push(json_to_sql(value));
}
} else {
for _ in &columns {
values.push(Box::new(None::<String>));
}
}
let params: Vec<&dyn ToSql> = values.iter().map(|v| v.as_ref()).collect();
stmt.execute(params_from_iter(params))?;
}
}
txn.commit()?;
conn.execute_batch(&format!("PRAGMA user_version={}", METADATA_SCHEMA_V1))?;
Ok(metadata.len())
}
pub fn update(index_path: &str, metadata: &[Value], doc_ids: &[i64]) -> Result<usize> {
if metadata.is_empty() {
return Ok(0);
}
if metadata.len() != doc_ids.len() {
return Err(Error::Filtering(format!(
"Metadata length ({}) must match doc_ids length ({})",
metadata.len(),
doc_ids.len()
)));
}
let db_path = get_db_path(index_path);
if !db_path.exists() {
return Err(Error::Filtering(
"Metadata database does not exist. Use create() first.".into(),
));
}
let mut conn = open_db_write(&db_path)?;
let mut existing_columns: Vec<String> = Vec::new();
{
let mut stmt = conn.prepare("PRAGMA table_info(METADATA)")?;
let rows = stmt.query_map([], |row| row.get::<_, String>(1))?;
for row in rows {
let col = row?;
if col != SUBSET_COLUMN {
existing_columns.push(col);
}
}
}
let existing_column_set: HashSet<&str> = existing_columns
.iter()
.map(|column| column.as_str())
.collect();
let has_new_columns = metadata.iter().any(|item| {
item.as_object().is_some_and(|obj| {
obj.keys()
.any(|key| !existing_column_set.contains(key.as_str()))
})
});
if !has_new_columns {
let fixed_columns: Vec<(&str, &str)> = existing_columns
.iter()
.map(|column| (column.as_str(), "TEXT"))
.collect();
return insert_fixed_metadata_rows(&mut conn, &fixed_columns, metadata, doc_ids);
}
let mut new_columns: Vec<String> = Vec::new();
let mut column_types: HashMap<String, &'static str> = HashMap::new();
for item in metadata {
if let Value::Object(obj) = item {
for (key, value) in obj {
if !existing_columns.contains(key) && !new_columns.contains(key) {
if !is_valid_column_name(key) {
return Err(Error::Filtering(format!(
"Invalid column name '{}'. Column names must start with a letter or \
underscore, followed by letters, digits, or underscores.",
key
)));
}
new_columns.push(key.clone());
}
if !value.is_null() && !column_types.contains_key(key) {
column_types.insert(key.clone(), infer_sql_type(value));
}
}
}
}
let txn = conn.transaction()?;
for col in &new_columns {
let sql_type = column_types.get(col).copied().unwrap_or("TEXT");
let alter_sql = format!("ALTER TABLE METADATA ADD COLUMN \"{}\" {}", col, sql_type);
txn.execute(&alter_sql, [])?;
}
let all_columns: Vec<String> = existing_columns.into_iter().chain(new_columns).collect();
let placeholders: Vec<&str> = std::iter::repeat_n("?", all_columns.len() + 1).collect();
let insert_sql = if all_columns.is_empty() {
format!("INSERT INTO METADATA (\"{}\") VALUES (?)", SUBSET_COLUMN,)
} else {
let col_names: Vec<String> = all_columns.iter().map(|c| format!("\"{}\"", c)).collect();
format!(
"INSERT INTO METADATA (\"{}\", {}) VALUES ({})",
SUBSET_COLUMN,
col_names.join(", "),
placeholders.join(", ")
)
};
{
let mut stmt = txn.prepare(&insert_sql)?;
for (i, item) in metadata.iter().enumerate() {
let mut values: Vec<Box<dyn ToSql>> = vec![Box::new(doc_ids[i])];
if let Value::Object(obj) = item {
for col in &all_columns {
let value = obj.get(col).unwrap_or(&Value::Null);
values.push(json_to_sql(value));
}
} else {
for _ in &all_columns {
values.push(Box::new(None::<String>));
}
}
let params: Vec<&dyn ToSql> = values.iter().map(|v| v.as_ref()).collect();
stmt.execute(params_from_iter(params))?;
}
}
txn.commit()?;
Ok(metadata.len())
}
pub fn delete(index_path: &str, subset: &[i64]) -> Result<usize> {
if subset.is_empty() {
return Ok(0);
}
let db_path = get_db_path(index_path);
if !db_path.exists() {
return Ok(0);
}
let conn = open_db_write(&db_path)?;
ensure_fast_subset_schema(&conn)?;
conn.execute("BEGIN", [])?;
let original_count: i64 = conn
.query_row(
&format!(
"SELECT COALESCE(MAX(\"{}\"), -1) FROM METADATA",
SUBSET_COLUMN
),
[],
|row| row.get::<_, i64>(0),
)
.unwrap_or(-1)
+ 1;
let (in_clause, in_params, temp_table) = crate::text_search::build_in_clause(&conn, subset)?;
let delete_sql = format!(
"DELETE FROM METADATA WHERE \"{}\" {}",
SUBSET_COLUMN, in_clause
);
let param_refs: Vec<&dyn ToSql> = in_params.iter().map(|v| v.as_ref()).collect();
let deleted = conn.execute(&delete_sql, params_from_iter(param_refs))?;
if let Some(ref name) = temp_table {
crate::text_search::drop_temp_table(&conn, name);
}
let mut sorted_ids: Vec<i64> = subset.to_vec();
sorted_ids.sort_unstable();
sorted_ids.dedup();
sorted_ids.retain(|&id| id >= 0 && id < original_count);
let max_id: i64 = conn
.query_row(
&format!(
"SELECT COALESCE(MAX(\"{}\"), -1) FROM METADATA",
SUBSET_COLUMN
),
[],
|row| row.get(0),
)
.unwrap_or(-1);
if max_id >= 0 && !sorted_ids.is_empty() {
let mut updates: Vec<(i64, i64, i64)> = Vec::new();
let mut i = 0;
while i < sorted_ids.len() {
let mut j = i + 1;
while j < sorted_ids.len() && sorted_ids[j] == sorted_ids[j - 1] + 1 {
j += 1;
}
let range_start = sorted_ids[j - 1] + 1;
let range_end = if j < sorted_ids.len() {
sorted_ids[j]
} else {
max_id + sorted_ids.len() as i64 + 1
};
if range_start < range_end {
updates.push((range_start, range_end, j as i64));
}
i = j;
}
for (from, to_excl, shift) in &updates {
conn.execute(
&format!(
"UPDATE METADATA SET \"{}\" = \"{}\" - ?1 WHERE \"{}\" >= ?2 AND \"{}\" < ?3",
SUBSET_COLUMN, SUBSET_COLUMN, SUBSET_COLUMN, SUBSET_COLUMN
),
rusqlite::params![shift, from, to_excl],
)?;
}
}
conn.execute("COMMIT", [])?;
Ok(deleted)
}
pub fn where_condition(
index_path: &str,
condition: &str,
parameters: &[Value],
) -> Result<Vec<i64>> {
let db_path = get_db_path(index_path);
if !db_path.exists() {
return Err(Error::Filtering(
"No metadata database found. Create it first by adding metadata during index creation."
.into(),
));
}
with_db_read(&db_path, |conn| {
let valid_columns = get_schema_columns(conn)?;
validate_condition(condition, &valid_columns)?;
let query = format!(
"SELECT \"{}\" FROM METADATA WHERE {}",
SUBSET_COLUMN, condition
);
let params: Vec<Box<dyn ToSql>> = parameters.iter().map(json_to_sql).collect();
let param_refs: Vec<&dyn ToSql> = params.iter().map(|v| v.as_ref()).collect();
let mut stmt = conn.prepare(&query)?;
let rows = stmt.query_map(params_from_iter(param_refs), |row| row.get::<_, i64>(0))?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
})
}
pub fn where_condition_regexp(
index_path: &str,
condition: &str,
parameters: &[Value],
) -> Result<Vec<i64>> {
let db_path = get_db_path(index_path);
if !db_path.exists() {
return Err(Error::Filtering(
"No metadata database found. Create it first by adding metadata during index creation."
.into(),
));
}
let regex_pattern = parameters
.first()
.and_then(|v| v.as_str())
.ok_or_else(|| Error::Filtering("REGEXP requires a pattern parameter".into()))?;
let compiled_regex = std::sync::Arc::new(
fancy_regex::RegexBuilder::new(regex_pattern)
.build()
.map_err(|e| {
Error::Filtering(format!("Invalid regex pattern '{}': {}", regex_pattern, e))
})?,
);
with_db_read(&db_path, |conn| {
let valid_columns = get_schema_columns(conn)?;
validate_condition(condition, &valid_columns)?;
let re = compiled_regex.clone();
conn.create_scalar_function(
"regexp",
2,
rusqlite::functions::FunctionFlags::SQLITE_UTF8
| rusqlite::functions::FunctionFlags::SQLITE_DETERMINISTIC,
move |ctx| {
let _pattern: String = ctx.get(0)?;
let text: String = ctx.get(1)?;
Ok(re.is_match(&text).unwrap_or(false))
},
)?;
let query = format!(
"SELECT \"{}\" FROM METADATA WHERE {}",
SUBSET_COLUMN, condition
);
let params: Vec<Box<dyn ToSql>> = parameters.iter().map(json_to_sql).collect();
let param_refs: Vec<&dyn ToSql> = params.iter().map(|v| v.as_ref()).collect();
let mut stmt = conn.prepare(&query)?;
let rows = stmt.query_map(params_from_iter(param_refs), |row| row.get::<_, i64>(0))?;
let mut result = Vec::new();
for row in rows {
result.push(row?);
}
Ok(result)
})
}
pub fn get_distinct_strings(index_path: &str, column: &str) -> Result<Vec<String>> {
let db_path = get_db_path(index_path);
if !db_path.exists() {
return Ok(Vec::new());
}
if !is_valid_column_name(column) {
return Err(Error::Filtering(format!(
"Invalid column name: '{}'",
column
)));
}
with_db_read(&db_path, |conn| {
let columns = get_schema_columns(conn)?;
if !columns.contains(column) {
return Ok(Vec::new());
}
let query = format!(
"SELECT DISTINCT \"{0}\" FROM METADATA WHERE \"{0}\" IS NOT NULL",
column
);
let mut stmt = conn.prepare(&query)?;
let rows = stmt.query_map([], |row| row.get::<_, Option<String>>(0))?;
let mut values: Vec<String> = Vec::new();
for row in rows {
if let Some(value) = row? {
values.push(value);
}
}
Ok(values)
})
}
pub fn get(
index_path: &str,
condition: Option<&str>,
parameters: &[Value],
subset: Option<&[i64]>,
) -> Result<Vec<Value>> {
if condition.is_some() && subset.is_some() {
return Err(Error::Filtering(
"Please provide either a 'condition' or a 'subset', not both.".into(),
));
}
let db_path = get_db_path(index_path);
if !db_path.exists() {
return Ok(Vec::new());
}
with_db_read(&db_path, |conn| {
if let Some(cond) = condition {
let valid_columns = get_schema_columns(conn)?;
validate_condition(cond, &valid_columns)?;
}
let mut columns: Vec<String> = Vec::new();
{
let mut stmt = conn.prepare("PRAGMA table_info(METADATA)")?;
let rows = stmt.query_map([], |row| row.get::<_, String>(1))?;
for row in rows {
columns.push(row?);
}
}
if let Some(ids) = subset {
if ids.is_empty() {
return Ok(Vec::new());
}
let mut results: Vec<Value> = Vec::new();
for chunk in ids.chunks(SQLITE_PARAM_LIMIT) {
let placeholders: Vec<&str> = std::iter::repeat_n("?", chunk.len()).collect();
let query = format!(
"SELECT * FROM METADATA WHERE \"{}\" IN ({})",
SUBSET_COLUMN,
placeholders.join(", ")
);
let params: Vec<Box<dyn ToSql>> = chunk
.iter()
.map(|&id| Box::new(id) as Box<dyn ToSql>)
.collect();
let param_refs: Vec<&dyn ToSql> = params.iter().map(|v| v.as_ref()).collect();
let mut stmt = conn.prepare(&query)?;
let mut rows = stmt.query(params_from_iter(param_refs))?;
while let Some(row) = rows.next()? {
let mut obj = serde_json::Map::new();
for (i, col) in columns.iter().enumerate() {
let value = row_to_json_value(row, i)?;
obj.insert(col.clone(), value);
}
results.push(Value::Object(obj));
}
}
let mut results_map: HashMap<i64, Value> = HashMap::new();
for result in results {
if let Some(id) = result.get(SUBSET_COLUMN).and_then(|v| v.as_i64()) {
results_map.insert(id, result);
}
}
return Ok(ids.iter().filter_map(|id| results_map.remove(id)).collect());
}
let (query, params): (String, Vec<Box<dyn ToSql>>) = if let Some(cond) = condition {
let query = format!(
"SELECT * FROM METADATA WHERE {} ORDER BY \"{}\"",
cond, SUBSET_COLUMN
);
let params = parameters.iter().map(json_to_sql).collect();
(query, params)
} else {
let query = format!("SELECT * FROM METADATA ORDER BY \"{}\"", SUBSET_COLUMN);
(query, Vec::new())
};
let param_refs: Vec<&dyn ToSql> = params.iter().map(|v| v.as_ref()).collect();
let mut stmt = conn.prepare(&query)?;
let mut rows = stmt.query(params_from_iter(param_refs))?;
let mut results: Vec<Value> = Vec::new();
while let Some(row) = rows.next()? {
let mut obj = serde_json::Map::new();
for (i, col) in columns.iter().enumerate() {
let value = row_to_json_value(row, i)?;
obj.insert(col.clone(), value);
}
results.push(Value::Object(obj));
}
Ok(results)
})
}
fn row_to_json_value(row: &rusqlite::Row, idx: usize) -> SqliteResult<Value> {
if let Ok(i) = row.get::<_, i64>(idx) {
return Ok(Value::Number(i.into()));
}
if let Ok(f) = row.get::<_, f64>(idx) {
return Ok(serde_json::Number::from_f64(f)
.map(Value::Number)
.unwrap_or(Value::Null));
}
if let Ok(s) = row.get::<_, String>(idx) {
return Ok(Value::String(s));
}
if let Ok(b) = row.get::<_, Vec<u8>>(idx) {
if let Ok(v) = serde_json::from_slice(&b) {
return Ok(v);
}
return Ok(Value::String(base64_encode(&b)));
}
Ok(Value::Null)
}
fn base64_encode(data: &[u8]) -> String {
let mut result = String::with_capacity(data.len() * 4 / 3 + 4);
const ALPHABET: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
for chunk in data.chunks(3) {
let b0 = chunk[0] as usize;
let b1 = chunk.get(1).copied().unwrap_or(0) as usize;
let b2 = chunk.get(2).copied().unwrap_or(0) as usize;
result.push(ALPHABET[b0 >> 2] as char);
result.push(ALPHABET[((b0 & 0x03) << 4) | (b1 >> 4)] as char);
if chunk.len() > 1 {
result.push(ALPHABET[((b1 & 0x0f) << 2) | (b2 >> 6)] as char);
} else {
result.push('=');
}
if chunk.len() > 2 {
result.push(ALPHABET[b2 & 0x3f] as char);
} else {
result.push('=');
}
}
result
}
pub fn update_where(
index_path: &str,
condition: &str,
parameters: &[Value],
updates: &Value,
) -> Result<usize> {
let db_path = get_db_path(index_path);
if !db_path.exists() {
return Err(Error::Filtering(
"No metadata database found. Create it first by adding metadata during index creation."
.into(),
));
}
let updates_obj = match updates {
Value::Object(obj) => obj,
_ => {
return Err(Error::Filtering("Updates must be a JSON object".into()));
}
};
if updates_obj.is_empty() {
return Ok(0);
}
let conn = open_db_write(&db_path)?;
let valid_columns = get_schema_columns(&conn)?;
validate_condition(condition, &valid_columns)?;
for col_name in updates_obj.keys() {
if col_name == SUBSET_COLUMN {
return Err(Error::Filtering("Cannot update the _subset_ column".into()));
}
if !is_valid_column_name(col_name) {
return Err(Error::Filtering(format!(
"Invalid column name '{}'. Column names must start with a letter or \
underscore, followed by letters, digits, or underscores.",
col_name
)));
}
let col_lower = col_name.to_lowercase();
let exists = valid_columns.iter().any(|c| c.to_lowercase() == col_lower);
if !exists {
return Err(Error::Filtering(format!(
"Unknown column '{}' in updates",
col_name
)));
}
}
let affected_ids: Vec<i64> = {
let affected_query = format!(
"SELECT \"{}\" FROM METADATA WHERE {}",
SUBSET_COLUMN, condition
);
let cond_params: Vec<Box<dyn ToSql>> = parameters.iter().map(json_to_sql).collect();
let cond_param_refs: Vec<&dyn ToSql> = cond_params.iter().map(|v| v.as_ref()).collect();
let mut affected_stmt = conn.prepare(&affected_query)?;
let rows = affected_stmt.query_map(params_from_iter(cond_param_refs), |row| {
row.get::<_, i64>(0)
})?;
rows.filter_map(|row| row.ok()).collect()
};
let set_parts: Vec<String> = updates_obj
.keys()
.map(|col| format!("\"{}\" = ?", col))
.collect();
let set_clause = set_parts.join(", ");
let query = format!("UPDATE METADATA SET {} WHERE {}", set_clause, condition);
let mut all_params: Vec<Box<dyn ToSql>> = updates_obj.values().map(json_to_sql).collect();
all_params.extend(parameters.iter().map(json_to_sql));
let param_refs: Vec<&dyn ToSql> = all_params.iter().map(|v| v.as_ref()).collect();
let updated = conn.execute(&query, params_from_iter(param_refs))?;
if updated > 0 && !affected_ids.is_empty() {
crate::text_search::update_rows(index_path, &affected_ids)?;
}
Ok(updated)
}
pub fn count(index_path: &str) -> Result<usize> {
let db_path = get_db_path(index_path);
if !db_path.exists() {
return Ok(0);
}
with_db_read(&db_path, |conn| {
let count: i64 = conn.query_row("SELECT COUNT(*) FROM METADATA", [], |row| row.get(0))?;
Ok(count as usize)
})
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use tempfile::TempDir;
fn setup_test_dir() -> TempDir {
TempDir::new().unwrap()
}
#[test]
fn test_create_empty() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let result = create(path, &[], &[]).unwrap();
assert_eq!(result, 0);
assert!(!exists(path));
}
#[test]
fn test_create_with_metadata() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![
json!({"name": "Alice", "age": 30, "score": 95.5}),
json!({"name": "Bob", "age": 25, "score": 87.0}),
json!({"name": "Charlie", "age": 35}),
];
let doc_ids: Vec<i64> = (0..3).collect();
let result = create(path, &metadata, &doc_ids).unwrap();
assert_eq!(result, 3);
assert!(exists(path));
assert_eq!(count(path).unwrap(), 3);
}
#[test]
fn test_create_invalid_column_name() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![json!({"valid_name": "Alice", "invalid name": 30})];
let doc_ids = vec![0];
let result = create(path, &metadata, &doc_ids);
assert!(result.is_err());
}
#[test]
fn test_where_condition() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![
json!({"name": "Alice", "category": "A", "score": 95}),
json!({"name": "Bob", "category": "B", "score": 87}),
json!({"name": "Charlie", "category": "A", "score": 92}),
];
let doc_ids: Vec<i64> = (0..3).collect();
create(path, &metadata, &doc_ids).unwrap();
let subset = where_condition(path, "category = ?", &[json!("A")]).unwrap();
assert_eq!(subset, vec![0, 2]);
let subset =
where_condition(path, "category = ? AND score > ?", &[json!("A"), json!(93)]).unwrap();
assert_eq!(subset, vec![0]);
}
#[test]
fn test_get_distinct_strings_returns_unique_values() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![
json!({"file": "src/a.rs", "code": "x"}),
json!({"file": "src/a.rs", "code": "y"}),
json!({"file": "src/b.rs", "code": "z"}),
];
let doc_ids: Vec<i64> = (0..3).collect();
create(path, &metadata, &doc_ids).unwrap();
let mut files = get_distinct_strings(path, "file").unwrap();
files.sort();
assert_eq!(files, vec!["src/a.rs".to_string(), "src/b.rs".to_string()]);
}
#[test]
fn test_get_distinct_strings_missing_db_returns_empty() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let files = get_distinct_strings(path, "file").unwrap();
assert!(files.is_empty());
}
#[test]
fn test_get_distinct_strings_unknown_column_returns_empty() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![json!({"file": "src/a.rs"})];
create(path, &metadata, &[0]).unwrap();
let values = get_distinct_strings(path, "not_a_column").unwrap();
assert!(values.is_empty());
}
#[test]
fn test_get_distinct_strings_rejects_invalid_column_name() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![json!({"file": "src/a.rs"})];
create(path, &metadata, &[0]).unwrap();
let result = get_distinct_strings(path, "file; DROP TABLE METADATA --");
assert!(result.is_err());
}
#[test]
fn test_get_all() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![
json!({"name": "Alice", "age": 30}),
json!({"name": "Bob", "age": 25}),
];
let doc_ids: Vec<i64> = (0..2).collect();
create(path, &metadata, &doc_ids).unwrap();
let results = get(path, None, &[], None).unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0]["name"], "Alice");
assert_eq!(results[1]["name"], "Bob");
}
#[test]
fn test_get_by_subset() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![
json!({"name": "Alice"}),
json!({"name": "Bob"}),
json!({"name": "Charlie"}),
];
let doc_ids: Vec<i64> = (0..3).collect();
create(path, &metadata, &doc_ids).unwrap();
let results = get(path, None, &[], Some(&[2, 0])).unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0]["name"], "Charlie");
assert_eq!(results[1]["name"], "Alice");
}
#[test]
fn test_update_adds_rows() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata1 = vec![json!({"name": "Alice"}), json!({"name": "Bob"})];
let doc_ids1: Vec<i64> = (0..2).collect();
create(path, &metadata1, &doc_ids1).unwrap();
assert_eq!(count(path).unwrap(), 2);
let metadata2 = vec![json!({"name": "Charlie"})];
let doc_ids2 = vec![2];
update(path, &metadata2, &doc_ids2).unwrap();
assert_eq!(count(path).unwrap(), 3);
let results = get(path, None, &[], None).unwrap();
assert_eq!(results[2]["_subset_"], 2);
assert_eq!(results[2]["name"], "Charlie");
}
#[test]
fn test_update_adds_columns() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata1 = vec![json!({"name": "Alice"})];
let doc_ids1 = vec![0];
create(path, &metadata1, &doc_ids1).unwrap();
let metadata2 = vec![json!({"name": "Bob", "age": 25, "city": "NYC"})];
let doc_ids2 = vec![1];
update(path, &metadata2, &doc_ids2).unwrap();
let results = get(path, None, &[], None).unwrap();
assert_eq!(results[0]["name"], "Alice");
assert!(results[0]["age"].is_null()); assert_eq!(results[1]["age"], 25);
assert_eq!(results[1]["city"], "NYC");
}
#[test]
fn test_delete_and_reindex() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![
json!({"name": "Alice"}),
json!({"name": "Bob"}),
json!({"name": "Charlie"}),
json!({"name": "Diana"}),
];
let doc_ids: Vec<i64> = (0..4).collect();
create(path, &metadata, &doc_ids).unwrap();
let deleted = delete(path, &[1, 2]).unwrap();
assert_eq!(deleted, 2);
assert_eq!(count(path).unwrap(), 2);
let results = get(path, None, &[], None).unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0]["_subset_"], 0);
assert_eq!(results[0]["name"], "Alice");
assert_eq!(results[1]["_subset_"], 1);
assert_eq!(results[1]["name"], "Diana");
}
#[test]
fn test_delete_resequence_ignores_out_of_range_and_negative_ids() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![
json!({"name": "Alice"}), json!({"name": "Bob"}), json!({"name": "Charlie"}), json!({"name": "Diana"}), json!({"name": "Eve"}), json!({"name": "Frank"}), ];
let doc_ids: Vec<i64> = (0..6).collect();
create(path, &metadata, &doc_ids).unwrap();
let deleted = delete(path, &[1, 3, -5, 999]).unwrap();
assert_eq!(deleted, 2, "only the two present ids are removed");
assert_eq!(count(path).unwrap(), 4);
let rows = get(path, None, &[], None).unwrap();
let got: Vec<(i64, String)> = rows
.iter()
.map(|r| {
(
r["_subset_"].as_i64().unwrap(),
r["name"].as_str().unwrap().to_string(),
)
})
.collect();
assert_eq!(
got,
vec![
(0, "Alice".into()),
(1, "Charlie".into()),
(2, "Eve".into()),
(3, "Frank".into()),
]
);
}
#[test]
fn test_where_with_like() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![
json!({"name": "Alice"}),
json!({"name": "Alex"}),
json!({"name": "Bob"}),
];
let doc_ids: Vec<i64> = (0..3).collect();
create(path, &metadata, &doc_ids).unwrap();
let subset = where_condition(path, "name LIKE ?", &[json!("Al%")]).unwrap();
assert_eq!(subset, vec![0, 1]);
}
#[test]
fn test_is_valid_column_name() {
assert!(is_valid_column_name("name"));
assert!(is_valid_column_name("_private"));
assert!(is_valid_column_name("column123"));
assert!(is_valid_column_name("Col_Name_2"));
assert!(!is_valid_column_name("123column")); assert!(!is_valid_column_name("column name")); assert!(!is_valid_column_name("column-name")); assert!(!is_valid_column_name("")); assert!(!is_valid_column_name("col;drop")); }
#[test]
fn test_type_inference() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![json!({
"int_val": 42,
"float_val": 3.125,
"str_val": "hello",
"bool_val": true,
"null_val": null
})];
let doc_ids = vec![0];
create(path, &metadata, &doc_ids).unwrap();
let results = get(path, None, &[], None).unwrap();
assert_eq!(results[0]["int_val"], 42);
assert!((results[0]["float_val"].as_f64().unwrap() - 3.125).abs() < 0.001);
assert_eq!(results[0]["str_val"], "hello");
assert_eq!(results[0]["bool_val"], 1); assert!(results[0]["null_val"].is_null());
}
fn test_columns() -> HashSet<String> {
["name", "category", "score", "status", "_subset_"]
.iter()
.map(|s| s.to_string())
.collect()
}
#[test]
fn test_validator_simple_equality() {
let cols = test_columns();
assert!(validate_condition("name = ?", &cols).is_ok());
assert!(validate_condition("score = ?", &cols).is_ok());
}
#[test]
fn test_validator_comparison_operators() {
let cols = test_columns();
assert!(validate_condition("score > ?", &cols).is_ok());
assert!(validate_condition("score >= ?", &cols).is_ok());
assert!(validate_condition("score < ?", &cols).is_ok());
assert!(validate_condition("score <= ?", &cols).is_ok());
assert!(validate_condition("score != ?", &cols).is_ok());
assert!(validate_condition("score <> ?", &cols).is_ok());
}
#[test]
fn test_validator_and_or() {
let cols = test_columns();
assert!(validate_condition("name = ? AND score > ?", &cols).is_ok());
assert!(validate_condition("category = ? OR status = ?", &cols).is_ok());
assert!(validate_condition("name = ? AND score > ? OR category = ?", &cols).is_ok());
}
#[test]
fn test_validator_like() {
let cols = test_columns();
assert!(validate_condition("name LIKE ?", &cols).is_ok());
assert!(validate_condition("name NOT LIKE ?", &cols).is_ok());
}
#[test]
fn test_validator_regexp() {
let cols = test_columns();
assert!(validate_condition("name REGEXP ?", &cols).is_ok());
assert!(validate_condition("name NOT REGEXP ?", &cols).is_ok());
}
#[test]
fn test_validator_between() {
let cols = test_columns();
assert!(validate_condition("score BETWEEN ? AND ?", &cols).is_ok());
assert!(validate_condition("score NOT BETWEEN ? AND ?", &cols).is_ok());
}
#[test]
fn test_validator_in() {
let cols = test_columns();
assert!(validate_condition("category IN (?)", &cols).is_ok());
assert!(validate_condition("category IN (?, ?)", &cols).is_ok());
assert!(validate_condition("category IN (?, ?, ?)", &cols).is_ok());
assert!(validate_condition("category NOT IN (?, ?)", &cols).is_ok());
}
#[test]
fn test_validator_is_null() {
let cols = test_columns();
assert!(validate_condition("name IS NULL", &cols).is_ok());
assert!(validate_condition("name IS NOT NULL", &cols).is_ok());
}
#[test]
fn test_validator_parentheses() {
let cols = test_columns();
assert!(validate_condition("(name = ?)", &cols).is_ok());
assert!(validate_condition("(name = ? AND score > ?)", &cols).is_ok());
assert!(validate_condition("(name = ? OR category = ?) AND score > ?", &cols).is_ok());
assert!(validate_condition("name = ? AND (category = ? OR status = ?)", &cols).is_ok());
}
#[test]
fn test_validator_not() {
let cols = test_columns();
assert!(validate_condition("NOT name = ?", &cols).is_ok());
assert!(validate_condition("NOT (name = ? AND score > ?)", &cols).is_ok());
}
#[test]
fn test_validator_quoted_identifiers() {
let cols = test_columns();
assert!(validate_condition("\"name\" = ?", &cols).is_ok());
assert!(validate_condition("\"score\" > ?", &cols).is_ok());
}
#[test]
fn test_validator_case_insensitive_keywords() {
let cols = test_columns();
assert!(validate_condition("name = ? and score > ?", &cols).is_ok());
assert!(validate_condition("name = ? AND score > ?", &cols).is_ok());
assert!(validate_condition("name LIKE ? or category = ?", &cols).is_ok());
assert!(validate_condition("score between ? and ?", &cols).is_ok());
}
#[test]
fn test_validator_allows_numeric_equality() {
let cols = test_columns();
assert!(validate_condition("1=1", &cols).is_ok());
assert!(validate_condition(" 1=1 ", &cols).is_ok()); assert!(validate_condition("0=0", &cols).is_ok());
assert!(validate_condition("1 = 1", &cols).is_ok()); assert!(validate_condition("42=42", &cols).is_ok());
assert!(validate_condition("1=0", &cols).is_ok()); }
#[test]
fn test_validator_rejects_semicolon() {
let cols = test_columns();
let result = validate_condition("name = ?; DROP TABLE METADATA", &cols);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("Semicolon"));
}
#[test]
fn test_validator_rejects_comments() {
let cols = test_columns();
assert!(validate_condition("name = ? -- comment", &cols).is_err());
assert!(validate_condition("name = ? /* comment */", &cols).is_err());
}
#[test]
fn test_validator_rejects_union() {
let cols = test_columns();
let result = validate_condition("name = ? UNION SELECT * FROM users", &cols);
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("UNION") || err_msg.contains("SELECT"),
"Expected error about UNION or SELECT, got: {}",
err_msg
);
}
#[test]
fn test_validator_rejects_subqueries() {
let cols = test_columns();
let result = validate_condition("name = (SELECT name FROM users)", &cols);
assert!(result.is_err());
}
#[test]
fn test_validator_rejects_ddl_keywords() {
let cols = test_columns();
assert!(validate_condition("DROP TABLE METADATA", &cols).is_err());
assert!(validate_condition("DELETE FROM METADATA", &cols).is_err());
assert!(validate_condition("INSERT INTO METADATA VALUES (?)", &cols).is_err());
assert!(validate_condition("UPDATE METADATA SET name = ?", &cols).is_err());
assert!(validate_condition("CREATE TABLE foo (id INT)", &cols).is_err());
assert!(validate_condition("ALTER TABLE METADATA ADD x INT", &cols).is_err());
assert!(validate_condition("TRUNCATE TABLE METADATA", &cols).is_err());
}
#[test]
fn test_validator_rejects_unknown_columns() {
let cols = test_columns();
let result = validate_condition("unknown_column = ?", &cols);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("Unknown column"));
}
#[test]
fn test_validator_rejects_string_literals() {
let cols = test_columns();
let result = validate_condition("name = 'Alice'", &cols);
assert!(result.is_err());
}
#[test]
fn test_validator_rejects_malformed_syntax() {
let cols = test_columns();
assert!(validate_condition("name =", &cols).is_err());
assert!(validate_condition("(name = ?", &cols).is_err());
assert!(validate_condition("name = ?)", &cols).is_err());
assert!(validate_condition("name = = ?", &cols).is_err());
assert!(validate_condition("= ?", &cols).is_err());
}
#[test]
fn test_validator_rejects_function_calls() {
let cols = test_columns();
let result = validate_condition("LENGTH(name) > ?", &cols);
assert!(result.is_err());
}
#[test]
fn test_validator_integration() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![
json!({"name": "Alice", "category": "A", "score": 95}),
json!({"name": "Bob", "category": "B", "score": 87}),
];
let doc_ids: Vec<i64> = (0..2).collect();
create(path, &metadata, &doc_ids).unwrap();
let result = where_condition(path, "category = ? AND score > ?", &[json!("A"), json!(90)]);
assert!(result.is_ok());
assert_eq!(result.unwrap(), vec![0]);
let result = where_condition(path, "category = ?; DROP TABLE METADATA", &[json!("A")]);
assert!(result.is_err());
let result = where_condition(path, "unknown = ?", &[json!("test")]);
assert!(result.is_err());
}
#[test]
fn test_validator_integration_get() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![
json!({"name": "Alice", "score": 95}),
json!({"name": "Bob", "score": 87}),
];
let doc_ids: Vec<i64> = (0..2).collect();
create(path, &metadata, &doc_ids).unwrap();
let result = get(path, Some("score > ?"), &[json!(90)], None);
assert!(result.is_ok());
assert_eq!(result.unwrap().len(), 1);
let result = get(path, Some("1=1 UNION SELECT * FROM users"), &[], None);
assert!(result.is_err());
}
#[test]
fn test_create_with_empty_metadata_objects() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![json!({}), json!({})];
let doc_ids: Vec<i64> = vec![0, 1];
let result = create(path, &metadata, &doc_ids).unwrap();
assert_eq!(result, 2);
assert!(exists(path));
assert_eq!(count(path).unwrap(), 2);
let all = get(path, None, &[], None).unwrap();
assert_eq!(all.len(), 2);
}
#[test]
fn test_update_with_empty_metadata_objects() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![json!({})];
let doc_ids: Vec<i64> = vec![0];
create(path, &metadata, &doc_ids).unwrap();
let new_metadata = vec![json!({})];
let new_doc_ids: Vec<i64> = vec![1];
let result = update(path, &new_metadata, &new_doc_ids).unwrap();
assert_eq!(result, 1);
assert_eq!(count(path).unwrap(), 2);
}
#[test]
fn test_create_with_mixed_empty_and_non_empty_metadata() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![
json!({"name": "Alice", "score": 95}),
json!({}),
json!({"name": "Charlie"}),
];
let doc_ids: Vec<i64> = vec![0, 1, 2];
let result = create(path, &metadata, &doc_ids).unwrap();
assert_eq!(result, 3);
assert_eq!(count(path).unwrap(), 3);
let with_name = get(path, Some("name IS NOT NULL"), &[], None).unwrap();
assert_eq!(with_name.len(), 2);
}
#[test]
fn test_update_with_mixed_empty_and_non_empty_metadata() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![json!({"name": "Alice"})];
let doc_ids: Vec<i64> = vec![0];
create(path, &metadata, &doc_ids).unwrap();
let new_metadata = vec![json!({})];
let new_doc_ids: Vec<i64> = vec![1];
let result = update(path, &new_metadata, &new_doc_ids).unwrap();
assert_eq!(result, 1);
assert_eq!(count(path).unwrap(), 2);
let with_name = get(path, Some("name IS NOT NULL"), &[], None).unwrap();
assert_eq!(with_name.len(), 1);
}
#[test]
fn test_read_only_helpers_work_with_query_only_connections() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata: Vec<Value> = (0..950)
.map(|i| {
json!({
"category": if i % 2 == 0 { "A" } else { "B" },
"source": format!("doc-{i}")
})
})
.collect();
let doc_ids: Vec<i64> = (0..950).collect();
create(path, &metadata, &doc_ids).unwrap();
assert_eq!(count(path).unwrap(), 950);
let mut sources = get_distinct_strings(path, "category").unwrap();
sources.sort();
assert_eq!(sources, vec!["A".to_string(), "B".to_string()]);
let filtered = where_condition(path, "category = ?", &[json!("A")]).unwrap();
assert_eq!(filtered.len(), 475);
let large_subset: Vec<i64> = (0..950).collect();
let rows = get(path, None, &[], Some(&large_subset)).unwrap();
assert_eq!(rows.len(), 950);
assert_eq!(rows[0]["_subset_"], json!(0));
assert_eq!(rows[949]["_subset_"], json!(949));
}
#[test]
fn test_update_fixed_schema_fast_path_reuses_connection() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let metadata = vec![json!({"category": "A", "source": "doc-0"})];
create(path, &metadata, &[0]).unwrap();
let new_metadata = vec![
json!({"category": "B", "source": "doc-1"}),
json!({"category": "A", "source": "doc-2"}),
];
let inserted = update(path, &new_metadata, &[1, 2]).unwrap();
assert_eq!(inserted, 2);
let rows = get(path, None, &[], Some(&[0, 1, 2])).unwrap();
assert_eq!(rows.len(), 3);
assert_eq!(rows[1]["source"], json!("doc-1"));
assert_eq!(count(path).unwrap(), 3);
}
#[test]
fn test_concurrent_metadata_reads_during_updates() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap().to_string();
let metadata: Vec<Value> = (0..20)
.map(|i| {
json!({
"category": if i % 2 == 0 { "A" } else { "B" },
"source": format!("doc-{i}")
})
})
.collect();
let doc_ids: Vec<i64> = (0..20).collect();
create(&path, &metadata, &doc_ids).unwrap();
let reader_count = 8;
let barrier = Arc::new(std::sync::Barrier::new(reader_count + 1));
std::thread::scope(|scope| {
for _ in 0..reader_count {
let path = path.clone();
let barrier = Arc::clone(&barrier);
scope.spawn(move || {
barrier.wait();
for _ in 0..100 {
let ids = where_condition(&path, "category = ?", &[json!("A")]).unwrap();
assert!(!ids.is_empty());
let subset_len = ids.len().min(3);
let rows = get(&path, None, &[], Some(&ids[..subset_len])).unwrap();
assert_eq!(rows.len(), subset_len);
assert!(count(&path).unwrap() >= 20);
}
});
}
let writer_path = path.clone();
let barrier = Arc::clone(&barrier);
scope.spawn(move || {
barrier.wait();
for i in 20..80 {
let metadata = vec![json!({
"category": if i % 2 == 0 { "A" } else { "B" },
"source": format!("doc-{i}")
})];
update(&writer_path, &metadata, &[i]).unwrap();
}
});
});
assert_eq!(count(&path).unwrap(), 80);
}
fn meta_db(path: &str) -> Connection {
Connection::open(std::path::Path::new(path).join(METADATA_DB_NAME)).unwrap()
}
fn user_version_of(path: &str) -> i64 {
meta_db(path)
.query_row("PRAGMA user_version", [], |r| r.get(0))
.unwrap()
}
fn subset_is_pk(path: &str) -> bool {
let c = meta_db(path);
let mut stmt = c.prepare("PRAGMA table_info(METADATA)").unwrap();
let rows = stmt
.query_map([], |row| {
Ok((row.get::<_, String>(1)?, row.get::<_, i64>(5)?))
})
.unwrap();
for r in rows {
let (name, pk) = r.unwrap();
if name == SUBSET_COLUMN && pk > 0 {
return true;
}
}
false
}
fn select_star_columns(path: &str) -> Vec<String> {
let rows = get(path, None, &[], None).unwrap();
rows[0].as_object().unwrap().keys().cloned().collect()
}
#[test]
fn test_create_uses_v1_layout() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let meta: Vec<serde_json::Value> = (0..5)
.map(|i| json!({"file": format!("f{i}.rs"), "code": format!("c{i}")}))
.collect();
create(path, &meta, &(0..5).collect::<Vec<i64>>()).unwrap();
assert_eq!(user_version_of(path), 1);
assert!(!subset_is_pk(path), "v1: _subset_ must not be the PK/rowid");
let cols = select_star_columns(path);
assert!(cols.iter().any(|c| c == SUBSET_COLUMN));
assert!(cols.iter().any(|c| c == "file") && cols.iter().any(|c| c == "code"));
assert!(!cols
.iter()
.any(|c| c == "rowid" || c.starts_with("_METADATA")));
}
#[test]
fn test_delete_resequences_dense() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
let meta: Vec<serde_json::Value> = (0..10)
.map(|i| json!({"file": format!("f{i}.rs")}))
.collect();
create(path, &meta, &(0..10).collect::<Vec<i64>>()).unwrap();
assert_eq!(delete(path, &[2, 5, 7]).unwrap(), 3);
let rows = get(path, None, &[], None).unwrap();
assert_eq!(rows.len(), 7);
let expected = [
"f0.rs", "f1.rs", "f3.rs", "f4.rs", "f6.rs", "f8.rs", "f9.rs",
];
for (i, row) in rows.iter().enumerate() {
assert_eq!(
row[SUBSET_COLUMN].as_i64().unwrap(),
i as i64,
"dense 0-based"
);
assert_eq!(
row["file"].as_str().unwrap(),
expected[i],
"survivor order preserved"
);
}
}
#[test]
fn test_legacy_v0_index_migrates_on_delete() {
let dir = setup_test_dir();
let path = dir.path().to_str().unwrap();
{
let c = meta_db(path);
c.execute_batch("PRAGMA user_version=0;").unwrap();
c.execute(
&format!(
"CREATE TABLE METADATA (\"{}\" INTEGER PRIMARY KEY, file TEXT, code TEXT)",
SUBSET_COLUMN
),
[],
)
.unwrap();
for i in 0..10i64 {
c.execute(
"INSERT INTO METADATA VALUES (?, ?, ?)",
rusqlite::params![i, format!("f{i}.rs"), format!("c{i}")],
)
.unwrap();
}
}
assert!(subset_is_pk(path), "precondition: legacy v0 layout");
assert_eq!(user_version_of(path), 0);
assert_eq!(delete(path, &[3]).unwrap(), 1);
assert_eq!(user_version_of(path), 1, "migrated to v1");
assert!(!subset_is_pk(path), "_subset_ demoted from PK");
let cols = select_star_columns(path);
assert!(!cols
.iter()
.any(|c| c == "rowid" || c.starts_with("_METADATA")));
assert!(cols.iter().any(|c| c == "file") && cols.iter().any(|c| c == "code"));
let rows = get(path, None, &[], None).unwrap();
assert_eq!(rows.len(), 9);
for (i, row) in rows.iter().enumerate() {
assert_eq!(row[SUBSET_COLUMN].as_i64().unwrap(), i as i64);
}
assert_eq!(rows[3]["file"].as_str().unwrap(), "f4.rs");
assert_eq!(rows[3]["code"].as_str().unwrap(), "c4");
}
}