Skip to main content

wp_knowledge/
loader.rs

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/// V2 KnowDB 配置:目录式 + 外置 SQL。仅支持单一数据文件:`<table_dir>/data.csv`,
20/// 或通过 `tables[n].data_file` 相对 `<table_dir>` 指定。
21#[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    /// `[fun.<name>]` — external named-query definitions.
36    #[serde(default)]
37    pub fun: HashMap<String, FunSpec>,
38
39    /// Raw provider config — `[provider.sqldb]` / `[provider.redis]`.
40    #[serde(default, rename = "provider")]
41    provider_raw: Option<ProviderConfig>,
42
43    /// `[intranet_nets]` — 内网网段知识配置
44    #[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// ---------------------------------------------------------------------------
55// Fun (external named query) config
56// ---------------------------------------------------------------------------
57
58#[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    /// Derive return type from the call (bf_exists/sismember → bool, hget/get → value).
71    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// ---------------------------------------------------------------------------
86// Cache config
87// ---------------------------------------------------------------------------
88
89#[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
109// ---------------------------------------------------------------------------
110// Provider configuration (new format: [provider.sqldb] / [provider.redis])
111// ---------------------------------------------------------------------------
112
113/// SQL 数据库 provider 的默认生效名:未显式 `name` 时使用。
114pub const DEFAULT_SQLDB_NAME: &str = "default";
115
116#[derive(Debug, Clone, Default, Deserialize)]
117pub struct ProviderConfig {
118    /// 支持 `[provider.sqldb]`(单个)或 `[[provider.sqldb]]`(多个)两种写法。
119    #[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    /// 返回全部 sqldb provider 的生效名列表(无显式 `name` 时为 `default`)。
127    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    /// 生效名称:显式 `name` 缺失时回退到 [`DEFAULT_SQLDB_NAME`]。
161    pub fn effective_name(&self) -> &str {
162        self.name.as_deref().unwrap_or(DEFAULT_SQLDB_NAME)
163    }
164
165    /// 校验 provider 名称字符集(仅 `[A-Za-z0-9_]`)与重名,返回生效名称集合。
166    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
185/// 同时接受 `[provider.sqldb]`(单表)与 `[[provider.sqldb]]`(数组)两种 TOML 写法。
186fn 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/// Config-level SQL provider kind — Postgres and Mysql only.
222/// The runtime-level [`ProviderKind`] additionally includes `SqliteAuthority`
223/// and `Redis` for the internal provider registry.
224#[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/// Runtime-level provider kind — used by the internal registry to identify
251/// the active provider. Includes built-in types (SqliteAuthority) and all
252/// external types (Postgres, Mysql, Redis).
253///
254/// For config-level SQL providers, see [`SqlProviderKind`].
255#[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
366/// 读取文本文件,返回字符串
367fn 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    // 1) 解析配置与 base_dir
396    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    // 2) 打开权威库
400    let db = open_authority(authority_uri)?;
401    // 3) 逐表加载(按配置顺序);不再处理显式依赖
402    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    // 注入内网网段配置(`[intranet_nets]` 节),供规则引擎消费
434    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    // 预注册内置 UDF 至权威库连接(注意:连接池可能返回不同连接,导入时也会再次注册)
447    let _ = db.with_conn(|conn| {
448        let _ = crate::sqlite_ext::register_builtin(conn);
449        Ok::<(), anyhow::Error>(())
450    });
451    Ok(db)
452}
453
454/// Kahn 拓扑排序:返回按依赖顺序的表索引列表。
455/// no topo_sort_tables: V2 简化版按配置顺序加载
456fn 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    // 目录与必须文件
474    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    // 建表与清理
490    db.with_conn(|conn| {
491        // 注册内置 UDF(导入连接)
492        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    // 数据源
500    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    // CSV 解析器
512    let mut rdr = build_csv_reader(csvd, &data_path)?;
513
514    // 列映射
515    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    // 导入(分批事务)
529    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        // 注册内置 UDF(用于 INSERT 绑定表达式)
534        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    // 行数校验
576    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        // 有显式 name 的生效名是 name;无 name 的生效名是 default
991        assert_eq!(specs[0].effective_name(), "geo");
992        assert_eq!(specs[1].effective_name(), DEFAULT_SQLDB_NAME);
993
994        // validate_specs 对混合配置应通过,且默认库规则为 "default" 优先
995        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        // No [cache] → defaults enabled=true, capacity=1024
1095        assert!(conf.cache.enabled);
1096        assert_eq!(conf.cache.capacity, 1024);
1097    }
1098
1099    // -----------------------------------------------------------------------
1100    // Fun (external named query) config tests
1101    // -----------------------------------------------------------------------
1102
1103    #[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); // default true
1164        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}