| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350 |
- # -*- 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()),
- }
|