Skip to main content

wp_knowledge/
loader.rs

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/// V2 KnowDB 配置:目录式 + 外置 SQL。仅支持单一数据文件:`<table_dir>/data.csv`,
18/// 或通过 `tables[n].data_file` 相对 `<table_dir>` 指定。
19#[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    /// `[fun.<name>]` — external named-query definitions.
34    #[serde(default)]
35    pub fun: HashMap<String, FunSpec>,
36
37    /// Raw provider config — `[provider.sqldb]` / `[provider.redis]`.
38    #[serde(default, rename = "provider")]
39    provider_raw: Option<ProviderConfig>,
40
41    /// `[intranet_nets]` — 内网网段知识配置
42    #[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// ---------------------------------------------------------------------------
53// Fun (external named query) config
54// ---------------------------------------------------------------------------
55
56#[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    /// Derive return type from the call (bf_exists/sismember → bool, hget/get → value).
69    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// ---------------------------------------------------------------------------
84// Cache config
85// ---------------------------------------------------------------------------
86
87#[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// ---------------------------------------------------------------------------
108// Provider configuration (new format: [provider.sqldb] / [provider.redis])
109// ---------------------------------------------------------------------------
110
111#[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/// Config-level SQL provider kind — Postgres and Mysql only.
137/// The runtime-level [`ProviderKind`] additionally includes `SqliteAuthority`
138/// and `Redis` for the internal provider registry.
139#[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/// Runtime-level provider kind — used by the internal registry to identify
166/// the active provider. Includes built-in types (SqliteAuthority) and all
167/// external types (Postgres, Mysql, Redis).
168///
169/// For config-level SQL providers, see [`SqlProviderKind`].
170#[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
281/// 读取文本文件,返回字符串
282fn 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    // 1) 解析配置与 base_dir
311    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    // 2) 打开权威库
315    let db = open_authority(authority_uri)?;
316    // 3) 逐表加载(按配置顺序);不再处理显式依赖
317    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    // 注入内网网段配置(`[intranet_nets]` 节),供规则引擎消费
349    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    // 预注册内置 UDF 至权威库连接(注意:连接池可能返回不同连接,导入时也会再次注册)
362    let _ = db.with_conn(|conn| {
363        let _ = crate::sqlite_ext::register_builtin(conn);
364        Ok::<(), anyhow::Error>(())
365    });
366    Ok(db)
367}
368
369/// Kahn 拓扑排序:返回按依赖顺序的表索引列表。
370/// no topo_sort_tables: V2 简化版按配置顺序加载
371fn 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    // 目录与必须文件
389    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    // 建表与清理
405    db.with_conn(|conn| {
406        // 注册内置 UDF(导入连接)
407        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    // 数据源
415    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    // CSV 解析器
427    let mut rdr = build_csv_reader(csvd, &data_path)?;
428
429    // 列映射
430    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    // 导入(分批事务)
444    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        // 注册内置 UDF(用于 INSERT 绑定表达式)
449        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    // 行数校验
491    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        // No [cache] → defaults enabled=true, capacity=1024
799        assert!(conf.cache.enabled);
800        assert_eq!(conf.cache.capacity, 1024);
801    }
802
803    // -----------------------------------------------------------------------
804    // Fun (external named query) config tests
805    // -----------------------------------------------------------------------
806
807    #[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); // default true
868        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}