loader.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350
  1. # -*- coding: utf-8 -*-
  2. """规则加载器(RULE LOADER)。
  3. 按 ARCHITECTURE.md Phase 7 设计,从 ``dl/rules/patterns/*.yaml`` 加载正则规则
  4. 和关键词分类规则,编译为 ``re.Pattern`` 对象并缓存,供业务层使用。
  5. 使用方式::
  6. from BiddingKG.dl.rules.loader import RuleLoader
  7. # 1. 获取编译后的正则对象(与 re.compile() 返回类型一致)
  8. pattern = RuleLoader.get_pattern("failure_keywords")
  9. if pattern.search(text):
  10. ...
  11. # 2. 使用 YAML 中指定的匹配方法(search/match/fullmatch)
  12. if RuleLoader.match("rank_first", text):
  13. return "win_tenderer"
  14. # 3. 关键词分类
  15. label = RuleLoader.classify("money_source", "财政拨款")
  16. # -> "财政资金,上级拨款"
  17. 依赖方向:rules 只依赖 common,不反向依赖业务层。
  18. """
  19. from __future__ import absolute_import
  20. import os
  21. import re
  22. import threading
  23. from typing import Dict, List, Optional, Any, Pattern as RegexPattern
  24. import yaml
  25. # ----------------------------------------------------------------------------
  26. # YAML Loader:on/off/yes/no 不解析为布尔(YAML 1.1 陷阱)
  27. # ----------------------------------------------------------------------------
  28. # PyYAML 按 YAML 1.1 把裸标量 on/off/yes/no 解析为 bool,导致 guard 条件的
  29. # `on: behind` 键变成 True、cond["on"] 取不到值而静默回退默认切片。
  30. # 此处仅保留 true/false 的布尔解析,on/off/yes/no 保持字符串。
  31. import re as _re
  32. class _RuleYamlLoader(yaml.SafeLoader):
  33. """规则 YAML 专用 Loader(on/off/yes/no 保持字符串)。"""
  34. _RuleYamlLoader.yaml_implicit_resolvers = {
  35. key: [(tag, regexp) for tag, regexp in resolvers
  36. if tag != 'tag:yaml.org,2002:bool']
  37. for key, resolvers in yaml.SafeLoader.yaml_implicit_resolvers.items()
  38. }
  39. _RuleYamlLoader.add_implicit_resolver(
  40. 'tag:yaml.org,2002:bool',
  41. _re.compile(r'^(?:true|True|TRUE|false|False|FALSE)$'),
  42. list('tTfF'),
  43. )
  44. def load_rule_yaml(stream):
  45. """加载规则 YAML(on/off/yes/no 保持字符串,其余语义同 safe_load)。"""
  46. return yaml.load(stream, Loader=_RuleYamlLoader)
  47. # ----------------------------------------------------------------------------
  48. # 路径常量
  49. # ----------------------------------------------------------------------------
  50. _RULES_DIR = os.path.dirname(os.path.abspath(__file__))
  51. _PATTERNS_DIR = os.path.join(_RULES_DIR, "patterns")
  52. # ----------------------------------------------------------------------------
  53. # 正则标志映射
  54. # ----------------------------------------------------------------------------
  55. _FLAG_MAP = {
  56. "IGNORECASE": re.IGNORECASE,
  57. "MULTILINE": re.MULTILINE,
  58. "DOTALL": re.DOTALL,
  59. "UNICODE": re.UNICODE,
  60. "VERBOSE": re.VERBOSE,
  61. }
  62. class RuleLoader:
  63. """规则加载器(单例缓存)。
  64. 所有方法均为类方法,通过类级缓存避免重复加载。
  65. 线程安全:首次加载使用锁保护。
  66. """
  67. # 类级缓存
  68. _patterns: Dict[str, RegexPattern] = {} # id → compiled regex
  69. _pattern_meta: Dict[str, dict] = {} # id → {method, flags, description, ...}
  70. _rules: List[dict] = [] # 全量规则 dict(Phase C 引擎消费,保持 YAML 顺序)
  71. _classifiers: Dict[str, dict] = {} # file_stem → {config, rules}
  72. _loaded: bool = False
  73. _lock = threading.Lock()
  74. # ------------------------------------------------------------------------
  75. # 加载
  76. # ------------------------------------------------------------------------
  77. @classmethod
  78. def _ensure_loaded(cls):
  79. """惰性加载所有 YAML 规则文件(线程安全)。"""
  80. if cls._loaded:
  81. return
  82. with cls._lock:
  83. if cls._loaded: # double-check
  84. return
  85. cls._load_all()
  86. cls._loaded = True
  87. @classmethod
  88. def _load_all(cls):
  89. """扫描 patterns/ 目录,加载所有 .yaml 文件。"""
  90. if not os.path.isdir(_PATTERNS_DIR):
  91. return
  92. for fname in sorted(os.listdir(_PATTERNS_DIR)):
  93. if not fname.endswith((".yaml", ".yml")):
  94. continue
  95. fpath = os.path.join(_PATTERNS_DIR, fname)
  96. stem = os.path.splitext(fname)[0]
  97. with open(fpath, encoding="utf-8") as f:
  98. data = load_rule_yaml(f)
  99. if not data:
  100. continue
  101. cls._load_patterns(data, stem)
  102. cls._load_classifier(data, stem)
  103. @classmethod
  104. def _load_patterns(cls, data: dict, stem: str):
  105. """加载正则规则。"""
  106. patterns = data.get("patterns") or []
  107. for p in patterns:
  108. pid = p["id"]
  109. pattern_str = p["pattern"]
  110. flags = 0
  111. for f in p.get("flags", []):
  112. flags |= _FLAG_MAP.get(f, 0)
  113. compiled = re.compile(pattern_str, flags)
  114. cls._patterns[pid] = compiled
  115. cls._pattern_meta[pid] = {
  116. "method": p.get("method", "search"),
  117. "description": p.get("description", ""),
  118. "category": data.get("category", "OTHER"),
  119. "source_file": stem,
  120. "pattern_string": pattern_str,
  121. }
  122. # Phase C:保留完整规则结构(stage/guard/set_prob 等)供引擎消费
  123. rule = dict(p)
  124. rule["pattern_compiled"] = compiled
  125. rule["category"] = data.get("category", "OTHER")
  126. rule["source_file"] = stem
  127. cls._rules.append(rule)
  128. @classmethod
  129. def _load_classifier(cls, data: dict, stem: str):
  130. """加载关键词分类规则。"""
  131. rules = data.get("keyword_rules")
  132. if not rules:
  133. return
  134. config = data.get("classifier_config", {})
  135. cls._classifiers[stem] = {
  136. "config": {
  137. "min_matches": config.get("min_matches", 1),
  138. "max_matches": config.get("max_matches", 99),
  139. "default_label": config.get("default_label", ""),
  140. "separator": config.get("separator", ","),
  141. "sort": config.get("sort", True),
  142. },
  143. "rules": rules,
  144. }
  145. # ------------------------------------------------------------------------
  146. # 公共 API:正则
  147. # ------------------------------------------------------------------------
  148. @classmethod
  149. def get_pattern(cls, pattern_id: str) -> RegexPattern:
  150. """获取编译后的正则对象。
  151. 返回 ``re.Pattern`` 实例,可直接调用 ``.search()``, ``.match()``,
  152. ``.findall()``, ``.finditer()`` 等方法,与 ``re.compile()`` 返回类型一致。
  153. :param pattern_id: 规则 ID(YAML 中的 ``id`` 字段)
  154. :return: 编译后的 ``re.Pattern`` 对象
  155. :raises KeyError: 规则 ID 不存在
  156. """
  157. cls._ensure_loaded()
  158. if pattern_id not in cls._patterns:
  159. raise KeyError("未找到规则 ID: %s(可用: %s)" % (
  160. pattern_id, ", ".join(sorted(cls._patterns.keys()))))
  161. return cls._patterns[pattern_id]
  162. @classmethod
  163. def match(cls, pattern_id: str, text: str) -> Optional[re.Match]:
  164. """使用 YAML 中指定的匹配方法进行匹配。
  165. 根据 YAML 中的 ``method`` 字段自动选择:
  166. - ``search`` → ``pattern.search(text)``
  167. - ``match`` → ``pattern.match(text)``
  168. - ``fullmatch`` → ``pattern.fullmatch(text)``
  169. :param pattern_id: 规则 ID
  170. :param text: 待匹配文本
  171. :return: ``re.Match`` 对象或 ``None``
  172. """
  173. cls._ensure_loaded()
  174. pattern = cls.get_pattern(pattern_id)
  175. method = cls._pattern_meta.get(pattern_id, {}).get("method", "search")
  176. if method == "match":
  177. return pattern.match(text)
  178. elif method == "fullmatch":
  179. return pattern.fullmatch(text)
  180. else:
  181. return pattern.search(text)
  182. @classmethod
  183. def search(cls, pattern_id: str, text: str) -> Optional[re.Match]:
  184. """强制使用 search 方法匹配。"""
  185. cls._ensure_loaded()
  186. return cls.get_pattern(pattern_id).search(text)
  187. @classmethod
  188. def get_all_pattern_ids(cls) -> List[str]:
  189. """返回所有已加载的规则 ID。"""
  190. cls._ensure_loaded()
  191. return sorted(cls._patterns.keys())
  192. @classmethod
  193. def get_pattern_meta(cls, pattern_id: str) -> dict:
  194. """返回规则的元信息(方法、描述、类别等)。"""
  195. cls._ensure_loaded()
  196. return cls._pattern_meta.get(pattern_id, {})
  197. # ------------------------------------------------------------------------
  198. # 公共 API:完整规则(Phase C RoleRuleEngine 消费)
  199. # ------------------------------------------------------------------------
  200. @classmethod
  201. def get_rules(cls, category: Optional[str] = None, staged: Optional[bool] = None) -> List[dict]:
  202. """返回完整规则列表(保持 YAML 文件内顺序)。
  203. 每条规则为 dict,除 YAML 原有字段(id/stage/model_labels/applies_to/
  204. window/guard/guard_not/or_trigger/set_label/set_prob/entity_filter/
  205. golden_examples 等)外,附加:
  206. - ``pattern_compiled`` : 编译后的 ``re.Pattern``
  207. - ``category`` : 规则文件类别(ROLE/MONEY/...)
  208. - ``source_file`` : 来源 YAML 文件名(去扩展名)
  209. :param category: 只返回该类别的规则(None = 全部)
  210. :param staged: True = 只返回带 ``stage`` 字段的规则(引擎修正规则),
  211. False = 只返回不带 stage 的规则,None = 全部
  212. """
  213. cls._ensure_loaded()
  214. rules = cls._rules
  215. if category is not None:
  216. rules = [r for r in rules if r.get("category") == category]
  217. if staged is not None:
  218. rules = [r for r in rules if ("stage" in r) == staged]
  219. return list(rules)
  220. @classmethod
  221. def get_rule(cls, pattern_id: str) -> dict:
  222. """按规则 ID 返回完整规则 dict(含 pattern_compiled)。
  223. :raises KeyError: 规则 ID 不存在
  224. """
  225. cls._ensure_loaded()
  226. for r in cls._rules:
  227. if r["id"] == pattern_id:
  228. return r
  229. raise KeyError("未找到规则 ID: %s" % pattern_id)
  230. # ------------------------------------------------------------------------
  231. # 公共 API:关键词分类
  232. # ------------------------------------------------------------------------
  233. @classmethod
  234. def classify(cls, classifier_id: str, text: str) -> str:
  235. """关键词分类。
  236. 根据 ``keyword_rules`` 规则列表,对文本进行多标签分类。
  237. 命中数在 ``min_matches`` 和 ``max_matches`` 之间时返回逗号分隔的标签,
  238. 否则返回 ``default_label``。
  239. :param classifier_id: 分类器 ID(YAML 文件名去掉扩展名,如 ``money_source``)
  240. :param text: 待分类文本
  241. :return: 分类标签字符串
  242. """
  243. cls._ensure_loaded()
  244. if classifier_id not in cls._classifiers:
  245. raise KeyError("未找到分类器 ID: %s(可用: %s)" % (
  246. classifier_id, ", ".join(sorted(cls._classifiers.keys()))))
  247. clf = cls._classifiers[classifier_id]
  248. config = clf["config"]
  249. rules = clf["rules"]
  250. matched_labels = []
  251. for rule in rules:
  252. # 检查排除词
  253. excluded = False
  254. for ex in rule.get("excludes", []):
  255. if re.search(ex, text):
  256. excluded = True
  257. break
  258. if excluded:
  259. continue
  260. # 检查正向关键词(任一命中即归类)
  261. for kw in rule["keywords"]:
  262. if re.search(kw, text):
  263. matched_labels.append(rule["label"])
  264. break
  265. if config["sort"]:
  266. matched_labels.sort()
  267. n = len(matched_labels)
  268. if n >= config["min_matches"] and n <= config["max_matches"]:
  269. return config["separator"].join(matched_labels)
  270. else:
  271. return config["default_label"]
  272. @classmethod
  273. def get_classifier_ids(cls) -> List[str]:
  274. """返回所有已加载的分类器 ID。"""
  275. cls._ensure_loaded()
  276. return sorted(cls._classifiers.keys())
  277. # ------------------------------------------------------------------------
  278. # 调试 / 测试辅助
  279. # ------------------------------------------------------------------------
  280. @classmethod
  281. def reload(cls):
  282. """清除缓存并重新加载(用于测试或热更新)。"""
  283. with cls._lock:
  284. cls._patterns.clear()
  285. cls._pattern_meta.clear()
  286. cls._rules[:] = []
  287. cls._classifiers.clear()
  288. cls._loaded = False
  289. @classmethod
  290. def get_stats(cls) -> dict:
  291. """返回加载统计信息。"""
  292. cls._ensure_loaded()
  293. return {
  294. "pattern_count": len(cls._patterns),
  295. "classifier_count": len(cls._classifiers),
  296. "pattern_ids": sorted(cls._patterns.keys()),
  297. "classifier_ids": sorted(cls._classifiers.keys()),
  298. }