1use std::collections::HashMap;
2use std::fs;
3use std::io::Read;
4use std::path::{Path, PathBuf};
5
6use orion_conf::EnvTomlLoad;
7use serde::Deserialize;
8use wp_log::info_ctrl;
9
10use crate::error::{KnowReason, KnowledgeResult};
11use crate::mem::memdb::MemDB;
12use orion_error::OperationContext;
13use orion_error::conversion::{SourceErr, SourceRawErr, ToStructError};
14use orion_variate::EnvDict;
15use rusqlite::OpenFlags;
16
17#[derive(Debug, Deserialize)]
20pub struct KnowDbConf {
21 pub version: u32,
22 #[serde(default = "default_dot")]
23 pub base_dir: String,
24 #[serde(default)]
25 pub default: OptLoadSpec,
26 #[serde(default)]
27 pub csv: CsvSpec,
28 #[serde(default)]
29 pub cache: CacheSpec,
30 #[serde(default)]
31 pub tables: Vec<TableSpec>,
32
33 #[serde(default)]
35 pub fun: HashMap<String, FunSpec>,
36
37 #[serde(default, rename = "provider")]
39 provider_raw: Option<ProviderConfig>,
40
41 #[serde(default)]
43 pub intranet_nets: Option<crate::intranet_nets::IntranetNetsConf>,
44}
45
46impl KnowDbConf {
47 pub fn provider(&self) -> Option<ProviderConfig> {
48 self.provider_raw.clone()
49 }
50}
51
52#[derive(Debug, Clone, Deserialize)]
57pub struct FunSpec {
58 pub call: FunCall,
59 #[serde(default)]
60 pub key: Option<String>,
61 #[serde(default = "default_true")]
62 pub cache: bool,
63 #[serde(default)]
64 pub ttl_ms: Option<u64>,
65}
66
67impl FunSpec {
68 pub fn returns_bool(&self) -> bool {
70 matches!(self.call, FunCall::BfExists | FunCall::Sismember)
71 }
72}
73
74#[derive(Debug, Clone, Deserialize, PartialEq)]
75#[serde(rename_all = "snake_case")]
76pub enum FunCall {
77 BfExists,
78 Sismember,
79 Hget,
80 Get,
81}
82
83#[derive(Debug, Clone, Deserialize)]
88pub struct CacheSpec {
89 #[serde(default = "default_true")]
90 pub enabled: bool,
91 #[serde(default = "default_result_cache_capacity")]
92 pub capacity: usize,
93 #[serde(default = "default_result_cache_ttl_ms")]
94 pub ttl_ms: u64,
95}
96
97impl Default for CacheSpec {
98 fn default() -> Self {
99 Self {
100 enabled: default_true(),
101 capacity: default_result_cache_capacity(),
102 ttl_ms: default_result_cache_ttl_ms(),
103 }
104 }
105}
106
107#[derive(Debug, Clone, Default, Deserialize)]
112pub struct ProviderConfig {
113 #[serde(default)]
114 pub sqldb: Option<SqlProviderSpec>,
115 #[serde(default)]
116 pub redis: Option<RedisProviderSpec>,
117}
118
119#[derive(Debug, Clone, Deserialize)]
120pub struct SqlProviderSpec {
121 #[serde(rename = "kind")]
122 pub kind: SqlProviderKind,
123 pub connection_uri: String,
124 #[serde(default)]
125 pub pool_size: Option<u32>,
126 #[serde(default)]
127 pub min_connections: Option<u32>,
128 #[serde(default)]
129 pub acquire_timeout_ms: Option<u64>,
130 #[serde(default)]
131 pub idle_timeout_ms: Option<u64>,
132 #[serde(default)]
133 pub max_lifetime_ms: Option<u64>,
134}
135
136#[derive(Debug, Clone, Deserialize)]
140#[serde(rename_all = "snake_case")]
141pub enum SqlProviderKind {
142 Postgres,
143 Mysql,
144}
145
146#[derive(Debug, Clone, Deserialize)]
147pub struct RedisProviderSpec {
148 pub connection_uri: String,
149 #[serde(default)]
150 pub pool_size: Option<usize>,
151 #[serde(default = "default_connect_timeout_ms")]
152 pub connect_timeout_ms: u64,
153 #[serde(default = "default_command_timeout_ms")]
154 pub command_timeout_ms: u64,
155}
156
157fn default_connect_timeout_ms() -> u64 {
158 3_000
159}
160
161fn default_command_timeout_ms() -> u64 {
162 100
163}
164
165#[derive(Debug, Clone, Deserialize)]
171#[serde(rename_all = "snake_case")]
172pub enum ProviderKind {
173 SqliteAuthority,
174 Postgres,
175 Mysql,
176 Redis,
177}
178
179#[derive(Debug, Clone, Deserialize)]
180pub struct OptLoadSpec {
181 #[serde(default = "default_true")]
182 pub transaction: bool,
183 #[serde(default = "default_batch")]
184 pub batch_size: usize,
185 #[serde(default = "default_on_error")]
186 pub on_error: OnError,
187}
188impl Default for OptLoadSpec {
189 fn default() -> Self {
190 Self {
191 transaction: true,
192 batch_size: default_batch(),
193 on_error: default_on_error(),
194 }
195 }
196}
197
198#[derive(Debug, Clone, Deserialize, Default)]
199#[serde(rename_all = "lowercase")]
200pub enum OnError {
201 #[default]
202 Fail,
203 Skip,
204}
205
206#[derive(Debug, Clone, Deserialize)]
207pub struct CsvSpec {
208 #[serde(default = "default_true")]
209 pub has_header: bool,
210 #[serde(default = "default_comma")]
211 pub delimiter: String,
212 #[serde(default = "default_utf8")]
213 pub encoding: String,
214 #[serde(default = "default_true")]
215 pub trim: bool,
216}
217impl Default for CsvSpec {
218 fn default() -> Self {
219 CsvSpec {
220 has_header: true,
221 delimiter: ",".into(),
222 encoding: "utf-8".into(),
223 trim: true,
224 }
225 }
226}
227
228#[derive(Debug, Clone, Deserialize)]
229pub struct TableSpec {
230 pub name: String,
231 #[serde(default)]
232 pub dir: Option<String>,
233 #[serde(default)]
234 pub data_file: Option<String>,
235 pub columns: ColumnsSpec,
236 #[serde(default)]
237 pub expected_rows: RowExpect,
238 #[serde(default = "default_true")]
239 pub enabled: bool,
240}
241
242#[derive(Debug, Clone, Deserialize)]
243pub struct ColumnsSpec {
244 #[serde(default)]
245 pub by_header: Vec<String>,
246 #[serde(default)]
247 pub by_index: Vec<usize>,
248}
249
250#[derive(Debug, Clone, Deserialize, Default)]
251pub struct RowExpect {
252 pub min: Option<usize>,
253 pub max: Option<usize>,
254}
255
256const fn default_true() -> bool {
257 true
258}
259const fn default_batch() -> usize {
260 2000
261}
262fn default_comma() -> String {
263 ",".to_string()
264}
265fn default_utf8() -> String {
266 "utf-8".to_string()
267}
268fn default_on_error() -> OnError {
269 OnError::Fail
270}
271fn default_dot() -> String {
272 ".".to_string()
273}
274const fn default_result_cache_capacity() -> usize {
275 1024
276}
277const fn default_result_cache_ttl_ms() -> u64 {
278 30_000
279}
280
281fn read_to_string(path: &Path) -> KnowledgeResult<String> {
283 let mut f = fs::File::open(path).source_raw_err(KnowReason::from_res(), "source error")?;
284 let mut buf = String::new();
285 f.read_to_string(&mut buf)
286 .source_raw_err(KnowReason::from_res(), "source error")?;
287 Ok(buf)
288}
289
290fn replace_table(sql: &str, table: &str) -> String {
291 sql.replace("{table}", table)
292}
293
294fn join_rel(base: &Path, rel: &str) -> PathBuf {
295 let p = Path::new(rel);
296 if p.is_absolute() {
297 p.to_path_buf()
298 } else {
299 base.join(p)
300 }
301}
302
303pub fn build_authority_from_knowdb(
304 root: &Path,
305 conf_path: &Path,
306 authority_uri: &str,
307 dict: &EnvDict,
308) -> KnowledgeResult<Vec<String>> {
309 let mut opx = OperationContext::doing("build authority from knowdb").with_auto_log();
310 let (conf, conf_abs, base_dir) = parse_knowdb_conf(root, conf_path, dict)?;
312 opx.record("conf", conf_abs.display());
313 opx.record("base_dir", base_dir.display());
314 let db = open_authority(authority_uri)?;
316 let mut loaded_names = Vec::new();
318 for t in &conf.tables {
319 if !t.enabled {
320 continue;
321 }
322 load_one_table(&db, &base_dir, t, &conf.csv, &conf.default)?;
323 info_ctrl!("load table {} suc!", base_dir.display(),);
324 loaded_names.push(t.name.clone());
325 }
326 opx.mark_suc();
327 Ok(loaded_names)
328}
329
330pub fn parse_knowdb_conf(
331 root: &Path,
332 conf_path: &Path,
333 dict: &EnvDict,
334) -> KnowledgeResult<(KnowDbConf, PathBuf, PathBuf)> {
335 let conf_abs = if conf_path.is_absolute() {
336 conf_path.to_path_buf()
337 } else {
338 root.join(conf_path)
339 };
340 let conf_txt = read_to_string(&conf_abs)?;
341 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(&conf_txt, dict)
342 .source_err(KnowReason::from_conf(), "parse knowdb config")?;
343 if conf.version != 2 {
344 return Err(KnowReason::from_conf()
345 .to_err()
346 .with_detail("unsupported knowdb.version"));
347 }
348 crate::intranet_nets::set_intranet_nets_conf(conf.intranet_nets.clone());
350 let conf_dir = conf_abs.parent().unwrap_or_else(|| Path::new("."));
351 let base_dir = join_rel(conf_dir, &conf.base_dir);
352 Ok((conf, conf_abs, base_dir))
353}
354
355fn open_authority(authority_uri: &str) -> KnowledgeResult<MemDB> {
356 ensure_parent_dir_for_file_uri(authority_uri);
357 let flags = OpenFlags::SQLITE_OPEN_READ_WRITE
358 | OpenFlags::SQLITE_OPEN_CREATE
359 | OpenFlags::SQLITE_OPEN_URI;
360 let db = MemDB::new_file(authority_uri, 1, flags)?;
361 let _ = db.with_conn(|conn| {
363 let _ = crate::sqlite_ext::register_builtin(conn);
364 Ok::<(), anyhow::Error>(())
365 });
366 Ok(db)
367}
368
369fn ensure_parent_dir_for_file_uri(uri: &str) {
372 if let Some(rest) = uri.strip_prefix("file:") {
373 let path_part = rest.split('?').next().unwrap_or(rest);
374 let p = Path::new(path_part);
375 if let Some(parent) = p.parent() {
376 let _ = fs::create_dir_all(parent);
377 }
378 }
379}
380
381fn load_one_table(
382 db: &MemDB,
383 base_dir: &Path,
384 t: &TableSpec,
385 csvd: &CsvSpec,
386 load: &OptLoadSpec,
387) -> KnowledgeResult<()> {
388 let mut opx = OperationContext::doing("load table to kdb")
390 .with_auto_log()
391 .with_mod_path("ctrl");
392 let dir_name: &str = t.dir.as_deref().unwrap_or(&t.name);
393 let table_dir = base_dir.join(dir_name);
394 opx.record("table_dir", table_dir.display());
395 let create_sql = replace_table(&read_to_string(&table_dir.join("create.sql"))?, &t.name);
396 let insert_sql = replace_table(&read_to_string(&table_dir.join("insert.sql"))?, &t.name);
397 let clean_path = table_dir.join("clean.sql");
398 let clean_sql = if clean_path.exists() {
399 replace_table(&read_to_string(&clean_path)?, &t.name)
400 } else {
401 format!("DELETE FROM {}", t.name)
402 };
403
404 db.with_conn(|conn| {
406 let _ = crate::sqlite_ext::register_builtin(conn);
408 conn.execute_batch(&create_sql)?;
409 conn.execute_batch(&clean_sql)?;
410 Ok::<(), anyhow::Error>(())
411 })
412 .source_err(KnowReason::from_res(), "prepare authority table")?;
413
414 let data_path = match &t.data_file {
416 Some(rel) => join_rel(&table_dir, rel),
417 None => table_dir.join("data.csv"),
418 };
419 if !data_path.exists() {
420 return Err(KnowReason::from_conf()
421 .to_err()
422 .with_detail("data.csv not found"));
423 }
424 opx.record("data_path", data_path.display());
425
426 let mut rdr = build_csv_reader(csvd, &data_path)?;
428
429 let col_indices: Vec<usize> = if !t.columns.by_header.is_empty() {
431 let headers = rdr
432 .headers()
433 .source_raw_err(KnowReason::from_res(), "source error")?;
434 select_indices_by_header(headers, &t.columns.by_header)?
435 } else if !t.columns.by_index.is_empty() {
436 t.columns.by_index.clone()
437 } else {
438 return Err(KnowReason::from_conf()
439 .to_err()
440 .with_detail("columns mapping required"));
441 };
442
443 let mut inserted: usize = 0;
445 let mut bad: usize = 0;
446 let mut batch_left = load.batch_size.max(1);
447 db.with_conn(|conn| {
448 let _ = crate::sqlite_ext::register_builtin(conn);
450 let mut tx = if load.transaction {
451 Some(conn.unchecked_transaction()?)
452 } else {
453 None
454 };
455 let mut stmt = conn.prepare(&insert_sql)?;
456 for rec in rdr.into_records() {
457 match rec {
458 Ok(record) => {
459 let refs = extract_row_refs(&record, &col_indices, &mut bad, load)?;
460 if let Some(refs) = refs {
461 stmt.execute(rusqlite::params_from_iter(refs))?;
462 inserted += 1;
463 if load.transaction {
464 batch_left -= 1;
465 if batch_left == 0 {
466 tx.take().unwrap().commit()?;
467 tx = Some(conn.unchecked_transaction()?);
468 batch_left = load.batch_size.max(1);
469 }
470 }
471 }
472 }
473 Err(_e) => {
474 if matches!(load.on_error, OnError::Skip) {
475 bad += 1;
476 continue;
477 } else {
478 anyhow::bail!("csv record parse error");
479 }
480 }
481 }
482 }
483 if let Some(tx) = tx {
484 tx.commit()?;
485 }
486 Ok::<(), anyhow::Error>(())
487 })
488 .source_err(KnowReason::from_res(), "load authority table data")?;
489
490 if let Some(min) = t.expected_rows.min
492 && inserted < min
493 {
494 return Err(KnowReason::from_conf()
495 .to_err()
496 .with_detail("table data less"));
497 }
498 if let Some(max) = t.expected_rows.max
499 && inserted > max
500 {
501 wp_log::warn_kdb!(
502 "table {} loaded rows {} exceed max {}",
503 &t.name,
504 inserted,
505 max
506 );
507 }
508 if bad > 0 {
509 wp_log::warn_kdb!("table {} skipped {} bad rows (on_error=skip)", &t.name, bad);
510 }
511 opx.mark_suc();
512 Ok(())
513}
514
515fn build_csv_reader(
516 csvd: &CsvSpec,
517 data_path: &Path,
518) -> KnowledgeResult<csv::Reader<std::fs::File>> {
519 if csvd.encoding.to_lowercase() != "utf-8" {
520 return Err(KnowReason::from_conf()
521 .to_err()
522 .with_detail("only utf-8 csv is supported"));
523 }
524 let mut rdr_b = csv::ReaderBuilder::new();
525 rdr_b.has_headers(csvd.has_header);
526 if csvd.delimiter.len() == 1 {
527 rdr_b.delimiter(csvd.delimiter.as_bytes()[0]);
528 }
529 if csvd.trim {
530 rdr_b.trim(csv::Trim::All);
531 }
532 rdr_b
533 .from_path(data_path)
534 .source_raw_err(KnowReason::from_res(), "source error")
535}
536
537fn select_indices_by_header(
538 headers: &csv::StringRecord,
539 wanted: &[String],
540) -> KnowledgeResult<Vec<usize>> {
541 let mut out = Vec::with_capacity(wanted.len());
542 for name in wanted {
543 let pos = headers.iter().position(|h| h == name).ok_or_else(|| {
544 KnowReason::from_conf()
545 .to_err()
546 .with_detail("header not found")
547 })?;
548 out.push(pos);
549 }
550 Ok(out)
551}
552
553fn extract_row_refs<'a>(
554 record: &'a csv::StringRecord,
555 col_indices: &[usize],
556 bad: &mut usize,
557 load: &OptLoadSpec,
558) -> anyhow::Result<Option<Vec<&'a str>>> {
559 let mut vs: Vec<&str> = Vec::with_capacity(col_indices.len());
560 for &idx in col_indices {
561 if idx >= record.len() {
562 if matches!(load.on_error, OnError::Skip) {
563 *bad += 1;
564 return Ok(None);
565 } else {
566 anyhow::bail!("missing column at index {}", idx);
567 }
568 }
569 vs.push(record.get(idx).unwrap_or(""));
570 }
571 Ok(Some(vs))
572}
573
574#[cfg(test)]
575mod tests {
576 use super::*;
577
578 #[test]
579 fn parse_new_style_sqldb_provider() {
580 let dict = EnvDict::default();
581 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
582 r#"
583version = 2
584
585[provider.sqldb]
586kind = "postgres"
587connection_uri = "postgres://demo:demo@127.0.0.1/demo"
588pool_size = 12
589"#,
590 &dict,
591 )
592 .expect("parse knowdb with sqldb provider");
593
594 let sqldb = conf
595 .provider()
596 .expect("provider")
597 .sqldb
598 .expect("sqldb provider");
599 assert!(matches!(sqldb.kind, SqlProviderKind::Postgres));
600 assert_eq!(sqldb.pool_size, Some(12));
601 }
602
603 #[test]
604 fn parse_new_style_redis_provider() {
605 let dict = EnvDict::default();
606 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
607 r#"
608version = 2
609
610[provider.redis]
611connection_uri = "redis://127.0.0.1:6379"
612pool_size = 16
613connect_timeout_ms = 5000
614command_timeout_ms = 200
615"#,
616 &dict,
617 )
618 .expect("parse knowdb with redis provider");
619
620 let redis_cfg = conf
621 .provider()
622 .expect("provider")
623 .redis
624 .expect("redis provider");
625 assert_eq!(redis_cfg.connection_uri, "redis://127.0.0.1:6379");
626 assert_eq!(redis_cfg.pool_size, Some(16));
627 assert_eq!(redis_cfg.connect_timeout_ms, 5000);
628 assert_eq!(redis_cfg.command_timeout_ms, 200);
629 }
630
631 #[test]
632 fn parse_redis_provider_with_default_timeouts() {
633 let dict = EnvDict::default();
634 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
635 r#"
636version = 2
637
638[provider.redis]
639connection_uri = "redis://127.0.0.1:6379"
640"#,
641 &dict,
642 )
643 .expect("parse knowdb with redis provider (no timeout fields)");
644
645 let redis_cfg = conf.provider().expect("provider").redis.expect("redis");
646 assert_eq!(redis_cfg.connect_timeout_ms, 3000);
647 assert_eq!(redis_cfg.command_timeout_ms, 100);
648 }
649
650 #[test]
651 fn parse_both_sqldb_and_redis_providers() {
652 let dict = EnvDict::default();
653 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
654 r#"
655version = 2
656
657[provider.sqldb]
658kind = "postgres"
659connection_uri = "postgres://demo:demo@127.0.0.1/demo"
660
661[provider.redis]
662connection_uri = "redis://10.0.0.1:6379"
663pool_size = 4
664"#,
665 &dict,
666 )
667 .expect("parse knowdb with both sqldb and redis");
668
669 let provider_cfg = conf.provider().expect("provider");
670 let sqldb = provider_cfg.sqldb.expect("sqldb");
671 let redis_cfg = provider_cfg.redis.expect("redis");
672 assert!(matches!(sqldb.kind, SqlProviderKind::Postgres));
673 assert_eq!(redis_cfg.connection_uri, "redis://10.0.0.1:6379");
674 assert_eq!(redis_cfg.pool_size, Some(4));
675 }
676
677 #[test]
678 fn parse_redis_only_without_sqldb() {
679 let dict = EnvDict::default();
680 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
681 r#"
682version = 2
683
684[provider.redis]
685connection_uri = "redis://127.0.0.1:6379"
686"#,
687 &dict,
688 )
689 .expect("parse knowdb with redis only");
690
691 let provider_cfg = conf.provider().expect("provider");
692 assert!(provider_cfg.sqldb.is_none());
693 assert!(provider_cfg.redis.is_some());
694 }
695
696 #[test]
697 fn parse_no_provider_section() {
698 let dict = EnvDict::default();
699 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
700 r#"
701version = 2
702"#,
703 &dict,
704 )
705 .expect("parse knowdb without provider");
706
707 assert!(conf.provider().is_none());
708 }
709
710 #[test]
711 fn new_style_sqldb_mysql_variant() {
712 let dict = EnvDict::default();
713 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
714 r#"
715version = 2
716
717[provider.sqldb]
718kind = "mysql"
719connection_uri = "mysql://user:pass@127.0.0.1:3306/db"
720pool_size = 8
721"#,
722 &dict,
723 )
724 .expect("parse new-style mysql sqldb");
725
726 let sqldb = conf.provider().expect("provider").sqldb.expect("sqldb");
727 assert!(matches!(sqldb.kind, SqlProviderKind::Mysql));
728 assert_eq!(sqldb.pool_size, Some(8));
729 }
730
731 #[test]
732 fn parse_cache_spec_with_defaults() {
733 let dict = EnvDict::default();
734 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
735 r#"
736version = 2
737"#,
738 &dict,
739 )
740 .expect("parse knowdb with default cache spec");
741
742 assert!(conf.cache.enabled);
743 assert_eq!(conf.cache.capacity, 1024);
744 assert_eq!(conf.cache.ttl_ms, 30_000);
745 }
746
747 #[test]
748 fn parse_cache_spec_from_toml() {
749 let dict = EnvDict::default();
750 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
751 r#"
752version = 2
753
754[cache]
755enabled = false
756capacity = 256
757ttl_ms = 1500
758"#,
759 &dict,
760 )
761 .expect("parse knowdb with cache spec");
762
763 assert!(!conf.cache.enabled);
764 assert_eq!(conf.cache.capacity, 256);
765 assert_eq!(conf.cache.ttl_ms, 1500);
766 }
767
768 #[test]
769 fn parse_redis_cache_spec() {
770 let dict = EnvDict::default();
771 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
772 r#"
773version = 2
774
775[cache]
776enabled = true
777capacity = 512
778"#,
779 &dict,
780 )
781 .expect("parse knowdb with cache");
782
783 assert!(conf.cache.enabled);
784 assert_eq!(conf.cache.capacity, 512);
785 }
786
787 #[test]
788 fn parse_redis_cache_defaults() {
789 let dict = EnvDict::default();
790 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
791 r#"
792version = 2
793"#,
794 &dict,
795 )
796 .expect("parse knowdb without redis.cache");
797
798 assert!(conf.cache.enabled);
800 assert_eq!(conf.cache.capacity, 1024);
801 }
802
803 #[test]
808 fn parse_fun_bool_services() {
809 let dict = EnvDict::default();
810 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
811 r#"
812version = 2
813
814[fun.password_check]
815call = "bf_exists"
816key = "weak_passwords"
817
818[fun.ip_whitelist]
819call = "sismember"
820key = "allowed_ips"
821"#,
822 &dict,
823 )
824 .expect("parse fun bool services");
825
826 let pw = conf.fun.get("password_check").expect("password_check");
827 assert_eq!(pw.call, FunCall::BfExists);
828 assert_eq!(pw.key.as_deref(), Some("weak_passwords"));
829 assert!(pw.returns_bool());
830
831 let ip = conf.fun.get("ip_whitelist").expect("ip_whitelist");
832 assert_eq!(ip.call, FunCall::Sismember);
833 assert_eq!(ip.key.as_deref(), Some("allowed_ips"));
834 assert!(ip.returns_bool());
835 }
836
837 #[test]
838 fn parse_fun_value_services() {
839 let dict = EnvDict::default();
840 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
841 r#"
842version = 2
843
844[fun.threat_actor]
845call = "hget"
846key = "threat_actors"
847cache = true
848ttl_ms = 60000
849
850[fun.user_tag]
851call = "get"
852"#,
853 &dict,
854 )
855 .expect("parse fun value services");
856
857 let ta = conf.fun.get("threat_actor").expect("threat_actor");
858 assert_eq!(ta.call, FunCall::Hget);
859 assert_eq!(ta.key.as_deref(), Some("threat_actors"));
860 assert!(ta.cache);
861 assert_eq!(ta.ttl_ms, Some(60000));
862 assert!(!ta.returns_bool());
863
864 let ut = conf.fun.get("user_tag").expect("user_tag");
865 assert_eq!(ut.call, FunCall::Get);
866 assert!(ut.key.is_none());
867 assert!(ut.cache); assert!(!ut.returns_bool());
869 }
870
871 #[test]
872 fn parse_fun_default_cache() {
873 let dict = EnvDict::default();
874 let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
875 r#"
876version = 2
877
878[fun.app_config]
879call = "get"
880key = "app_config"
881"#,
882 &dict,
883 )
884 .expect("parse fun default cache");
885
886 let spec = conf.fun.get("app_config").expect("app_config");
887 assert!(spec.cache);
888 assert!(spec.ttl_ms.is_none());
889 }
890}