use nedb_engine::Db;
use nedb_engine::nql;
use serde_json::json;
fn tmpdir(tag: &str) -> std::path::PathBuf {
let mut p = std::env::temp_dir();
p.push(format!(
"nedb-cast-test-{}-{}",
tag,
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
std::fs::create_dir_all(&p).unwrap();
p
}
fn seeded_db(tag: &str) -> (Db, std::path::PathBuf) {
let dir = tmpdir(tag);
let db = seeded_db_unflushed_at(&dir);
db.flush_all();
(db, dir)
}
fn seeded_db_unflushed_at(dir: &std::path::Path) -> Db {
let db = Db::open(dir, None).expect("open db");
db.put("orders", "o1", json!({"total": 150, "status": "paid"}), vec![], None, None)
.expect("put o1");
db.put("orders", "o2", json!({"total": 40, "status": "pending"}), vec![], None, None)
.expect("put o2");
db.put("orders", "o3", json!({"total": 900, "status": "paid"}), vec![], None, None)
.expect("put o3");
db
}
#[test]
fn collections_visible_without_explicit_flush() {
let dir = tmpdir("noflush");
let db = seeded_db_unflushed_at(&dir);
let colls = db.id_index.collections();
assert!(
colls.contains(&"orders".to_string()),
"collection invisible before flush — /cast would 422 a collection that \
exists. got {colls:?}"
);
let (_rows, count) = nql::query(&db, "FROM orders").expect("query orders");
assert_eq!(count, 3, "expected 3 seeded orders, got {count}");
let _ = std::fs::remove_dir_all(dir);
}
#[test]
fn parse_accepts_what_execute_accepts() {
let (db, dir) = seeded_db("agree");
let good = [
"FROM orders",
"FROM orders WHERE total > 100",
"FROM orders WHERE status = \"paid\" AND total >= 100",
"FROM orders ORDER BY total DESC LIMIT 2",
"FROM orders GROUP BY status COUNT",
"FROM orders AS OF 1",
];
for q in good {
let parsed = nql::parse(q);
let executed = nql::execute(&db, q);
assert_eq!(
parsed.is_ok(),
executed.is_ok(),
"parse and execute disagree on {q:?}: parse={:?} execute={:?}",
parsed.as_ref().err().map(|e| e.to_string()),
executed.as_ref().err().map(|e| e.to_string()),
);
}
let _ = std::fs::remove_dir_all(dir);
}
#[test]
fn parse_rejects_garbage() {
for bad in ["FROM orders WHERE bogus ~~ 3", "SELECT * FROM orders", "WHERE x = 1"] {
assert!(nql::parse(bad).is_err(), "should have rejected {bad:?}");
}
}
#[test]
fn parse_has_no_side_effects() {
let (db, dir) = seeded_db("noside");
let before = db.seq.load(std::sync::atomic::Ordering::SeqCst);
let _ = nql::parse("FROM orders WHERE total > 100");
let _ = nql::parse("garbage that will not parse");
let after = db.seq.load(std::sync::atomic::Ordering::SeqCst);
assert_eq!(before, after, "parse() must not touch the database");
let _ = std::fs::remove_dir_all(dir);
}
#[cfg(feature = "cast")]
mod with_model {
use super::*;
use nedb_engine::cast::Caster;
fn require_model() -> bool {
std::env::var("CAST_REQUIRE_MODEL").map(|v| v == "1").unwrap_or(false)
}
fn caster(dir: &std::path::Path) -> Option<Caster> {
match Caster::load(dir) {
Ok(c) => Some(c),
Err(e) => {
if require_model() {
panic!("CAST_REQUIRE_MODEL=1 but no model could be loaded: {e}");
}
eprintln!("skip: {e}");
None
}
}
}
#[test]
fn model_loads_and_reports_shape() {
let dir = tmpdir("load");
let Some(c) = caster(&dir) else { return };
eprintln!(
" loaded {:.2}M params, vocab {}, from {}",
c.n_params() as f64 / 1e6,
c.vocab_size(),
c.source()
);
assert!(c.vocab_size() > 100, "implausible vocab {}", c.vocab_size());
assert!(c.n_params() > 1_000_000, "implausible params {}", c.n_params());
let _ = std::fs::remove_dir_all(dir);
}
#[test]
fn generated_nql_parses() {
let dir = tmpdir("parses");
let Some(c) = caster(&dir) else { return };
let prompts = [
"show me all orders",
"orders over 100",
"paid orders sorted by total descending",
"top 5 orders",
];
let mut invalid = Vec::new();
for p in prompts {
let nqls = c.cast(p);
eprintln!(" {p:?} -> {nqls}");
if nql::parse(&nqls).is_err() {
invalid.push(format!("{p:?} -> {nqls:?}"));
}
}
assert!(invalid.is_empty(), "unparseable output:\n {}", invalid.join("\n "));
let _ = std::fs::remove_dir_all(dir);
}
#[test]
fn cast_then_execute_returns_correct_rows() {
let (db, dir) = seeded_db("e2e");
let Some(c) = caster(&dir) else {
let _ = std::fs::remove_dir_all(dir);
return;
};
let nqls = c.cast("orders over 100");
eprintln!(" cast -> {nqls}");
let parsed = nql::parse(&nqls);
assert!(parsed.is_ok(), "generated NQL did not parse: {nqls:?}");
let (rows, count) = nql::query(&db, &nqls).expect("execute generated NQL");
eprintln!(" executed -> {count} rows");
for r in &rows {
let total = r.get("total").and_then(|v| v.as_f64()).unwrap_or(0.0);
assert!(total > 100.0, "row leaked through the filter: {r}");
}
assert_eq!(count, 2, "expected o1 and o3, got {count} rows: {rows:?}");
let _ = std::fs::remove_dir_all(dir);
}
#[test]
fn unknown_collection_is_detected_not_silently_empty() {
let (db, dir) = seeded_db("unknown");
let Some(c) = caster(&dir) else {
let _ = std::fs::remove_dir_all(dir);
return;
};
let collections = db.id_index.collections();
assert!(collections.contains(&"orders".to_string()));
let res = c.cast_checked("show me all stylists", &collections);
eprintln!(
" cast -> {} | collection={:?} known={}",
res.nql, res.collection, res.collection_known
);
if res.collection.as_deref() == Some("stylists") {
assert!(
!res.collection_known,
"stylists is not in {collections:?} but was reported as known"
);
}
let res2 = c.cast_checked("show me all orders", &collections);
if res2.collection.as_deref() == Some("orders") {
assert!(res2.collection_known, "orders IS in {collections:?}");
}
let _ = std::fs::remove_dir_all(dir);
}
#[test]
fn drift_surfaces_through_cast_checked() {
let (db, dir) = seeded_db("drift");
let Some(c) = caster(&dir) else {
let _ = std::fs::remove_dir_all(dir);
return;
};
let collections = db.id_index.collections();
let res = c.cast_checked("memories about pricing", &collections);
eprintln!(" cast -> {} | drift={:?}", res.nql, res.drift);
if res.nql.contains("SEARCH") {
let lit = res.nql.split('"').nth(1).unwrap_or("");
if !lit.is_empty() && !"memories about pricing".contains(lit) {
assert!(
res.drift.is_some(),
"emitted literal {lit:?} absent from the prompt but drift was None"
);
}
}
let clean = c.cast_checked("show me all orders", &collections);
eprintln!(" cast -> {} | drift={:?}", clean.nql, clean.drift);
assert!(clean.drift.is_none(), "false positive on {:?}", clean.nql);
let _ = std::fs::remove_dir_all(dir);
}
#[test]
fn decoding_is_deterministic() {
let dir = tmpdir("determinism");
let Some(c) = caster(&dir) else { return };
let a = c.cast("orders over 100");
let b = c.cast("orders over 100");
assert_eq!(a, b, "greedy decode is not deterministic");
let _ = std::fs::remove_dir_all(dir);
}
}