role_rule_engine.py 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386
  1. # -*- coding: utf-8 -*-
  2. """RoleRuleEngine — YAML 驱动的角色/金额模型后修正引擎(角色分类流程优化 Phase C Stage 2)。
  3. 职责
  4. ====
  5. 消费 ``dl/rules/patterns/role_context_fix.yaml``(ROLE)与
  6. ``role_money_fix.yaml``(MONEY)中带 ``stage`` 字段的机器可读修正规则,
  7. 按「阶段序 + 阶段内首条命中」的语义执行,精确复刻原
  8. ``prem.predict_role`` / ``predict_money`` / ``correct_money_by_rule``
  9. 中内联 if/elif 修正链的行为(行为等价重构,Phase D 双跑回归验证)。
  10. 阶段序(与原代码 if/elif 短路语义一一对应)
  11. ============================================
  12. 角色(correct_role)::
  13. global_pre(非终止:命中后 label 改写继续走后续判定)
  14. → 阈值过滤(label∈[0,4] 且 values[label] < P_MODEL_THRESHOLD → label=5,终止)
  15. → seq_check(label∈[2,3,4] 命中 → label=5,终止)
  16. → 分支:label=0→l0;label=2→l2;label=1→l1;label∈[3,4]→l34;
  17. label=5→l5_win_yes → notify → l5(逐阶段首条命中即止)
  18. 金额(correct_money)::
  19. 阈值过滤(label∈[0,1] 且 values[label] < P_MODEL_THRESHOLD → label=2,终止)
  20. → 分支:label=1→m1;label=0→m0;label=2→m_bid
  21. 金额标题类别(correct_money_by_doc)::
  22. doc_title 阶段独立执行(entity_filter 过滤实体,不依赖上下文切片),
  23. 替代原 correct_money_by_rule。
  24. 依赖方向说明
  25. ============
  26. 本模块位于 CORE(predictors),只读消费 ``dl/rules`` 的 patterns 运行时数据
  27. 与 loader 框架,不依赖 ``dl/rules/generated/``(ARCHITECTURE.md §4.3)。
  28. ``is_agency`` 谓词来自同包 ``predictors/_common.py``(原 prem.py 即如此使用)。
  29. 规则条件 DSL 见 ``dl/rules/schema/pattern_schema.json`` definitions.condition。
  30. """
  31. from __future__ import absolute_import
  32. import re
  33. from typing import Any, Dict, List, Optional, Tuple
  34. from BiddingKG.dl.predictors._common import is_agency
  35. from BiddingKG.dl.predictors.role_context import P_MODEL_THRESHOLD
  36. from BiddingKG.dl.rules.loader import RuleLoader
  37. __all__ = ["RoleRuleEngine"]
  38. #: 角色分支阶段调度:当前 label → 依次尝试的阶段列表
  39. _ROLE_BRANCH_STAGES = {
  40. 0: ("l0",),
  41. 2: ("l2",),
  42. 1: ("l1",),
  43. 3: ("l34",),
  44. 4: ("l34",),
  45. 5: ("l5_win_yes", "notify", "l5"),
  46. }
  47. #: 金额分支阶段调度:模型原始 label(阈值过滤后)→ 分支阶段
  48. _MONEY_BRANCH_STAGES = {
  49. 1: ("m1",),
  50. 0: ("m0",),
  51. 2: ("m_bid",),
  52. }
  53. #: 条件求值缺省的 title_content_head / content_head 截断长度
  54. _DEFAULT_HEAD_SIZE = 100
  55. class RoleRuleEngine(object):
  56. """YAML 驱动的角色/金额模型后修正引擎。
  57. 用法(prem.py)::
  58. engine = RoleRuleEngine()
  59. label = engine.correct_role(entity, label, values, front, middle, behind)
  60. entity.set_Role(label, values)
  61. values 被就地修改(与原代码一致);label 由返回值带回。
  62. """
  63. def __init__(self, agency_checker=None):
  64. """
  65. :param agency_checker: is_agency 谓词(可注入替换,缺省用
  66. ``predictors/_common.is_agency``,与原 prem.py 行为一致)
  67. """
  68. self._agency_checker = agency_checker or is_agency
  69. self._role_stages = self._group_stages("ROLE")
  70. self._money_stages = self._group_stages("MONEY")
  71. # ------------------------------------------------------------------
  72. # 规则装载
  73. # ------------------------------------------------------------------
  74. @staticmethod
  75. def _group_stages(category):
  76. """按 stage 分组规则(组内保持 YAML 顺序)。"""
  77. stages = {}
  78. for rule in RuleLoader.get_rules(category=category, staged=True):
  79. stages.setdefault(rule["stage"], []).append(rule)
  80. return stages
  81. # ------------------------------------------------------------------
  82. # 角色:替代 prem.predict_role 内联修正链
  83. # ------------------------------------------------------------------
  84. def correct_role(self, entity, label, values, front, middle, behind):
  85. """执行角色模型后修正(行为等价原 prem.predict_role if/elif 链)。
  86. :param entity: 实体(entity_text 用于 is_agency 谓词)
  87. :param label: 模型预测 label(int)
  88. :param values: 模型输出概率数组(就地修改)
  89. :param front/middle/behind: role_model_text 三元组(前文23/实体/后文25)
  90. :return: 修正后的 label
  91. """
  92. env = {
  93. "front": front,
  94. "behind": behind,
  95. "whole": front[-10:] + middle + behind[:10],
  96. "entity_text": entity.entity_text,
  97. }
  98. # ---- global_pre:非终止,命中后以新 label 继续走阈值/分支 ----
  99. rule = self._first_match("global_pre", label, env, entity, values)
  100. if rule is not None:
  101. label = self._apply(rule, label, values)
  102. # ---- 阈值过滤(程序逻辑,终止)----
  103. if label in (0, 1, 2, 3, 4) and values[label] < P_MODEL_THRESHOLD:
  104. return 5
  105. # ---- seq_check(终止)----
  106. if label in (2, 3, 4):
  107. rule = self._first_match("seq_check", label, env, entity, values)
  108. if rule is not None:
  109. return 5
  110. # ---- 分支阶段 ----
  111. for stage in _ROLE_BRANCH_STAGES.get(label, ()):
  112. rule = self._first_match(stage, label, env, entity, values)
  113. if rule is not None:
  114. label = self._apply(rule, label, values)
  115. break
  116. return label
  117. # ------------------------------------------------------------------
  118. # 金额:替代 prem.predict_money 内联修正链
  119. # ------------------------------------------------------------------
  120. def correct_money(self, entity, label, values, front, middle, behind):
  121. """执行金额模型后修正(行为等价原 prem.predict_money if/elif 链)。
  122. :param front/middle/behind: money_model_text 三元组(前文13/实体/后文15)
  123. :return: 修正后的 label
  124. """
  125. env = {
  126. "front": front,
  127. "behind": behind,
  128. "entity_text": entity.entity_text,
  129. }
  130. # ---- 阈值过滤(程序逻辑,终止)----
  131. if label in (0, 1) and values[label] < P_MODEL_THRESHOLD:
  132. return 2
  133. for stage in _MONEY_BRANCH_STAGES.get(label, ()):
  134. rule = self._first_match(stage, label, env, entity, values)
  135. if rule is not None:
  136. label = self._apply(rule, label, values)
  137. break
  138. return label
  139. # ------------------------------------------------------------------
  140. # 金额标题类别:替代 prem.correct_money_by_rule
  141. # ------------------------------------------------------------------
  142. def correct_money_by_doc(self, title, content, list_entitys):
  143. """按公告标题/正文类别批量修正金额实体(doc_title 阶段)。
  144. :param title: 公告标题
  145. :param content: 首篇文章正文(原代码 list_articles[0].content)
  146. :param list_entitys: 与原 correct_money_by_rule 相同的嵌套实体列表
  147. """
  148. rules = self._money_stages.get("doc_title", ())
  149. if not rules:
  150. return
  151. for list_entity in list_entitys:
  152. for entity in list_entity:
  153. env = {
  154. "title": title,
  155. "content": content,
  156. "entity_text": entity.entity_text,
  157. }
  158. for rule in rules:
  159. if not self._entity_filter_ok(entity, rule):
  160. continue
  161. if self._fires(rule, env, entity, entity.values):
  162. new_label = self._apply(rule, entity.label, entity.values)
  163. entity.set_Money(new_label, entity.values)
  164. break
  165. # ------------------------------------------------------------------
  166. # 内部:规则匹配与动作
  167. # ------------------------------------------------------------------
  168. def _first_match(self, stage, label, env, entity, values):
  169. """阶段内按 YAML 顺序找首条命中规则(对应原 elif 链短路)。"""
  170. for rule in list(self._role_stages.get(stage, ())) + list(self._money_stages.get(stage, ())):
  171. if "model_labels" in rule and label not in rule["model_labels"]:
  172. continue
  173. if self._fires(rule, env, entity, values):
  174. return rule
  175. return None
  176. def _fires(self, rule, env, entity, values):
  177. """规则触发判定:trigger(pattern 或 or_trigger)且 guard 且无 guard_not。"""
  178. # ---- trigger ----
  179. trig = False
  180. pattern = rule.get("pattern_compiled")
  181. if pattern is not None:
  182. text = self._apply_slice(
  183. rule.get("applies_to", "front"),
  184. rule.get("window", {}).get("size"),
  185. env,
  186. )
  187. if text is not None:
  188. method = rule.get("method", "search")
  189. if method == "match":
  190. trig = pattern.match(text) is not None
  191. elif method == "fullmatch":
  192. trig = pattern.fullmatch(text) is not None
  193. else:
  194. trig = pattern.search(text) is not None
  195. if not trig:
  196. for cond in rule.get("or_trigger", ()):
  197. if self._eval_cond(cond, env, entity, values):
  198. trig = True
  199. break
  200. if not trig:
  201. return False
  202. # ---- guard:AND ----
  203. for cond in rule.get("guard", ()):
  204. if not self._eval_cond(cond, env, entity, values):
  205. return False
  206. # ---- guard_not:OR 排除 ----
  207. for cond in rule.get("guard_not", ()):
  208. if self._eval_cond(cond, env, entity, values):
  209. return False
  210. return True
  211. @staticmethod
  212. def _apply(rule, label, values):
  213. """应用 set_label / set_prob,返回新 label。values 就地修改。"""
  214. new_label = rule.get("set_label", label)
  215. set_prob = rule.get("set_prob")
  216. if set_prob is not None:
  217. index = set_prob.get("index", rule.get("set_label", label))
  218. values[index] = set_prob["value"]
  219. return new_label
  220. # ------------------------------------------------------------------
  221. # 内部:切片与条件求值
  222. # ------------------------------------------------------------------
  223. @staticmethod
  224. def _apply_slice(on, size, env):
  225. """解析条件目标切片;front/behind 按 size 截断(front 取末段)。"""
  226. if on == "front":
  227. text = env.get("front")
  228. if text is None:
  229. return None
  230. return text[-size:] if size else text
  231. if on == "behind":
  232. text = env.get("behind")
  233. if text is None:
  234. return None
  235. return text[:size] if size else text
  236. if on == "whole":
  237. return env.get("whole")
  238. if on == "entity_text":
  239. return env.get("entity_text")
  240. if on == "title":
  241. return env.get("title")
  242. if on == "content":
  243. return env.get("content")
  244. if on == "content_head":
  245. content = env.get("content")
  246. if content is None:
  247. return None
  248. return content[: size or _DEFAULT_HEAD_SIZE]
  249. if on == "title_content_head":
  250. title = env.get("title")
  251. content = env.get("content")
  252. if title is None or content is None:
  253. return None
  254. return title + content[: size or _DEFAULT_HEAD_SIZE]
  255. return None
  256. def _eval_cond(self, cond, env, entity, values):
  257. """求值结构化条件(schema definitions.condition)。"""
  258. ctype = cond["type"]
  259. if ctype == "regex":
  260. text = self._apply_slice(cond.get("on", "front"), cond.get("size"), env)
  261. if text is None:
  262. return False
  263. return re_search(cond["pattern"], text) is not None
  264. if ctype == "regex_not":
  265. text = self._apply_slice(cond.get("on", "front"), cond.get("size"), env)
  266. if text is None:
  267. return False
  268. return re_search(cond["pattern"], text) is None
  269. if ctype == "match":
  270. text = self._apply_slice(cond.get("on", "front"), cond.get("size"), env)
  271. if text is None:
  272. return False
  273. return re_match(cond["pattern"], text) is not None
  274. if ctype == "is_agency":
  275. return bool(self._agency_checker(env.get("entity_text") or ""))
  276. if ctype == "not_is_agency":
  277. return not bool(self._agency_checker(env.get("entity_text") or ""))
  278. if ctype == "value_lt":
  279. return values[cond["index"]] < cond["threshold"]
  280. if ctype == "notes_in":
  281. return getattr(entity, "notes", None) in cond["values"]
  282. if ctype == "any":
  283. return any(self._eval_cond(c, env, entity, values) for c in cond["of"])
  284. if ctype == "all":
  285. return all(self._eval_cond(c, env, entity, values) for c in cond["of"])
  286. if ctype == "count_eq":
  287. text = self._apply_slice(cond.get("on", "title"), cond.get("size"), env)
  288. if text is None:
  289. return False
  290. return len(re_findall(cond["pattern"], text)) == cond["equals"]
  291. raise ValueError("未知条件类型: %s" % ctype)
  292. @staticmethod
  293. def _entity_filter_ok(entity, rule):
  294. """doc_title 阶段实体过滤(读规则的 entity_filter 字段)。"""
  295. flt = rule.get("entity_filter")
  296. if not flt:
  297. return True
  298. if "entity_type" in flt and entity.entity_type != flt["entity_type"]:
  299. return False
  300. if "notes_in" in flt and getattr(entity, "notes", None) not in flt["notes_in"]:
  301. return False
  302. if "label" in flt and entity.label != flt["label"]:
  303. return False
  304. return True
  305. # ----------------------------------------------------------------------------
  306. # 条件正则缓存(guard/match/count_eq 用,避免每次求值重复编译)
  307. # ----------------------------------------------------------------------------
  308. _COND_CACHE: Dict[str, Any] = {}
  309. def re_search(pattern, text):
  310. key = "s:" + pattern
  311. pat = _COND_CACHE.get(key)
  312. if pat is None:
  313. pat = _COND_CACHE[key] = re.compile(pattern)
  314. return pat.search(text)
  315. def re_match(pattern, text):
  316. key = "m:" + pattern
  317. pat = _COND_CACHE.get(key)
  318. if pat is None:
  319. pat = _COND_CACHE[key] = re.compile(pattern)
  320. return pat.match(text)
  321. def re_findall(pattern, text):
  322. key = "f:" + pattern
  323. pat = _COND_CACHE.get(key)
  324. if pat is None:
  325. pat = _COND_CACHE[key] = re.compile(pattern)
  326. return pat.findall(text)