use std::collections::{HashMap, HashSet};
use std::path::PathBuf;
use std::sync::{Arc, RwLock};
use std::time::Duration;
use tokio::sync::mpsc;
use crate::error::{KnowReason, KnowledgeResult};
use crate::mem::RowData;
use orion_error::conversion::ToStructError;
pub const DEFAULT_SIGNAL_CAPACITY: usize = 1024;
#[derive(Debug, Clone)]
pub enum RefreshSource {
Authority {
root: PathBuf,
conf: PathBuf,
authority_uri: String,
table: String,
},
NamedSql {
provider: String,
sql: String,
code: String,
},
}
#[derive(Debug, Clone)]
pub struct RefreshSpec {
pub name: String,
pub interval: Duration,
pub source: RefreshSource,
}
#[derive(Debug)]
pub struct TableData {
pub name: String,
pub rows: Vec<RowData>,
}
#[derive(Default)]
pub struct TableStore {
current: RwLock<HashMap<String, Arc<TableData>>>,
}
impl TableStore {
pub fn snapshot(&self, name: &str) -> Option<Arc<TableData>> {
self.current
.read()
.expect("table store lock poisoned")
.get(name)
.cloned()
}
pub fn insert(&self, data: Arc<TableData>) {
self.current
.write()
.expect("table store lock poisoned")
.insert(data.name.clone(), data);
}
}
#[derive(Debug)]
pub struct RefreshSignal {
pub name: String,
}
pub struct RefreshService {
pub store: Arc<TableStore>,
pub signals: mpsc::Receiver<RefreshSignal>,
handles: Vec<tokio::task::AbortHandle>,
}
impl RefreshService {
pub fn spawn(specs: Vec<RefreshSpec>) -> Self {
Self::spawn_with_store(specs, Arc::new(TableStore::default()))
}
pub fn spawn_with_store(specs: Vec<RefreshSpec>, store: Arc<TableStore>) -> Self {
let (tx, signals) = mpsc::channel(DEFAULT_SIGNAL_CAPACITY);
let mut handles = Vec::with_capacity(specs.len());
let mut seen = HashSet::new();
for spec in specs {
if !seen.insert(spec.name.clone()) {
log::warn!(
"knowdb refresh: 重复规格 {} 已忽略(一表一条刷新任务)",
spec.name
);
continue;
}
let tx = tx.clone();
let store = Arc::clone(&store);
handles.push(tokio::spawn(run_spec(spec, store, tx)).abort_handle());
}
drop(tx);
Self {
store,
signals,
handles,
}
}
pub fn shutdown(&mut self) {
for h in self.handles.drain(..) {
h.abort();
}
}
}
impl Drop for RefreshService {
fn drop(&mut self) {
self.shutdown();
}
}
pub fn load_rows(spec: &RefreshSpec) -> KnowledgeResult<Vec<RowData>> {
match &spec.source {
RefreshSource::NamedSql {
provider,
sql,
code,
} => {
let sql = crate::vel::render(sql, code, crate::vel::current_wall_nanos())?;
crate::facade::query_for(provider, &sql)
}
RefreshSource::Authority {
root,
conf,
authority_uri,
table,
} => crate::loader::reload_table_rows(
root,
conf,
authority_uri,
table,
&orion_variate::EnvDict::default(),
),
}
}
async fn run_spec(spec: RefreshSpec, store: Arc<TableStore>, tx: mpsc::Sender<RefreshSignal>) {
if spec.interval.is_zero() {
log::warn!("refresh spec {:?} interval is zero; skipped", spec.name);
return;
}
let mut ticker = tokio::time::interval(spec.interval);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
ticker.tick().await;
loop {
ticker.tick().await;
let rows = match reload(&spec).await {
Ok(rows) => rows,
Err(e) => {
log::warn!("knowdb refresh {:?} reload failed: {e}", spec.name);
continue;
}
};
store.insert(Arc::new(TableData {
name: spec.name.clone(),
rows,
}));
if let Err(e) = tx.try_send(RefreshSignal {
name: spec.name.clone(),
}) {
match e {
mpsc::error::TrySendError::Full(_) => {
log::warn!(
"knowdb refresh {:?} signal dropped (channel full; store 已换代)",
spec.name
);
}
mpsc::error::TrySendError::Closed(_) => {
log::debug!("knowdb refresh {:?} receiver closed; exit", spec.name);
return;
}
}
}
}
}
async fn reload(spec: &RefreshSpec) -> KnowledgeResult<Vec<RowData>> {
match &spec.source {
RefreshSource::NamedSql {
provider,
sql,
code,
} => {
let sql = crate::vel::render(sql, code, crate::vel::current_wall_nanos())?;
crate::facade::query_async_for(provider, &sql).await
}
RefreshSource::Authority {
root,
conf,
authority_uri,
table,
} => {
let root = root.clone();
let conf = conf.clone();
let authority_uri = authority_uri.clone();
let table = table.clone();
tokio::task::spawn_blocking(move || {
crate::loader::reload_table_rows(
&root,
&conf,
&authority_uri,
&table,
&orion_variate::EnvDict::default(),
)
})
.await
.map_err(|join| {
KnowReason::from_res()
.to_err()
.with_detail(format!("refresh task join failed: {join}"))
})?
}
}
}
#[cfg(test)]
mod tests {
use std::path::PathBuf;
use std::time::Duration;
use super::*;
fn fixture_spec(table: &str, tag: &str) -> RefreshSpec {
RefreshSpec {
name: table.to_string(),
interval: Duration::from_millis(80),
source: RefreshSource::Authority {
root: PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("knowdb"),
conf: PathBuf::from("knowdb.toml"),
authority_uri: format!(
"file:{}/refresh_fixture_{}_{}_{}.sqlite",
std::env::temp_dir().display(),
table,
tag,
std::process::id()
),
table: table.to_string(),
},
}
}
async fn collect(
service: &mut RefreshService,
n: usize,
timeout: Duration,
) -> Vec<RefreshSignal> {
let mut out = Vec::new();
let deadline = tokio::time::Instant::now() + timeout;
while out.len() < n {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
break;
}
match tokio::time::timeout(remaining, service.signals.recv()).await {
Ok(Some(sig)) => out.push(sig),
_ => break,
}
}
out
}
#[test]
fn store_snapshot_missing_is_none_and_insert_returns_current_generation() {
let store = TableStore::default();
assert!(store.snapshot("nope").is_none(), "无装载 → None");
let gen1 = Arc::new(TableData {
name: "t".into(),
rows: Vec::new(),
});
store.insert(Arc::clone(&gen1));
let snap = store.snapshot("t").expect("装载后应有当前代");
assert!(Arc::ptr_eq(&snap, &gen1), "snapshot 应零复制共享同一 Arc");
}
#[test]
fn store_swap_keeps_old_generation_valid_for_holders() {
let store = TableStore::default();
let gen1 = Arc::new(TableData {
name: "t".into(),
rows: Vec::new(),
});
store.insert(Arc::clone(&gen1));
let holder = store.snapshot("t").expect("gen1");
let gen2 = Arc::new(TableData {
name: "t".into(),
rows: Vec::new(),
});
store.insert(Arc::clone(&gen2));
assert!(
Arc::ptr_eq(&store.snapshot("t").unwrap(), &gen2),
"store 已换代"
);
assert!(
Arc::ptr_eq(&holder, &gen1),
"旧 Arc 持有者仍指向完整的 gen1(不可变)"
);
}
#[tokio::test]
async fn tick_swaps_store_and_signals_without_payload() {
let mut service = RefreshService::spawn(vec![fixture_spec("address", "a")]);
let signals = collect(&mut service, 2, Duration::from_millis(1500)).await;
assert!(
signals.len() >= 2,
"期望 ≥2 次周期信号,实际 {}",
signals.len()
);
for sig in &signals {
assert_eq!(sig.name, "address");
}
let data = service
.store
.snapshot("address")
.expect("换代的表应在 store 中有当前代");
assert_eq!(data.rows.len(), 10, "address 表应重灌 10 行");
let field = &data.rows[0][0];
assert_eq!(field.get_name(), "value");
}
#[tokio::test]
async fn multiple_specs_notify_concurrently_and_independently() {
let mut service = RefreshService::spawn(vec![
fixture_spec("address", "m1"),
fixture_spec("example", "m2"),
]);
let signals = collect(&mut service, 4, Duration::from_millis(2000)).await;
let mut names: Vec<String> = signals.iter().map(|s| s.name.clone()).collect();
names.sort();
names.dedup();
assert!(
names.contains(&"address".to_string()) && names.contains(&"example".to_string()),
"两表应各自独立换代出信号: {names:?}"
);
assert!(signals.len() >= 4, "期望 ≥4 次信号,实际 {}", signals.len());
for name in ["address", "example"] {
assert!(
service.store.snapshot(name).is_some(),
"{name} 换代后 store 应有当前代"
);
}
}
#[tokio::test]
async fn zero_interval_spec_is_skipped() {
let mut spec = fixture_spec("address", "z");
spec.interval = Duration::ZERO;
let mut service = RefreshService::spawn(vec![spec]);
let signals = collect(&mut service, 1, Duration::from_millis(200)).await;
assert!(signals.is_empty(), "零周期规格不应出信号");
}
#[tokio::test]
async fn first_signal_not_before_first_interval() {
let mut spec = fixture_spec("address", "skip1");
spec.interval = Duration::from_millis(250);
let mut service = RefreshService::spawn(vec![spec]);
let early = collect(&mut service, 1, Duration::from_millis(150)).await;
assert!(early.is_empty(), "首 interval 前不应出信号");
assert!(
service.store.snapshot("address").is_none(),
"首 tick 前 store 尚无当前代(等待调用者 seed / 首 interval 换代)"
);
let signals = collect(&mut service, 1, Duration::from_millis(800)).await;
assert_eq!(signals.len(), 1, "首个 interval 后应恰好出 1 次信号");
assert_eq!(signals[0].name, "address");
assert_eq!(
service.store.snapshot("address").expect("换代").rows.len(),
10
);
}
#[tokio::test]
async fn reload_failure_is_skipped_and_shutdown_closes_channel() {
let mut spec = fixture_spec("ghost_table", "f");
spec.interval = Duration::from_millis(60);
let mut service = RefreshService::spawn(vec![spec]);
let signals = collect(&mut service, 1, Duration::from_millis(400)).await;
assert!(signals.is_empty(), "失败表不应出信号");
assert!(
service.store.snapshot("ghost_table").is_none(),
"失败表不应换代"
);
service.shutdown();
match tokio::time::timeout(Duration::from_millis(300), service.signals.recv()).await {
Ok(None) => {}
other => panic!("shutdown 后信号通道应关闭并 recv None,实际 {other:?}"),
}
}
#[tokio::test]
async fn drop_aborts_tasks_and_closes_channel() {
let mut service = RefreshService::spawn(vec![fixture_spec("address", "d")]);
let signals = collect(&mut service, 1, Duration::from_millis(1500)).await;
assert_eq!(signals.len(), 1);
drop(service);
}
#[tokio::test(flavor = "current_thread")]
async fn named_sql_spec_substitutes_code_before_each_query() {
let _guard = crate::runtime::runtime_test_guard().lock_async().await;
let db = crate::mem::memdb::MemDB::instance();
db.execute("CREATE TABLE refresh_vars_t (k TEXT, v TEXT)")
.expect("create");
db.execute("INSERT INTO refresh_vars_t VALUES ('a', '1'), ('b', '2')")
.expect("seed");
crate::facade::init_mem_provider(db).expect("init mem provider");
let mut service = RefreshService::spawn(vec![RefreshSpec {
name: "vars_t".into(),
interval: Duration::from_millis(80),
source: RefreshSource::NamedSql {
provider: "default".to_string(),
sql: "SELECT v FROM refresh_vars_t WHERE k = '$cur'".to_string(),
code: "$cur = \"b\"".to_string(),
},
}]);
let signals = collect(&mut service, 1, Duration::from_millis(1500)).await;
assert_eq!(signals.len(), 1, "应出 1 次刷新信号");
assert_eq!(signals[0].name, "vars_t");
let data = service
.store
.snapshot("vars_t")
.expect("刷新换代后应有当前代");
assert_eq!(data.rows.len(), 1, "$cur→'b' 过滤后应只回 1 行");
let field = &data.rows[0][0];
assert_eq!(field.get_name(), "v");
assert_eq!(field.to_string(), "chars(2)");
}
#[test]
fn sync_load_rows_substitutes_code_before_query() {
let _guard = crate::runtime::runtime_test_guard().lock();
let db = crate::mem::memdb::MemDB::instance();
db.execute("CREATE TABLE sync_load_t (k TEXT, v TEXT)")
.expect("create");
db.execute("INSERT INTO sync_load_t VALUES ('a', '1'), ('b', '2')")
.expect("seed");
crate::facade::init_mem_provider(db).expect("init mem provider");
let spec = RefreshSpec {
name: "sync_t".into(),
interval: Duration::from_millis(80),
source: RefreshSource::NamedSql {
provider: "default".to_string(),
sql: "SELECT v FROM sync_load_t WHERE k = '$cur'".to_string(),
code: "$cur = \"a\"".to_string(),
},
};
let rows = load_rows(&spec).expect("同步装载应成功");
assert_eq!(rows.len(), 1, "$cur→'a' 过滤后应只回 1 行");
let field = &rows[0][0];
assert_eq!(field.get_name(), "v");
}
#[test]
fn sync_load_rows_authority_reloads_typed_rows() {
let spec = fixture_spec("address", "sync_auth");
let rows = load_rows(&spec).expect("同步 Authority 装载应成功");
assert_eq!(rows.len(), 10, "address 表应重灌 10 行");
let field = &rows[0][0];
assert_eq!(field.get_name(), "value");
}
#[tokio::test]
async fn spawn_with_store_replaces_seeded_generation_on_first_tick() {
let store = Arc::new(TableStore::default());
let seed = Arc::new(TableData {
name: "address".into(),
rows: Vec::new(),
});
store.insert(Arc::clone(&seed));
let mut service = RefreshService::spawn_with_store(
vec![fixture_spec("address", "shared")],
Arc::clone(&store),
);
let signals = collect(&mut service, 1, Duration::from_millis(1500)).await;
assert_eq!(signals.len(), 1, "首 tick 后应出信号");
let cur = store.snapshot("address").expect("tick 后应有当前代");
assert_eq!(cur.rows.len(), 10, "fixture address 表 10 行");
assert!(!Arc::ptr_eq(&cur, &seed), "tick 换代应替换启动 seed 代");
assert!(seed.rows.is_empty(), "旧 seed 代不可变(持有者视角完整)");
}
#[tokio::test(flavor = "current_thread")]
async fn failed_reload_keeps_last_generation_and_no_signal() {
let _guard = crate::runtime::runtime_test_guard().lock_async().await;
let db = crate::mem::memdb::MemDB::instance();
db.execute("CREATE TABLE refresh_fail_keep_t (k TEXT)")
.expect("create");
crate::facade::init_mem_provider(db).expect("init mem provider");
let store = Arc::new(TableStore::default());
let seed = Arc::new(TableData {
name: "missing_t".into(),
rows: Vec::new(),
});
store.insert(Arc::clone(&seed));
let mut service = RefreshService::spawn_with_store(
vec![RefreshSpec {
name: "missing_t".into(),
interval: Duration::from_millis(50),
source: RefreshSource::NamedSql {
provider: "default".to_string(),
sql: "SELECT * FROM refresh_missing_xyz_t".to_string(),
code: String::new(),
},
}],
Arc::clone(&store),
);
let signals = collect(&mut service, 1, Duration::from_millis(400)).await;
assert!(signals.is_empty(), "失败表不应发信号");
let cur = store.snapshot("missing_t").expect("seed 仍在");
assert!(
Arc::ptr_eq(&cur, &seed),
"失败 tick 不应换代(保留最后一代)"
);
}
#[tokio::test]
async fn invalid_vel_code_skips_tick_and_load_rows_errors() {
let spec = RefreshSpec {
name: "velbad".into(),
interval: Duration::from_millis(40),
source: RefreshSource::NamedSql {
provider: "default".to_string(),
sql: "SELECT 1 WHERE '$bad'".to_string(),
code: "$bad = nope(1)".to_string(),
},
};
assert!(
load_rows(&spec).is_err(),
"坏 code 应在同步装载期(渲染)报错"
);
let mut service = RefreshService::spawn(vec![spec]);
let signals = collect(&mut service, 1, Duration::from_millis(300)).await;
assert!(signals.is_empty(), "渲染失败不应发信号");
assert!(
service.store.snapshot("velbad").is_none(),
"渲染失败不应换代"
);
}
#[tokio::test]
async fn duplicate_spec_name_keeps_only_first_ticker() {
let mut first_long = fixture_spec("address", "dup_long");
first_long.interval = Duration::from_secs(1);
let mut second_short = fixture_spec("address", "dup_short");
second_short.interval = Duration::from_millis(200);
let mut service = RefreshService::spawn(vec![first_long, second_short]);
let signals = collect(&mut service, 3, Duration::from_millis(700)).await;
assert!(
signals.is_empty(),
"重复同名 spec 应被去重:仅首个(1s)运行,短间隔第二个不得出信号"
);
assert!(
service.store.snapshot("address").is_none(),
"首 tick 未到不应换代"
);
let mut first_short = fixture_spec("address", "dup_short2");
first_short.interval = Duration::from_millis(200);
let mut second_long = fixture_spec("address", "dup_long2");
second_long.interval = Duration::from_secs(1);
let mut service = RefreshService::spawn(vec![first_short, second_long]);
let signals = collect(&mut service, 1, Duration::from_millis(1500)).await;
assert_eq!(signals.len(), 1, "保留首个(短间隔)应出信号");
assert_eq!(signals[0].name, "address");
assert!(
service.store.snapshot("address").is_some(),
"短间隔任务应已换代"
);
}
#[test]
fn store_concurrent_swap_and_snapshot_consistent() {
use std::sync::atomic::{AtomicBool, Ordering};
use std::thread;
let store = Arc::new(TableStore::default());
let gen_a = Arc::new(TableData {
name: "t".into(),
rows: Vec::new(),
});
let gen_b = Arc::new(TableData {
name: "t".into(),
rows: Vec::new(),
});
store.insert(Arc::clone(&gen_a));
let stop = Arc::new(AtomicBool::new(false));
let w_store = Arc::clone(&store);
let w_a = Arc::clone(&gen_a);
let w_b = Arc::clone(&gen_b);
let w_stop = Arc::clone(&stop);
let writer = thread::spawn(move || {
for i in 0..5000u32 {
w_store.insert(if i % 2 == 0 {
Arc::clone(&w_a)
} else {
Arc::clone(&w_b)
});
}
w_stop.store(true, Ordering::SeqCst);
});
let mut readers = Vec::new();
for _ in 0..4 {
let r_store = Arc::clone(&store);
let r_a = Arc::clone(&gen_a);
let r_b = Arc::clone(&gen_b);
let r_stop = Arc::clone(&stop);
readers.push(thread::spawn(move || {
while !r_stop.load(Ordering::SeqCst) {
let cur = r_store.snapshot("t").expect("首次 insert 后始终有当前代");
assert!(
Arc::ptr_eq(&cur, &r_a) || Arc::ptr_eq(&cur, &r_b),
"snapshot 应为完整一代(A/B 之一)"
);
}
}));
}
writer.join().expect("writer panicked");
for r in readers {
r.join().expect("reader panicked");
}
assert!(store.snapshot("t").is_some(), "结束后仍有当前代");
}
}