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/// PostgreSQL `plan_cache_mode` 枚举值(`[provider.sqldb.postgres_session]`)。
117#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Deserialize)]
118#[serde(rename_all = "snake_case")]
119pub enum PlanCacheMode {
120    #[default]
121    Auto,
122    ForceGenericPlan,
123    ForceCustomPlan,
124}
125
126impl PlanCacheMode {
127    /// PostgreSQL `SET plan_cache_mode` 接受的取值。
128    pub fn to_sql_value(self) -> &'static str {
129        match self {
130            PlanCacheMode::Auto => "auto",
131            PlanCacheMode::ForceGenericPlan => "force_generic_plan",
132            PlanCacheMode::ForceCustomPlan => "force_custom_plan",
133        }
134    }
135}
136
137/// PostgreSQL 连接级 session 初始化配置(稳定执行计划用)。
138///
139/// 对应 `[provider.sqldb.postgres_session]`;仅 `kind = "postgres"` 生效,
140/// 其他 kind 配置该字段会在 [`SqlProviderSpec::validate_specs`] 被拒绝。
141/// 未配置(`None`)时保持默认行为;每一项 `None` 表示该项不下发。
142/// 不提供任意 after_connect_sql 入口。
143#[derive(Debug, Clone, Default, PartialEq, Deserialize)]
144#[serde(deny_unknown_fields)]
145pub struct PostgresSessionSpec {
146    #[serde(default)]
147    pub plan_cache_mode: Option<PlanCacheMode>,
148    /// `None` = 不下发 `SET jit`(保持数据库默认)。
149    #[serde(default)]
150    pub jit: Option<bool>,
151    /// `None` = 不下发 `SET application_name`。长度 ≤ 63 字节(PG `NAMEDATALEN-1`),无控制字符。
152    #[serde(default)]
153    pub application_name: Option<String>,
154}
155
156#[derive(Debug, Clone, Default, Deserialize)]
157pub struct ProviderConfig {
158    /// 支持 `[provider.sqldb]`(单个)或 `[[provider.sqldb]]`(多个)两种写法。
159    #[serde(default, deserialize_with = "deserialize_sqldb")]
160    pub sqldb: Option<Vec<SqlProviderSpec>>,
161    #[serde(default)]
162    pub redis: Option<RedisProviderSpec>,
163}
164
165impl ProviderConfig {
166    /// 返回全部 sqldb provider 的生效名列表(无显式 `name` 时为 `default`)。
167    pub fn sqldb_names(&self) -> Vec<String> {
168        self.sqldb
169            .as_ref()
170            .map(|specs| {
171                specs
172                    .iter()
173                    .map(|spec| spec.effective_name().to_string())
174                    .collect()
175            })
176            .unwrap_or_default()
177    }
178}
179
180#[derive(Debug, Clone, Deserialize)]
181pub struct SqlProviderSpec {
182    #[serde(default)]
183    pub name: Option<String>,
184    #[serde(rename = "kind")]
185    pub kind: SqlProviderKind,
186    pub connection_uri: String,
187    #[serde(default)]
188    pub pool_size: Option<u32>,
189    #[serde(default)]
190    pub min_connections: Option<u32>,
191    #[serde(default)]
192    pub acquire_timeout_ms: Option<u64>,
193    #[serde(default)]
194    pub idle_timeout_ms: Option<u64>,
195    #[serde(default)]
196    pub max_lifetime_ms: Option<u64>,
197    /// PostgreSQL 专属:连接级 session 初始化(稳定执行计划);仅 `kind = "postgres"` 消费。
198    #[serde(default)]
199    pub postgres_session: Option<PostgresSessionSpec>,
200}
201
202impl SqlProviderSpec {
203    /// 生效名称:显式 `name` 缺失时回退到 [`DEFAULT_SQLDB_NAME`]。
204    pub fn effective_name(&self) -> &str {
205        self.name.as_deref().unwrap_or(DEFAULT_SQLDB_NAME)
206    }
207
208    /// 校验 provider 名称字符集(仅 `[A-Za-z0-9_]`)与重名,返回生效名称集合。
209    pub fn validate_specs(specs: &[SqlProviderSpec]) -> KnowledgeResult<HashSet<String>> {
210        let mut seen = HashSet::with_capacity(specs.len());
211        for spec in specs {
212            let name = spec.effective_name();
213            if !name.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
214                return Err(KnowReason::from_conf().to_err().with_detail(format!(
215                    "invalid sqldb provider name '{name}' (allowed [A-Za-z0-9_])"
216                )));
217            }
218            if !seen.insert(name.to_string()) {
219                return Err(KnowReason::from_conf()
220                    .to_err()
221                    .with_detail(format!("duplicate sqldb provider name '{name}'")));
222            }
223            if let Some(session) = &spec.postgres_session {
224                if !matches!(spec.kind, SqlProviderKind::Postgres) {
225                    return Err(KnowReason::from_conf().to_err().with_detail(format!(
226                        "provider '{name}': postgres_session is only valid for kind = \"postgres\""
227                    )));
228                }
229                if let Some(app) = &session.application_name {
230                    if app.len() > 63 || app.chars().any(|c| c.is_control()) {
231                        return Err(KnowReason::from_conf().to_err().with_detail(format!(
232                            "provider '{name}': postgres_session.application_name invalid \
233                             (≤ 63 bytes, no control characters)"
234                        )));
235                    }
236                }
237            }
238        }
239        Ok(seen)
240    }
241}
242
243/// 同时接受 `[provider.sqldb]`(单表)与 `[[provider.sqldb]]`(数组)两种 TOML 写法。
244fn deserialize_sqldb<'de, D>(deserializer: D) -> Result<Option<Vec<SqlProviderSpec>>, D::Error>
245where
246    D: Deserializer<'de>,
247{
248    struct SqlDbVisitor;
249
250    impl<'de> Visitor<'de> for SqlDbVisitor {
251        type Value = Option<Vec<SqlProviderSpec>>;
252
253        fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
254            formatter
255                .write_str("a `[provider.sqldb]` table or a `[[provider.sqldb]]` array of tables")
256        }
257
258        fn visit_map<A>(self, map: A) -> Result<Self::Value, A::Error>
259        where
260            A: MapAccess<'de>,
261        {
262            let spec = SqlProviderSpec::deserialize(de::value::MapAccessDeserializer::new(map))?;
263            Ok(Some(vec![spec]))
264        }
265
266        fn visit_seq<A>(self, seq: A) -> Result<Self::Value, A::Error>
267        where
268            A: SeqAccess<'de>,
269        {
270            let specs =
271                Vec::<SqlProviderSpec>::deserialize(de::value::SeqAccessDeserializer::new(seq))?;
272            Ok(Some(specs))
273        }
274    }
275
276    deserializer.deserialize_any(SqlDbVisitor)
277}
278
279/// Config-level SQL provider kind — Postgres and Mysql only.
280/// The runtime-level [`ProviderKind`] additionally includes `SqliteAuthority`
281/// and `Redis` for the internal provider registry.
282#[derive(Debug, Clone, Deserialize)]
283#[serde(rename_all = "snake_case")]
284pub enum SqlProviderKind {
285    Postgres,
286    Mysql,
287}
288
289#[derive(Debug, Clone, Deserialize)]
290pub struct RedisProviderSpec {
291    pub connection_uri: String,
292    #[serde(default)]
293    pub pool_size: Option<usize>,
294    #[serde(default = "default_connect_timeout_ms")]
295    pub connect_timeout_ms: u64,
296    #[serde(default = "default_command_timeout_ms")]
297    pub command_timeout_ms: u64,
298}
299
300fn default_connect_timeout_ms() -> u64 {
301    3_000
302}
303
304fn default_command_timeout_ms() -> u64 {
305    100
306}
307
308/// Runtime-level provider kind — used by the internal registry to identify
309/// the active provider. Includes built-in types (SqliteAuthority) and all
310/// external types (Postgres, Mysql, Redis).
311///
312/// For config-level SQL providers, see [`SqlProviderKind`].
313#[derive(Debug, Clone, Deserialize)]
314#[serde(rename_all = "snake_case")]
315pub enum ProviderKind {
316    SqliteAuthority,
317    Postgres,
318    Mysql,
319    Redis,
320}
321
322#[derive(Debug, Clone, Deserialize)]
323pub struct OptLoadSpec {
324    #[serde(default = "default_true")]
325    pub transaction: bool,
326    #[serde(default = "default_batch")]
327    pub batch_size: usize,
328    #[serde(default = "default_on_error")]
329    pub on_error: OnError,
330}
331impl Default for OptLoadSpec {
332    fn default() -> Self {
333        Self {
334            transaction: true,
335            batch_size: default_batch(),
336            on_error: default_on_error(),
337        }
338    }
339}
340
341#[derive(Debug, Clone, Deserialize, Default)]
342#[serde(rename_all = "lowercase")]
343pub enum OnError {
344    #[default]
345    Fail,
346    Skip,
347}
348
349#[derive(Debug, Clone, Deserialize)]
350pub struct CsvSpec {
351    #[serde(default = "default_true")]
352    pub has_header: bool,
353    #[serde(default = "default_comma")]
354    pub delimiter: String,
355    #[serde(default = "default_utf8")]
356    pub encoding: String,
357    #[serde(default = "default_true")]
358    pub trim: bool,
359}
360impl Default for CsvSpec {
361    fn default() -> Self {
362        CsvSpec {
363            has_header: true,
364            delimiter: ",".into(),
365            encoding: "utf-8".into(),
366            trim: true,
367        }
368    }
369}
370
371#[derive(Debug, Clone, Deserialize)]
372pub struct TableSpec {
373    pub name: String,
374    #[serde(default)]
375    pub dir: Option<String>,
376    #[serde(default)]
377    pub data_file: Option<String>,
378    pub columns: ColumnsSpec,
379    #[serde(default)]
380    pub expected_rows: RowExpect,
381    #[serde(default = "default_true")]
382    pub enabled: bool,
383}
384
385#[derive(Debug, Clone, Deserialize)]
386pub struct ColumnsSpec {
387    #[serde(default)]
388    pub by_header: Vec<String>,
389    #[serde(default)]
390    pub by_index: Vec<usize>,
391}
392
393#[derive(Debug, Clone, Deserialize, Default)]
394pub struct RowExpect {
395    pub min: Option<usize>,
396    pub max: Option<usize>,
397}
398
399const fn default_true() -> bool {
400    true
401}
402const fn default_batch() -> usize {
403    2000
404}
405fn default_comma() -> String {
406    ",".to_string()
407}
408fn default_utf8() -> String {
409    "utf-8".to_string()
410}
411fn default_on_error() -> OnError {
412    OnError::Fail
413}
414fn default_dot() -> String {
415    ".".to_string()
416}
417const fn default_result_cache_capacity() -> usize {
418    1024
419}
420const fn default_result_cache_ttl_ms() -> u64 {
421    30_000
422}
423
424/// 读取文本文件,返回字符串
425fn read_to_string(path: &Path) -> KnowledgeResult<String> {
426    let mut f = fs::File::open(path).source_raw_err(KnowReason::from_res(), "source error")?;
427    let mut buf = String::new();
428    f.read_to_string(&mut buf)
429        .source_raw_err(KnowReason::from_res(), "source error")?;
430    Ok(buf)
431}
432
433fn replace_table(sql: &str, table: &str) -> String {
434    sql.replace("{table}", table)
435}
436
437fn join_rel(base: &Path, rel: &str) -> PathBuf {
438    let p = Path::new(rel);
439    if p.is_absolute() {
440        p.to_path_buf()
441    } else {
442        base.join(p)
443    }
444}
445
446pub fn build_authority_from_knowdb(
447    root: &Path,
448    conf_path: &Path,
449    authority_uri: &str,
450    dict: &EnvDict,
451) -> KnowledgeResult<Vec<String>> {
452    let mut opx = OperationContext::doing("build authority from knowdb").with_auto_log();
453    // 1) 解析配置与 base_dir
454    let (conf, conf_abs, base_dir) = parse_knowdb_conf(root, conf_path, dict)?;
455    opx.record("conf", conf_abs.display());
456    opx.record("base_dir", base_dir.display());
457    // 2) 打开权威库
458    let db = open_authority(authority_uri)?;
459    // 3) 逐表加载(按配置顺序);不再处理显式依赖
460    let mut loaded_names = Vec::new();
461    for t in &conf.tables {
462        if !t.enabled {
463            continue;
464        }
465        load_one_table(&db, &base_dir, t, &conf.csv, &conf.default)?;
466        info_ctrl!("load table {} suc!", base_dir.display(),);
467        loaded_names.push(t.name.clone());
468    }
469    opx.mark_suc();
470    Ok(loaded_names)
471}
472
473pub fn parse_knowdb_conf(
474    root: &Path,
475    conf_path: &Path,
476    dict: &EnvDict,
477) -> KnowledgeResult<(KnowDbConf, PathBuf, PathBuf)> {
478    let conf_abs = if conf_path.is_absolute() {
479        conf_path.to_path_buf()
480    } else {
481        root.join(conf_path)
482    };
483    let conf_txt = read_to_string(&conf_abs)?;
484    let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(&conf_txt, dict)
485        .source_err(KnowReason::from_conf(), "parse knowdb config")?;
486    if conf.version != 2 {
487        return Err(KnowReason::from_conf()
488            .to_err()
489            .with_detail("unsupported knowdb.version"));
490    }
491    // 注入内网网段配置(`[intranet_nets]` 节),供规则引擎消费
492    crate::intranet_nets::set_intranet_nets_conf(conf.intranet_nets.clone());
493    let conf_dir = conf_abs.parent().unwrap_or_else(|| Path::new("."));
494    let base_dir = join_rel(conf_dir, &conf.base_dir);
495    Ok((conf, conf_abs, base_dir))
496}
497
498fn open_authority(authority_uri: &str) -> KnowledgeResult<MemDB> {
499    ensure_parent_dir_for_file_uri(authority_uri);
500    let flags = OpenFlags::SQLITE_OPEN_READ_WRITE
501        | OpenFlags::SQLITE_OPEN_CREATE
502        | OpenFlags::SQLITE_OPEN_URI;
503    let db = MemDB::new_file(authority_uri, 1, flags)?;
504    // 预注册内置 UDF 至权威库连接(注意:连接池可能返回不同连接,导入时也会再次注册)
505    let _ = db.with_conn(|conn| {
506        let _ = crate::sqlite_ext::register_builtin(conn);
507        Ok::<(), anyhow::Error>(())
508    });
509    Ok(db)
510}
511
512/// Kahn 拓扑排序:返回按依赖顺序的表索引列表。
513/// no topo_sort_tables: V2 简化版按配置顺序加载
514fn ensure_parent_dir_for_file_uri(uri: &str) {
515    if let Some(rest) = uri.strip_prefix("file:") {
516        let path_part = rest.split('?').next().unwrap_or(rest);
517        let p = Path::new(path_part);
518        if let Some(parent) = p.parent() {
519            let _ = fs::create_dir_all(parent);
520        }
521    }
522}
523
524fn load_one_table(
525    db: &MemDB,
526    base_dir: &Path,
527    t: &TableSpec,
528    csvd: &CsvSpec,
529    load: &OptLoadSpec,
530) -> KnowledgeResult<()> {
531    // 目录与必须文件
532    let mut opx = OperationContext::doing("load table to kdb")
533        .with_auto_log()
534        .with_mod_path("ctrl");
535    let dir_name: &str = t.dir.as_deref().unwrap_or(&t.name);
536    let table_dir = base_dir.join(dir_name);
537    opx.record("table_dir", table_dir.display());
538    let create_sql = replace_table(&read_to_string(&table_dir.join("create.sql"))?, &t.name);
539    let insert_sql = replace_table(&read_to_string(&table_dir.join("insert.sql"))?, &t.name);
540    let clean_path = table_dir.join("clean.sql");
541    let clean_sql = if clean_path.exists() {
542        replace_table(&read_to_string(&clean_path)?, &t.name)
543    } else {
544        format!("DELETE FROM {}", t.name)
545    };
546
547    // 建表与清理
548    db.with_conn(|conn| {
549        // 注册内置 UDF(导入连接)
550        let _ = crate::sqlite_ext::register_builtin(conn);
551        conn.execute_batch(&create_sql)?;
552        conn.execute_batch(&clean_sql)?;
553        Ok::<(), anyhow::Error>(())
554    })
555    .source_err(KnowReason::from_res(), "prepare authority table")?;
556
557    // 数据源
558    let data_path = match &t.data_file {
559        Some(rel) => join_rel(&table_dir, rel),
560        None => table_dir.join("data.csv"),
561    };
562    if !data_path.exists() {
563        return Err(KnowReason::from_conf()
564            .to_err()
565            .with_detail("data.csv not found"));
566    }
567    opx.record("data_path", data_path.display());
568
569    // CSV 解析器
570    let mut rdr = build_csv_reader(csvd, &data_path)?;
571
572    // 列映射
573    let col_indices: Vec<usize> = if !t.columns.by_header.is_empty() {
574        let headers = rdr
575            .headers()
576            .source_raw_err(KnowReason::from_res(), "source error")?;
577        select_indices_by_header(headers, &t.columns.by_header)?
578    } else if !t.columns.by_index.is_empty() {
579        t.columns.by_index.clone()
580    } else {
581        return Err(KnowReason::from_conf()
582            .to_err()
583            .with_detail("columns mapping required"));
584    };
585
586    // 导入(分批事务)
587    let mut inserted: usize = 0;
588    let mut bad: usize = 0;
589    let mut batch_left = load.batch_size.max(1);
590    db.with_conn(|conn| {
591        // 注册内置 UDF(用于 INSERT 绑定表达式)
592        let _ = crate::sqlite_ext::register_builtin(conn);
593        let mut tx = if load.transaction {
594            Some(conn.unchecked_transaction()?)
595        } else {
596            None
597        };
598        let mut stmt = conn.prepare(&insert_sql)?;
599        for rec in rdr.into_records() {
600            match rec {
601                Ok(record) => {
602                    let refs = extract_row_refs(&record, &col_indices, &mut bad, load)?;
603                    if let Some(refs) = refs {
604                        stmt.execute(rusqlite::params_from_iter(refs))?;
605                        inserted += 1;
606                        if load.transaction {
607                            batch_left -= 1;
608                            if batch_left == 0 {
609                                tx.take().unwrap().commit()?;
610                                tx = Some(conn.unchecked_transaction()?);
611                                batch_left = load.batch_size.max(1);
612                            }
613                        }
614                    }
615                }
616                Err(_e) => {
617                    if matches!(load.on_error, OnError::Skip) {
618                        bad += 1;
619                        continue;
620                    } else {
621                        anyhow::bail!("csv record parse error");
622                    }
623                }
624            }
625        }
626        if let Some(tx) = tx {
627            tx.commit()?;
628        }
629        Ok::<(), anyhow::Error>(())
630    })
631    .source_err(KnowReason::from_res(), "load authority table data")?;
632
633    // 行数校验
634    if let Some(min) = t.expected_rows.min
635        && inserted < min
636    {
637        return Err(KnowReason::from_conf()
638            .to_err()
639            .with_detail("table data less"));
640    }
641    if let Some(max) = t.expected_rows.max
642        && inserted > max
643    {
644        wp_log::warn_kdb!(
645            "table {} loaded rows {} exceed max {}",
646            &t.name,
647            inserted,
648            max
649        );
650    }
651    if bad > 0 {
652        wp_log::warn_kdb!("table {} skipped {} bad rows (on_error=skip)", &t.name, bad);
653    }
654    opx.mark_suc();
655    Ok(())
656}
657
658fn build_csv_reader(
659    csvd: &CsvSpec,
660    data_path: &Path,
661) -> KnowledgeResult<csv::Reader<std::fs::File>> {
662    if csvd.encoding.to_lowercase() != "utf-8" {
663        return Err(KnowReason::from_conf()
664            .to_err()
665            .with_detail("only utf-8 csv is supported"));
666    }
667    let mut rdr_b = csv::ReaderBuilder::new();
668    rdr_b.has_headers(csvd.has_header);
669    if csvd.delimiter.len() == 1 {
670        rdr_b.delimiter(csvd.delimiter.as_bytes()[0]);
671    }
672    if csvd.trim {
673        rdr_b.trim(csv::Trim::All);
674    }
675    rdr_b
676        .from_path(data_path)
677        .source_raw_err(KnowReason::from_res(), "source error")
678}
679
680fn select_indices_by_header(
681    headers: &csv::StringRecord,
682    wanted: &[String],
683) -> KnowledgeResult<Vec<usize>> {
684    let mut out = Vec::with_capacity(wanted.len());
685    for name in wanted {
686        let pos = headers.iter().position(|h| h == name).ok_or_else(|| {
687            KnowReason::from_conf()
688                .to_err()
689                .with_detail("header not found")
690        })?;
691        out.push(pos);
692    }
693    Ok(out)
694}
695
696fn extract_row_refs<'a>(
697    record: &'a csv::StringRecord,
698    col_indices: &[usize],
699    bad: &mut usize,
700    load: &OptLoadSpec,
701) -> anyhow::Result<Option<Vec<&'a str>>> {
702    let mut vs: Vec<&str> = Vec::with_capacity(col_indices.len());
703    for &idx in col_indices {
704        if idx >= record.len() {
705            if matches!(load.on_error, OnError::Skip) {
706                *bad += 1;
707                return Ok(None);
708            } else {
709                anyhow::bail!("missing column at index {}", idx);
710            }
711        }
712        vs.push(record.get(idx).unwrap_or(""));
713    }
714    Ok(Some(vs))
715}
716
717#[cfg(test)]
718mod tests {
719    use super::*;
720
721    #[test]
722    fn parse_new_style_sqldb_provider() {
723        let dict = EnvDict::default();
724        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
725            r#"
726version = 2
727
728[provider.sqldb]
729kind = "postgres"
730connection_uri = "postgres://demo:demo@127.0.0.1/demo"
731pool_size = 12
732"#,
733            &dict,
734        )
735        .expect("parse knowdb with sqldb provider");
736
737        let sqldb = conf
738            .provider()
739            .expect("provider")
740            .sqldb
741            .expect("sqldb provider");
742        assert_eq!(sqldb.len(), 1);
743        let sqldb = &sqldb[0];
744        assert!(matches!(sqldb.kind, SqlProviderKind::Postgres));
745        assert_eq!(sqldb.pool_size, Some(12));
746        assert_eq!(sqldb.effective_name(), DEFAULT_SQLDB_NAME);
747    }
748
749    #[test]
750    fn parse_new_style_redis_provider() {
751        let dict = EnvDict::default();
752        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
753            r#"
754version = 2
755
756[provider.redis]
757connection_uri = "redis://127.0.0.1:6379"
758pool_size = 16
759connect_timeout_ms = 5000
760command_timeout_ms = 200
761"#,
762            &dict,
763        )
764        .expect("parse knowdb with redis provider");
765
766        let redis_cfg = conf
767            .provider()
768            .expect("provider")
769            .redis
770            .expect("redis provider");
771        assert_eq!(redis_cfg.connection_uri, "redis://127.0.0.1:6379");
772        assert_eq!(redis_cfg.pool_size, Some(16));
773        assert_eq!(redis_cfg.connect_timeout_ms, 5000);
774        assert_eq!(redis_cfg.command_timeout_ms, 200);
775    }
776
777    #[test]
778    fn parse_redis_provider_with_default_timeouts() {
779        let dict = EnvDict::default();
780        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
781            r#"
782version = 2
783
784[provider.redis]
785connection_uri = "redis://127.0.0.1:6379"
786"#,
787            &dict,
788        )
789        .expect("parse knowdb with redis provider (no timeout fields)");
790
791        let redis_cfg = conf.provider().expect("provider").redis.expect("redis");
792        assert_eq!(redis_cfg.connect_timeout_ms, 3000);
793        assert_eq!(redis_cfg.command_timeout_ms, 100);
794    }
795
796    #[test]
797    fn parse_both_sqldb_and_redis_providers() {
798        let dict = EnvDict::default();
799        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
800            r#"
801version = 2
802
803[provider.sqldb]
804kind = "postgres"
805connection_uri = "postgres://demo:demo@127.0.0.1/demo"
806
807[provider.redis]
808connection_uri = "redis://10.0.0.1:6379"
809pool_size = 4
810"#,
811            &dict,
812        )
813        .expect("parse knowdb with both sqldb and redis");
814
815        let provider_cfg = conf.provider().expect("provider");
816        let sqldb = &provider_cfg.sqldb.expect("sqldb")[0];
817        let redis_cfg = provider_cfg.redis.expect("redis");
818        assert!(matches!(sqldb.kind, SqlProviderKind::Postgres));
819        assert_eq!(redis_cfg.connection_uri, "redis://10.0.0.1:6379");
820        assert_eq!(redis_cfg.pool_size, Some(4));
821    }
822
823    #[test]
824    fn parse_redis_only_without_sqldb() {
825        let dict = EnvDict::default();
826        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
827            r#"
828version = 2
829
830[provider.redis]
831connection_uri = "redis://127.0.0.1:6379"
832"#,
833            &dict,
834        )
835        .expect("parse knowdb with redis only");
836
837        let provider_cfg = conf.provider().expect("provider");
838        assert!(provider_cfg.sqldb.is_none());
839        assert!(provider_cfg.redis.is_some());
840    }
841
842    #[test]
843    fn parse_no_provider_section() {
844        let dict = EnvDict::default();
845        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
846            r#"
847version = 2
848"#,
849            &dict,
850        )
851        .expect("parse knowdb without provider");
852
853        assert!(conf.provider().is_none());
854    }
855
856    #[test]
857    fn new_style_sqldb_mysql_variant() {
858        let dict = EnvDict::default();
859        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
860            r#"
861version = 2
862
863[provider.sqldb]
864kind = "mysql"
865connection_uri = "mysql://user:pass@127.0.0.1:3306/db"
866pool_size = 8
867"#,
868            &dict,
869        )
870        .expect("parse new-style mysql sqldb");
871
872        let sqldb = &conf.provider().expect("provider").sqldb.expect("sqldb")[0];
873        assert!(matches!(sqldb.kind, SqlProviderKind::Mysql));
874        assert_eq!(sqldb.pool_size, Some(8));
875    }
876
877    #[test]
878    fn parse_array_style_multi_sqldb_providers() {
879        let dict = EnvDict::default();
880        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
881            r#"
882version = 2
883
884[[provider.sqldb]]
885name = "geo"
886kind = "postgres"
887connection_uri = "postgres://demo@127.0.0.1:5432/geo_db"
888pool_size = 8
889
890[[provider.sqldb]]
891name = "asset"
892kind = "postgres"
893connection_uri = "postgres://demo@127.0.0.1:5432/asset_db"
894pool_size = 12
895"#,
896            &dict,
897        )
898        .expect("parse knowdb with multiple sqldb providers");
899
900        let specs = conf
901            .provider()
902            .expect("provider")
903            .sqldb
904            .expect("sqldb providers");
905        assert_eq!(specs.len(), 2);
906        assert_eq!(specs[0].name.as_deref(), Some("geo"));
907        assert_eq!(specs[0].effective_name(), "geo");
908        assert_eq!(
909            specs[0].connection_uri,
910            "postgres://demo@127.0.0.1:5432/geo_db"
911        );
912        assert_eq!(specs[0].pool_size, Some(8));
913        assert_eq!(specs[1].name.as_deref(), Some("asset"));
914        assert_eq!(specs[1].effective_name(), "asset");
915        assert_eq!(
916            specs[1].connection_uri,
917            "postgres://demo@127.0.0.1:5432/asset_db"
918        );
919        assert_eq!(specs[1].pool_size, Some(12));
920    }
921
922    #[test]
923    fn parse_single_sqldb_with_explicit_name() {
924        let dict = EnvDict::default();
925        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
926            r#"
927version = 2
928
929[provider.sqldb]
930name = "geo"
931kind = "postgres"
932connection_uri = "postgres://demo@127.0.0.1/geo_db"
933"#,
934            &dict,
935        )
936        .expect("parse single named sqldb");
937
938        let spec = &conf.provider().expect("provider").sqldb.expect("sqldb")[0];
939        assert_eq!(spec.effective_name(), "geo");
940    }
941
942    #[test]
943    fn sqldb_names_applies_default_name() {
944        let dict = EnvDict::default();
945        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
946            r#"
947version = 2
948
949[provider.sqldb]
950kind = "postgres"
951connection_uri = "postgres://demo@127.0.0.1/demo"
952"#,
953            &dict,
954        )
955        .expect("parse unnamed sqldb");
956
957        let names = conf.provider().expect("provider").sqldb_names();
958        assert_eq!(names, vec!["default".to_string()]);
959    }
960
961    #[test]
962    fn validate_specs_rejects_duplicate_effective_names() {
963        let dict = EnvDict::default();
964        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
965            r#"
966version = 2
967
968[[provider.sqldb]]
969kind = "postgres"
970connection_uri = "postgres://demo@127.0.0.1/db1"
971
972[[provider.sqldb]]
973kind = "postgres"
974connection_uri = "postgres://demo@127.0.0.1/db2"
975"#,
976            &dict,
977        )
978        .expect("parse two unnamed sqldb");
979
980        let specs = conf.provider().expect("provider").sqldb.expect("sqldb");
981        let err = SqlProviderSpec::validate_specs(&specs).expect_err("duplicate 'default' name");
982        assert!(
983            err.to_string()
984                .contains("duplicate sqldb provider name 'default'")
985        );
986    }
987
988    #[test]
989    fn validate_specs_rejects_invalid_name_charset() {
990        let specs = [SqlProviderSpec {
991            name: Some("bad name".to_string()),
992            kind: SqlProviderKind::Postgres,
993            connection_uri: "postgres://demo@127.0.0.1/db".to_string(),
994            pool_size: None,
995            min_connections: None,
996            acquire_timeout_ms: None,
997            idle_timeout_ms: None,
998            max_lifetime_ms: None,
999            postgres_session: None,
1000        }];
1001        let err = SqlProviderSpec::validate_specs(&specs).expect_err("invalid name charset");
1002        assert!(err.to_string().contains("invalid sqldb provider name"));
1003    }
1004
1005    #[test]
1006    fn parse_sqldb_postgres_session() {
1007        let dict = EnvDict::default();
1008        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1009            r#"
1010version = 2
1011
1012[provider.sqldb]
1013name = "geo"
1014kind = "postgres"
1015connection_uri = "postgres://demo@127.0.0.1/geo_db"
1016
1017[provider.sqldb.postgres_session]
1018plan_cache_mode = "force_generic_plan"
1019jit = false
1020application_name = "ip_geo_service"
1021"#,
1022            &dict,
1023        )
1024        .expect("parse postgres_session config");
1025
1026        let spec = &conf.provider().expect("provider").sqldb.expect("sqldb")[0];
1027        let session = spec.postgres_session.as_ref().expect("postgres_session");
1028        assert_eq!(
1029            session.plan_cache_mode,
1030            Some(PlanCacheMode::ForceGenericPlan)
1031        );
1032        assert_eq!(session.jit, Some(false));
1033        assert_eq!(session.application_name.as_deref(), Some("ip_geo_service"));
1034    }
1035
1036    #[test]
1037    fn parse_sqldb_postgres_session_unknown_key_rejected() {
1038        let dict = EnvDict::default();
1039        let r = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1040            r#"
1041version = 2
1042
1043[provider.sqldb]
1044kind = "postgres"
1045connection_uri = "postgres://demo@127.0.0.1/demo"
1046
1047[provider.sqldb.postgres_session]
1048plan_cache_mode = "auto"
1049after_connect_sql = "SET x = 1"
1050"#,
1051            &dict,
1052        );
1053        assert!(
1054            r.is_err(),
1055            "unknown postgres_session key (after_connect_sql) should be rejected"
1056        );
1057    }
1058
1059    #[test]
1060    fn validate_specs_rejects_session_on_non_postgres() {
1061        let specs = [SqlProviderSpec {
1062            name: Some("geo".to_string()),
1063            kind: SqlProviderKind::Mysql,
1064            connection_uri: "mysql://demo@127.0.0.1/db".to_string(),
1065            pool_size: None,
1066            min_connections: None,
1067            acquire_timeout_ms: None,
1068            idle_timeout_ms: None,
1069            max_lifetime_ms: None,
1070            postgres_session: Some(PostgresSessionSpec {
1071                plan_cache_mode: Some(PlanCacheMode::ForceGenericPlan),
1072                jit: None,
1073                application_name: None,
1074            }),
1075        }];
1076        let err = SqlProviderSpec::validate_specs(&specs).expect_err("mysql with postgres_session");
1077        assert!(err.to_string().contains("postgres_session"));
1078    }
1079
1080    #[test]
1081    fn validate_specs_rejects_application_name_too_long() {
1082        let specs = [SqlProviderSpec {
1083            name: Some("geo".to_string()),
1084            kind: SqlProviderKind::Postgres,
1085            connection_uri: "postgres://demo@127.0.0.1/db".to_string(),
1086            pool_size: None,
1087            min_connections: None,
1088            acquire_timeout_ms: None,
1089            idle_timeout_ms: None,
1090            max_lifetime_ms: None,
1091            postgres_session: Some(PostgresSessionSpec {
1092                plan_cache_mode: None,
1093                jit: None,
1094                application_name: Some("x".repeat(64)),
1095            }),
1096        }];
1097        let err = SqlProviderSpec::validate_specs(&specs).expect_err("application_name too long");
1098        assert!(err.to_string().contains("application_name"));
1099    }
1100
1101    #[test]
1102    fn parse_empty_sqldb_array_is_empty() {
1103        let dict = EnvDict::default();
1104        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1105            r#"
1106version = 2
1107
1108[provider]
1109sqldb = []
1110"#,
1111            &dict,
1112        )
1113        .expect("parse empty sqldb array");
1114
1115        let specs = conf
1116            .provider()
1117            .expect("provider")
1118            .sqldb
1119            .expect("sqldb present");
1120        assert!(specs.is_empty());
1121    }
1122
1123    #[test]
1124    fn mixed_named_and_unnamed_sqldb_applies_default_name() {
1125        let dict = EnvDict::default();
1126        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1127            r#"
1128version = 2
1129
1130[[provider.sqldb]]
1131name = "geo"
1132kind = "postgres"
1133connection_uri = "postgres://demo@127.0.0.1/geo_db"
1134
1135[[provider.sqldb]]
1136kind = "postgres"
1137connection_uri = "postgres://demo@127.0.0.1/asset_db"
1138"#,
1139            &dict,
1140        )
1141        .expect("parse mixed named/unnamed sqldb");
1142
1143        let specs = conf.provider().expect("provider").sqldb.expect("sqldb");
1144        assert_eq!(specs.len(), 2);
1145        // 有显式 name 的生效名是 name;无 name 的生效名是 default
1146        assert_eq!(specs[0].effective_name(), "geo");
1147        assert_eq!(specs[1].effective_name(), DEFAULT_SQLDB_NAME);
1148
1149        // validate_specs 对混合配置应通过,且默认库规则为 "default" 优先
1150        let expected = SqlProviderSpec::validate_specs(&specs).expect("valid specs");
1151        assert!(expected.contains(DEFAULT_SQLDB_NAME));
1152    }
1153
1154    #[test]
1155    fn explicit_default_name_is_recognized() {
1156        let dict = EnvDict::default();
1157        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1158            r#"
1159version = 2
1160
1161[[provider.sqldb]]
1162name = "default"
1163kind = "postgres"
1164connection_uri = "postgres://demo@127.0.0.1/main_db"
1165
1166[[provider.sqldb]]
1167name = "geo"
1168kind = "postgres"
1169connection_uri = "postgres://demo@127.0.0.1/geo_db"
1170"#,
1171            &dict,
1172        )
1173        .expect("parse explicit default name");
1174
1175        let specs = conf.provider().expect("provider").sqldb.expect("sqldb");
1176        assert_eq!(specs[0].effective_name(), DEFAULT_SQLDB_NAME);
1177        assert_eq!(specs[1].effective_name(), "geo");
1178        let expected = SqlProviderSpec::validate_specs(&specs).expect("valid specs");
1179        assert!(expected.contains(DEFAULT_SQLDB_NAME));
1180    }
1181
1182    #[test]
1183    fn parse_cache_spec_with_defaults() {
1184        let dict = EnvDict::default();
1185        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1186            r#"
1187version = 2
1188"#,
1189            &dict,
1190        )
1191        .expect("parse knowdb with default cache spec");
1192
1193        assert!(conf.cache.enabled);
1194        assert_eq!(conf.cache.capacity, 1024);
1195        assert_eq!(conf.cache.ttl_ms, 30_000);
1196    }
1197
1198    #[test]
1199    fn parse_cache_spec_from_toml() {
1200        let dict = EnvDict::default();
1201        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1202            r#"
1203version = 2
1204
1205[cache]
1206enabled = false
1207capacity = 256
1208ttl_ms = 1500
1209"#,
1210            &dict,
1211        )
1212        .expect("parse knowdb with cache spec");
1213
1214        assert!(!conf.cache.enabled);
1215        assert_eq!(conf.cache.capacity, 256);
1216        assert_eq!(conf.cache.ttl_ms, 1500);
1217    }
1218
1219    #[test]
1220    fn parse_redis_cache_spec() {
1221        let dict = EnvDict::default();
1222        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1223            r#"
1224version = 2
1225
1226[cache]
1227enabled = true
1228capacity = 512
1229"#,
1230            &dict,
1231        )
1232        .expect("parse knowdb with cache");
1233
1234        assert!(conf.cache.enabled);
1235        assert_eq!(conf.cache.capacity, 512);
1236    }
1237
1238    #[test]
1239    fn parse_redis_cache_defaults() {
1240        let dict = EnvDict::default();
1241        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1242            r#"
1243version = 2
1244"#,
1245            &dict,
1246        )
1247        .expect("parse knowdb without redis.cache");
1248
1249        // No [cache] → defaults enabled=true, capacity=1024
1250        assert!(conf.cache.enabled);
1251        assert_eq!(conf.cache.capacity, 1024);
1252    }
1253
1254    // -----------------------------------------------------------------------
1255    // Fun (external named query) config tests
1256    // -----------------------------------------------------------------------
1257
1258    #[test]
1259    fn parse_fun_bool_services() {
1260        let dict = EnvDict::default();
1261        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1262            r#"
1263version = 2
1264
1265[fun.password_check]
1266call = "bf_exists"
1267key = "weak_passwords"
1268
1269[fun.ip_whitelist]
1270call = "sismember"
1271key = "allowed_ips"
1272"#,
1273            &dict,
1274        )
1275        .expect("parse fun bool services");
1276
1277        let pw = conf.fun.get("password_check").expect("password_check");
1278        assert_eq!(pw.call, FunCall::BfExists);
1279        assert_eq!(pw.key.as_deref(), Some("weak_passwords"));
1280        assert!(pw.returns_bool());
1281
1282        let ip = conf.fun.get("ip_whitelist").expect("ip_whitelist");
1283        assert_eq!(ip.call, FunCall::Sismember);
1284        assert_eq!(ip.key.as_deref(), Some("allowed_ips"));
1285        assert!(ip.returns_bool());
1286    }
1287
1288    #[test]
1289    fn parse_fun_value_services() {
1290        let dict = EnvDict::default();
1291        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1292            r#"
1293version = 2
1294
1295[fun.threat_actor]
1296call = "hget"
1297key = "threat_actors"
1298cache = true
1299ttl_ms = 60000
1300
1301[fun.user_tag]
1302call = "get"
1303"#,
1304            &dict,
1305        )
1306        .expect("parse fun value services");
1307
1308        let ta = conf.fun.get("threat_actor").expect("threat_actor");
1309        assert_eq!(ta.call, FunCall::Hget);
1310        assert_eq!(ta.key.as_deref(), Some("threat_actors"));
1311        assert!(ta.cache);
1312        assert_eq!(ta.ttl_ms, Some(60000));
1313        assert!(!ta.returns_bool());
1314
1315        let ut = conf.fun.get("user_tag").expect("user_tag");
1316        assert_eq!(ut.call, FunCall::Get);
1317        assert!(ut.key.is_none());
1318        assert!(ut.cache); // default true
1319        assert!(!ut.returns_bool());
1320    }
1321
1322    #[test]
1323    fn parse_fun_default_cache() {
1324        let dict = EnvDict::default();
1325        let conf: KnowDbConf = <KnowDbConf as EnvTomlLoad<KnowDbConf>>::env_parse_toml(
1326            r#"
1327version = 2
1328
1329[fun.app_config]
1330call = "get"
1331key = "app_config"
1332"#,
1333            &dict,
1334        )
1335        .expect("parse fun default cache");
1336
1337        let spec = conf.fun.get("app_config").expect("app_config");
1338        assert!(spec.cache);
1339        assert!(spec.ttl_ms.is_none());
1340    }
1341}