#![cfg(feature = "persistence")]
use std::collections::{BTreeSet, HashMap};
use std::sync::Arc;
use parking_lot::RwLock;
use proptest::prelude::*;
use serde_json::json;
use tempfile::TempDir;
use velesdb_core::velesql::{Condition, Parser};
use velesdb_core::{
AccessDecision, AccessScope, Database, DatabaseObserver, DistanceMetric, Point,
QueryAccessContext,
};
const CATEGORIES: [&str; 4] = ["tech", "science", "art", "food"];
struct ScopingObserver {
scope: RwLock<Option<Condition>>,
}
impl ScopingObserver {
fn new() -> Self {
Self {
scope: RwLock::new(None),
}
}
fn set_allow(&self) {
*self.scope.write() = None;
}
fn set_scope(&self, cond: Condition) {
*self.scope.write() = Some(cond);
}
}
impl DatabaseObserver for ScopingObserver {
fn on_query_request(&self, _ctx: &QueryAccessContext) -> velesdb_core::Result<AccessDecision> {
match self.scope.read().clone() {
None => Ok(AccessDecision::Allow),
Some(cond) => {
#[allow(clippy::field_reassign_with_default)]
let mut scope = AccessScope::default();
scope.filter = Some(cond);
Ok(AccessDecision::AllowWithScope(scope))
}
}
}
}
fn setup(observer: Arc<dyn DatabaseObserver>, points: Vec<Point>) -> (TempDir, Database) {
let dir = TempDir::new().expect("test: tempdir");
let db = Database::open_with_observer(dir.path(), observer).expect("test: open db");
db.create_vector_collection("items", 4, DistanceMetric::Cosine)
.expect("test: create collection");
let collection = db
.get_vector_collection("items")
.expect("test: get collection");
collection.upsert(points).expect("test: upsert items");
(dir, db)
}
fn where_fragment(kind: u8, price: i64, category: &str) -> Option<String> {
match kind {
1 => Some(format!("price > {price}")),
2 => Some(format!("price <= {price}")),
3 => Some(format!("category = '{category}'")),
4 => Some(format!("price > {price} AND category = '{category}'")),
_ => None,
}
}
fn run_ids(db: &Database, sql: &str) -> BTreeSet<u64> {
let query = Parser::parse(sql).expect("test: parse query");
db.execute_query(&query, &HashMap::new())
.expect("test: execute query")
.into_iter()
.map(|result| result.point.id)
.collect()
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(200))]
#[test]
fn allow_with_scope_and_composes_and_never_widens(
data in prop::collection::vec((0usize..4usize, 0i64..=200i64), 1..=15),
base_kind in 0u8..5,
base_price in 0i64..=200,
base_cat in 0usize..4,
scope_kind in 1u8..5,
scope_price in 0i64..=200,
scope_cat in 0usize..4,
) {
let points: Vec<Point> = data
.iter()
.enumerate()
.map(|(index, (cat_idx, price))| {
let id = u64::try_from(index).expect("test: index fits u64") + 1;
Point::new(
id,
vec![1.0, 0.0, 0.0, 0.0],
Some(json!({
"_labels": ["Item"],
"category": CATEGORIES[*cat_idx],
"price": price,
})),
)
})
.collect();
let observer = Arc::new(ScopingObserver::new());
let (_dir, db) = setup(observer.clone() as Arc<dyn DatabaseObserver>, points);
let base_frag = where_fragment(base_kind, base_price, CATEGORIES[base_cat]);
let base_sql = match &base_frag {
Some(frag) => format!("SELECT * FROM items WHERE {frag} LIMIT 1000"),
None => "SELECT * FROM items LIMIT 1000".to_string(),
};
let scope_frag = where_fragment(scope_kind, scope_price, CATEGORIES[scope_cat])
.expect("test: scope kind 1..=4 always yields a fragment");
let scope_cond = Parser::parse(&format!("SELECT * FROM items WHERE {scope_frag}"))
.expect("test: parse scope condition")
.select
.where_clause
.expect("test: scope query has a WHERE clause");
let intersect_sql = match &base_frag {
Some(frag) => {
format!("SELECT * FROM items WHERE ({frag}) AND ({scope_frag}) LIMIT 1000")
}
None => format!("SELECT * FROM items WHERE {scope_frag} LIMIT 1000"),
};
observer.set_allow();
let r = run_ids(&db, &base_sql);
observer.set_scope(scope_cond);
let r_scoped = run_ids(&db, &base_sql);
observer.set_allow();
let r_intersect = run_ids(&db, &intersect_sql);
prop_assert!(
r_scoped.is_subset(&r),
"AllowWithScope widened the result set: R'={r_scoped:?} not a subset of R={r:?}"
);
prop_assert_eq!(
&r_scoped,
&r_intersect,
"R' must equal R ∩ F (base_sql={}, scope={})",
base_sql,
scope_frag
);
}
}