From 0dad895aac024b5c0547e5fec54cefa903453cec Mon Sep 17 00:00:00 2001 From: tzt <14718231+flying-travel@user.noreply.gitee.com> Date: Sat, 19 Sep 2026 09:24:56 +0800 Subject: [PATCH] =?UTF-8?q?feat(v1):=20T-R2=20=E9=87=87=E7=BA=B3=20llmrout?= =?UTF-8?q?er=E3=80=8C=E8=A7=84=E5=88=99=E6=96=87=E6=A1=A3=E5=8D=B3?= =?UTF-8?q?=E9=85=8D=E7=BD=AE=E3=80=8D=E2=80=94=E2=80=94=E5=88=86=E7=B1=BB?= =?UTF-8?q?=E8=A7=84=E5=88=99=E5=A4=96=E7=BD=AE=20config/routes.json?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - classifier.load_domain_rules():支持 .json(必有)/.yaml(pyyaml 可选,与 config.py 同一可选依赖纪律);文件按 domain 整域覆盖内置 DOMAIN_RULES (与 v2 知识库'文件按 id 覆盖'同一惯例);缺失/格式非法/条目非法整体 安全回退内置(llmrouter 失败安全思想,规则文档编辑错误不打垮路由) - build_classifier:cfg.rules_file 显式指定,未指定时约定路径 config/routes.json 存在即自动加载(约定优于配置,ROUTES.md 精髓:改文档即改行为,可 review 可版本化) - 新增 tests/test_rules_external.py 5 项;全量 38 passed(33+5) --- router_system/classifier.py | 67 ++++++++++++++++++++++++++++++++++-- tests/test_rules_external.py | 53 ++++++++++++++++++++++++++++ 2 files changed, 117 insertions(+), 3 deletions(-) create mode 100644 tests/test_rules_external.py diff --git a/router_system/classifier.py b/router_system/classifier.py index f275300..099aad3 100644 --- a/router_system/classifier.py +++ b/router_system/classifier.py @@ -6,11 +6,18 @@ 置信度设计:每个领域有一组 (关键词, 权重)。命中权重求和得原始分 s, confidence = 1 - exp(-s),保证 s=1 -> 0.63,s=2 -> 0.86,s=3 -> 0.95。 无领域命中(或最高分领域为 general)时置信度低,触发 should_fallback。 + +规则外置(T-R2,采纳 llmrouter「规则文档即配置」设计):领域规则支持从 +config/routes.json(或 .yaml,pyyaml 可选)加载,文件按 domain 覆盖内置 +DOMAIN_RULES(与 v2 知识库"文件按 id 覆盖"同一惯例);文件缺失/格式非法 +整体安全回退内置(llmrouter 失败安全思想)。 """ from __future__ import annotations +import json import math -from typing import Dict, List, Tuple +from pathlib import Path +from typing import Dict, List, Optional, Tuple from .difficulty import estimate_difficulty from .models import Classification @@ -82,6 +89,50 @@ _STOPWORDS = { "or", "do", "does", "can", "could", "would", "should", "please", "me", "my", } +# 外置规则默认路径(约定优于配置:文件存在即自动加载,T-R2) +DEFAULT_RULES_FILE = Path(__file__).resolve().parent.parent / "config" / "routes.json" + + +def load_domain_rules(path: Optional[str | Path] = None + ) -> Optional[Dict[str, List[Tuple[str, float]]]]: + """加载外置分类规则文件,返回 {domain: [(关键词, 权重), ...]}。 + + 文件缺失 / 格式非法 / 条目非法时返回 None(调用方安全回退内置规则, + 不抛异常——规则文档可被人工编辑,编辑错误不应打垮路由)。 + 支持 .json(必有)与 .yaml(pyyaml 可选依赖)。 + """ + p = Path(path) if path else DEFAULT_RULES_FILE + if not p.exists(): + return None + try: + if p.suffix.lower() in (".yaml", ".yml"): + try: + import yaml # type: ignore + except ImportError: + return None + with open(p, "r", encoding="utf-8") as f: + data = yaml.safe_load(f) + else: + with open(p, "r", encoding="utf-8") as f: + data = json.load(f) + except (OSError, ValueError, Exception): # noqa: BLE001 解析失败一律回退 + return None + if not isinstance(data, dict): + return None + rules: Dict[str, List[Tuple[str, float]]] = {} + for domain, entries in data.items(): + if not isinstance(domain, str) or not isinstance(entries, list): + return None + pairs: List[Tuple[str, float]] = [] + for entry in entries: + if (not isinstance(entry, (list, tuple)) or len(entry) != 2 + or not isinstance(entry[0], str) + or not isinstance(entry[1], (int, float))): + return None + pairs.append((entry[0], float(entry[1]))) + rules[domain] = pairs + return rules + class BaseClassifier: def classify(self, query: str) -> Classification: @@ -198,11 +249,21 @@ class HuggingFaceClassifier(BaseClassifier): def build_classifier(cfg: Dict) -> BaseClassifier: - """根据配置构建分类器。cfg 为 classifier 段配置。""" + """根据配置构建分类器。cfg 为 classifier 段配置。 + + T-R2:cfg.rules_file 指定外置规则文件;未指定时若约定路径 + config/routes.json 存在则自动加载。文件按 domain 覆盖内置规则。 + """ ctype = cfg.get("type", "rule") floor = cfg.get("confidence_floor", 0.55) if ctype == "rule": - return RuleClassifier(confidence_floor=floor) + clf = RuleClassifier(confidence_floor=floor) + rules_file = cfg.get("rules_file") + external = load_domain_rules(rules_file) if rules_file \ + else (load_domain_rules() if DEFAULT_RULES_FILE.exists() else None) + if external: + clf.rules = {**DOMAIN_RULES, **external} + return clf if ctype == "hf": return HuggingFaceClassifier(cfg.get("model", "Qwen/Qwen3-0.6B"), confidence_floor=floor) raise ValueError(f"未知分类器类型: {ctype}(支持 rule | hf)") diff --git a/tests/test_rules_external.py b/tests/test_rules_external.py new file mode 100644 index 0000000..20d57dc --- /dev/null +++ b/tests/test_rules_external.py @@ -0,0 +1,53 @@ +"""分类规则外置单元测试(T-R2,采纳 llmrouter「规则文档即配置」设计)。""" +import json + +from router_system.classifier import (DOMAIN_RULES, RuleClassifier, build_classifier, + load_domain_rules) + + +def test_load_missing_file_returns_none(tmp_path): + """文件缺失:返回 None(调用方回退内置),不抛异常。""" + assert load_domain_rules(tmp_path / "nope.json") is None + + +def test_external_rules_override_domain(tmp_path): + """外置文件按 domain 覆盖内置:改写 code 规则即改变路由行为。""" + f = tmp_path / "routes.json" + f.write_text(json.dumps({"code": [["速排", 2.0]]}, ensure_ascii=False), + encoding="utf-8") + rules = load_domain_rules(f) + assert rules is not None and rules["code"] == [("速排", 2.0)] + + clf = RuleClassifier() + clf.rules = {**DOMAIN_RULES, **rules} + r = clf.classify("讲讲速排的思路") + assert r.domain == "code" + + +def test_build_classifier_merges_over_builtin(tmp_path): + """build_classifier 合并语义:文件只写 legal,其余领域保持内置。""" + f = tmp_path / "routes.json" + f.write_text(json.dumps({"legal": [["劳动合同", 3.0]]}, ensure_ascii=False), + encoding="utf-8") + clf = build_classifier({"type": "rule", "rules_file": str(f)}) + assert clf.rules["legal"] == [("劳动合同", 3.0)] + assert clf.rules["code"] == DOMAIN_RULES["code"] # 未覆盖领域保持内置 + r = clf.classify("劳动合同到期不续签需要支付经济补偿吗") + assert r.domain == "legal" + + +def test_malformed_file_falls_back_to_builtin(tmp_path): + """格式非法:整体回退内置(失败安全,编辑错误不打垮路由)。""" + f = tmp_path / "routes.json" + f.write_text('{"code": [["坏数据", "不是数字"]]}', encoding="utf-8") + assert load_domain_rules(f) is None + clf = build_classifier({"type": "rule", "rules_file": str(f)}) + assert clf.rules == DOMAIN_RULES + r = clf.classify("用 Python 写一个快速排序函数") + assert r.domain == "code" + + +def test_explicit_rules_file_missing_uses_builtin(tmp_path): + """显式路径不存在:静默回退内置,分类行为不变。""" + clf = build_classifier({"type": "rule", "rules_file": str(tmp_path / "no.json")}) + assert clf.rules == DOMAIN_RULES