1use std::collections::{HashMap, HashSet};
32use std::path::PathBuf;
33use std::sync::{Arc, RwLock};
34use std::time::Duration;
35
36use tokio::sync::mpsc;
37
38use crate::error::{KnowReason, KnowledgeResult};
39use crate::mem::RowData;
40use orion_error::conversion::ToStructError;
41
42pub const DEFAULT_SIGNAL_CAPACITY: usize = 1024;
44
45#[derive(Debug, Clone)]
47pub enum RefreshSource {
48 Authority {
51 root: PathBuf,
53 conf: PathBuf,
55 authority_uri: String,
57 table: String,
59 },
60 NamedSql {
62 provider: String,
64 sql: String,
67 code: String,
69 },
70}
71
72#[derive(Debug, Clone)]
74pub struct RefreshSpec {
75 pub name: String,
77 pub interval: Duration,
79 pub source: RefreshSource,
81}
82
83#[derive(Debug)]
88pub struct TableData {
89 pub name: String,
91 pub rows: Vec<RowData>,
93}
94
95#[derive(Default)]
102pub struct TableStore {
103 current: RwLock<HashMap<String, Arc<TableData>>>,
104}
105
106impl TableStore {
107 pub fn snapshot(&self, name: &str) -> Option<Arc<TableData>> {
109 self.current
110 .read()
111 .expect("table store lock poisoned")
112 .get(name)
113 .cloned()
114 }
115
116 pub fn insert(&self, data: Arc<TableData>) {
118 self.current
119 .write()
120 .expect("table store lock poisoned")
121 .insert(data.name.clone(), data);
122 }
123}
124
125#[derive(Debug)]
128pub struct RefreshSignal {
129 pub name: String,
131}
132
133pub struct RefreshService {
135 pub store: Arc<TableStore>,
137 pub signals: mpsc::Receiver<RefreshSignal>,
139 handles: Vec<tokio::task::AbortHandle>,
140}
141
142impl RefreshService {
143 pub fn spawn(specs: Vec<RefreshSpec>) -> Self {
147 Self::spawn_with_store(specs, Arc::new(TableStore::default()))
148 }
149
150 pub fn spawn_with_store(specs: Vec<RefreshSpec>, store: Arc<TableStore>) -> Self {
156 let (tx, signals) = mpsc::channel(DEFAULT_SIGNAL_CAPACITY);
157 let mut handles = Vec::with_capacity(specs.len());
158 let mut seen = HashSet::new();
159 for spec in specs {
160 if !seen.insert(spec.name.clone()) {
161 log::warn!(
162 "knowdb refresh: 重复规格 {} 已忽略(一表一条刷新任务)",
163 spec.name
164 );
165 continue;
166 }
167 let tx = tx.clone();
168 let store = Arc::clone(&store);
169 handles.push(tokio::spawn(run_spec(spec, store, tx)).abort_handle());
170 }
171 drop(tx);
172 Self {
173 store,
174 signals,
175 handles,
176 }
177 }
178
179 pub fn shutdown(&mut self) {
181 for h in self.handles.drain(..) {
182 h.abort();
183 }
184 }
185}
186
187impl Drop for RefreshService {
188 fn drop(&mut self) {
189 self.shutdown();
190 }
191}
192
193pub fn load_rows(spec: &RefreshSpec) -> KnowledgeResult<Vec<RowData>> {
196 match &spec.source {
197 RefreshSource::NamedSql {
198 provider,
199 sql,
200 code,
201 } => {
202 let sql = crate::vel::render(sql, code, crate::vel::current_wall_nanos())?;
204 crate::facade::query_for(provider, &sql)
205 }
206 RefreshSource::Authority {
207 root,
208 conf,
209 authority_uri,
210 table,
211 } => crate::loader::reload_table_rows(
212 root,
213 conf,
214 authority_uri,
215 table,
216 &orion_variate::EnvDict::default(),
217 ),
218 }
219}
220
221async fn run_spec(spec: RefreshSpec, store: Arc<TableStore>, tx: mpsc::Sender<RefreshSignal>) {
222 if spec.interval.is_zero() {
223 log::warn!("refresh spec {:?} interval is zero; skipped", spec.name);
224 return;
225 }
226 let mut ticker = tokio::time::interval(spec.interval);
227 ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
228 ticker.tick().await;
230 loop {
231 ticker.tick().await;
232 let rows = match reload(&spec).await {
233 Ok(rows) => rows,
234 Err(e) => {
235 log::warn!("knowdb refresh {:?} reload failed: {e}", spec.name);
236 continue;
237 }
238 };
239 store.insert(Arc::new(TableData {
241 name: spec.name.clone(),
242 rows,
243 }));
244 if let Err(e) = tx.try_send(RefreshSignal {
245 name: spec.name.clone(),
246 }) {
247 match e {
248 mpsc::error::TrySendError::Full(_) => {
249 log::warn!(
250 "knowdb refresh {:?} signal dropped (channel full; store 已换代)",
251 spec.name
252 );
253 }
254 mpsc::error::TrySendError::Closed(_) => {
255 log::debug!("knowdb refresh {:?} receiver closed; exit", spec.name);
256 return;
257 }
258 }
259 }
260 }
261}
262
263async fn reload(spec: &RefreshSpec) -> KnowledgeResult<Vec<RowData>> {
264 match &spec.source {
265 RefreshSource::NamedSql {
266 provider,
267 sql,
268 code,
269 } => {
270 let sql = crate::vel::render(sql, code, crate::vel::current_wall_nanos())?;
273 crate::facade::query_async_for(provider, &sql).await
274 }
275 RefreshSource::Authority {
276 root,
277 conf,
278 authority_uri,
279 table,
280 } => {
281 let root = root.clone();
282 let conf = conf.clone();
283 let authority_uri = authority_uri.clone();
284 let table = table.clone();
285 tokio::task::spawn_blocking(move || {
286 crate::loader::reload_table_rows(
287 &root,
288 &conf,
289 &authority_uri,
290 &table,
291 &orion_variate::EnvDict::default(),
292 )
293 })
294 .await
295 .map_err(|join| {
296 KnowReason::from_res()
297 .to_err()
298 .with_detail(format!("refresh task join failed: {join}"))
299 })?
300 }
301 }
302}
303
304#[cfg(test)]
305mod tests {
306 use std::path::PathBuf;
307 use std::time::Duration;
308
309 use super::*;
310
311 fn fixture_spec(table: &str, tag: &str) -> RefreshSpec {
312 RefreshSpec {
313 name: table.to_string(),
314 interval: Duration::from_millis(80),
315 source: RefreshSource::Authority {
316 root: PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("knowdb"),
317 conf: PathBuf::from("knowdb.toml"),
318 authority_uri: format!(
319 "file:{}/refresh_fixture_{}_{}_{}.sqlite",
320 std::env::temp_dir().display(),
321 table,
322 tag,
323 std::process::id()
324 ),
325 table: table.to_string(),
326 },
327 }
328 }
329
330 async fn collect(
331 service: &mut RefreshService,
332 n: usize,
333 timeout: Duration,
334 ) -> Vec<RefreshSignal> {
335 let mut out = Vec::new();
336 let deadline = tokio::time::Instant::now() + timeout;
337 while out.len() < n {
338 let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
339 if remaining.is_zero() {
340 break;
341 }
342 match tokio::time::timeout(remaining, service.signals.recv()).await {
343 Ok(Some(sig)) => out.push(sig),
344 _ => break,
345 }
346 }
347 out
348 }
349
350 #[test]
355 fn store_snapshot_missing_is_none_and_insert_returns_current_generation() {
356 let store = TableStore::default();
357 assert!(store.snapshot("nope").is_none(), "无装载 → None");
358
359 let gen1 = Arc::new(TableData {
360 name: "t".into(),
361 rows: Vec::new(),
362 });
363 store.insert(Arc::clone(&gen1));
364 let snap = store.snapshot("t").expect("装载后应有当前代");
365 assert!(Arc::ptr_eq(&snap, &gen1), "snapshot 应零复制共享同一 Arc");
366 }
367
368 #[test]
369 fn store_swap_keeps_old_generation_valid_for_holders() {
370 let store = TableStore::default();
372 let gen1 = Arc::new(TableData {
373 name: "t".into(),
374 rows: Vec::new(),
375 });
376 store.insert(Arc::clone(&gen1));
377 let holder = store.snapshot("t").expect("gen1");
378
379 let gen2 = Arc::new(TableData {
380 name: "t".into(),
381 rows: Vec::new(),
382 });
383 store.insert(Arc::clone(&gen2));
384 assert!(
385 Arc::ptr_eq(&store.snapshot("t").unwrap(), &gen2),
386 "store 已换代"
387 );
388 assert!(
389 Arc::ptr_eq(&holder, &gen1),
390 "旧 Arc 持有者仍指向完整的 gen1(不可变)"
391 );
392 }
393
394 #[tokio::test]
399 async fn tick_swaps_store_and_signals_without_payload() {
400 let mut service = RefreshService::spawn(vec![fixture_spec("address", "a")]);
401 let signals = collect(&mut service, 2, Duration::from_millis(1500)).await;
402 assert!(
403 signals.len() >= 2,
404 "期望 ≥2 次周期信号,实际 {}",
405 signals.len()
406 );
407 for sig in &signals {
408 assert_eq!(sig.name, "address");
409 }
410 let data = service
412 .store
413 .snapshot("address")
414 .expect("换代的表应在 store 中有当前代");
415 assert_eq!(data.rows.len(), 10, "address 表应重灌 10 行");
416 let field = &data.rows[0][0];
417 assert_eq!(field.get_name(), "value");
418 }
419
420 #[tokio::test]
421 async fn multiple_specs_notify_concurrently_and_independently() {
422 let mut service = RefreshService::spawn(vec![
423 fixture_spec("address", "m1"),
424 fixture_spec("example", "m2"),
425 ]);
426 let signals = collect(&mut service, 4, Duration::from_millis(2000)).await;
427 let mut names: Vec<String> = signals.iter().map(|s| s.name.clone()).collect();
428 names.sort();
429 names.dedup();
430 assert!(
431 names.contains(&"address".to_string()) && names.contains(&"example".to_string()),
432 "两表应各自独立换代出信号: {names:?}"
433 );
434 assert!(signals.len() >= 4, "期望 ≥4 次信号,实际 {}", signals.len());
435 for name in ["address", "example"] {
436 assert!(
437 service.store.snapshot(name).is_some(),
438 "{name} 换代后 store 应有当前代"
439 );
440 }
441 }
442
443 #[tokio::test]
444 async fn zero_interval_spec_is_skipped() {
445 let mut spec = fixture_spec("address", "z");
446 spec.interval = Duration::ZERO;
447 let mut service = RefreshService::spawn(vec![spec]);
448 let signals = collect(&mut service, 1, Duration::from_millis(200)).await;
449 assert!(signals.is_empty(), "零周期规格不应出信号");
450 }
451
452 #[tokio::test]
453 async fn first_signal_not_before_first_interval() {
454 let mut spec = fixture_spec("address", "skip1");
457 spec.interval = Duration::from_millis(250);
458 let mut service = RefreshService::spawn(vec![spec]);
459
460 let early = collect(&mut service, 1, Duration::from_millis(150)).await;
461 assert!(early.is_empty(), "首 interval 前不应出信号");
462 assert!(
463 service.store.snapshot("address").is_none(),
464 "首 tick 前 store 尚无当前代(等待调用者 seed / 首 interval 换代)"
465 );
466
467 let signals = collect(&mut service, 1, Duration::from_millis(800)).await;
468 assert_eq!(signals.len(), 1, "首个 interval 后应恰好出 1 次信号");
469 assert_eq!(signals[0].name, "address");
470 assert_eq!(
471 service.store.snapshot("address").expect("换代").rows.len(),
472 10
473 );
474 }
475
476 #[tokio::test]
477 async fn reload_failure_is_skipped_and_shutdown_closes_channel() {
478 let mut spec = fixture_spec("ghost_table", "f");
481 spec.interval = Duration::from_millis(60);
482 let mut service = RefreshService::spawn(vec![spec]);
483
484 let signals = collect(&mut service, 1, Duration::from_millis(400)).await;
485 assert!(signals.is_empty(), "失败表不应出信号");
486 assert!(
487 service.store.snapshot("ghost_table").is_none(),
488 "失败表不应换代"
489 );
490
491 service.shutdown();
492 match tokio::time::timeout(Duration::from_millis(300), service.signals.recv()).await {
493 Ok(None) => {}
494 other => panic!("shutdown 后信号通道应关闭并 recv None,实际 {other:?}"),
495 }
496 }
497
498 #[tokio::test]
499 async fn drop_aborts_tasks_and_closes_channel() {
500 let mut service = RefreshService::spawn(vec![fixture_spec("address", "d")]);
501 let signals = collect(&mut service, 1, Duration::from_millis(1500)).await;
503 assert_eq!(signals.len(), 1);
504 drop(service);
505 }
508
509 #[tokio::test(flavor = "current_thread")]
510 async fn named_sql_spec_substitutes_code_before_each_query() {
511 let _guard = crate::runtime::runtime_test_guard().lock_async().await;
512 let db = crate::mem::memdb::MemDB::instance();
514 db.execute("CREATE TABLE refresh_vars_t (k TEXT, v TEXT)")
515 .expect("create");
516 db.execute("INSERT INTO refresh_vars_t VALUES ('a', '1'), ('b', '2')")
517 .expect("seed");
518 crate::facade::init_mem_provider(db).expect("init mem provider");
519 let mut service = RefreshService::spawn(vec![RefreshSpec {
520 name: "vars_t".into(),
521 interval: Duration::from_millis(80),
522 source: RefreshSource::NamedSql {
523 provider: "default".to_string(),
524 sql: "SELECT v FROM refresh_vars_t WHERE k = '$cur'".to_string(),
525 code: "$cur = \"b\"".to_string(),
526 },
527 }]);
528 let signals = collect(&mut service, 1, Duration::from_millis(1500)).await;
529 assert_eq!(signals.len(), 1, "应出 1 次刷新信号");
530 assert_eq!(signals[0].name, "vars_t");
531 let data = service
532 .store
533 .snapshot("vars_t")
534 .expect("刷新换代后应有当前代");
535 assert_eq!(data.rows.len(), 1, "$cur→'b' 过滤后应只回 1 行");
536 let field = &data.rows[0][0];
537 assert_eq!(field.get_name(), "v");
538 assert_eq!(field.to_string(), "chars(2)");
539 }
540
541 #[test]
546 fn sync_load_rows_substitutes_code_before_query() {
547 let _guard = crate::runtime::runtime_test_guard().lock();
548 let db = crate::mem::memdb::MemDB::instance();
550 db.execute("CREATE TABLE sync_load_t (k TEXT, v TEXT)")
551 .expect("create");
552 db.execute("INSERT INTO sync_load_t VALUES ('a', '1'), ('b', '2')")
553 .expect("seed");
554 crate::facade::init_mem_provider(db).expect("init mem provider");
555 let spec = RefreshSpec {
556 name: "sync_t".into(),
557 interval: Duration::from_millis(80),
558 source: RefreshSource::NamedSql {
559 provider: "default".to_string(),
560 sql: "SELECT v FROM sync_load_t WHERE k = '$cur'".to_string(),
561 code: "$cur = \"a\"".to_string(),
562 },
563 };
564 let rows = load_rows(&spec).expect("同步装载应成功");
565 assert_eq!(rows.len(), 1, "$cur→'a' 过滤后应只回 1 行");
566 let field = &rows[0][0];
567 assert_eq!(field.get_name(), "v");
568 }
569
570 #[test]
571 fn sync_load_rows_authority_reloads_typed_rows() {
572 let spec = fixture_spec("address", "sync_auth");
575 let rows = load_rows(&spec).expect("同步 Authority 装载应成功");
576 assert_eq!(rows.len(), 10, "address 表应重灌 10 行");
577 let field = &rows[0][0];
578 assert_eq!(field.get_name(), "value");
579 }
580
581 #[tokio::test]
586 async fn spawn_with_store_replaces_seeded_generation_on_first_tick() {
587 let store = Arc::new(TableStore::default());
590 let seed = Arc::new(TableData {
591 name: "address".into(),
592 rows: Vec::new(),
593 });
594 store.insert(Arc::clone(&seed));
595 let mut service = RefreshService::spawn_with_store(
596 vec![fixture_spec("address", "shared")],
597 Arc::clone(&store),
598 );
599
600 let signals = collect(&mut service, 1, Duration::from_millis(1500)).await;
601 assert_eq!(signals.len(), 1, "首 tick 后应出信号");
602 let cur = store.snapshot("address").expect("tick 后应有当前代");
603 assert_eq!(cur.rows.len(), 10, "fixture address 表 10 行");
604 assert!(!Arc::ptr_eq(&cur, &seed), "tick 换代应替换启动 seed 代");
605 assert!(seed.rows.is_empty(), "旧 seed 代不可变(持有者视角完整)");
606 }
607
608 #[tokio::test(flavor = "current_thread")]
609 async fn failed_reload_keeps_last_generation_and_no_signal() {
610 let _guard = crate::runtime::runtime_test_guard().lock_async().await;
611 let db = crate::mem::memdb::MemDB::instance();
614 db.execute("CREATE TABLE refresh_fail_keep_t (k TEXT)")
615 .expect("create");
616 crate::facade::init_mem_provider(db).expect("init mem provider");
617
618 let store = Arc::new(TableStore::default());
619 let seed = Arc::new(TableData {
620 name: "missing_t".into(),
621 rows: Vec::new(),
622 });
623 store.insert(Arc::clone(&seed));
624 let mut service = RefreshService::spawn_with_store(
625 vec![RefreshSpec {
626 name: "missing_t".into(),
627 interval: Duration::from_millis(50),
628 source: RefreshSource::NamedSql {
629 provider: "default".to_string(),
630 sql: "SELECT * FROM refresh_missing_xyz_t".to_string(),
631 code: String::new(),
632 },
633 }],
634 Arc::clone(&store),
635 );
636
637 let signals = collect(&mut service, 1, Duration::from_millis(400)).await;
638 assert!(signals.is_empty(), "失败表不应发信号");
639 let cur = store.snapshot("missing_t").expect("seed 仍在");
640 assert!(
641 Arc::ptr_eq(&cur, &seed),
642 "失败 tick 不应换代(保留最后一代)"
643 );
644 }
645
646 #[tokio::test]
647 async fn invalid_vel_code_skips_tick_and_load_rows_errors() {
648 let spec = RefreshSpec {
651 name: "velbad".into(),
652 interval: Duration::from_millis(40),
653 source: RefreshSource::NamedSql {
654 provider: "default".to_string(),
655 sql: "SELECT 1 WHERE '$bad'".to_string(),
656 code: "$bad = nope(1)".to_string(),
657 },
658 };
659 assert!(
660 load_rows(&spec).is_err(),
661 "坏 code 应在同步装载期(渲染)报错"
662 );
663 let mut service = RefreshService::spawn(vec![spec]);
664 let signals = collect(&mut service, 1, Duration::from_millis(300)).await;
665 assert!(signals.is_empty(), "渲染失败不应发信号");
666 assert!(
667 service.store.snapshot("velbad").is_none(),
668 "渲染失败不应换代"
669 );
670 }
671
672 #[tokio::test]
673 async fn duplicate_spec_name_keeps_only_first_ticker() {
674 let mut first_long = fixture_spec("address", "dup_long");
678 first_long.interval = Duration::from_secs(1);
679 let mut second_short = fixture_spec("address", "dup_short");
680 second_short.interval = Duration::from_millis(200);
681 let mut service = RefreshService::spawn(vec![first_long, second_short]);
682
683 let signals = collect(&mut service, 3, Duration::from_millis(700)).await;
684 assert!(
685 signals.is_empty(),
686 "重复同名 spec 应被去重:仅首个(1s)运行,短间隔第二个不得出信号"
687 );
688 assert!(
689 service.store.snapshot("address").is_none(),
690 "首 tick 未到不应换代"
691 );
692
693 let mut first_short = fixture_spec("address", "dup_short2");
695 first_short.interval = Duration::from_millis(200);
696 let mut second_long = fixture_spec("address", "dup_long2");
697 second_long.interval = Duration::from_secs(1);
698 let mut service = RefreshService::spawn(vec![first_short, second_long]);
699 let signals = collect(&mut service, 1, Duration::from_millis(1500)).await;
700 assert_eq!(signals.len(), 1, "保留首个(短间隔)应出信号");
701 assert_eq!(signals[0].name, "address");
702 assert!(
703 service.store.snapshot("address").is_some(),
704 "短间隔任务应已换代"
705 );
706 }
707
708 #[test]
713 fn store_concurrent_swap_and_snapshot_consistent() {
714 use std::sync::atomic::{AtomicBool, Ordering};
715 use std::thread;
716
717 let store = Arc::new(TableStore::default());
718 let gen_a = Arc::new(TableData {
719 name: "t".into(),
720 rows: Vec::new(),
721 });
722 let gen_b = Arc::new(TableData {
723 name: "t".into(),
724 rows: Vec::new(),
725 });
726 store.insert(Arc::clone(&gen_a));
727 let stop = Arc::new(AtomicBool::new(false));
728
729 let w_store = Arc::clone(&store);
730 let w_a = Arc::clone(&gen_a);
731 let w_b = Arc::clone(&gen_b);
732 let w_stop = Arc::clone(&stop);
733 let writer = thread::spawn(move || {
734 for i in 0..5000u32 {
735 w_store.insert(if i % 2 == 0 {
736 Arc::clone(&w_a)
737 } else {
738 Arc::clone(&w_b)
739 });
740 }
741 w_stop.store(true, Ordering::SeqCst);
742 });
743
744 let mut readers = Vec::new();
745 for _ in 0..4 {
746 let r_store = Arc::clone(&store);
747 let r_a = Arc::clone(&gen_a);
748 let r_b = Arc::clone(&gen_b);
749 let r_stop = Arc::clone(&stop);
750 readers.push(thread::spawn(move || {
751 while !r_stop.load(Ordering::SeqCst) {
752 let cur = r_store.snapshot("t").expect("首次 insert 后始终有当前代");
754 assert!(
755 Arc::ptr_eq(&cur, &r_a) || Arc::ptr_eq(&cur, &r_b),
756 "snapshot 应为完整一代(A/B 之一)"
757 );
758 }
759 }));
760 }
761
762 writer.join().expect("writer panicked");
763 for r in readers {
764 r.join().expect("reader panicked");
765 }
766 assert!(store.snapshot("t").is_some(), "结束后仍有当前代");
767 }
768}