use std::collections::BTreeMap;
use std::fs;
use std::path::{Path, PathBuf};
use anyhow::{Context, Result, anyhow, bail};
use surrealdb::Surreal;
use surrealdb::engine::any::Any;
use crate::constants::seed_dir;
use crate::core::{display, exec_surql, sha256_hex};
use crate::schema_state::{canonicalise_keys, folder_relative_key};
use crate::variables::TemplateVars;
const ENSURE_SEED_TABLE: &str = "\
DEFINE TABLE IF NOT EXISTS __seed SCHEMAFULL PERMISSIONS NONE; \
DEFINE FIELD IF NOT EXISTS key ON __seed TYPE string; \
DEFINE FIELD IF NOT EXISTS hash ON __seed TYPE string; \
DEFINE FIELD IF NOT EXISTS applied_at ON __seed TYPE datetime DEFAULT time::now(); \
DEFINE INDEX IF NOT EXISTS by_seed_key ON __seed FIELDS key UNIQUE;";
pub struct EmbeddedSeedFile {
pub path: &'static str,
pub sql: &'static str,
}
enum SeedSource<'a> {
Embedded(&'a [EmbeddedSeedFile]),
Dir(String),
}
pub struct Seed<'a> {
source: SeedSource<'a>,
vars: TemplateVars,
force: bool,
}
impl<'a> Seed<'a> {
pub fn embedded(files: &'a [EmbeddedSeedFile]) -> Self {
Self {
source: SeedSource::Embedded(files),
vars: TemplateVars::default(),
force: false,
}
}
pub fn from_dir(folder: impl Into<String>) -> Self {
Self {
source: SeedSource::Dir(folder.into()),
vars: TemplateVars::default(),
force: false,
}
}
pub fn vars(mut self, vars: TemplateVars) -> Self {
self.vars = vars;
self
}
pub fn force(mut self, force: bool) -> Self {
self.force = force;
self
}
pub async fn run(self, db: &Surreal<Any>) -> Result<()> {
let Seed {
source,
vars,
force,
} = self;
match source {
SeedSource::Embedded(files) => {
let keys: Vec<String> = files.iter().map(|f| f.path.to_string()).collect();
let tracked = migrate_legacy_seed_keys(db, &keys).await?;
let mut stats = SeedStats::default();
for f in files {
apply_seed(db, f.path, f.sql, &tracked, force, &vars, &mut stats).await?;
}
stats.report();
Ok(())
}
SeedSource::Dir(folder) => {
let dir = seed_dir(&folder);
if !dir.is_dir() {
bail!(
"no seed directory at {}.\n\
Create it and put your .surql files inside:\n\
\x20 mkdir -p {0}",
display(&dir)
);
}
run_dir(db, Some(folder.as_str()), &dir, &vars, force).await
}
}
}
}
pub async fn seed(db: &Surreal<Any>, folder: &str, vars: &TemplateVars) -> Result<()> {
Seed::from_dir(folder).vars(vars.clone()).run(db).await
}
#[doc(hidden)]
pub async fn seed_from_dir(db: &Surreal<Any>, dir: &Path, vars: &TemplateVars) -> Result<()> {
run_dir(db, None, dir, vars, false).await
}
#[derive(Default)]
struct SeedStats {
executed: usize,
skipped: usize,
}
impl SeedStats {
fn report(&self) {
log::info!("Seeded {} file(s); {} unchanged", self.executed, self.skipped);
}
}
async fn apply_seed(
db: &Surreal<Any>,
key: &str,
raw_sql: &str,
tracked: &BTreeMap<String, String>,
force: bool,
vars: &TemplateVars,
stats: &mut SeedStats,
) -> Result<()> {
let hash = sha256_hex(raw_sql.as_bytes());
if !force && tracked.get(key).is_some_and(|prev| prev == &hash) {
log::info!(" skipping {key} (unchanged)");
stats.skipped += 1;
return Ok(());
}
log::info!(" executing {key}");
let sql =
vars.apply(raw_sql).with_context(|| format!("applying template variables in {key}"))?;
exec_surql(db, &sql).await.with_context(|| format!("executing {key}"))?;
store_seed_hash(db, key, &hash).await?;
stats.executed += 1;
Ok(())
}
async fn run_dir(
db: &Surreal<Any>,
root: Option<&str>,
dir: &Path,
vars: &TemplateVars,
force: bool,
) -> Result<()> {
let mut files: Vec<PathBuf> = fs::read_dir(dir)
.with_context(|| format!("reading directory {}", display(dir)))?
.filter_map(|entry| {
let path = entry.ok()?.path();
(path.extension().and_then(|e| e.to_str()) == Some("surql")).then_some(path)
})
.collect();
if files.is_empty() {
return Err(anyhow!("no .surql files found in {}", display(dir)));
}
files.sort();
log::info!("Seeding from {} ({} files found)", display(dir), files.len());
let keys: Vec<String> = files
.iter()
.map(|path| match root {
Some(root) => folder_relative_key(root, path),
None => Ok(display(path)),
})
.collect::<Result<_>>()?;
let tracked = migrate_legacy_seed_keys(db, &keys).await?;
let mut stats = SeedStats::default();
for (path, key) in files.iter().zip(&keys) {
let raw = fs::read_to_string(path).with_context(|| format!("reading {}", display(path)))?;
apply_seed(db, key, &raw, &tracked, force, vars, &mut stats).await?;
}
stats.report();
Ok(())
}
async fn migrate_legacy_seed_keys(
db: &Surreal<Any>,
canonical: &[String],
) -> Result<BTreeMap<String, String>> {
let stored = load_seed_hashes(db).await?;
let (tracked, re_keyed) = canonicalise_keys(&stored, canonical);
if re_keyed.is_empty() {
return Ok(tracked);
}
log::info!(
"re-keyed {} tracked seed file(s) to folder-relative paths (e.g. {} -> {})",
re_keyed.len(),
re_keyed[0].0,
re_keyed[0].1
);
for (legacy, target) in &re_keyed {
let hash = tracked.get(target).cloned().unwrap_or_default();
store_seed_hash(db, target, &hash).await?;
delete_seed_hash(db, legacy).await?;
}
Ok(tracked)
}
async fn delete_seed_hash(db: &Surreal<Any>, key: &str) -> Result<()> {
db.query("DELETE __seed WHERE key = $key;").bind(("key", key.to_string())).await?.check()?;
Ok(())
}
async fn load_seed_hashes(db: &Surreal<Any>) -> Result<BTreeMap<String, String>> {
let rows: Vec<serde_json::Value> = match db.query("SELECT key, hash FROM __seed;").await {
Ok(mut resp) => resp.take(0).unwrap_or_default(),
Err(_) => Vec::new(),
};
let mut out = BTreeMap::new();
for row in rows {
let key = row.get("key").and_then(|v| v.as_str()).map(str::to_string);
let hash = row.get("hash").and_then(|v| v.as_str()).map(str::to_string);
if let (Some(key), Some(hash)) = (key, hash) {
out.insert(key, hash);
}
}
Ok(out)
}
async fn store_seed_hash(db: &Surreal<Any>, key: &str, hash: &str) -> Result<()> {
let sql = format!(
"{ENSURE_SEED_TABLE} \
DELETE __seed WHERE key = $key; \
CREATE __seed CONTENT {{ key: $key, hash: $hash, applied_at: time::now() }};",
);
db.query(sql).bind(("key", key.to_string())).bind(("hash", hash.to_string())).await?.check()?;
Ok(())
}
#[cfg(test)]
mod tests {
#[tokio::test]
async fn missing_seed_directory_names_the_path_and_the_fix() {
let tmp = tempfile::TempDir::new().expect("tmpdir");
let folder = tmp.path().to_string_lossy().to_string();
let db = surrealdb::engine::any::connect((
"mem://",
surrealdb::opt::Config::new()
.capabilities(surrealdb::opt::capabilities::Capabilities::all()),
))
.await
.expect("mem db");
db.use_ns("t").use_db("t").await.expect("use");
let err = Seed::from_dir(&folder).run(&db).await.unwrap_err().to_string();
assert!(err.contains("no seed directory"), "got: {err}");
assert!(err.contains("mkdir"), "error should say how to fix it: {err}");
}
use surrealdb::engine::any::connect;
use surrealdb::opt::Config;
use surrealdb::opt::capabilities::Capabilities;
use tempfile::TempDir;
use super::*;
use crate::variables::TemplateVars;
async fn mem_db() -> Surreal<Any> {
let config = Config::new().capabilities(Capabilities::all());
let db = connect(("mem://", config)).await.expect("connect mem://");
db.use_ns("test").use_db("seed_test").await.expect("use_ns/use_db");
db
}
#[tokio::test]
async fn seed_dir_runs_files_in_alphabetical_order() {
let tmp = TempDir::new().unwrap();
fs::write(tmp.path().join("02_b.surql"), "CREATE ordered:2 SET step = 2;").unwrap();
fs::write(tmp.path().join("01_a.surql"), "CREATE ordered:1 SET step = 1;").unwrap();
let db = mem_db().await;
seed_from_dir(&db, tmp.path(), &TemplateVars::default()).await.unwrap();
let count: Option<serde_json::Value> =
db.query("SELECT count() FROM ordered GROUP ALL").await.unwrap().take(0).unwrap();
let n = count.and_then(|v| v["count"].as_u64()).unwrap_or(0);
assert_eq!(n, 2, "both files should have been seeded");
}
#[tokio::test]
async fn seed_dir_ignores_non_surql_files() {
let tmp = TempDir::new().unwrap();
fs::write(tmp.path().join("data.surql"), "CREATE kept:1;").unwrap();
fs::write(tmp.path().join("README.md"), "# not SQL").unwrap();
fs::write(tmp.path().join("data.sql"), "CREATE ignored:1;").unwrap();
let db = mem_db().await;
seed_from_dir(&db, tmp.path(), &TemplateVars::default()).await.unwrap();
let kept: Vec<serde_json::Value> =
db.query("SELECT * FROM kept").await.unwrap().take(0).unwrap();
assert_eq!(kept.len(), 1);
let tables: Option<serde_json::Value> =
db.query("INFO FOR DB").await.unwrap().take(0).unwrap();
let table_names = tables
.as_ref()
.and_then(|v| v["tables"].as_object())
.map(|m| m.keys().cloned().collect::<Vec<_>>())
.unwrap_or_default();
assert!(!table_names.contains(&"ignored".to_string()));
}
#[tokio::test]
async fn seed_dir_errors_when_no_surql_files_present() {
let tmp = TempDir::new().unwrap();
fs::write(tmp.path().join("notes.txt"), "nothing here").unwrap();
let db = mem_db().await;
let err = seed_from_dir(&db, tmp.path(), &TemplateVars::default()).await.unwrap_err();
assert!(err.to_string().contains("no .surql files found"), "unexpected error: {err}");
}
#[tokio::test]
async fn seed_dir_error_includes_failing_file_name() {
let tmp = TempDir::new().unwrap();
fs::write(tmp.path().join("01_good.surql"), "CREATE good:1;").unwrap();
fs::write(tmp.path().join("02_bad.surql"), "THIS IS NOT VALID SURQL @@@").unwrap();
let db = mem_db().await;
let err = seed_from_dir(&db, tmp.path(), &TemplateVars::default()).await.unwrap_err();
assert!(
err.to_string().contains("02_bad.surql"),
"error should name the failing file, got: {err}"
);
}
#[tokio::test]
async fn seed_dir_handles_many_files_without_oom() {
let tmp = TempDir::new().unwrap();
let file_count = 50;
let records_per_file = 100;
for i in 0..file_count {
let sql: String = (0..records_per_file)
.map(|j| {
format!("CREATE chunk_{}:{} SET n = {};\n", i, j, i * records_per_file + j)
})
.collect();
fs::write(tmp.path().join(format!("{:03}_chunk.surql", i)), sql).unwrap();
}
let db = mem_db().await;
seed_from_dir(&db, tmp.path(), &TemplateVars::default()).await.unwrap();
let count: Option<serde_json::Value> =
db.query("SELECT count() FROM chunk_0 GROUP ALL").await.unwrap().take(0).unwrap();
let n = count.and_then(|v| v["count"].as_u64()).unwrap_or(0);
assert_eq!(n, records_per_file as u64);
}
async fn seed_count(db: &Surreal<Any>) -> u64 {
count_rows(db, "__seed").await
}
async fn count_rows(db: &Surreal<Any>, table: &str) -> u64 {
let q = format!("SELECT count() FROM {table} GROUP ALL");
let count: Option<serde_json::Value> = db.query(q).await.unwrap().take(0).unwrap();
count.and_then(|v| v["count"].as_u64()).unwrap_or(0)
}
#[tokio::test]
async fn embedded_seed_runs_once_then_skips_unchanged() {
static SEEDS: &[EmbeddedSeedFile] = &[EmbeddedSeedFile {
path: "database/seed/people.surql",
sql: "CREATE marker SET at = time::now();",
}];
let db = mem_db().await;
Seed::embedded(SEEDS).run(&db).await.unwrap();
Seed::embedded(SEEDS).run(&db).await.unwrap();
assert_eq!(count_rows(&db, "marker").await, 1, "unchanged seed should run exactly once");
assert_eq!(seed_count(&db).await, 1, "one __seed row tracked");
}
#[tokio::test]
async fn embedded_seed_reruns_when_content_changes() {
let db = mem_db().await;
static V1: &[EmbeddedSeedFile] = &[EmbeddedSeedFile {
path: "database/seed/people.surql",
sql: "CREATE marker SET at = time::now();",
}];
static V2: &[EmbeddedSeedFile] = &[EmbeddedSeedFile {
path: "database/seed/people.surql",
sql: "CREATE marker SET at = time::now(); -- v2",
}];
Seed::embedded(V1).run(&db).await.unwrap();
Seed::embedded(V2).run(&db).await.unwrap();
assert_eq!(count_rows(&db, "marker").await, 2, "changed content should re-run");
}
#[tokio::test]
async fn force_reruns_unchanged_seed() {
static SEEDS: &[EmbeddedSeedFile] = &[EmbeddedSeedFile {
path: "database/seed/people.surql",
sql: "CREATE marker SET at = time::now();",
}];
let db = mem_db().await;
Seed::embedded(SEEDS).run(&db).await.unwrap();
Seed::embedded(SEEDS).force(true).run(&db).await.unwrap();
assert_eq!(count_rows(&db, "marker").await, 2, "force should re-run even when unchanged");
}
#[tokio::test]
async fn dir_seed_is_idempotent_across_runs() {
let tmp = TempDir::new().unwrap();
fs::write(tmp.path().join("01.surql"), "CREATE once:1 SET n = 1;").unwrap();
let db = mem_db().await;
seed_from_dir(&db, tmp.path(), &TemplateVars::default()).await.unwrap();
seed_from_dir(&db, tmp.path(), &TemplateVars::default()).await.unwrap();
assert_eq!(seed_count(&db).await, 1, "one tracked seed file");
}
}