1use std::collections::{HashMap, HashSet};
2use std::fmt;
3use std::fs;
4use std::io::Read;
5use std::path::{Path, PathBuf};
6
7use orion_conf::EnvTomlLoad;
8use serde::Deserialize;
9use serde::de::{self, Deserializer, MapAccess, SeqAccess, Visitor};
10use wp_log::info_ctrl;
11
12use crate::error::{KnowReason, KnowledgeResult};
13use crate::mem::memdb::MemDB;
14use crate::mem::{DBQuery, RowData};
15use orion_error::OperationContext;
16use orion_error::conversion::{SourceErr, SourceRawErr, ToStructError};
17use orion_variate::EnvDict;
18use rusqlite::OpenFlags;
19
20#[derive(Debug, Deserialize)]
23pub struct KnowDbConf {
24 pub version: u32,
25 #[serde(default = "default_dot")]
26 pub base_dir: String,
27 #[serde(default)]
28 pub default: OptLoadSpec,
29 #[serde(default)]
30 pub csv: CsvSpec,
31 #[serde(default)]
32 pub cache: CacheSpec,
33 #[serde(default)]
34 pub tables: Vec<TableSpec>,
35
36 #[serde(default)]
38 pub fun: HashMap<String, FunSpec>,
39
40 #[serde(default, rename = "provider")]
42 provider_raw: Option<ProviderConfig>,
43
44 #[serde(default)]
46 pub intranet_nets: Option<crate::intranet_nets::IntranetNetsConf>,
47}
48
49impl KnowDbConf {
50 pub fn provider(&self) -> Option<ProviderConfig> {
51 self.provider_raw.clone()
52 }
53}
54
55#[derive(Debug, Clone, Deserialize)]
60pub struct FunSpec {
61 pub call: FunCall,
62 #[serde(default)]
63 pub key: Option<String>,
64 #[serde(default = "default_true")]
65 pub cache: bool,
66 #[serde(default)]
67 pub ttl_ms: Option<u64>,
68}
69
70impl FunSpec {
71 pub fn returns_bool(&self) -> bool {
73 matches!(self.call, FunCall::BfExists | FunCall::Sismember)
74 }
75}
76
77#[derive(Debug, Clone, Deserialize, PartialEq)]
78#[serde(rename_all = "snake_case")]
79pub enum FunCall {
80 BfExists,
81 Sismember,
82 Hget,
83 Get,
84}
85
86#[derive(Debug, Clone, Deserialize)]
91pub struct CacheSpec {
92 #[serde(default = "default_true")]
93 pub enabled: bool,
94 #[serde(default = "default_result_cache_capacity")]
95 pub capacity: usize,
96 #[serde(default = "default_result_cache_ttl_ms")]
97 pub ttl_ms: u64,
98}
99
100impl Default for CacheSpec {
101 fn default() -> Self {
102 Self {
103 enabled: default_true(),
104 capacity: default_result_cache_capacity(),
105 ttl_ms: default_result_cache_ttl_ms(),
106 }
107 }
108}
109
110pub const DEFAULT_SQLDB_NAME: &str = "default";
116
117#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Deserialize)]
119#[serde(rename_all = "snake_case")]
120pub enum PlanCacheMode {
121 #[default]
122 Auto,
123 ForceGenericPlan,
124 ForceCustomPlan,
125}
126
127impl PlanCacheMode {
128 pub fn to_sql_value(self) -> &'static str {
130 match self {
131 PlanCacheMode::Auto => "auto",
132 PlanCacheMode::ForceGenericPlan => "force_generic_plan",
133 PlanCacheMode::ForceCustomPlan => "force_custom_plan",
134 }
135 }
136}
137
138#[derive(Debug, Clone, Default, PartialEq, Deserialize)]
145#[serde(deny_unknown_fields)]
146pub struct PostgresSessionSpec {
147 #[serde(default)]
148 pub plan_cache_mode: Option<PlanCacheMode>,
149 #[serde(default)]
151 pub jit: Option<bool>,
152 #[serde(default)]
154 pub application_name: Option<String>,
155}
156
157#[derive(Debug, Clone, Default, Deserialize)]
158pub struct ProviderConfig {
159 #[serde(default, deserialize_with = "deserialize_sqldb")]
161 pub sqldb: Option<Vec<SqlProviderSpec>>,
162 #[serde(default)]
163 pub redis: Option<RedisProviderSpec>,
164}
165
166impl ProviderConfig {
167 pub fn sqldb_names(&self) -> Vec<String> {
169 self.sqldb
170 .as_ref()
171 .map(|specs| {
172 specs
173 .iter()
174 .map(|spec| spec.effective_name().to_string())
175 .collect()
176 })
177 .unwrap_or_default()
178 }
179}
180
181#[derive(Debug, Clone, Deserialize)]
182pub struct SqlProviderSpec {
183 #[serde(default)]
184 pub name: Option<String>,
185 #[serde(rename = "kind")]
186 pub kind: SqlProviderKind,
187 pub connection_uri: String,
188 #[serde(default)]
189 pub pool_size: Option<u32>,
190 #[serde(default)]
191 pub min_connections: Option<u32>,
192 #[serde(default)]
193 pub acquire_timeout_ms: Option<u64>,
194 #[serde(default)]
195 pub idle_timeout_ms: Option<u64>,
196 #[serde(default)]
197 pub max_lifetime_ms: Option<u64>,
198 #[serde(default)]
200 pub postgres_session: Option<PostgresSessionSpec>,
201}
202
203impl SqlProviderSpec {
204 pub fn effective_name(&self) -> &str {
206 self.name.as_deref().unwrap_or(DEFAULT_SQLDB_NAME)
207 }
208
209 pub fn validate_specs(specs: &[SqlProviderSpec]) -> KnowledgeResult<HashSet<String>> {
211 let mut seen = HashSet::with_capacity(specs.len());
212 for spec in specs {
213 let name = spec.effective_name();
214 if !name.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
215 return Err(KnowReason::from_conf().to_err().with_detail(format!(
216 "invalid sqldb provider name '{name}' (allowed [A-Za-z0-9_])"
217 )));
218 }
219 if !seen.insert(name.to_string()) {
220 return Err(KnowReason::from_conf()
221 .to_err()
222 .with_detail(format!("duplicate sqldb provider name '{name}'")));
223 }
224 if let Some(session) = &spec.postgres_session {
225 if !matches!(spec.kind, SqlProviderKind::Postgres) {
226 return Err(KnowReason::from_conf().to_err().with_detail(format!(
227 "provider '{name}': postgres_session is only valid for kind = \"postgres\""
228 )));
229 }
230 if let Some(app) = &session.application_name
231 && (app.len() > 63 || app.chars().any(|c| c.is_control()))
232 {
233 return Err(KnowReason::from_conf().to_err().with_detail(format!(
234 "provider '{name}': postgres_session.application_name invalid \
235 (≤ 63 bytes, no control characters)"
236 )));
237 }
238 }
239 }
240 Ok(seen)
241 }
242}
243
244fn deserialize_sqldb<'de, D>(deserializer: D) -> Result<Option<Vec<SqlProviderSpec>>, D::Error>
246where
247 D: Deserializer<'de>,
248{
249 struct SqlDbVisitor;
250
251 impl<'de> Visitor<'de> for SqlDbVisitor {
252 type Value = Option<Vec<SqlProviderSpec>>;
253
254 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
255 formatter
256 .write_str("a `[provider.sqldb]` table or a `[[provider.sqldb]]` array of tables")
257 }
258
259 fn visit_map<A>(self, map: A) -> Result<Self::Value, A::Error>
260 where
261 A: MapAccess<'de>,
262 {
263 let spec = SqlProviderSpec::deserialize(de::value::MapAccessDeserializer::new(map))?;
264 Ok(Some(vec![spec]))
265 }
266
267 fn visit_seq<A>(self, seq: A) -> Result<Self::Value, A::Error>
268 where
269 A: SeqAccess<'de>,
270 {
271 let specs =
272 Vec::<SqlProviderSpec>::deserialize(de::value::SeqAccessDeserializer::new(seq))?;
273 Ok(Some(specs))
274 }
275 }
276
277 deserializer.deserialize_any(SqlDbVisitor)
278}
279
280#[derive(Debug, Clone, Deserialize)]
284#[serde(rename_all = "snake_case")]
285pub enum SqlProviderKind {
286 Postgres,
287 Mysql,
288}
289
290#[derive(Debug, Clone, Deserialize)]
291pub struct RedisProviderSpec {
292 pub connection_uri: String,
293 #[serde(default)]
294 pub pool_size: Option<usize>,
295 #[serde(default = "default_connect_timeout_ms")]
296 pub connect_timeout_ms: u64,
297 #[serde(default = "default_command_timeout_ms")]
298 pub command_timeout_ms: u64,
299}
300
301fn default_connect_timeout_ms() -> u64 {
302 3_000
303}
304
305fn default_command_timeout_ms() -> u64 {
306 100
307}
308
309#[derive(Debug, Clone, Deserialize)]
315#[serde(rename_all = "snake_case")]
316pub enum ProviderKind {
317 SqliteAuthority,
318 Postgres,
319 Mysql,
320 Redis,
321}
322
323#[derive(Debug, Clone, Deserialize)]
324pub struct OptLoadSpec {
325 #[serde(default = "default_true")]
326 pub transaction: bool,
327 #[serde(default = "default_batch")]
328 pub batch_size: usize,
329 #[serde(default = "default_on_error")]
330 pub on_error: OnError,
331}
332impl Default for OptLoadSpec {
333 fn default() -> Self {
334 Self {
335 transaction: true,
336 batch_size: default_batch(),
337 on_error: default_on_error(),
338 }
339 }
340}
341
342#[derive(Debug, Clone, Deserialize, Default)]
343#[serde(rename_all = "lowercase")]
344pub enum OnError {
345 #[default]
346 Fail,
347 Skip,
348}
349
350#[derive(Debug, Clone, Deserialize)]
351pub struct CsvSpec {
352 #[serde(default = "default_true")]
353 pub has_header: bool,
354 #[serde(default = "default_comma")]
355 pub delimiter: String,
356 #[serde(default = "default_utf8")]
357 pub encoding: String,
358 #[serde(default = "default_true")]
359 pub trim: bool,
360}
361impl Default for CsvSpec {
362 fn default() -> Self {
363 CsvSpec {
364 has_header: true,
365 delimiter: ",".into(),
366 encoding: "utf-8".into(),
367 trim: true,
368 }
369 }
370}
371
372#[derive(Debug, Clone, Deserialize)]
373pub struct TableSpec {
374 pub name: String,
375 #[serde(default)]
376 pub dir: Option<String>,
377 #[serde(default)]
378 pub data_file: Option<String>,
379 pub columns: ColumnsSpec,
380 #[serde(default)]
381 pub expected_rows: RowExpect,
382 #[serde(default = "default_true")]
383 pub enabled: bool,
384}
385
386#[derive(Debug, Clone, Deserialize)]
387pub struct ColumnsSpec {
388 #[serde(default)]
389 pub by_header: Vec<String>,
390 #[serde(default)]
391 pub by_index: Vec<usize>,
392}
393
394#[derive(Debug, Clone, Deserialize, Default)]
395pub struct RowExpect {
396 pub min: Option<usize>,
397 pub max: Option<usize>,
398}
399
400const fn default_true() -> bool {
401 true
402}
403const fn default_batch() -> usize {
404 2000
405}
406fn default_comma() -> String {
407 ",".to_string()
408}
409fn default_utf8() -> String {
410 "utf-8".to_string()
411}
412fn default_on_error() -> OnError {
413 OnError::Fail
414}
415fn default_dot() -> String {
416 ".".to_string()
417}
418const fn default_result_cache_capacity() -> usize {
419 1024
420}
421const fn default_result_cache_ttl_ms() -> u64 {
422 30_000
423}
424
425fn read_to_string(path: &Path) -> KnowledgeResult<String> {
427 let mut f = fs::File::open(path).source_raw_err(KnowReason::from_res(), "source error")?;
428 let mut buf = String::new();
429 f.read_to_string(&mut buf)
430 .source_raw_err(KnowReason::from_res(), "source error")?;
431 Ok(buf)
432}
433
434fn replace_table(sql: &str, table: &str) -> String {
435 sql.replace("{table}", table)
436}
437
438fn join_rel(base: &Path, rel: &str) -> PathBuf {
439 let p = Path::new(rel);
440 if p.is_absolute() {
441 p.to_path_buf()
442 } else {
443 base.join(p)
444 }
445}
446
447pub fn build_authority_from_knowdb(
448 root: &Path,
449 conf_path: &Path,
450 authority_uri: &str,
451 dict: &EnvDict,
452) -> KnowledgeResult<Vec<String>> {
453 let mut opx = OperationContext::doing("build authority from knowdb").with_auto_log();
454 let (conf, conf_abs, base_dir) = parse_knowdb_conf(root, conf_path, dict)?;
456 opx.record("conf", conf_abs.display());
457 opx.record("base_dir", base_dir.display());
458 let db = open_authority(authority_uri)?;
460 let mut loaded_names = Vec::new();
462 for t in &conf.tables {
463 if !t.enabled {
464 continue;
465 }
466 load_one_table(&db, &base_dir, t, &conf.csv, &conf.default)?;
467 info_ctrl!("load table {} suc!", base_dir.display(),);
468 loaded_names.push(t.name.clone());
469 }
470 opx.mark_suc();
471 Ok(loaded_names)
472}
473
474pub fn reload_table_rows(
485 root: &Path,
486 conf_path: &Path,
487 authority_uri: &str,
488 table: &str,
489 dict: &EnvDict,
490) -> KnowledgeResult<Vec<RowData>> {
491 let (conf, _conf_abs, base_dir) = parse_knowdb_conf(root, conf_path, dict)?;
492 let spec = conf
493 .tables
494 .iter()
495 .find(|t| t.enabled && t.name == table)
496 .ok_or_else(|| {
497 KnowReason::from_conf()
498 .to_err()
499 .with_detail(format!("refresh table {table:?} not found or disabled"))
500 })?;
501 if spec.columns.by_header.is_empty() {
502 return Err(KnowReason::from_conf().to_err().with_detail(format!(
503 "refresh table {table:?} requires columns.by_header for row projection"
504 )));
505 }
506 let db = open_authority(authority_uri)?;
507 load_one_table(&db, &base_dir, spec, &conf.csv, &conf.default)?;
508 let cols = spec.columns.by_header.join(", ");
509 let sql = format!("SELECT {cols} FROM {}", spec.name);
510 db.query(&sql)
511}
512
513pub fn parse_knowdb_conf(
514 root: &Path,
515 conf_path: &Path,
516 dict: &EnvDict,
517) -> KnowledgeResult<(KnowDbConf, PathBuf, PathBuf)> {
518 let conf_abs = if conf_path.is_absolute() {
519 conf_path.to_path_buf()
520 } else {
521 root.join(conf_path)
522 };
523 let conf_txt = read_to_string(&conf_abs)?;
524 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(&conf_txt, dict)
525 .source_err(KnowReason::from_conf(), "parse knowdb config")?;
526 if conf.version != 2 {
527 return Err(KnowReason::from_conf()
528 .to_err()
529 .with_detail("unsupported knowdb.version"));
530 }
531 crate::intranet_nets::set_intranet_nets_conf(conf.intranet_nets.clone());
533 let conf_dir = conf_abs.parent().unwrap_or_else(|| Path::new("."));
534 let base_dir = join_rel(conf_dir, &conf.base_dir);
535 Ok((conf, conf_abs, base_dir))
536}
537
538fn open_authority(authority_uri: &str) -> KnowledgeResult<MemDB> {
539 ensure_parent_dir_for_file_uri(authority_uri);
540 let flags = OpenFlags::SQLITE_OPEN_READ_WRITE
541 | OpenFlags::SQLITE_OPEN_CREATE
542 | OpenFlags::SQLITE_OPEN_URI;
543 let db = MemDB::new_file(authority_uri, 1, flags)?;
544 let _ = db.with_conn(|conn| {
546 let _ = crate::sqlite_ext::register_builtin(conn);
547 Ok::<(), anyhow::Error>(())
548 });
549 Ok(db)
550}
551
552fn ensure_parent_dir_for_file_uri(uri: &str) {
555 if let Some(rest) = uri.strip_prefix("file:") {
556 let path_part = rest.split('?').next().unwrap_or(rest);
557 let p = Path::new(path_part);
558 if let Some(parent) = p.parent() {
559 let _ = fs::create_dir_all(parent);
560 }
561 }
562}
563
564fn load_one_table(
565 db: &MemDB,
566 base_dir: &Path,
567 t: &TableSpec,
568 csvd: &CsvSpec,
569 load: &OptLoadSpec,
570) -> KnowledgeResult<()> {
571 let mut opx = OperationContext::doing("load table to kdb")
573 .with_auto_log()
574 .with_mod_path("ctrl");
575 let dir_name: &str = t.dir.as_deref().unwrap_or(&t.name);
576 let table_dir = base_dir.join(dir_name);
577 opx.record("table_dir", table_dir.display());
578 let create_sql = replace_table(&read_to_string(&table_dir.join("create.sql"))?, &t.name);
579 let insert_sql = replace_table(&read_to_string(&table_dir.join("insert.sql"))?, &t.name);
580 let clean_path = table_dir.join("clean.sql");
581 let clean_sql = if clean_path.exists() {
582 replace_table(&read_to_string(&clean_path)?, &t.name)
583 } else {
584 format!("DELETE FROM {}", t.name)
585 };
586
587 db.with_conn(|conn| {
589 let _ = crate::sqlite_ext::register_builtin(conn);
591 conn.execute_batch(&create_sql)?;
592 conn.execute_batch(&clean_sql)?;
593 Ok::<(), anyhow::Error>(())
594 })
595 .source_err(KnowReason::from_res(), "prepare authority table")?;
596
597 let data_path = match &t.data_file {
599 Some(rel) => join_rel(&table_dir, rel),
600 None => table_dir.join("data.csv"),
601 };
602 if !data_path.exists() {
603 return Err(KnowReason::from_conf()
604 .to_err()
605 .with_detail("data.csv not found"));
606 }
607 opx.record("data_path", data_path.display());
608
609 let mut rdr = build_csv_reader(csvd, &data_path)?;
611
612 let col_indices: Vec<usize> = if !t.columns.by_header.is_empty() {
614 let headers = rdr
615 .headers()
616 .source_raw_err(KnowReason::from_res(), "source error")?;
617 select_indices_by_header(headers, &t.columns.by_header)?
618 } else if !t.columns.by_index.is_empty() {
619 t.columns.by_index.clone()
620 } else {
621 return Err(KnowReason::from_conf()
622 .to_err()
623 .with_detail("columns mapping required"));
624 };
625
626 let mut inserted: usize = 0;
628 let mut bad: usize = 0;
629 let mut batch_left = load.batch_size.max(1);
630 db.with_conn(|conn| {
631 let _ = crate::sqlite_ext::register_builtin(conn);
633 let mut tx = if load.transaction {
634 Some(conn.unchecked_transaction()?)
635 } else {
636 None
637 };
638 let mut stmt = conn.prepare(&insert_sql)?;
639 for rec in rdr.into_records() {
640 match rec {
641 Ok(record) => {
642 let refs = extract_row_refs(&record, &col_indices, &mut bad, load)?;
643 if let Some(refs) = refs {
644 stmt.execute(rusqlite::params_from_iter(refs))?;
645 inserted += 1;
646 if load.transaction {
647 batch_left -= 1;
648 if batch_left == 0 {
649 tx.take().unwrap().commit()?;
650 tx = Some(conn.unchecked_transaction()?);
651 batch_left = load.batch_size.max(1);
652 }
653 }
654 }
655 }
656 Err(_e) => {
657 if matches!(load.on_error, OnError::Skip) {
658 bad += 1;
659 continue;
660 } else {
661 anyhow::bail!("csv record parse error");
662 }
663 }
664 }
665 }
666 if let Some(tx) = tx {
667 tx.commit()?;
668 }
669 Ok::<(), anyhow::Error>(())
670 })
671 .source_err(KnowReason::from_res(), "load authority table data")?;
672
673 if let Some(min) = t.expected_rows.min
675 && inserted < min
676 {
677 return Err(KnowReason::from_conf()
678 .to_err()
679 .with_detail("table data less"));
680 }
681 if let Some(max) = t.expected_rows.max
682 && inserted > max
683 {
684 wp_log::warn_kdb!(
685 "table {} loaded rows {} exceed max {}",
686 &t.name,
687 inserted,
688 max
689 );
690 }
691 if bad > 0 {
692 wp_log::warn_kdb!("table {} skipped {} bad rows (on_error=skip)", &t.name, bad);
693 }
694 opx.mark_suc();
695 Ok(())
696}
697
698fn build_csv_reader(
699 csvd: &CsvSpec,
700 data_path: &Path,
701) -> KnowledgeResult<csv::Reader<std::fs::File>> {
702 if csvd.encoding.to_lowercase() != "utf-8" {
703 return Err(KnowReason::from_conf()
704 .to_err()
705 .with_detail("only utf-8 csv is supported"));
706 }
707 let mut rdr_b = csv::ReaderBuilder::new();
708 rdr_b.has_headers(csvd.has_header);
709 if csvd.delimiter.len() == 1 {
710 rdr_b.delimiter(csvd.delimiter.as_bytes()[0]);
711 }
712 if csvd.trim {
713 rdr_b.trim(csv::Trim::All);
714 }
715 rdr_b
716 .from_path(data_path)
717 .source_raw_err(KnowReason::from_res(), "source error")
718}
719
720fn select_indices_by_header(
721 headers: &csv::StringRecord,
722 wanted: &[String],
723) -> KnowledgeResult<Vec<usize>> {
724 let mut out = Vec::with_capacity(wanted.len());
725 for name in wanted {
726 let pos = headers.iter().position(|h| h == name).ok_or_else(|| {
727 KnowReason::from_conf()
728 .to_err()
729 .with_detail("header not found")
730 })?;
731 out.push(pos);
732 }
733 Ok(out)
734}
735
736fn extract_row_refs<'a>(
737 record: &'a csv::StringRecord,
738 col_indices: &[usize],
739 bad: &mut usize,
740 load: &OptLoadSpec,
741) -> anyhow::Result<Option<Vec<&'a str>>> {
742 let mut vs: Vec<&str> = Vec::with_capacity(col_indices.len());
743 for &idx in col_indices {
744 if idx >= record.len() {
745 if matches!(load.on_error, OnError::Skip) {
746 *bad += 1;
747 return Ok(None);
748 } else {
749 anyhow::bail!("missing column at index {}", idx);
750 }
751 }
752 vs.push(record.get(idx).unwrap_or(""));
753 }
754 Ok(Some(vs))
755}
756
757#[cfg(test)]
758mod tests {
759 use super::*;
760
761 #[test]
762 fn parse_new_style_sqldb_provider() {
763 let dict = EnvDict::default();
764 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
765 r#"
766version = 2
767
768[provider.sqldb]
769kind = "postgres"
770connection_uri = "postgres://demo:demo@127.0.0.1/demo"
771pool_size = 12
772"#,
773 &dict,
774 )
775 .expect("parse knowdb with sqldb provider");
776
777 let sqldb = conf
778 .provider()
779 .expect("provider")
780 .sqldb
781 .expect("sqldb provider");
782 assert_eq!(sqldb.len(), 1);
783 let sqldb = &sqldb[0];
784 assert!(matches!(sqldb.kind, SqlProviderKind::Postgres));
785 assert_eq!(sqldb.pool_size, Some(12));
786 assert_eq!(sqldb.effective_name(), DEFAULT_SQLDB_NAME);
787 }
788
789 #[test]
790 fn parse_new_style_redis_provider() {
791 let dict = EnvDict::default();
792 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
793 r#"
794version = 2
795
796[provider.redis]
797connection_uri = "redis://127.0.0.1:6379"
798pool_size = 16
799connect_timeout_ms = 5000
800command_timeout_ms = 200
801"#,
802 &dict,
803 )
804 .expect("parse knowdb with redis provider");
805
806 let redis_cfg = conf
807 .provider()
808 .expect("provider")
809 .redis
810 .expect("redis provider");
811 assert_eq!(redis_cfg.connection_uri, "redis://127.0.0.1:6379");
812 assert_eq!(redis_cfg.pool_size, Some(16));
813 assert_eq!(redis_cfg.connect_timeout_ms, 5000);
814 assert_eq!(redis_cfg.command_timeout_ms, 200);
815 }
816
817 #[test]
818 fn parse_redis_provider_with_default_timeouts() {
819 let dict = EnvDict::default();
820 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
821 r#"
822version = 2
823
824[provider.redis]
825connection_uri = "redis://127.0.0.1:6379"
826"#,
827 &dict,
828 )
829 .expect("parse knowdb with redis provider (no timeout fields)");
830
831 let redis_cfg = conf.provider().expect("provider").redis.expect("redis");
832 assert_eq!(redis_cfg.connect_timeout_ms, 3000);
833 assert_eq!(redis_cfg.command_timeout_ms, 100);
834 }
835
836 #[test]
837 fn parse_both_sqldb_and_redis_providers() {
838 let dict = EnvDict::default();
839 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
840 r#"
841version = 2
842
843[provider.sqldb]
844kind = "postgres"
845connection_uri = "postgres://demo:demo@127.0.0.1/demo"
846
847[provider.redis]
848connection_uri = "redis://10.0.0.1:6379"
849pool_size = 4
850"#,
851 &dict,
852 )
853 .expect("parse knowdb with both sqldb and redis");
854
855 let provider_cfg = conf.provider().expect("provider");
856 let sqldb = &provider_cfg.sqldb.expect("sqldb")[0];
857 let redis_cfg = provider_cfg.redis.expect("redis");
858 assert!(matches!(sqldb.kind, SqlProviderKind::Postgres));
859 assert_eq!(redis_cfg.connection_uri, "redis://10.0.0.1:6379");
860 assert_eq!(redis_cfg.pool_size, Some(4));
861 }
862
863 #[test]
864 fn parse_redis_only_without_sqldb() {
865 let dict = EnvDict::default();
866 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
867 r#"
868version = 2
869
870[provider.redis]
871connection_uri = "redis://127.0.0.1:6379"
872"#,
873 &dict,
874 )
875 .expect("parse knowdb with redis only");
876
877 let provider_cfg = conf.provider().expect("provider");
878 assert!(provider_cfg.sqldb.is_none());
879 assert!(provider_cfg.redis.is_some());
880 }
881
882 #[test]
883 fn parse_no_provider_section() {
884 let dict = EnvDict::default();
885 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
886 r#"
887version = 2
888"#,
889 &dict,
890 )
891 .expect("parse knowdb without provider");
892
893 assert!(conf.provider().is_none());
894 }
895
896 #[test]
897 fn new_style_sqldb_mysql_variant() {
898 let dict = EnvDict::default();
899 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
900 r#"
901version = 2
902
903[provider.sqldb]
904kind = "mysql"
905connection_uri = "mysql://user:pass@127.0.0.1:3306/db"
906pool_size = 8
907"#,
908 &dict,
909 )
910 .expect("parse new-style mysql sqldb");
911
912 let sqldb = &conf.provider().expect("provider").sqldb.expect("sqldb")[0];
913 assert!(matches!(sqldb.kind, SqlProviderKind::Mysql));
914 assert_eq!(sqldb.pool_size, Some(8));
915 }
916
917 #[test]
918 fn parse_array_style_multi_sqldb_providers() {
919 let dict = EnvDict::default();
920 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
921 r#"
922version = 2
923
924[[provider.sqldb]]
925name = "geo"
926kind = "postgres"
927connection_uri = "postgres://demo@127.0.0.1:5432/geo_db"
928pool_size = 8
929
930[[provider.sqldb]]
931name = "asset"
932kind = "postgres"
933connection_uri = "postgres://demo@127.0.0.1:5432/asset_db"
934pool_size = 12
935"#,
936 &dict,
937 )
938 .expect("parse knowdb with multiple sqldb providers");
939
940 let specs = conf
941 .provider()
942 .expect("provider")
943 .sqldb
944 .expect("sqldb providers");
945 assert_eq!(specs.len(), 2);
946 assert_eq!(specs[0].name.as_deref(), Some("geo"));
947 assert_eq!(specs[0].effective_name(), "geo");
948 assert_eq!(
949 specs[0].connection_uri,
950 "postgres://demo@127.0.0.1:5432/geo_db"
951 );
952 assert_eq!(specs[0].pool_size, Some(8));
953 assert_eq!(specs[1].name.as_deref(), Some("asset"));
954 assert_eq!(specs[1].effective_name(), "asset");
955 assert_eq!(
956 specs[1].connection_uri,
957 "postgres://demo@127.0.0.1:5432/asset_db"
958 );
959 assert_eq!(specs[1].pool_size, Some(12));
960 }
961
962 #[test]
963 fn parse_single_sqldb_with_explicit_name() {
964 let dict = EnvDict::default();
965 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
966 r#"
967version = 2
968
969[provider.sqldb]
970name = "geo"
971kind = "postgres"
972connection_uri = "postgres://demo@127.0.0.1/geo_db"
973"#,
974 &dict,
975 )
976 .expect("parse single named sqldb");
977
978 let spec = &conf.provider().expect("provider").sqldb.expect("sqldb")[0];
979 assert_eq!(spec.effective_name(), "geo");
980 }
981
982 #[test]
983 fn sqldb_names_applies_default_name() {
984 let dict = EnvDict::default();
985 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
986 r#"
987version = 2
988
989[provider.sqldb]
990kind = "postgres"
991connection_uri = "postgres://demo@127.0.0.1/demo"
992"#,
993 &dict,
994 )
995 .expect("parse unnamed sqldb");
996
997 let names = conf.provider().expect("provider").sqldb_names();
998 assert_eq!(names, vec!["default".to_string()]);
999 }
1000
1001 #[test]
1002 fn validate_specs_rejects_duplicate_effective_names() {
1003 let dict = EnvDict::default();
1004 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1005 r#"
1006version = 2
1007
1008[[provider.sqldb]]
1009kind = "postgres"
1010connection_uri = "postgres://demo@127.0.0.1/db1"
1011
1012[[provider.sqldb]]
1013kind = "postgres"
1014connection_uri = "postgres://demo@127.0.0.1/db2"
1015"#,
1016 &dict,
1017 )
1018 .expect("parse two unnamed sqldb");
1019
1020 let specs = conf.provider().expect("provider").sqldb.expect("sqldb");
1021 let err = SqlProviderSpec::validate_specs(&specs).expect_err("duplicate 'default' name");
1022 assert!(
1023 err.to_string()
1024 .contains("duplicate sqldb provider name 'default'")
1025 );
1026 }
1027
1028 #[test]
1029 fn validate_specs_rejects_invalid_name_charset() {
1030 let specs = [SqlProviderSpec {
1031 name: Some("bad name".to_string()),
1032 kind: SqlProviderKind::Postgres,
1033 connection_uri: "postgres://demo@127.0.0.1/db".to_string(),
1034 pool_size: None,
1035 min_connections: None,
1036 acquire_timeout_ms: None,
1037 idle_timeout_ms: None,
1038 max_lifetime_ms: None,
1039 postgres_session: None,
1040 }];
1041 let err = SqlProviderSpec::validate_specs(&specs).expect_err("invalid name charset");
1042 assert!(err.to_string().contains("invalid sqldb provider name"));
1043 }
1044
1045 #[test]
1046 fn parse_sqldb_postgres_session() {
1047 let dict = EnvDict::default();
1048 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1049 r#"
1050version = 2
1051
1052[provider.sqldb]
1053name = "geo"
1054kind = "postgres"
1055connection_uri = "postgres://demo@127.0.0.1/geo_db"
1056
1057[provider.sqldb.postgres_session]
1058plan_cache_mode = "force_generic_plan"
1059jit = false
1060application_name = "ip_geo_service"
1061"#,
1062 &dict,
1063 )
1064 .expect("parse postgres_session config");
1065
1066 let spec = &conf.provider().expect("provider").sqldb.expect("sqldb")[0];
1067 let session = spec.postgres_session.as_ref().expect("postgres_session");
1068 assert_eq!(
1069 session.plan_cache_mode,
1070 Some(PlanCacheMode::ForceGenericPlan)
1071 );
1072 assert_eq!(session.jit, Some(false));
1073 assert_eq!(session.application_name.as_deref(), Some("ip_geo_service"));
1074 }
1075
1076 #[test]
1077 fn parse_sqldb_postgres_session_unknown_key_rejected() {
1078 let dict = EnvDict::default();
1079 let r = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1080 r#"
1081version = 2
1082
1083[provider.sqldb]
1084kind = "postgres"
1085connection_uri = "postgres://demo@127.0.0.1/demo"
1086
1087[provider.sqldb.postgres_session]
1088plan_cache_mode = "auto"
1089after_connect_sql = "SET x = 1"
1090"#,
1091 &dict,
1092 );
1093 assert!(
1094 r.is_err(),
1095 "unknown postgres_session key (after_connect_sql) should be rejected"
1096 );
1097 }
1098
1099 #[test]
1100 fn validate_specs_rejects_session_on_non_postgres() {
1101 let specs = [SqlProviderSpec {
1102 name: Some("geo".to_string()),
1103 kind: SqlProviderKind::Mysql,
1104 connection_uri: "mysql://demo@127.0.0.1/db".to_string(),
1105 pool_size: None,
1106 min_connections: None,
1107 acquire_timeout_ms: None,
1108 idle_timeout_ms: None,
1109 max_lifetime_ms: None,
1110 postgres_session: Some(PostgresSessionSpec {
1111 plan_cache_mode: Some(PlanCacheMode::ForceGenericPlan),
1112 jit: None,
1113 application_name: None,
1114 }),
1115 }];
1116 let err = SqlProviderSpec::validate_specs(&specs).expect_err("mysql with postgres_session");
1117 assert!(err.to_string().contains("postgres_session"));
1118 }
1119
1120 #[test]
1121 fn validate_specs_rejects_application_name_too_long() {
1122 let specs = [SqlProviderSpec {
1123 name: Some("geo".to_string()),
1124 kind: SqlProviderKind::Postgres,
1125 connection_uri: "postgres://demo@127.0.0.1/db".to_string(),
1126 pool_size: None,
1127 min_connections: None,
1128 acquire_timeout_ms: None,
1129 idle_timeout_ms: None,
1130 max_lifetime_ms: None,
1131 postgres_session: Some(PostgresSessionSpec {
1132 plan_cache_mode: None,
1133 jit: None,
1134 application_name: Some("x".repeat(64)),
1135 }),
1136 }];
1137 let err = SqlProviderSpec::validate_specs(&specs).expect_err("application_name too long");
1138 assert!(err.to_string().contains("application_name"));
1139 }
1140
1141 #[test]
1142 fn parse_empty_sqldb_array_is_empty() {
1143 let dict = EnvDict::default();
1144 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1145 r#"
1146version = 2
1147
1148[provider]
1149sqldb = []
1150"#,
1151 &dict,
1152 )
1153 .expect("parse empty sqldb array");
1154
1155 let specs = conf
1156 .provider()
1157 .expect("provider")
1158 .sqldb
1159 .expect("sqldb present");
1160 assert!(specs.is_empty());
1161 }
1162
1163 #[test]
1164 fn mixed_named_and_unnamed_sqldb_applies_default_name() {
1165 let dict = EnvDict::default();
1166 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1167 r#"
1168version = 2
1169
1170[[provider.sqldb]]
1171name = "geo"
1172kind = "postgres"
1173connection_uri = "postgres://demo@127.0.0.1/geo_db"
1174
1175[[provider.sqldb]]
1176kind = "postgres"
1177connection_uri = "postgres://demo@127.0.0.1/asset_db"
1178"#,
1179 &dict,
1180 )
1181 .expect("parse mixed named/unnamed sqldb");
1182
1183 let specs = conf.provider().expect("provider").sqldb.expect("sqldb");
1184 assert_eq!(specs.len(), 2);
1185 assert_eq!(specs[0].effective_name(), "geo");
1187 assert_eq!(specs[1].effective_name(), DEFAULT_SQLDB_NAME);
1188
1189 let expected = SqlProviderSpec::validate_specs(&specs).expect("valid specs");
1191 assert!(expected.contains(DEFAULT_SQLDB_NAME));
1192 }
1193
1194 #[test]
1195 fn explicit_default_name_is_recognized() {
1196 let dict = EnvDict::default();
1197 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1198 r#"
1199version = 2
1200
1201[[provider.sqldb]]
1202name = "default"
1203kind = "postgres"
1204connection_uri = "postgres://demo@127.0.0.1/main_db"
1205
1206[[provider.sqldb]]
1207name = "geo"
1208kind = "postgres"
1209connection_uri = "postgres://demo@127.0.0.1/geo_db"
1210"#,
1211 &dict,
1212 )
1213 .expect("parse explicit default name");
1214
1215 let specs = conf.provider().expect("provider").sqldb.expect("sqldb");
1216 assert_eq!(specs[0].effective_name(), DEFAULT_SQLDB_NAME);
1217 assert_eq!(specs[1].effective_name(), "geo");
1218 let expected = SqlProviderSpec::validate_specs(&specs).expect("valid specs");
1219 assert!(expected.contains(DEFAULT_SQLDB_NAME));
1220 }
1221
1222 #[test]
1223 fn parse_cache_spec_with_defaults() {
1224 let dict = EnvDict::default();
1225 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1226 r#"
1227version = 2
1228"#,
1229 &dict,
1230 )
1231 .expect("parse knowdb with default cache spec");
1232
1233 assert!(conf.cache.enabled);
1234 assert_eq!(conf.cache.capacity, 1024);
1235 assert_eq!(conf.cache.ttl_ms, 30_000);
1236 }
1237
1238 #[test]
1239 fn parse_cache_spec_from_toml() {
1240 let dict = EnvDict::default();
1241 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1242 r#"
1243version = 2
1244
1245[cache]
1246enabled = false
1247capacity = 256
1248ttl_ms = 1500
1249"#,
1250 &dict,
1251 )
1252 .expect("parse knowdb with cache spec");
1253
1254 assert!(!conf.cache.enabled);
1255 assert_eq!(conf.cache.capacity, 256);
1256 assert_eq!(conf.cache.ttl_ms, 1500);
1257 }
1258
1259 #[test]
1260 fn parse_redis_cache_spec() {
1261 let dict = EnvDict::default();
1262 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1263 r#"
1264version = 2
1265
1266[cache]
1267enabled = true
1268capacity = 512
1269"#,
1270 &dict,
1271 )
1272 .expect("parse knowdb with cache");
1273
1274 assert!(conf.cache.enabled);
1275 assert_eq!(conf.cache.capacity, 512);
1276 }
1277
1278 #[test]
1279 fn parse_redis_cache_defaults() {
1280 let dict = EnvDict::default();
1281 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1282 r#"
1283version = 2
1284"#,
1285 &dict,
1286 )
1287 .expect("parse knowdb without redis.cache");
1288
1289 assert!(conf.cache.enabled);
1291 assert_eq!(conf.cache.capacity, 1024);
1292 }
1293
1294 #[test]
1299 fn parse_fun_bool_services() {
1300 let dict = EnvDict::default();
1301 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1302 r#"
1303version = 2
1304
1305[fun.password_check]
1306call = "bf_exists"
1307key = "weak_passwords"
1308
1309[fun.ip_whitelist]
1310call = "sismember"
1311key = "allowed_ips"
1312"#,
1313 &dict,
1314 )
1315 .expect("parse fun bool services");
1316
1317 let pw = conf.fun.get("password_check").expect("password_check");
1318 assert_eq!(pw.call, FunCall::BfExists);
1319 assert_eq!(pw.key.as_deref(), Some("weak_passwords"));
1320 assert!(pw.returns_bool());
1321
1322 let ip = conf.fun.get("ip_whitelist").expect("ip_whitelist");
1323 assert_eq!(ip.call, FunCall::Sismember);
1324 assert_eq!(ip.key.as_deref(), Some("allowed_ips"));
1325 assert!(ip.returns_bool());
1326 }
1327
1328 #[test]
1329 fn parse_fun_value_services() {
1330 let dict = EnvDict::default();
1331 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1332 r#"
1333version = 2
1334
1335[fun.threat_actor]
1336call = "hget"
1337key = "threat_actors"
1338cache = true
1339ttl_ms = 60000
1340
1341[fun.user_tag]
1342call = "get"
1343"#,
1344 &dict,
1345 )
1346 .expect("parse fun value services");
1347
1348 let ta = conf.fun.get("threat_actor").expect("threat_actor");
1349 assert_eq!(ta.call, FunCall::Hget);
1350 assert_eq!(ta.key.as_deref(), Some("threat_actors"));
1351 assert!(ta.cache);
1352 assert_eq!(ta.ttl_ms, Some(60000));
1353 assert!(!ta.returns_bool());
1354
1355 let ut = conf.fun.get("user_tag").expect("user_tag");
1356 assert_eq!(ut.call, FunCall::Get);
1357 assert!(ut.key.is_none());
1358 assert!(ut.cache); assert!(!ut.returns_bool());
1360 }
1361
1362 #[test]
1363 fn parse_fun_default_cache() {
1364 let dict = EnvDict::default();
1365 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1366 r#"
1367version = 2
1368
1369[fun.app_config]
1370call = "get"
1371key = "app_config"
1372"#,
1373 &dict,
1374 )
1375 .expect("parse fun default cache");
1376
1377 let spec = conf.fun.get("app_config").expect("app_config");
1378 assert!(spec.cache);
1379 assert!(spec.ttl_ms.is_none());
1380 }
1381
1382 fn scaffold_reload_fixture(root: &std::path::Path, rows_txt: &str) {
1389 let table_dir = root.join("address");
1390 fs::create_dir_all(&table_dir).expect("create table dir");
1391 let fixture = std::path::Path::new(env!("CARGO_MANIFEST_DIR")).join("knowdb/address");
1392 fs::copy(fixture.join("create.sql"), table_dir.join("create.sql")).expect("copy create");
1393 fs::copy(fixture.join("insert.sql"), table_dir.join("insert.sql")).expect("copy insert");
1394 fs::write(table_dir.join("data.csv"), rows_txt).expect("write data.csv");
1395 fs::write(
1396 root.join("knowdb.toml"),
1397 r#"version = 2
1398base_dir = "."
1399
1400[default]
1401on_error = "fail"
1402
1403[[tables]]
1404name = "address"
1405dir = "address"
1406enabled = true
1407columns.by_header = ["value"]
1408"#,
1409 )
1410 .expect("write conf");
1411 }
1412
1413 fn reload_uri(tag: &str) -> String {
1414 format!(
1415 "file:{}/wf_loader_reload_{}_{}.sqlite",
1416 std::env::temp_dir().display(),
1417 tag,
1418 std::process::id()
1419 )
1420 }
1421
1422 #[test]
1423 fn reload_table_rows_reflects_csv_overwrite_and_projects_typed_rows() {
1424 let root =
1425 std::env::temp_dir().join(format!("wf_loader_reload_root_a_{}", std::process::id()));
1426 let _ = fs::remove_dir_all(&root);
1427 fs::create_dir_all(&root).expect("create root");
1428 let dict = EnvDict::default();
1429 let conf = PathBuf::from("knowdb.toml");
1430
1431 scaffold_reload_fixture(&root, "value\nv1\nv2\nv3\n");
1433 let rows = reload_table_rows(&root, &conf, &reload_uri("a"), "address", &dict)
1434 .expect("first load");
1435 assert_eq!(rows.len(), 3);
1436 let field = &rows[0][0];
1437 assert_eq!(field.get_name(), "value");
1438 assert_eq!(field.get_meta(), &wp_model_core::model::DataType::Chars);
1439 assert!(matches!(
1440 field.get_value(),
1441 wp_model_core::model::Value::Chars(_)
1442 ));
1443
1444 scaffold_reload_fixture(&root, "value\nv9\n");
1446 let rows2 = reload_table_rows(&root, &conf, &reload_uri("a"), "address", &dict)
1447 .expect("reload after overwrite");
1448 assert_eq!(rows2.len(), 1, "重载应反映覆盖后的文件");
1449 let field2 = &rows2[0][0];
1450 assert_eq!(field2.get_name(), "value");
1451 assert_eq!(field2.get_meta(), &wp_model_core::model::DataType::Chars);
1452
1453 let _ = fs::remove_dir_all(&root);
1454 }
1455
1456 #[test]
1457 fn reload_table_rows_errors_on_unknown_disabled_and_by_index_only_tables() {
1458 let root =
1459 std::env::temp_dir().join(format!("wf_loader_reload_root_b_{}", std::process::id()));
1460 let _ = fs::remove_dir_all(&root);
1461 fs::create_dir_all(&root).expect("create root");
1462 fs::write(
1463 root.join("knowdb.toml"),
1464 r#"version = 2
1465base_dir = "."
1466
1467[[tables]]
1468name = "on"
1469dir = "address"
1470enabled = true
1471columns.by_header = ["value"]
1472
1473[[tables]]
1474name = "off"
1475dir = "address"
1476enabled = false
1477columns.by_header = ["value"]
1478
1479[[tables]]
1480name = "by_index"
1481dir = "address"
1482enabled = true
1483columns.by_index = [0]
1484"#,
1485 )
1486 .expect("write conf");
1487 let dict = EnvDict::default();
1488 let conf = PathBuf::from("knowdb.toml");
1489
1490 let err_unknown = reload_table_rows(&root, &conf, &reload_uri("b"), "ghost", &dict)
1491 .expect_err("未知表应报错");
1492 let msg = format!("{err_unknown}");
1493 assert!(msg.contains("not found"), "unknown: {msg}");
1494
1495 let err_disabled = reload_table_rows(&root, &conf, &reload_uri("b"), "off", &dict)
1496 .expect_err("禁用表应报错");
1497 let msg = format!("{err_disabled}");
1498 assert!(msg.contains("not found"), "disabled: {msg}");
1499
1500 let err_by_index = reload_table_rows(&root, &conf, &reload_uri("b"), "by_index", &dict)
1501 .expect_err("纯 by_index 表应报错");
1502 let msg = format!("{err_by_index}");
1503 assert!(msg.contains("columns.by_header"), "by_index only: {msg}");
1504
1505 let _ = fs::remove_dir_all(&root);
1506 }
1507}