# -*- coding: utf-8 -*- """RoleRuleEngine — YAML 驱动的角色/金额模型后修正引擎(角色分类流程优化 Phase C Stage 2)。 职责 ==== 消费 ``dl/rules/patterns/role_context_fix.yaml``(ROLE)与 ``role_money_fix.yaml``(MONEY)中带 ``stage`` 字段的机器可读修正规则, 按「阶段序 + 阶段内首条命中」的语义执行,精确复刻原 ``prem.predict_role`` / ``predict_money`` / ``correct_money_by_rule`` 中内联 if/elif 修正链的行为(行为等价重构,Phase D 双跑回归验证)。 阶段序(与原代码 if/elif 短路语义一一对应) ============================================ 角色(correct_role):: global_pre(非终止:命中后 label 改写继续走后续判定) → 阈值过滤(label∈[0,4] 且 values[label] < P_MODEL_THRESHOLD → label=5,终止) → seq_check(label∈[2,3,4] 命中 → label=5,终止) → 分支:label=0→l0;label=2→l2;label=1→l1;label∈[3,4]→l34; label=5→l5_win_yes → notify → l5(逐阶段首条命中即止) 金额(correct_money):: 阈值过滤(label∈[0,1] 且 values[label] < P_MODEL_THRESHOLD → label=2,终止) → 分支:label=1→m1;label=0→m0;label=2→m_bid 金额标题类别(correct_money_by_doc):: doc_title 阶段独立执行(entity_filter 过滤实体,不依赖上下文切片), 替代原 correct_money_by_rule。 依赖方向说明 ============ 本模块位于 CORE(predictors),只读消费 ``dl/rules`` 的 patterns 运行时数据 与 loader 框架,不依赖 ``dl/rules/generated/``(ARCHITECTURE.md §4.3)。 ``is_agency`` 谓词来自同包 ``predictors/_common.py``(原 prem.py 即如此使用)。 规则条件 DSL 见 ``dl/rules/schema/pattern_schema.json`` definitions.condition。 """ from __future__ import absolute_import import re from typing import Any, Dict, List, Optional, Tuple from BiddingKG.dl.predictors._common import is_agency from BiddingKG.dl.predictors.role_context import P_MODEL_THRESHOLD from BiddingKG.dl.rules.loader import RuleLoader __all__ = ["RoleRuleEngine"] #: 角色分支阶段调度:当前 label → 依次尝试的阶段列表 _ROLE_BRANCH_STAGES = { 0: ("l0",), 2: ("l2",), 1: ("l1",), 3: ("l34",), 4: ("l34",), 5: ("l5_win_yes", "notify", "l5"), } #: 金额分支阶段调度:模型原始 label(阈值过滤后)→ 分支阶段 _MONEY_BRANCH_STAGES = { 1: ("m1",), 0: ("m0",), 2: ("m_bid",), } #: 条件求值缺省的 title_content_head / content_head 截断长度 _DEFAULT_HEAD_SIZE = 100 class RoleRuleEngine(object): """YAML 驱动的角色/金额模型后修正引擎。 用法(prem.py):: engine = RoleRuleEngine() label = engine.correct_role(entity, label, values, front, middle, behind) entity.set_Role(label, values) values 被就地修改(与原代码一致);label 由返回值带回。 """ def __init__(self, agency_checker=None): """ :param agency_checker: is_agency 谓词(可注入替换,缺省用 ``predictors/_common.is_agency``,与原 prem.py 行为一致) """ self._agency_checker = agency_checker or is_agency self._role_stages = self._group_stages("ROLE") self._money_stages = self._group_stages("MONEY") # ------------------------------------------------------------------ # 规则装载 # ------------------------------------------------------------------ @staticmethod def _group_stages(category): """按 stage 分组规则(组内保持 YAML 顺序)。""" stages = {} for rule in RuleLoader.get_rules(category=category, staged=True): stages.setdefault(rule["stage"], []).append(rule) return stages # ------------------------------------------------------------------ # 角色:替代 prem.predict_role 内联修正链 # ------------------------------------------------------------------ def correct_role(self, entity, label, values, front, middle, behind): """执行角色模型后修正(行为等价原 prem.predict_role if/elif 链)。 :param entity: 实体(entity_text 用于 is_agency 谓词) :param label: 模型预测 label(int) :param values: 模型输出概率数组(就地修改) :param front/middle/behind: role_model_text 三元组(前文23/实体/后文25) :return: 修正后的 label """ env = { "front": front, "behind": behind, "whole": front[-10:] + middle + behind[:10], "entity_text": entity.entity_text, } # ---- global_pre:非终止,命中后以新 label 继续走阈值/分支 ---- rule = self._first_match("global_pre", label, env, entity, values) if rule is not None: label = self._apply(rule, label, values) # ---- 阈值过滤(程序逻辑,终止)---- if label in (0, 1, 2, 3, 4) and values[label] < P_MODEL_THRESHOLD: return 5 # ---- seq_check(终止)---- if label in (2, 3, 4): rule = self._first_match("seq_check", label, env, entity, values) if rule is not None: return 5 # ---- 分支阶段 ---- for stage in _ROLE_BRANCH_STAGES.get(label, ()): rule = self._first_match(stage, label, env, entity, values) if rule is not None: label = self._apply(rule, label, values) break return label # ------------------------------------------------------------------ # 金额:替代 prem.predict_money 内联修正链 # ------------------------------------------------------------------ def correct_money(self, entity, label, values, front, middle, behind): """执行金额模型后修正(行为等价原 prem.predict_money if/elif 链)。 :param front/middle/behind: money_model_text 三元组(前文13/实体/后文15) :return: 修正后的 label """ env = { "front": front, "behind": behind, "entity_text": entity.entity_text, } # ---- 阈值过滤(程序逻辑,终止)---- if label in (0, 1) and values[label] < P_MODEL_THRESHOLD: return 2 for stage in _MONEY_BRANCH_STAGES.get(label, ()): rule = self._first_match(stage, label, env, entity, values) if rule is not None: label = self._apply(rule, label, values) break return label # ------------------------------------------------------------------ # 金额标题类别:替代 prem.correct_money_by_rule # ------------------------------------------------------------------ def correct_money_by_doc(self, title, content, list_entitys): """按公告标题/正文类别批量修正金额实体(doc_title 阶段)。 :param title: 公告标题 :param content: 首篇文章正文(原代码 list_articles[0].content) :param list_entitys: 与原 correct_money_by_rule 相同的嵌套实体列表 """ rules = self._money_stages.get("doc_title", ()) if not rules: return for list_entity in list_entitys: for entity in list_entity: env = { "title": title, "content": content, "entity_text": entity.entity_text, } for rule in rules: if not self._entity_filter_ok(entity, rule): continue if self._fires(rule, env, entity, entity.values): new_label = self._apply(rule, entity.label, entity.values) entity.set_Money(new_label, entity.values) break # ------------------------------------------------------------------ # 内部:规则匹配与动作 # ------------------------------------------------------------------ def _first_match(self, stage, label, env, entity, values): """阶段内按 YAML 顺序找首条命中规则(对应原 elif 链短路)。""" for rule in list(self._role_stages.get(stage, ())) + list(self._money_stages.get(stage, ())): if "model_labels" in rule and label not in rule["model_labels"]: continue if self._fires(rule, env, entity, values): return rule return None def _fires(self, rule, env, entity, values): """规则触发判定:trigger(pattern 或 or_trigger)且 guard 且无 guard_not。""" # ---- trigger ---- trig = False pattern = rule.get("pattern_compiled") if pattern is not None: text = self._apply_slice( rule.get("applies_to", "front"), rule.get("window", {}).get("size"), env, ) if text is not None: method = rule.get("method", "search") if method == "match": trig = pattern.match(text) is not None elif method == "fullmatch": trig = pattern.fullmatch(text) is not None else: trig = pattern.search(text) is not None if not trig: for cond in rule.get("or_trigger", ()): if self._eval_cond(cond, env, entity, values): trig = True break if not trig: return False # ---- guard:AND ---- for cond in rule.get("guard", ()): if not self._eval_cond(cond, env, entity, values): return False # ---- guard_not:OR 排除 ---- for cond in rule.get("guard_not", ()): if self._eval_cond(cond, env, entity, values): return False return True @staticmethod def _apply(rule, label, values): """应用 set_label / set_prob,返回新 label。values 就地修改。""" new_label = rule.get("set_label", label) set_prob = rule.get("set_prob") if set_prob is not None: index = set_prob.get("index", rule.get("set_label", label)) values[index] = set_prob["value"] return new_label # ------------------------------------------------------------------ # 内部:切片与条件求值 # ------------------------------------------------------------------ @staticmethod def _apply_slice(on, size, env): """解析条件目标切片;front/behind 按 size 截断(front 取末段)。""" if on == "front": text = env.get("front") if text is None: return None return text[-size:] if size else text if on == "behind": text = env.get("behind") if text is None: return None return text[:size] if size else text if on == "whole": return env.get("whole") if on == "entity_text": return env.get("entity_text") if on == "title": return env.get("title") if on == "content": return env.get("content") if on == "content_head": content = env.get("content") if content is None: return None return content[: size or _DEFAULT_HEAD_SIZE] if on == "title_content_head": title = env.get("title") content = env.get("content") if title is None or content is None: return None return title + content[: size or _DEFAULT_HEAD_SIZE] return None def _eval_cond(self, cond, env, entity, values): """求值结构化条件(schema definitions.condition)。""" ctype = cond["type"] if ctype == "regex": text = self._apply_slice(cond.get("on", "front"), cond.get("size"), env) if text is None: return False return re_search(cond["pattern"], text) is not None if ctype == "regex_not": text = self._apply_slice(cond.get("on", "front"), cond.get("size"), env) if text is None: return False return re_search(cond["pattern"], text) is None if ctype == "match": text = self._apply_slice(cond.get("on", "front"), cond.get("size"), env) if text is None: return False return re_match(cond["pattern"], text) is not None if ctype == "is_agency": return bool(self._agency_checker(env.get("entity_text") or "")) if ctype == "not_is_agency": return not bool(self._agency_checker(env.get("entity_text") or "")) if ctype == "value_lt": return values[cond["index"]] < cond["threshold"] if ctype == "notes_in": return getattr(entity, "notes", None) in cond["values"] if ctype == "any": return any(self._eval_cond(c, env, entity, values) for c in cond["of"]) if ctype == "all": return all(self._eval_cond(c, env, entity, values) for c in cond["of"]) if ctype == "count_eq": text = self._apply_slice(cond.get("on", "title"), cond.get("size"), env) if text is None: return False return len(re_findall(cond["pattern"], text)) == cond["equals"] raise ValueError("未知条件类型: %s" % ctype) @staticmethod def _entity_filter_ok(entity, rule): """doc_title 阶段实体过滤(读规则的 entity_filter 字段)。""" flt = rule.get("entity_filter") if not flt: return True if "entity_type" in flt and entity.entity_type != flt["entity_type"]: return False if "notes_in" in flt and getattr(entity, "notes", None) not in flt["notes_in"]: return False if "label" in flt and entity.label != flt["label"]: return False return True # ---------------------------------------------------------------------------- # 条件正则缓存(guard/match/count_eq 用,避免每次求值重复编译) # ---------------------------------------------------------------------------- _COND_CACHE: Dict[str, Any] = {} def re_search(pattern, text): key = "s:" + pattern pat = _COND_CACHE.get(key) if pat is None: pat = _COND_CACHE[key] = re.compile(pattern) return pat.search(text) def re_match(pattern, text): key = "m:" + pattern pat = _COND_CACHE.get(key) if pat is None: pat = _COND_CACHE[key] = re.compile(pattern) return pat.match(text) def re_findall(pattern, text): key = "f:" + pattern pat = _COND_CACHE.get(key) if pat is None: pat = _COND_CACHE[key] = re.compile(pattern) return pat.findall(text)