sleep_ai_demo / db /src /sleep_db /rules.py
lokework's picture
Upload sleep full-agent demo
6025aa5 verified
Raw
History Blame Contribute Delete
3.26 kB
from __future__ import annotations
import operator
from pathlib import Path
from typing import Any
import yaml
from sleep_db.hashing import stable_file_hash
from sleep_db.schema import ALLOWED_LABELS, ALLOWED_REL_TYPES, REQUIRED_RULE_FIELDS
OPS = {
">": operator.gt,
">=": operator.ge,
"<": operator.lt,
"<=": operator.le,
"==": operator.eq,
}
def load_rule_file(path: Path) -> tuple[dict[str, Any], str]:
text = path.read_text(encoding="utf-8")
rule = yaml.safe_load(text) or {}
validate_rule(rule, path)
return rule, stable_file_hash(text)
def load_rules_directory(rules_dir: Path) -> list[tuple[dict[str, Any], str]]:
return [load_rule_file(path) for path in sorted(rules_dir.rglob("*.yaml"))]
def validate_rule(rule: dict[str, Any], path: Path | None = None) -> None:
missing = REQUIRED_RULE_FIELDS - set(rule)
if missing:
raise ValueError(f"{path}: missing rule fields: {sorted(missing)}")
if not isinstance(rule["source_doc_ids"], list) or not rule["source_doc_ids"]:
raise ValueError(f"{path}: source_doc_ids must be a non-empty list")
for node in rule.get("nodes", []):
label = node.get("label")
node_id = node.get("id")
if label not in ALLOWED_LABELS:
raise ValueError(f"{path}: unsupported node label {label!r}")
if not node_id:
raise ValueError(f"{path}: node is missing id")
node_ids = {node["id"] for node in rule.get("nodes", [])}
for edge in rule.get("edges", []):
rel_type = edge.get("type")
if rel_type not in ALLOWED_REL_TYPES:
raise ValueError(f"{path}: unsupported relationship type {rel_type!r}")
if edge.get("from") not in node_ids or edge.get("to") not in node_ids:
raise ValueError(f"{path}: edge references a missing node")
def eval_condition(condition: dict[str, Any], user_metrics: dict[str, Any], baseline_metrics: dict[str, Any]) -> bool:
metric_key = condition["metric_key"]
current = user_metrics.get(metric_key)
if current is None:
return False
threshold_type = condition["threshold_type"]
if threshold_type == "relative":
baseline = _baseline_value(condition, user_metrics, baseline_metrics)
if baseline is None:
return False
value = float(current) - float(baseline)
threshold = float(condition["offset"])
elif threshold_type == "absolute":
value = float(current)
threshold = float(condition["threshold"])
elif threshold_type == "relative_ratio":
baseline = _baseline_value(condition, user_metrics, baseline_metrics)
if baseline in (None, 0):
return False
value = float(current) / float(baseline)
threshold = float(condition["ratio"])
else:
return False
op = OPS.get(condition["operator"])
if not op:
return False
return bool(op(value, threshold))
def _baseline_value(
condition: dict[str, Any],
user_metrics: dict[str, Any],
baseline_metrics: dict[str, Any],
) -> Any:
baseline_key = condition.get("baseline_key")
if not baseline_key:
return None
return baseline_metrics.get(baseline_key, user_metrics.get(baseline_key))