1use std::collections::{HashMap, HashSet};
2use std::fmt;
3use std::fs;
4use std::io::Read;
5use std::path::{Path, PathBuf};
6
7use orion_conf::EnvTomlLoad;
8use serde::Deserialize;
9use serde::de::{self, Deserializer, MapAccess, SeqAccess, Visitor};
10use wp_log::info_ctrl;
11
12use crate::error::{KnowReason, KnowledgeResult};
13use crate::mem::memdb::MemDB;
14use orion_error::OperationContext;
15use orion_error::conversion::{SourceErr, SourceRawErr, ToStructError};
16use orion_variate::EnvDict;
17use rusqlite::OpenFlags;
18
19#[derive(Debug, Deserialize)]
22pub struct KnowDbConf {
23 pub version: u32,
24 #[serde(default = "default_dot")]
25 pub base_dir: String,
26 #[serde(default)]
27 pub default: OptLoadSpec,
28 #[serde(default)]
29 pub csv: CsvSpec,
30 #[serde(default)]
31 pub cache: CacheSpec,
32 #[serde(default)]
33 pub tables: Vec<TableSpec>,
34
35 #[serde(default)]
37 pub fun: HashMap<String, FunSpec>,
38
39 #[serde(default, rename = "provider")]
41 provider_raw: Option<ProviderConfig>,
42
43 #[serde(default)]
45 pub intranet_nets: Option<crate::intranet_nets::IntranetNetsConf>,
46}
47
48impl KnowDbConf {
49 pub fn provider(&self) -> Option<ProviderConfig> {
50 self.provider_raw.clone()
51 }
52}
53
54#[derive(Debug, Clone, Deserialize)]
59pub struct FunSpec {
60 pub call: FunCall,
61 #[serde(default)]
62 pub key: Option<String>,
63 #[serde(default = "default_true")]
64 pub cache: bool,
65 #[serde(default)]
66 pub ttl_ms: Option<u64>,
67}
68
69impl FunSpec {
70 pub fn returns_bool(&self) -> bool {
72 matches!(self.call, FunCall::BfExists | FunCall::Sismember)
73 }
74}
75
76#[derive(Debug, Clone, Deserialize, PartialEq)]
77#[serde(rename_all = "snake_case")]
78pub enum FunCall {
79 BfExists,
80 Sismember,
81 Hget,
82 Get,
83}
84
85#[derive(Debug, Clone, Deserialize)]
90pub struct CacheSpec {
91 #[serde(default = "default_true")]
92 pub enabled: bool,
93 #[serde(default = "default_result_cache_capacity")]
94 pub capacity: usize,
95 #[serde(default = "default_result_cache_ttl_ms")]
96 pub ttl_ms: u64,
97}
98
99impl Default for CacheSpec {
100 fn default() -> Self {
101 Self {
102 enabled: default_true(),
103 capacity: default_result_cache_capacity(),
104 ttl_ms: default_result_cache_ttl_ms(),
105 }
106 }
107}
108
109pub const DEFAULT_SQLDB_NAME: &str = "default";
115
116#[derive(Debug, Clone, 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 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#[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 #[serde(default)]
150 pub jit: Option<bool>,
151 #[serde(default)]
153 pub application_name: Option<String>,
154}
155
156#[derive(Debug, Clone, Default, Deserialize)]
157pub struct ProviderConfig {
158 #[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 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 #[serde(default)]
199 pub postgres_session: Option<PostgresSessionSpec>,
200}
201
202impl SqlProviderSpec {
203 pub fn effective_name(&self) -> &str {
205 self.name.as_deref().unwrap_or(DEFAULT_SQLDB_NAME)
206 }
207
208 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 && (app.len() > 63 || app.chars().any(|c| c.is_control()))
231 {
232 return Err(KnowReason::from_conf().to_err().with_detail(format!(
233 "provider '{name}': postgres_session.application_name invalid \
234 (≤ 63 bytes, no control characters)"
235 )));
236 }
237 }
238 }
239 Ok(seen)
240 }
241}
242
243fn 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#[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#[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
424fn 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 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 let db = open_authority(authority_uri)?;
459 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 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 let _ = db.with_conn(|conn| {
506 let _ = crate::sqlite_ext::register_builtin(conn);
507 Ok::<(), anyhow::Error>(())
508 });
509 Ok(db)
510}
511
512fn 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 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 db.with_conn(|conn| {
549 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 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 let mut rdr = build_csv_reader(csvd, &data_path)?;
571
572 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 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 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 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 assert_eq!(specs[0].effective_name(), "geo");
1147 assert_eq!(specs[1].effective_name(), DEFAULT_SQLDB_NAME);
1148
1149 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 assert!(conf.cache.enabled);
1251 assert_eq!(conf.cache.capacity, 1024);
1252 }
1253
1254 #[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); 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}