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 orion_error::OperationContext;
15use orion_error::conversion::{SourceErr, SourceRawErr, ToStructError};
16use orion_variate::EnvDict;
17use rusqlite::OpenFlags;
18
19#[derive(Debug, Deserialize)]
22pub struct KnowDbConf {
23 pub version: u32,
24 #[serde(default = "default_dot")]
25 pub base_dir: String,
26 #[serde(default)]
27 pub default: OptLoadSpec,
28 #[serde(default)]
29 pub csv: CsvSpec,
30 #[serde(default)]
31 pub cache: CacheSpec,
32 #[serde(default)]
33 pub tables: Vec<TableSpec>,
34
35 #[serde(default)]
37 pub fun: HashMap<String, FunSpec>,
38
39 #[serde(default, rename = "provider")]
41 provider_raw: Option<ProviderConfig>,
42
43 #[serde(default)]
45 pub intranet_nets: Option<crate::intranet_nets::IntranetNetsConf>,
46}
47
48impl KnowDbConf {
49 pub fn provider(&self) -> Option<ProviderConfig> {
50 self.provider_raw.clone()
51 }
52}
53
54#[derive(Debug, Clone, Deserialize)]
59pub struct FunSpec {
60 pub call: FunCall,
61 #[serde(default)]
62 pub key: Option<String>,
63 #[serde(default = "default_true")]
64 pub cache: bool,
65 #[serde(default)]
66 pub ttl_ms: Option<u64>,
67}
68
69impl FunSpec {
70 pub fn returns_bool(&self) -> bool {
72 matches!(self.call, FunCall::BfExists | FunCall::Sismember)
73 }
74}
75
76#[derive(Debug, Clone, Deserialize, PartialEq)]
77#[serde(rename_all = "snake_case")]
78pub enum FunCall {
79 BfExists,
80 Sismember,
81 Hget,
82 Get,
83}
84
85#[derive(Debug, Clone, Deserialize)]
90pub struct CacheSpec {
91 #[serde(default = "default_true")]
92 pub enabled: bool,
93 #[serde(default = "default_result_cache_capacity")]
94 pub capacity: usize,
95 #[serde(default = "default_result_cache_ttl_ms")]
96 pub ttl_ms: u64,
97}
98
99impl Default for CacheSpec {
100 fn default() -> Self {
101 Self {
102 enabled: default_true(),
103 capacity: default_result_cache_capacity(),
104 ttl_ms: default_result_cache_ttl_ms(),
105 }
106 }
107}
108
109pub const DEFAULT_SQLDB_NAME: &str = "default";
115
116#[derive(Debug, Clone, Default, Deserialize)]
117pub struct ProviderConfig {
118 #[serde(default, deserialize_with = "deserialize_sqldb")]
120 pub sqldb: Option<Vec<SqlProviderSpec>>,
121 #[serde(default)]
122 pub redis: Option<RedisProviderSpec>,
123}
124
125impl ProviderConfig {
126 pub fn sqldb_names(&self) -> Vec<String> {
128 self.sqldb
129 .as_ref()
130 .map(|specs| {
131 specs
132 .iter()
133 .map(|spec| spec.effective_name().to_string())
134 .collect()
135 })
136 .unwrap_or_default()
137 }
138}
139
140#[derive(Debug, Clone, Deserialize)]
141pub struct SqlProviderSpec {
142 #[serde(default)]
143 pub name: Option<String>,
144 #[serde(rename = "kind")]
145 pub kind: SqlProviderKind,
146 pub connection_uri: String,
147 #[serde(default)]
148 pub pool_size: Option<u32>,
149 #[serde(default)]
150 pub min_connections: Option<u32>,
151 #[serde(default)]
152 pub acquire_timeout_ms: Option<u64>,
153 #[serde(default)]
154 pub idle_timeout_ms: Option<u64>,
155 #[serde(default)]
156 pub max_lifetime_ms: Option<u64>,
157}
158
159impl SqlProviderSpec {
160 pub fn effective_name(&self) -> &str {
162 self.name.as_deref().unwrap_or(DEFAULT_SQLDB_NAME)
163 }
164
165 pub fn validate_specs(specs: &[SqlProviderSpec]) -> KnowledgeResult<HashSet<String>> {
167 let mut seen = HashSet::with_capacity(specs.len());
168 for spec in specs {
169 let name = spec.effective_name();
170 if !name.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
171 return Err(KnowReason::from_conf().to_err().with_detail(format!(
172 "invalid sqldb provider name '{name}' (allowed [A-Za-z0-9_])"
173 )));
174 }
175 if !seen.insert(name.to_string()) {
176 return Err(KnowReason::from_conf()
177 .to_err()
178 .with_detail(format!("duplicate sqldb provider name '{name}'")));
179 }
180 }
181 Ok(seen)
182 }
183}
184
185fn deserialize_sqldb<'de, D>(deserializer: D) -> Result<Option<Vec<SqlProviderSpec>>, D::Error>
187where
188 D: Deserializer<'de>,
189{
190 struct SqlDbVisitor;
191
192 impl<'de> Visitor<'de> for SqlDbVisitor {
193 type Value = Option<Vec<SqlProviderSpec>>;
194
195 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
196 formatter
197 .write_str("a `[provider.sqldb]` table or a `[[provider.sqldb]]` array of tables")
198 }
199
200 fn visit_map<A>(self, map: A) -> Result<Self::Value, A::Error>
201 where
202 A: MapAccess<'de>,
203 {
204 let spec = SqlProviderSpec::deserialize(de::value::MapAccessDeserializer::new(map))?;
205 Ok(Some(vec![spec]))
206 }
207
208 fn visit_seq<A>(self, seq: A) -> Result<Self::Value, A::Error>
209 where
210 A: SeqAccess<'de>,
211 {
212 let specs =
213 Vec::<SqlProviderSpec>::deserialize(de::value::SeqAccessDeserializer::new(seq))?;
214 Ok(Some(specs))
215 }
216 }
217
218 deserializer.deserialize_any(SqlDbVisitor)
219}
220
221#[derive(Debug, Clone, Deserialize)]
225#[serde(rename_all = "snake_case")]
226pub enum SqlProviderKind {
227 Postgres,
228 Mysql,
229}
230
231#[derive(Debug, Clone, Deserialize)]
232pub struct RedisProviderSpec {
233 pub connection_uri: String,
234 #[serde(default)]
235 pub pool_size: Option<usize>,
236 #[serde(default = "default_connect_timeout_ms")]
237 pub connect_timeout_ms: u64,
238 #[serde(default = "default_command_timeout_ms")]
239 pub command_timeout_ms: u64,
240}
241
242fn default_connect_timeout_ms() -> u64 {
243 3_000
244}
245
246fn default_command_timeout_ms() -> u64 {
247 100
248}
249
250#[derive(Debug, Clone, Deserialize)]
256#[serde(rename_all = "snake_case")]
257pub enum ProviderKind {
258 SqliteAuthority,
259 Postgres,
260 Mysql,
261 Redis,
262}
263
264#[derive(Debug, Clone, Deserialize)]
265pub struct OptLoadSpec {
266 #[serde(default = "default_true")]
267 pub transaction: bool,
268 #[serde(default = "default_batch")]
269 pub batch_size: usize,
270 #[serde(default = "default_on_error")]
271 pub on_error: OnError,
272}
273impl Default for OptLoadSpec {
274 fn default() -> Self {
275 Self {
276 transaction: true,
277 batch_size: default_batch(),
278 on_error: default_on_error(),
279 }
280 }
281}
282
283#[derive(Debug, Clone, Deserialize, Default)]
284#[serde(rename_all = "lowercase")]
285pub enum OnError {
286 #[default]
287 Fail,
288 Skip,
289}
290
291#[derive(Debug, Clone, Deserialize)]
292pub struct CsvSpec {
293 #[serde(default = "default_true")]
294 pub has_header: bool,
295 #[serde(default = "default_comma")]
296 pub delimiter: String,
297 #[serde(default = "default_utf8")]
298 pub encoding: String,
299 #[serde(default = "default_true")]
300 pub trim: bool,
301}
302impl Default for CsvSpec {
303 fn default() -> Self {
304 CsvSpec {
305 has_header: true,
306 delimiter: ",".into(),
307 encoding: "utf-8".into(),
308 trim: true,
309 }
310 }
311}
312
313#[derive(Debug, Clone, Deserialize)]
314pub struct TableSpec {
315 pub name: String,
316 #[serde(default)]
317 pub dir: Option<String>,
318 #[serde(default)]
319 pub data_file: Option<String>,
320 pub columns: ColumnsSpec,
321 #[serde(default)]
322 pub expected_rows: RowExpect,
323 #[serde(default = "default_true")]
324 pub enabled: bool,
325}
326
327#[derive(Debug, Clone, Deserialize)]
328pub struct ColumnsSpec {
329 #[serde(default)]
330 pub by_header: Vec<String>,
331 #[serde(default)]
332 pub by_index: Vec<usize>,
333}
334
335#[derive(Debug, Clone, Deserialize, Default)]
336pub struct RowExpect {
337 pub min: Option<usize>,
338 pub max: Option<usize>,
339}
340
341const fn default_true() -> bool {
342 true
343}
344const fn default_batch() -> usize {
345 2000
346}
347fn default_comma() -> String {
348 ",".to_string()
349}
350fn default_utf8() -> String {
351 "utf-8".to_string()
352}
353fn default_on_error() -> OnError {
354 OnError::Fail
355}
356fn default_dot() -> String {
357 ".".to_string()
358}
359const fn default_result_cache_capacity() -> usize {
360 1024
361}
362const fn default_result_cache_ttl_ms() -> u64 {
363 30_000
364}
365
366fn read_to_string(path: &Path) -> KnowledgeResult<String> {
368 let mut f = fs::File::open(path).source_raw_err(KnowReason::from_res(), "source error")?;
369 let mut buf = String::new();
370 f.read_to_string(&mut buf)
371 .source_raw_err(KnowReason::from_res(), "source error")?;
372 Ok(buf)
373}
374
375fn replace_table(sql: &str, table: &str) -> String {
376 sql.replace("{table}", table)
377}
378
379fn join_rel(base: &Path, rel: &str) -> PathBuf {
380 let p = Path::new(rel);
381 if p.is_absolute() {
382 p.to_path_buf()
383 } else {
384 base.join(p)
385 }
386}
387
388pub fn build_authority_from_knowdb(
389 root: &Path,
390 conf_path: &Path,
391 authority_uri: &str,
392 dict: &EnvDict,
393) -> KnowledgeResult<Vec<String>> {
394 let mut opx = OperationContext::doing("build authority from knowdb").with_auto_log();
395 let (conf, conf_abs, base_dir) = parse_knowdb_conf(root, conf_path, dict)?;
397 opx.record("conf", conf_abs.display());
398 opx.record("base_dir", base_dir.display());
399 let db = open_authority(authority_uri)?;
401 let mut loaded_names = Vec::new();
403 for t in &conf.tables {
404 if !t.enabled {
405 continue;
406 }
407 load_one_table(&db, &base_dir, t, &conf.csv, &conf.default)?;
408 info_ctrl!("load table {} suc!", base_dir.display(),);
409 loaded_names.push(t.name.clone());
410 }
411 opx.mark_suc();
412 Ok(loaded_names)
413}
414
415pub fn parse_knowdb_conf(
416 root: &Path,
417 conf_path: &Path,
418 dict: &EnvDict,
419) -> KnowledgeResult<(KnowDbConf, PathBuf, PathBuf)> {
420 let conf_abs = if conf_path.is_absolute() {
421 conf_path.to_path_buf()
422 } else {
423 root.join(conf_path)
424 };
425 let conf_txt = read_to_string(&conf_abs)?;
426 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(&conf_txt, dict)
427 .source_err(KnowReason::from_conf(), "parse knowdb config")?;
428 if conf.version != 2 {
429 return Err(KnowReason::from_conf()
430 .to_err()
431 .with_detail("unsupported knowdb.version"));
432 }
433 crate::intranet_nets::set_intranet_nets_conf(conf.intranet_nets.clone());
435 let conf_dir = conf_abs.parent().unwrap_or_else(|| Path::new("."));
436 let base_dir = join_rel(conf_dir, &conf.base_dir);
437 Ok((conf, conf_abs, base_dir))
438}
439
440fn open_authority(authority_uri: &str) -> KnowledgeResult<MemDB> {
441 ensure_parent_dir_for_file_uri(authority_uri);
442 let flags = OpenFlags::SQLITE_OPEN_READ_WRITE
443 | OpenFlags::SQLITE_OPEN_CREATE
444 | OpenFlags::SQLITE_OPEN_URI;
445 let db = MemDB::new_file(authority_uri, 1, flags)?;
446 let _ = db.with_conn(|conn| {
448 let _ = crate::sqlite_ext::register_builtin(conn);
449 Ok::<(), anyhow::Error>(())
450 });
451 Ok(db)
452}
453
454fn ensure_parent_dir_for_file_uri(uri: &str) {
457 if let Some(rest) = uri.strip_prefix("file:") {
458 let path_part = rest.split('?').next().unwrap_or(rest);
459 let p = Path::new(path_part);
460 if let Some(parent) = p.parent() {
461 let _ = fs::create_dir_all(parent);
462 }
463 }
464}
465
466fn load_one_table(
467 db: &MemDB,
468 base_dir: &Path,
469 t: &TableSpec,
470 csvd: &CsvSpec,
471 load: &OptLoadSpec,
472) -> KnowledgeResult<()> {
473 let mut opx = OperationContext::doing("load table to kdb")
475 .with_auto_log()
476 .with_mod_path("ctrl");
477 let dir_name: &str = t.dir.as_deref().unwrap_or(&t.name);
478 let table_dir = base_dir.join(dir_name);
479 opx.record("table_dir", table_dir.display());
480 let create_sql = replace_table(&read_to_string(&table_dir.join("create.sql"))?, &t.name);
481 let insert_sql = replace_table(&read_to_string(&table_dir.join("insert.sql"))?, &t.name);
482 let clean_path = table_dir.join("clean.sql");
483 let clean_sql = if clean_path.exists() {
484 replace_table(&read_to_string(&clean_path)?, &t.name)
485 } else {
486 format!("DELETE FROM {}", t.name)
487 };
488
489 db.with_conn(|conn| {
491 let _ = crate::sqlite_ext::register_builtin(conn);
493 conn.execute_batch(&create_sql)?;
494 conn.execute_batch(&clean_sql)?;
495 Ok::<(), anyhow::Error>(())
496 })
497 .source_err(KnowReason::from_res(), "prepare authority table")?;
498
499 let data_path = match &t.data_file {
501 Some(rel) => join_rel(&table_dir, rel),
502 None => table_dir.join("data.csv"),
503 };
504 if !data_path.exists() {
505 return Err(KnowReason::from_conf()
506 .to_err()
507 .with_detail("data.csv not found"));
508 }
509 opx.record("data_path", data_path.display());
510
511 let mut rdr = build_csv_reader(csvd, &data_path)?;
513
514 let col_indices: Vec<usize> = if !t.columns.by_header.is_empty() {
516 let headers = rdr
517 .headers()
518 .source_raw_err(KnowReason::from_res(), "source error")?;
519 select_indices_by_header(headers, &t.columns.by_header)?
520 } else if !t.columns.by_index.is_empty() {
521 t.columns.by_index.clone()
522 } else {
523 return Err(KnowReason::from_conf()
524 .to_err()
525 .with_detail("columns mapping required"));
526 };
527
528 let mut inserted: usize = 0;
530 let mut bad: usize = 0;
531 let mut batch_left = load.batch_size.max(1);
532 db.with_conn(|conn| {
533 let _ = crate::sqlite_ext::register_builtin(conn);
535 let mut tx = if load.transaction {
536 Some(conn.unchecked_transaction()?)
537 } else {
538 None
539 };
540 let mut stmt = conn.prepare(&insert_sql)?;
541 for rec in rdr.into_records() {
542 match rec {
543 Ok(record) => {
544 let refs = extract_row_refs(&record, &col_indices, &mut bad, load)?;
545 if let Some(refs) = refs {
546 stmt.execute(rusqlite::params_from_iter(refs))?;
547 inserted += 1;
548 if load.transaction {
549 batch_left -= 1;
550 if batch_left == 0 {
551 tx.take().unwrap().commit()?;
552 tx = Some(conn.unchecked_transaction()?);
553 batch_left = load.batch_size.max(1);
554 }
555 }
556 }
557 }
558 Err(_e) => {
559 if matches!(load.on_error, OnError::Skip) {
560 bad += 1;
561 continue;
562 } else {
563 anyhow::bail!("csv record parse error");
564 }
565 }
566 }
567 }
568 if let Some(tx) = tx {
569 tx.commit()?;
570 }
571 Ok::<(), anyhow::Error>(())
572 })
573 .source_err(KnowReason::from_res(), "load authority table data")?;
574
575 if let Some(min) = t.expected_rows.min
577 && inserted < min
578 {
579 return Err(KnowReason::from_conf()
580 .to_err()
581 .with_detail("table data less"));
582 }
583 if let Some(max) = t.expected_rows.max
584 && inserted > max
585 {
586 wp_log::warn_kdb!(
587 "table {} loaded rows {} exceed max {}",
588 &t.name,
589 inserted,
590 max
591 );
592 }
593 if bad > 0 {
594 wp_log::warn_kdb!("table {} skipped {} bad rows (on_error=skip)", &t.name, bad);
595 }
596 opx.mark_suc();
597 Ok(())
598}
599
600fn build_csv_reader(
601 csvd: &CsvSpec,
602 data_path: &Path,
603) -> KnowledgeResult<csv::Reader<std::fs::File>> {
604 if csvd.encoding.to_lowercase() != "utf-8" {
605 return Err(KnowReason::from_conf()
606 .to_err()
607 .with_detail("only utf-8 csv is supported"));
608 }
609 let mut rdr_b = csv::ReaderBuilder::new();
610 rdr_b.has_headers(csvd.has_header);
611 if csvd.delimiter.len() == 1 {
612 rdr_b.delimiter(csvd.delimiter.as_bytes()[0]);
613 }
614 if csvd.trim {
615 rdr_b.trim(csv::Trim::All);
616 }
617 rdr_b
618 .from_path(data_path)
619 .source_raw_err(KnowReason::from_res(), "source error")
620}
621
622fn select_indices_by_header(
623 headers: &csv::StringRecord,
624 wanted: &[String],
625) -> KnowledgeResult<Vec<usize>> {
626 let mut out = Vec::with_capacity(wanted.len());
627 for name in wanted {
628 let pos = headers.iter().position(|h| h == name).ok_or_else(|| {
629 KnowReason::from_conf()
630 .to_err()
631 .with_detail("header not found")
632 })?;
633 out.push(pos);
634 }
635 Ok(out)
636}
637
638fn extract_row_refs<'a>(
639 record: &'a csv::StringRecord,
640 col_indices: &[usize],
641 bad: &mut usize,
642 load: &OptLoadSpec,
643) -> anyhow::Result<Option<Vec<&'a str>>> {
644 let mut vs: Vec<&str> = Vec::with_capacity(col_indices.len());
645 for &idx in col_indices {
646 if idx >= record.len() {
647 if matches!(load.on_error, OnError::Skip) {
648 *bad += 1;
649 return Ok(None);
650 } else {
651 anyhow::bail!("missing column at index {}", idx);
652 }
653 }
654 vs.push(record.get(idx).unwrap_or(""));
655 }
656 Ok(Some(vs))
657}
658
659#[cfg(test)]
660mod tests {
661 use super::*;
662
663 #[test]
664 fn parse_new_style_sqldb_provider() {
665 let dict = EnvDict::default();
666 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
667 r#"
668version = 2
669
670[provider.sqldb]
671kind = "postgres"
672connection_uri = "postgres://demo:demo@127.0.0.1/demo"
673pool_size = 12
674"#,
675 &dict,
676 )
677 .expect("parse knowdb with sqldb provider");
678
679 let sqldb = conf
680 .provider()
681 .expect("provider")
682 .sqldb
683 .expect("sqldb provider");
684 assert_eq!(sqldb.len(), 1);
685 let sqldb = &sqldb[0];
686 assert!(matches!(sqldb.kind, SqlProviderKind::Postgres));
687 assert_eq!(sqldb.pool_size, Some(12));
688 assert_eq!(sqldb.effective_name(), DEFAULT_SQLDB_NAME);
689 }
690
691 #[test]
692 fn parse_new_style_redis_provider() {
693 let dict = EnvDict::default();
694 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
695 r#"
696version = 2
697
698[provider.redis]
699connection_uri = "redis://127.0.0.1:6379"
700pool_size = 16
701connect_timeout_ms = 5000
702command_timeout_ms = 200
703"#,
704 &dict,
705 )
706 .expect("parse knowdb with redis provider");
707
708 let redis_cfg = conf
709 .provider()
710 .expect("provider")
711 .redis
712 .expect("redis provider");
713 assert_eq!(redis_cfg.connection_uri, "redis://127.0.0.1:6379");
714 assert_eq!(redis_cfg.pool_size, Some(16));
715 assert_eq!(redis_cfg.connect_timeout_ms, 5000);
716 assert_eq!(redis_cfg.command_timeout_ms, 200);
717 }
718
719 #[test]
720 fn parse_redis_provider_with_default_timeouts() {
721 let dict = EnvDict::default();
722 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
723 r#"
724version = 2
725
726[provider.redis]
727connection_uri = "redis://127.0.0.1:6379"
728"#,
729 &dict,
730 )
731 .expect("parse knowdb with redis provider (no timeout fields)");
732
733 let redis_cfg = conf.provider().expect("provider").redis.expect("redis");
734 assert_eq!(redis_cfg.connect_timeout_ms, 3000);
735 assert_eq!(redis_cfg.command_timeout_ms, 100);
736 }
737
738 #[test]
739 fn parse_both_sqldb_and_redis_providers() {
740 let dict = EnvDict::default();
741 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
742 r#"
743version = 2
744
745[provider.sqldb]
746kind = "postgres"
747connection_uri = "postgres://demo:demo@127.0.0.1/demo"
748
749[provider.redis]
750connection_uri = "redis://10.0.0.1:6379"
751pool_size = 4
752"#,
753 &dict,
754 )
755 .expect("parse knowdb with both sqldb and redis");
756
757 let provider_cfg = conf.provider().expect("provider");
758 let sqldb = &provider_cfg.sqldb.expect("sqldb")[0];
759 let redis_cfg = provider_cfg.redis.expect("redis");
760 assert!(matches!(sqldb.kind, SqlProviderKind::Postgres));
761 assert_eq!(redis_cfg.connection_uri, "redis://10.0.0.1:6379");
762 assert_eq!(redis_cfg.pool_size, Some(4));
763 }
764
765 #[test]
766 fn parse_redis_only_without_sqldb() {
767 let dict = EnvDict::default();
768 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
769 r#"
770version = 2
771
772[provider.redis]
773connection_uri = "redis://127.0.0.1:6379"
774"#,
775 &dict,
776 )
777 .expect("parse knowdb with redis only");
778
779 let provider_cfg = conf.provider().expect("provider");
780 assert!(provider_cfg.sqldb.is_none());
781 assert!(provider_cfg.redis.is_some());
782 }
783
784 #[test]
785 fn parse_no_provider_section() {
786 let dict = EnvDict::default();
787 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
788 r#"
789version = 2
790"#,
791 &dict,
792 )
793 .expect("parse knowdb without provider");
794
795 assert!(conf.provider().is_none());
796 }
797
798 #[test]
799 fn new_style_sqldb_mysql_variant() {
800 let dict = EnvDict::default();
801 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
802 r#"
803version = 2
804
805[provider.sqldb]
806kind = "mysql"
807connection_uri = "mysql://user:pass@127.0.0.1:3306/db"
808pool_size = 8
809"#,
810 &dict,
811 )
812 .expect("parse new-style mysql sqldb");
813
814 let sqldb = &conf.provider().expect("provider").sqldb.expect("sqldb")[0];
815 assert!(matches!(sqldb.kind, SqlProviderKind::Mysql));
816 assert_eq!(sqldb.pool_size, Some(8));
817 }
818
819 #[test]
820 fn parse_array_style_multi_sqldb_providers() {
821 let dict = EnvDict::default();
822 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
823 r#"
824version = 2
825
826[[provider.sqldb]]
827name = "geo"
828kind = "postgres"
829connection_uri = "postgres://demo@127.0.0.1:5432/geo_db"
830pool_size = 8
831
832[[provider.sqldb]]
833name = "asset"
834kind = "postgres"
835connection_uri = "postgres://demo@127.0.0.1:5432/asset_db"
836pool_size = 12
837"#,
838 &dict,
839 )
840 .expect("parse knowdb with multiple sqldb providers");
841
842 let specs = conf
843 .provider()
844 .expect("provider")
845 .sqldb
846 .expect("sqldb providers");
847 assert_eq!(specs.len(), 2);
848 assert_eq!(specs[0].name.as_deref(), Some("geo"));
849 assert_eq!(specs[0].effective_name(), "geo");
850 assert_eq!(
851 specs[0].connection_uri,
852 "postgres://demo@127.0.0.1:5432/geo_db"
853 );
854 assert_eq!(specs[0].pool_size, Some(8));
855 assert_eq!(specs[1].name.as_deref(), Some("asset"));
856 assert_eq!(specs[1].effective_name(), "asset");
857 assert_eq!(
858 specs[1].connection_uri,
859 "postgres://demo@127.0.0.1:5432/asset_db"
860 );
861 assert_eq!(specs[1].pool_size, Some(12));
862 }
863
864 #[test]
865 fn parse_single_sqldb_with_explicit_name() {
866 let dict = EnvDict::default();
867 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
868 r#"
869version = 2
870
871[provider.sqldb]
872name = "geo"
873kind = "postgres"
874connection_uri = "postgres://demo@127.0.0.1/geo_db"
875"#,
876 &dict,
877 )
878 .expect("parse single named sqldb");
879
880 let spec = &conf.provider().expect("provider").sqldb.expect("sqldb")[0];
881 assert_eq!(spec.effective_name(), "geo");
882 }
883
884 #[test]
885 fn sqldb_names_applies_default_name() {
886 let dict = EnvDict::default();
887 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
888 r#"
889version = 2
890
891[provider.sqldb]
892kind = "postgres"
893connection_uri = "postgres://demo@127.0.0.1/demo"
894"#,
895 &dict,
896 )
897 .expect("parse unnamed sqldb");
898
899 let names = conf.provider().expect("provider").sqldb_names();
900 assert_eq!(names, vec!["default".to_string()]);
901 }
902
903 #[test]
904 fn validate_specs_rejects_duplicate_effective_names() {
905 let dict = EnvDict::default();
906 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
907 r#"
908version = 2
909
910[[provider.sqldb]]
911kind = "postgres"
912connection_uri = "postgres://demo@127.0.0.1/db1"
913
914[[provider.sqldb]]
915kind = "postgres"
916connection_uri = "postgres://demo@127.0.0.1/db2"
917"#,
918 &dict,
919 )
920 .expect("parse two unnamed sqldb");
921
922 let specs = conf.provider().expect("provider").sqldb.expect("sqldb");
923 let err = SqlProviderSpec::validate_specs(&specs).expect_err("duplicate 'default' name");
924 assert!(
925 err.to_string()
926 .contains("duplicate sqldb provider name 'default'")
927 );
928 }
929
930 #[test]
931 fn validate_specs_rejects_invalid_name_charset() {
932 let specs = [SqlProviderSpec {
933 name: Some("bad name".to_string()),
934 kind: SqlProviderKind::Postgres,
935 connection_uri: "postgres://demo@127.0.0.1/db".to_string(),
936 pool_size: None,
937 min_connections: None,
938 acquire_timeout_ms: None,
939 idle_timeout_ms: None,
940 max_lifetime_ms: None,
941 }];
942 let err = SqlProviderSpec::validate_specs(&specs).expect_err("invalid name charset");
943 assert!(err.to_string().contains("invalid sqldb provider name"));
944 }
945
946 #[test]
947 fn parse_empty_sqldb_array_is_empty() {
948 let dict = EnvDict::default();
949 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
950 r#"
951version = 2
952
953[provider]
954sqldb = []
955"#,
956 &dict,
957 )
958 .expect("parse empty sqldb array");
959
960 let specs = conf
961 .provider()
962 .expect("provider")
963 .sqldb
964 .expect("sqldb present");
965 assert!(specs.is_empty());
966 }
967
968 #[test]
969 fn mixed_named_and_unnamed_sqldb_applies_default_name() {
970 let dict = EnvDict::default();
971 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
972 r#"
973version = 2
974
975[[provider.sqldb]]
976name = "geo"
977kind = "postgres"
978connection_uri = "postgres://demo@127.0.0.1/geo_db"
979
980[[provider.sqldb]]
981kind = "postgres"
982connection_uri = "postgres://demo@127.0.0.1/asset_db"
983"#,
984 &dict,
985 )
986 .expect("parse mixed named/unnamed sqldb");
987
988 let specs = conf.provider().expect("provider").sqldb.expect("sqldb");
989 assert_eq!(specs.len(), 2);
990 assert_eq!(specs[0].effective_name(), "geo");
992 assert_eq!(specs[1].effective_name(), DEFAULT_SQLDB_NAME);
993
994 let expected = SqlProviderSpec::validate_specs(&specs).expect("valid specs");
996 assert!(expected.contains(DEFAULT_SQLDB_NAME));
997 }
998
999 #[test]
1000 fn explicit_default_name_is_recognized() {
1001 let dict = EnvDict::default();
1002 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1003 r#"
1004version = 2
1005
1006[[provider.sqldb]]
1007name = "default"
1008kind = "postgres"
1009connection_uri = "postgres://demo@127.0.0.1/main_db"
1010
1011[[provider.sqldb]]
1012name = "geo"
1013kind = "postgres"
1014connection_uri = "postgres://demo@127.0.0.1/geo_db"
1015"#,
1016 &dict,
1017 )
1018 .expect("parse explicit default name");
1019
1020 let specs = conf.provider().expect("provider").sqldb.expect("sqldb");
1021 assert_eq!(specs[0].effective_name(), DEFAULT_SQLDB_NAME);
1022 assert_eq!(specs[1].effective_name(), "geo");
1023 let expected = SqlProviderSpec::validate_specs(&specs).expect("valid specs");
1024 assert!(expected.contains(DEFAULT_SQLDB_NAME));
1025 }
1026
1027 #[test]
1028 fn parse_cache_spec_with_defaults() {
1029 let dict = EnvDict::default();
1030 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1031 r#"
1032version = 2
1033"#,
1034 &dict,
1035 )
1036 .expect("parse knowdb with default cache spec");
1037
1038 assert!(conf.cache.enabled);
1039 assert_eq!(conf.cache.capacity, 1024);
1040 assert_eq!(conf.cache.ttl_ms, 30_000);
1041 }
1042
1043 #[test]
1044 fn parse_cache_spec_from_toml() {
1045 let dict = EnvDict::default();
1046 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1047 r#"
1048version = 2
1049
1050[cache]
1051enabled = false
1052capacity = 256
1053ttl_ms = 1500
1054"#,
1055 &dict,
1056 )
1057 .expect("parse knowdb with cache spec");
1058
1059 assert!(!conf.cache.enabled);
1060 assert_eq!(conf.cache.capacity, 256);
1061 assert_eq!(conf.cache.ttl_ms, 1500);
1062 }
1063
1064 #[test]
1065 fn parse_redis_cache_spec() {
1066 let dict = EnvDict::default();
1067 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1068 r#"
1069version = 2
1070
1071[cache]
1072enabled = true
1073capacity = 512
1074"#,
1075 &dict,
1076 )
1077 .expect("parse knowdb with cache");
1078
1079 assert!(conf.cache.enabled);
1080 assert_eq!(conf.cache.capacity, 512);
1081 }
1082
1083 #[test]
1084 fn parse_redis_cache_defaults() {
1085 let dict = EnvDict::default();
1086 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1087 r#"
1088version = 2
1089"#,
1090 &dict,
1091 )
1092 .expect("parse knowdb without redis.cache");
1093
1094 assert!(conf.cache.enabled);
1096 assert_eq!(conf.cache.capacity, 1024);
1097 }
1098
1099 #[test]
1104 fn parse_fun_bool_services() {
1105 let dict = EnvDict::default();
1106 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1107 r#"
1108version = 2
1109
1110[fun.password_check]
1111call = "bf_exists"
1112key = "weak_passwords"
1113
1114[fun.ip_whitelist]
1115call = "sismember"
1116key = "allowed_ips"
1117"#,
1118 &dict,
1119 )
1120 .expect("parse fun bool services");
1121
1122 let pw = conf.fun.get("password_check").expect("password_check");
1123 assert_eq!(pw.call, FunCall::BfExists);
1124 assert_eq!(pw.key.as_deref(), Some("weak_passwords"));
1125 assert!(pw.returns_bool());
1126
1127 let ip = conf.fun.get("ip_whitelist").expect("ip_whitelist");
1128 assert_eq!(ip.call, FunCall::Sismember);
1129 assert_eq!(ip.key.as_deref(), Some("allowed_ips"));
1130 assert!(ip.returns_bool());
1131 }
1132
1133 #[test]
1134 fn parse_fun_value_services() {
1135 let dict = EnvDict::default();
1136 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1137 r#"
1138version = 2
1139
1140[fun.threat_actor]
1141call = "hget"
1142key = "threat_actors"
1143cache = true
1144ttl_ms = 60000
1145
1146[fun.user_tag]
1147call = "get"
1148"#,
1149 &dict,
1150 )
1151 .expect("parse fun value services");
1152
1153 let ta = conf.fun.get("threat_actor").expect("threat_actor");
1154 assert_eq!(ta.call, FunCall::Hget);
1155 assert_eq!(ta.key.as_deref(), Some("threat_actors"));
1156 assert!(ta.cache);
1157 assert_eq!(ta.ttl_ms, Some(60000));
1158 assert!(!ta.returns_bool());
1159
1160 let ut = conf.fun.get("user_tag").expect("user_tag");
1161 assert_eq!(ut.call, FunCall::Get);
1162 assert!(ut.key.is_none());
1163 assert!(ut.cache); assert!(!ut.returns_bool());
1165 }
1166
1167 #[test]
1168 fn parse_fun_default_cache() {
1169 let dict = EnvDict::default();
1170 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1171 r#"
1172version = 2
1173
1174[fun.app_config]
1175call = "get"
1176key = "app_config"
1177"#,
1178 &dict,
1179 )
1180 .expect("parse fun default cache");
1181
1182 let spec = conf.fun.get("app_config").expect("app_config");
1183 assert!(spec.cache);
1184 assert!(spec.ttl_ms.is_none());
1185 }
1186}