# -*- coding: utf-8 -*- """规则加载器(RULE LOADER)。 按 ARCHITECTURE.md Phase 7 设计,从 ``dl/rules/patterns/*.yaml`` 加载正则规则 和关键词分类规则,编译为 ``re.Pattern`` 对象并缓存,供业务层使用。 使用方式:: from BiddingKG.dl.rules.loader import RuleLoader # 1. 获取编译后的正则对象(与 re.compile() 返回类型一致) pattern = RuleLoader.get_pattern("failure_keywords") if pattern.search(text): ... # 2. 使用 YAML 中指定的匹配方法(search/match/fullmatch) if RuleLoader.match("rank_first", text): return "win_tenderer" # 3. 关键词分类 label = RuleLoader.classify("money_source", "财政拨款") # -> "财政资金,上级拨款" 依赖方向:rules 只依赖 common,不反向依赖业务层。 """ from __future__ import absolute_import import os import re import threading from typing import Dict, List, Optional, Any, Pattern as RegexPattern import yaml # ---------------------------------------------------------------------------- # YAML Loader:on/off/yes/no 不解析为布尔(YAML 1.1 陷阱) # ---------------------------------------------------------------------------- # PyYAML 按 YAML 1.1 把裸标量 on/off/yes/no 解析为 bool,导致 guard 条件的 # `on: behind` 键变成 True、cond["on"] 取不到值而静默回退默认切片。 # 此处仅保留 true/false 的布尔解析,on/off/yes/no 保持字符串。 import re as _re class _RuleYamlLoader(yaml.SafeLoader): """规则 YAML 专用 Loader(on/off/yes/no 保持字符串)。""" _RuleYamlLoader.yaml_implicit_resolvers = { key: [(tag, regexp) for tag, regexp in resolvers if tag != 'tag:yaml.org,2002:bool'] for key, resolvers in yaml.SafeLoader.yaml_implicit_resolvers.items() } _RuleYamlLoader.add_implicit_resolver( 'tag:yaml.org,2002:bool', _re.compile(r'^(?:true|True|TRUE|false|False|FALSE)$'), list('tTfF'), ) def load_rule_yaml(stream): """加载规则 YAML(on/off/yes/no 保持字符串,其余语义同 safe_load)。""" return yaml.load(stream, Loader=_RuleYamlLoader) # ---------------------------------------------------------------------------- # 路径常量 # ---------------------------------------------------------------------------- _RULES_DIR = os.path.dirname(os.path.abspath(__file__)) _PATTERNS_DIR = os.path.join(_RULES_DIR, "patterns") # ---------------------------------------------------------------------------- # 正则标志映射 # ---------------------------------------------------------------------------- _FLAG_MAP = { "IGNORECASE": re.IGNORECASE, "MULTILINE": re.MULTILINE, "DOTALL": re.DOTALL, "UNICODE": re.UNICODE, "VERBOSE": re.VERBOSE, } class RuleLoader: """规则加载器(单例缓存)。 所有方法均为类方法,通过类级缓存避免重复加载。 线程安全:首次加载使用锁保护。 """ # 类级缓存 _patterns: Dict[str, RegexPattern] = {} # id → compiled regex _pattern_meta: Dict[str, dict] = {} # id → {method, flags, description, ...} _rules: List[dict] = [] # 全量规则 dict(Phase C 引擎消费,保持 YAML 顺序) _classifiers: Dict[str, dict] = {} # file_stem → {config, rules} _loaded: bool = False _lock = threading.Lock() # ------------------------------------------------------------------------ # 加载 # ------------------------------------------------------------------------ @classmethod def _ensure_loaded(cls): """惰性加载所有 YAML 规则文件(线程安全)。""" if cls._loaded: return with cls._lock: if cls._loaded: # double-check return cls._load_all() cls._loaded = True @classmethod def _load_all(cls): """扫描 patterns/ 目录,加载所有 .yaml 文件。""" if not os.path.isdir(_PATTERNS_DIR): return for fname in sorted(os.listdir(_PATTERNS_DIR)): if not fname.endswith((".yaml", ".yml")): continue fpath = os.path.join(_PATTERNS_DIR, fname) stem = os.path.splitext(fname)[0] with open(fpath, encoding="utf-8") as f: data = load_rule_yaml(f) if not data: continue cls._load_patterns(data, stem) cls._load_classifier(data, stem) @classmethod def _load_patterns(cls, data: dict, stem: str): """加载正则规则。""" patterns = data.get("patterns") or [] for p in patterns: pid = p["id"] pattern_str = p["pattern"] flags = 0 for f in p.get("flags", []): flags |= _FLAG_MAP.get(f, 0) compiled = re.compile(pattern_str, flags) cls._patterns[pid] = compiled cls._pattern_meta[pid] = { "method": p.get("method", "search"), "description": p.get("description", ""), "category": data.get("category", "OTHER"), "source_file": stem, "pattern_string": pattern_str, } # Phase C:保留完整规则结构(stage/guard/set_prob 等)供引擎消费 rule = dict(p) rule["pattern_compiled"] = compiled rule["category"] = data.get("category", "OTHER") rule["source_file"] = stem cls._rules.append(rule) @classmethod def _load_classifier(cls, data: dict, stem: str): """加载关键词分类规则。""" rules = data.get("keyword_rules") if not rules: return config = data.get("classifier_config", {}) cls._classifiers[stem] = { "config": { "min_matches": config.get("min_matches", 1), "max_matches": config.get("max_matches", 99), "default_label": config.get("default_label", ""), "separator": config.get("separator", ","), "sort": config.get("sort", True), }, "rules": rules, } # ------------------------------------------------------------------------ # 公共 API:正则 # ------------------------------------------------------------------------ @classmethod def get_pattern(cls, pattern_id: str) -> RegexPattern: """获取编译后的正则对象。 返回 ``re.Pattern`` 实例,可直接调用 ``.search()``, ``.match()``, ``.findall()``, ``.finditer()`` 等方法,与 ``re.compile()`` 返回类型一致。 :param pattern_id: 规则 ID(YAML 中的 ``id`` 字段) :return: 编译后的 ``re.Pattern`` 对象 :raises KeyError: 规则 ID 不存在 """ cls._ensure_loaded() if pattern_id not in cls._patterns: raise KeyError("未找到规则 ID: %s(可用: %s)" % ( pattern_id, ", ".join(sorted(cls._patterns.keys())))) return cls._patterns[pattern_id] @classmethod def match(cls, pattern_id: str, text: str) -> Optional[re.Match]: """使用 YAML 中指定的匹配方法进行匹配。 根据 YAML 中的 ``method`` 字段自动选择: - ``search`` → ``pattern.search(text)`` - ``match`` → ``pattern.match(text)`` - ``fullmatch`` → ``pattern.fullmatch(text)`` :param pattern_id: 规则 ID :param text: 待匹配文本 :return: ``re.Match`` 对象或 ``None`` """ cls._ensure_loaded() pattern = cls.get_pattern(pattern_id) method = cls._pattern_meta.get(pattern_id, {}).get("method", "search") if method == "match": return pattern.match(text) elif method == "fullmatch": return pattern.fullmatch(text) else: return pattern.search(text) @classmethod def search(cls, pattern_id: str, text: str) -> Optional[re.Match]: """强制使用 search 方法匹配。""" cls._ensure_loaded() return cls.get_pattern(pattern_id).search(text) @classmethod def get_all_pattern_ids(cls) -> List[str]: """返回所有已加载的规则 ID。""" cls._ensure_loaded() return sorted(cls._patterns.keys()) @classmethod def get_pattern_meta(cls, pattern_id: str) -> dict: """返回规则的元信息(方法、描述、类别等)。""" cls._ensure_loaded() return cls._pattern_meta.get(pattern_id, {}) # ------------------------------------------------------------------------ # 公共 API:完整规则(Phase C RoleRuleEngine 消费) # ------------------------------------------------------------------------ @classmethod def get_rules(cls, category: Optional[str] = None, staged: Optional[bool] = None) -> List[dict]: """返回完整规则列表(保持 YAML 文件内顺序)。 每条规则为 dict,除 YAML 原有字段(id/stage/model_labels/applies_to/ window/guard/guard_not/or_trigger/set_label/set_prob/entity_filter/ golden_examples 等)外,附加: - ``pattern_compiled`` : 编译后的 ``re.Pattern`` - ``category`` : 规则文件类别(ROLE/MONEY/...) - ``source_file`` : 来源 YAML 文件名(去扩展名) :param category: 只返回该类别的规则(None = 全部) :param staged: True = 只返回带 ``stage`` 字段的规则(引擎修正规则), False = 只返回不带 stage 的规则,None = 全部 """ cls._ensure_loaded() rules = cls._rules if category is not None: rules = [r for r in rules if r.get("category") == category] if staged is not None: rules = [r for r in rules if ("stage" in r) == staged] return list(rules) @classmethod def get_rule(cls, pattern_id: str) -> dict: """按规则 ID 返回完整规则 dict(含 pattern_compiled)。 :raises KeyError: 规则 ID 不存在 """ cls._ensure_loaded() for r in cls._rules: if r["id"] == pattern_id: return r raise KeyError("未找到规则 ID: %s" % pattern_id) # ------------------------------------------------------------------------ # 公共 API:关键词分类 # ------------------------------------------------------------------------ @classmethod def classify(cls, classifier_id: str, text: str) -> str: """关键词分类。 根据 ``keyword_rules`` 规则列表,对文本进行多标签分类。 命中数在 ``min_matches`` 和 ``max_matches`` 之间时返回逗号分隔的标签, 否则返回 ``default_label``。 :param classifier_id: 分类器 ID(YAML 文件名去掉扩展名,如 ``money_source``) :param text: 待分类文本 :return: 分类标签字符串 """ cls._ensure_loaded() if classifier_id not in cls._classifiers: raise KeyError("未找到分类器 ID: %s(可用: %s)" % ( classifier_id, ", ".join(sorted(cls._classifiers.keys())))) clf = cls._classifiers[classifier_id] config = clf["config"] rules = clf["rules"] matched_labels = [] for rule in rules: # 检查排除词 excluded = False for ex in rule.get("excludes", []): if re.search(ex, text): excluded = True break if excluded: continue # 检查正向关键词(任一命中即归类) for kw in rule["keywords"]: if re.search(kw, text): matched_labels.append(rule["label"]) break if config["sort"]: matched_labels.sort() n = len(matched_labels) if n >= config["min_matches"] and n <= config["max_matches"]: return config["separator"].join(matched_labels) else: return config["default_label"] @classmethod def get_classifier_ids(cls) -> List[str]: """返回所有已加载的分类器 ID。""" cls._ensure_loaded() return sorted(cls._classifiers.keys()) # ------------------------------------------------------------------------ # 调试 / 测试辅助 # ------------------------------------------------------------------------ @classmethod def reload(cls): """清除缓存并重新加载(用于测试或热更新)。""" with cls._lock: cls._patterns.clear() cls._pattern_meta.clear() cls._rules[:] = [] cls._classifiers.clear() cls._loaded = False @classmethod def get_stats(cls) -> dict: """返回加载统计信息。""" cls._ensure_loaded() return { "pattern_count": len(cls._patterns), "classifier_count": len(cls._classifiers), "pattern_ids": sorted(cls._patterns.keys()), "classifier_ids": sorted(cls._classifiers.keys()), }