DCE / decision /engine /config_loader.py
That guy James Bond :)
Deploy Medical Intent Escalation API
af61b34
Raw
History Blame Contribute Delete
18 kB
"""
Configuration loader for the Decision Engine.
Loads and validates all YAML configuration from the decision/ folder:
- global/global_rules.yaml
- taxonomy/*/triggers.yaml, rules.yaml, flows/*.yaml
- primitives/*.yaml
- orchestrator/*.yaml
- context/fhir_mappings.yaml
- tenants/*/tenant.yaml
- journeys/*.yaml
"""
from __future__ import annotations
import logging
import os
import re
from pathlib import Path
from typing import Any, Dict, List, Optional, Set
import yaml
logger = logging.getLogger("decision.config_loader")
class DecisionConfigLoader:
"""Loads and caches all decision engine configuration from YAML files."""
def __init__(self, decision_dir: str | Path):
self.decision_dir = Path(decision_dir)
if not self.decision_dir.is_dir():
raise FileNotFoundError(
f"Decision config directory not found: {self.decision_dir}"
)
# Loaded config caches
self._global_rules: Optional[Dict[str, Any]] = None
self._taxonomy_triggers: Optional[Dict[str, Dict[str, Any]]] = None
self._taxonomy_rules: Optional[Dict[str, Dict[str, Any]]] = None
self._taxonomy_flows: Optional[Dict[str, Dict[str, Dict[str, Any]]]] = None
self._primitives: Optional[Dict[str, Dict[str, Any]]] = None
self._orchestrator: Optional[Dict[str, Dict[str, Any]]] = None
self._fhir_mappings: Optional[Dict[str, Any]] = None
self._tenants: Optional[Dict[str, Dict[str, Any]]] = None
self._journeys: Optional[Dict[str, Dict[str, Any]]] = None
# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
def load_all(self) -> None:
"""Load all configuration files. Call once at startup."""
logger.info("Loading decision engine configuration from %s", self.decision_dir)
self._load_global_rules()
self._load_taxonomy()
self._load_primitives()
self._load_orchestrator()
self._load_fhir_mappings()
self._load_tenants()
self._load_journeys()
self._validate()
logger.info(
"Decision config loaded: %d domains, %d primitives, %d tenants, %d journeys",
len(self._taxonomy_triggers or {}),
len(self._primitives or {}),
len(self._tenants or {}),
len(self._journeys or {}),
)
@property
def global_rules(self) -> Dict[str, Any]:
if self._global_rules is None:
self._load_global_rules()
return self._global_rules
@property
def taxonomy_triggers(self) -> Dict[str, Dict[str, Any]]:
"""domain_name -> triggers.yaml content"""
if self._taxonomy_triggers is None:
self._load_taxonomy()
return self._taxonomy_triggers
@property
def taxonomy_rules(self) -> Dict[str, Dict[str, Any]]:
"""domain_name -> rules.yaml content"""
if self._taxonomy_rules is None:
self._load_taxonomy()
return self._taxonomy_rules
@property
def taxonomy_flows(self) -> Dict[str, Dict[str, Dict[str, Any]]]:
"""domain_name -> {flow_id -> flow.yaml content}"""
if self._taxonomy_flows is None:
self._load_taxonomy()
return self._taxonomy_flows
@property
def primitives(self) -> Dict[str, Dict[str, Any]]:
if self._primitives is None:
self._load_primitives()
return self._primitives
@property
def orchestrator(self) -> Dict[str, Dict[str, Any]]:
if self._orchestrator is None:
self._load_orchestrator()
return self._orchestrator
@property
def fhir_mappings(self) -> Dict[str, Any]:
if self._fhir_mappings is None:
self._load_fhir_mappings()
return self._fhir_mappings
@property
def tenants(self) -> Dict[str, Dict[str, Any]]:
if self._tenants is None:
self._load_tenants()
return self._tenants
@property
def journeys(self) -> Dict[str, Dict[str, Any]]:
if self._journeys is None:
self._load_journeys()
return self._journeys
# ------------------------------------------------------------------
# Derived accessors
# ------------------------------------------------------------------
@property
def safety_precedence(self) -> List[str]:
"""Ordered list of domains by clinical acuity (highest first)."""
return self.global_rules.get("safety_precedence", [])
@property
def hard_escalate_domains(self) -> Set[str]:
"""Domains that always escalate regardless of confidence.
Includes both explicitly listed domains from global_rules AND
any taxonomy domain whose rules.yaml declares target_risk_class: R3.
"""
explicit = set(self.global_rules.get("hard_escalate_domains", []))
# Auto-include every domain with target_risk_class R3
for domain_name, rules in (self._taxonomy_rules or {}).items():
if rules.get("target_risk_class") == "R3":
explicit.add(domain_name)
return explicit
@property
def risk_escalation_rules(self) -> List[Dict[str, Any]]:
"""Risk escalation rule definitions from global_rules."""
return self.global_rules.get("risk_escalation_rules", [])
@property
def domain_suppression_rules(self) -> List[Dict[str, Any]]:
"""Domain suppression rules from global_rules."""
return self.global_rules.get("domain_suppression", [])
@property
def risk_class_descriptions(self) -> Dict[str, str]:
"""R0/R1/R2/R3 descriptions."""
return self.global_rules.get("risk_classes", {})
@property
def domain_names(self) -> List[str]:
"""All known domain names from taxonomy."""
return sorted(self.taxonomy_triggers.keys())
def get_tenant(self, tenant_id: str) -> Optional[Dict[str, Any]]:
"""Get tenant config by ID."""
return self.tenants.get(tenant_id)
def get_domain_priority(self, domain: str) -> int:
"""Get priority for a domain (higher = more urgent). Default 0."""
triggers = self.taxonomy_triggers.get(domain, {})
return triggers.get("priority", 0)
def get_domain_target_risk(self, domain: str) -> Optional[str]:
"""Get default target risk class for a domain from rules.yaml."""
rules = self.taxonomy_rules.get(domain, {})
return rules.get("target_risk_class")
def get_domain_decision_rules(self, domain: str) -> List[Dict[str, Any]]:
"""Get decision rules for a domain."""
rules = self.taxonomy_rules.get(domain, {})
return rules.get("decision_rules", [])
# ------------------------------------------------------------------
# Internal loaders
# ------------------------------------------------------------------
def _safe_load_yaml(self, path: Path) -> Dict[str, Any]:
"""Load a YAML file safely, returning empty dict on error."""
try:
with open(path, "r", encoding="utf-8") as f:
data = yaml.safe_load(f)
return data if isinstance(data, dict) else {}
except Exception as e:
logger.warning("Failed to load YAML %s: %s", path, e)
return {}
def _load_global_rules(self) -> None:
path = self.decision_dir / "global" / "global_rules.yaml"
self._global_rules = self._safe_load_yaml(path)
if not self._global_rules:
logger.error("CRITICAL: global_rules.yaml is empty or missing at %s", path)
def _load_taxonomy(self) -> None:
self._taxonomy_triggers = {}
self._taxonomy_rules = {}
self._taxonomy_flows = {}
taxonomy_dir = self.decision_dir / "taxonomy"
if not taxonomy_dir.is_dir():
logger.warning("Taxonomy directory not found: %s", taxonomy_dir)
return
for domain_dir in sorted(taxonomy_dir.iterdir()):
if not domain_dir.is_dir():
continue
domain_name = domain_dir.name
# Load triggers.yaml
triggers_path = domain_dir / "triggers.yaml"
if triggers_path.is_file():
triggers = self._safe_load_yaml(triggers_path)
if triggers:
self._taxonomy_triggers[domain_name] = triggers
# Pre-compile regex patterns for performance
self._compile_trigger_regex(domain_name, triggers)
# Load rules.yaml
rules_path = domain_dir / "rules.yaml"
if rules_path.is_file():
rules = self._safe_load_yaml(rules_path)
if rules:
self._taxonomy_rules[domain_name] = rules
# Load flows/*.yaml
flows_dir = domain_dir / "flows"
if flows_dir.is_dir():
domain_flows = {}
for flow_file in sorted(flows_dir.iterdir()):
if flow_file.suffix in (".yaml", ".yml"):
flow_data = self._safe_load_yaml(flow_file)
if flow_data:
flow_id = flow_data.get("flow_id", flow_file.stem)
domain_flows[flow_id] = flow_data
if domain_flows:
self._taxonomy_flows[domain_name] = domain_flows
def _compile_trigger_regex(
self, domain_name: str, triggers: Dict[str, Any]
) -> None:
"""Pre-compile regex patterns in trigger definitions for performance."""
lexical = triggers.get("lexical_signals", {})
for tier_name in ("high", "medium", "low"):
tier = lexical.get(tier_name, {})
raw_patterns = tier.get("regex", [])
compiled = []
for pattern in raw_patterns:
try:
compiled.append(re.compile(pattern, re.IGNORECASE))
except re.error as e:
logger.warning(
"Invalid regex in %s/%s: %r → %s",
domain_name,
tier_name,
pattern,
e,
)
tier["_compiled_regex"] = compiled
# Compile negation patterns
negation = triggers.get("negation_handling", {})
neg_patterns = negation.get("patterns", [])
negation["_compiled_patterns"] = [
p.lower() for p in neg_patterns
]
def load_additional_taxonomy(self, taxonomy_dir: str | Path) -> int:
"""
Load additional taxonomy domains from a secondary directory.
Only loads domains that are NOT already present in the primary taxonomy.
Returns the number of new domains added.
"""
taxonomy_dir = Path(taxonomy_dir)
if not taxonomy_dir.is_dir():
logger.warning("Additional taxonomy directory not found: %s", taxonomy_dir)
return 0
if self._taxonomy_triggers is None:
self._taxonomy_triggers = {}
if self._taxonomy_rules is None:
self._taxonomy_rules = {}
if self._taxonomy_flows is None:
self._taxonomy_flows = {}
added = 0
for domain_dir in sorted(taxonomy_dir.iterdir()):
if not domain_dir.is_dir():
continue
domain_name = domain_dir.name
# Skip domains already loaded from primary taxonomy
if domain_name in self._taxonomy_triggers:
continue
# Load triggers.yaml
triggers_path = domain_dir / "triggers.yaml"
if triggers_path.is_file():
triggers = self._safe_load_yaml(triggers_path)
if triggers:
self._taxonomy_triggers[domain_name] = triggers
self._compile_trigger_regex(domain_name, triggers)
# Load rules.yaml
rules_path = domain_dir / "rules.yaml"
if rules_path.is_file():
rules = self._safe_load_yaml(rules_path)
if rules:
self._taxonomy_rules[domain_name] = rules
# Load flows/*.yaml
flows_dir = domain_dir / "flows"
if flows_dir.is_dir():
domain_flows = {}
for flow_file in sorted(flows_dir.iterdir()):
if flow_file.suffix in (".yaml", ".yml"):
flow_data = self._safe_load_yaml(flow_file)
if flow_data:
flow_id = flow_data.get("flow_id", flow_file.stem)
domain_flows[flow_id] = flow_data
if domain_flows:
self._taxonomy_flows[domain_name] = domain_flows
added += 1
if added:
logger.info(
"Loaded %d additional taxonomy domains from %s (total: %d)",
added, taxonomy_dir, len(self._taxonomy_triggers),
)
return added
def _load_primitives(self) -> None:
self._primitives = {}
prim_dir = self.decision_dir / "primitives"
if not prim_dir.is_dir():
return
for f in sorted(prim_dir.iterdir()):
if f.suffix in (".yaml", ".yml"):
data = self._safe_load_yaml(f)
if data:
flow_id = data.get("flow_id", f.stem)
self._primitives[flow_id] = data
def _load_orchestrator(self) -> None:
self._orchestrator = {}
orch_dir = self.decision_dir / "orchestrator"
if not orch_dir.is_dir():
return
for f in sorted(orch_dir.iterdir()):
if f.suffix in (".yaml", ".yml"):
data = self._safe_load_yaml(f)
if data:
self._orchestrator[f.stem] = data
def _load_fhir_mappings(self) -> None:
path = self.decision_dir / "context" / "fhir_mappings.yaml"
self._fhir_mappings = self._safe_load_yaml(path) if path.is_file() else {}
def _load_tenants(self) -> None:
self._tenants = {}
tenants_dir = self.decision_dir / "tenants"
if not tenants_dir.is_dir():
return
for tenant_dir in sorted(tenants_dir.iterdir()):
if not tenant_dir.is_dir():
continue
tenant_path = tenant_dir / "tenant.yaml"
if tenant_path.is_file():
data = self._safe_load_yaml(tenant_path)
if data:
tid = data.get("tenant_id", tenant_dir.name)
self._tenants[tid] = data
def _load_journeys(self) -> None:
self._journeys = {}
journeys_dir = self.decision_dir / "journeys"
if not journeys_dir.is_dir():
return
for f in sorted(journeys_dir.iterdir()):
if f.suffix in (".yaml", ".yml"):
data = self._safe_load_yaml(f)
if data:
wt = data.get("wedge_type", f.stem)
self._journeys[wt] = data
# ------------------------------------------------------------------
# Validation
# ------------------------------------------------------------------
def _validate(self) -> None:
"""Run basic validation checks on loaded configuration."""
errors = []
# Check global rules has required keys
required_global_keys = [
"safety_precedence",
"risk_classes",
"risk_escalation_rules",
"hard_escalate_domains",
]
for key in required_global_keys:
if key not in self.global_rules:
errors.append(f"global_rules.yaml missing required key: {key}")
# Check every domain in safety_precedence has triggers
for domain in self.safety_precedence:
if domain not in self.taxonomy_triggers:
errors.append(
f"safety_precedence domain '{domain}' has no triggers.yaml"
)
# Check hard_escalate domains exist
for domain in self.hard_escalate_domains:
if domain not in self.taxonomy_triggers:
errors.append(
f"hard_escalate domain '{domain}' has no triggers.yaml"
)
# Check risk escalation rules reference valid domains
for rule in self.risk_escalation_rules:
rule_domain = rule.get("if", {}).get("domain")
if rule_domain and rule_domain not in self.taxonomy_triggers:
errors.append(
f"risk rule '{rule.get('id')}' references unknown domain "
f"'{rule_domain}'"
)
# Check tenant wedge references
for tid, tenant in self.tenants.items():
for wedge in tenant.get("enabled_wedges", []):
# Map short names to journey keys
if wedge not in self.journeys:
# Try mapping common abbreviations
mapping = {
"pde": "pde",
"pre_op": "pre_op",
"awv": "awv",
"gap_closure": "gap_closure",
"hra": "hra",
}
if wedge not in mapping:
logger.debug(
"Tenant '%s' references wedge '%s' not in journeys",
tid,
wedge,
)
if errors:
for err in errors:
logger.warning("Config validation: %s", err)
logger.warning(
"Decision config has %d validation warning(s)", len(errors)
)
else:
logger.info("Decision config validation passed")