331 lines
11 KiB
Python
331 lines
11 KiB
Python
#!/usr/bin/env python3
|
||
"""干预数据脚本的公共能力。
|
||
|
||
这里尽量只做确定性处理:
|
||
- 读取历史干预 TSV。
|
||
- 归纳设备枚举和 code 到下发 domain 的多数映射。
|
||
- 复用 label-master 校验 target 语法。
|
||
- 校验候选干预条目是否和历史数据重复或冲突。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import csv
|
||
import importlib.util
|
||
import json
|
||
import re
|
||
import sys
|
||
from collections import Counter, defaultdict
|
||
from dataclasses import dataclass
|
||
from pathlib import Path
|
||
from typing import Any
|
||
|
||
|
||
SCRIPT_DIR = Path(__file__).resolve().parent
|
||
SKILL_ROOT = SCRIPT_DIR.parent
|
||
REPO_ROOT = SKILL_ROOT.parent.parent
|
||
DEFAULT_EXACT_SOURCE = REPO_ROOT / "干预skill" / "20260134"
|
||
DEFAULT_REGEX_SOURCE = REPO_ROOT / "干预skill" / "20260207"
|
||
LABEL_VALIDATOR_PATH = REPO_ROOT / "skills" / "label-master" / "scripts" / "validate_label_output.py"
|
||
CONFIG_MAPPING_PATH = SKILL_ROOT / "knowledge" / "type_match_agent.json"
|
||
AGENT_TARGET_RE = re.compile(r'^Agent\(\s*tag\s*=\s*["\'](.+?)["\']\s*\)$')
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class ExactRecord:
|
||
device: str
|
||
query: str
|
||
target: str
|
||
dispatch_domain: str
|
||
line_no: int
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class RegexRecord:
|
||
device: str
|
||
pattern: str
|
||
target: str
|
||
line_no: int
|
||
|
||
|
||
def normalize_device(device: str) -> str:
|
||
"""归一化历史数据里的设备写法,但保留常见大小写。"""
|
||
|
||
value = device.strip()
|
||
if value.startswith("#"):
|
||
value = value[1:].strip()
|
||
aliases = {
|
||
"micar": "miCar",
|
||
"miai": "miai",
|
||
"tVV2": "tvV2",
|
||
"tvv2": "tvV2",
|
||
"unify": "unify",
|
||
}
|
||
return aliases.get(value, value)
|
||
|
||
|
||
def normalize_target(target: str) -> str:
|
||
return target.replace("“", '"').replace("”", '"').replace("‘", "'").replace("’", "'").strip()
|
||
|
||
|
||
def read_tsv_rows(path: Path) -> list[list[str]]:
|
||
with path.open("r", encoding="utf-8-sig", newline="") as handle:
|
||
return list(csv.reader(handle, delimiter="\t"))
|
||
|
||
|
||
def load_exact_records(path: Path = DEFAULT_EXACT_SOURCE) -> list[ExactRecord]:
|
||
rows = read_tsv_rows(path)
|
||
records: list[ExactRecord] = []
|
||
for index, row in enumerate(rows, start=1):
|
||
if index == 1 and row[:4] == ["#设备名", "query", "code", "下发domain"]:
|
||
continue
|
||
if not row or all(not cell.strip() for cell in row):
|
||
continue
|
||
padded = [*row, "", "", "", ""]
|
||
records.append(
|
||
ExactRecord(
|
||
device=normalize_device(padded[0]),
|
||
query=padded[1].strip(),
|
||
target=normalize_target(padded[2]),
|
||
dispatch_domain=padded[3].strip(),
|
||
line_no=index,
|
||
)
|
||
)
|
||
return records
|
||
|
||
|
||
def load_regex_records(path: Path = DEFAULT_REGEX_SOURCE) -> list[RegexRecord]:
|
||
rows = read_tsv_rows(path)
|
||
records: list[RegexRecord] = []
|
||
for index, row in enumerate(rows, start=1):
|
||
if not row or all(not cell.strip() for cell in row):
|
||
continue
|
||
padded = [*row, "", ""]
|
||
records.append(
|
||
RegexRecord(
|
||
device=normalize_device(padded[0]),
|
||
pattern=padded[1].strip(),
|
||
target=normalize_target(padded[2]),
|
||
line_no=index,
|
||
)
|
||
)
|
||
return records
|
||
|
||
|
||
def build_source_summary(
|
||
exact_path: Path = DEFAULT_EXACT_SOURCE,
|
||
regex_path: Path = DEFAULT_REGEX_SOURCE,
|
||
) -> dict[str, Any]:
|
||
exact_records = load_exact_records(exact_path) if exact_path.exists() else []
|
||
regex_records = load_regex_records(regex_path) if regex_path.exists() else []
|
||
|
||
exact_devices = Counter(record.device for record in exact_records)
|
||
regex_devices = Counter(record.device for record in regex_records)
|
||
devices = sorted(set(exact_devices) | set(regex_devices))
|
||
|
||
code_domain_counts: dict[str, Counter[str]] = defaultdict(Counter)
|
||
for record in exact_records:
|
||
if record.target and record.dispatch_domain:
|
||
code_domain_counts[record.target][record.dispatch_domain] += 1
|
||
|
||
code_domain_mapping: dict[str, dict[str, Any]] = {}
|
||
for target, counts in code_domain_counts.items():
|
||
domain, count = counts.most_common(1)[0]
|
||
total = sum(counts.values())
|
||
code_domain_mapping[target] = {
|
||
"dispatch_domain": domain,
|
||
"count": count,
|
||
"total": total,
|
||
"confidence": round(count / total, 4) if total else 0,
|
||
"alternatives": [
|
||
{"dispatch_domain": item_domain, "count": item_count}
|
||
for item_domain, item_count in counts.most_common(5)
|
||
],
|
||
}
|
||
|
||
config_mapping = load_type_match_agent_mapping()
|
||
|
||
return {
|
||
"source_files": {
|
||
"exact": str(exact_path),
|
||
"regex": str(regex_path),
|
||
},
|
||
"counts": {
|
||
"exact_records": len(exact_records),
|
||
"regex_records": len(regex_records),
|
||
},
|
||
"devices": devices,
|
||
"device_counts": {
|
||
"exact": dict(exact_devices.most_common()),
|
||
"regex": dict(regex_devices.most_common()),
|
||
},
|
||
"code_domain_mapping": code_domain_mapping,
|
||
"configured_code_domain_mapping": config_mapping,
|
||
"historical_targets": sorted({record.target for record in exact_records + regex_records if record.target}),
|
||
"top_exact_targets": dict(Counter(record.target for record in exact_records).most_common(50)),
|
||
"top_regex_targets": dict(Counter(record.target for record in regex_records).most_common(50)),
|
||
}
|
||
|
||
|
||
def historical_targets(summary: dict[str, Any]) -> set[str]:
|
||
targets = set(summary.get("top_exact_targets", {})) | set(summary.get("top_regex_targets", {}))
|
||
targets.update(summary.get("code_domain_mapping", {}))
|
||
targets.update(summary.get("historical_targets", []))
|
||
return {normalize_target(target) for target in targets if target}
|
||
|
||
|
||
def infer_dispatch_domain(target: str, summary: dict[str, Any]) -> dict[str, Any] | None:
|
||
normalized = normalize_target(target)
|
||
historical = summary.get("code_domain_mapping", {}).get(normalized)
|
||
if historical:
|
||
return {**historical, "source": "historical"}
|
||
configured = infer_configured_dispatch_domain(normalized, "unify", summary)
|
||
if configured:
|
||
return {
|
||
"dispatch_domain": configured["dispatch_domain"],
|
||
"count": 0,
|
||
"total": 0,
|
||
"confidence": None,
|
||
"alternatives": [],
|
||
"source": "type_match_agent",
|
||
"matched_key": configured["matched_key"],
|
||
}
|
||
return None
|
||
|
||
|
||
def infer_dispatch_domain_for_device(target: str, device: str, summary: dict[str, Any]) -> dict[str, Any] | None:
|
||
"""按设备推断推荐下发 domain。
|
||
|
||
优先级:
|
||
1. `标签#设备` 这种明确配置。
|
||
2. 历史精确干预表里的全局多数映射。
|
||
3. `标签` 通用配置。
|
||
"""
|
||
|
||
normalized = normalize_target(target)
|
||
configured = infer_configured_dispatch_domain(normalized, device, summary)
|
||
if configured and "#" in configured["matched_key"]:
|
||
return {
|
||
"dispatch_domain": configured["dispatch_domain"],
|
||
"count": 0,
|
||
"total": 0,
|
||
"confidence": None,
|
||
"alternatives": [],
|
||
"source": "type_match_agent",
|
||
"matched_key": configured["matched_key"],
|
||
}
|
||
|
||
historical = summary.get("code_domain_mapping", {}).get(normalized)
|
||
if historical:
|
||
return {**historical, "source": "historical"}
|
||
|
||
if configured:
|
||
return {
|
||
"dispatch_domain": configured["dispatch_domain"],
|
||
"count": 0,
|
||
"total": 0,
|
||
"confidence": None,
|
||
"alternatives": [],
|
||
"source": "type_match_agent",
|
||
"matched_key": configured["matched_key"],
|
||
}
|
||
return None
|
||
|
||
|
||
def load_type_match_agent_mapping(path: Path = CONFIG_MAPPING_PATH) -> dict[str, str]:
|
||
if not path.exists():
|
||
return {}
|
||
data = json.loads(path.read_text(encoding="utf-8"))
|
||
return {str(key): str(value) for key, value in data.items()}
|
||
|
||
|
||
def extract_agent_tag(target: str) -> str | None:
|
||
match = AGENT_TARGET_RE.match(normalize_target(target))
|
||
if not match:
|
||
return None
|
||
return match.group(1)
|
||
|
||
|
||
def infer_configured_dispatch_domain(target: str, device: str, summary: dict[str, Any]) -> dict[str, str] | None:
|
||
tag = extract_agent_tag(target)
|
||
if not tag:
|
||
return None
|
||
mapping = summary.get("configured_code_domain_mapping", {})
|
||
normalized_device = normalize_device(device)
|
||
for key in (f"{tag}#{normalized_device}", tag):
|
||
if key in mapping:
|
||
return {"matched_key": key, "dispatch_domain": mapping[key]}
|
||
return None
|
||
|
||
|
||
def split_domain(domain: str) -> list[str]:
|
||
return [part for part in domain.split("|") if part]
|
||
|
||
|
||
def domain_compatible(actual: str, expected: str) -> bool:
|
||
"""判断实际下发 domain 是否和配置基线兼容。
|
||
|
||
下发表里有些是宽泛入口,例如 `音乐 -> contentCopilot`;
|
||
历史干预里常见更细路径,例如 `contentCopilot|music`。
|
||
只要二者一方的管道 token 是另一方的子集,就认为兼容。
|
||
"""
|
||
|
||
actual_parts = set(split_domain(actual))
|
||
expected_parts = set(split_domain(expected))
|
||
if not actual_parts or not expected_parts:
|
||
return False
|
||
return actual_parts.issubset(expected_parts) or expected_parts.issubset(actual_parts)
|
||
|
||
|
||
def load_label_validator() -> Any:
|
||
spec = importlib.util.spec_from_file_location("label_master_validate_label_output", LABEL_VALIDATOR_PATH)
|
||
if spec is None or spec.loader is None:
|
||
raise RuntimeError(f"无法加载 label-master 校验器:{LABEL_VALIDATOR_PATH}")
|
||
module = importlib.util.module_from_spec(spec)
|
||
sys.modules[spec.name] = module
|
||
spec.loader.exec_module(module)
|
||
manifest_path = module.DEFAULT_OUTPUT
|
||
validator = module.LabelOutputValidator(module.load_manifest(manifest_path))
|
||
return validator
|
||
|
||
|
||
def validate_target(target: str) -> dict[str, Any]:
|
||
validator = load_label_validator()
|
||
# 历史 TSV 中多行 function/Agent 目标常以字面量 `\n` 存在,
|
||
# label-master 校验器需要真实换行才能按多行程序解析。
|
||
return validator.validate(normalize_target(target).replace("\\n", "\n"))
|
||
|
||
|
||
def escape_regex_literal(text: str) -> str:
|
||
return re.escape(text.strip())
|
||
|
||
|
||
def build_query_pattern(query: str, before_queries: list[str] | None = None) -> str:
|
||
before_queries = before_queries or []
|
||
parts: list[str] = []
|
||
for before_query in before_queries:
|
||
if before_query.strip():
|
||
parts.append(f"beforeQuery#{escape_regex_literal(before_query)}")
|
||
parts.append(f"query#{escape_regex_literal(query)}")
|
||
return "^" + "#".join(parts) + "$"
|
||
|
||
|
||
def exact_tsv_line(device: str, query: str, target: str, dispatch_domain: str) -> str:
|
||
return "\t".join([normalize_device(device), query.strip(), normalize_target(target), dispatch_domain.strip()])
|
||
|
||
|
||
def regex_tsv_line(device: str, pattern: str, target: str) -> str:
|
||
return "\t".join([normalize_device(device), pattern.strip(), normalize_target(target)])
|
||
|
||
|
||
def parse_candidate_lines(text: str) -> list[list[str]]:
|
||
rows = []
|
||
for line in text.splitlines():
|
||
if line.strip():
|
||
rows.append(next(csv.reader([line], delimiter="\t")))
|
||
return rows
|
||
|
||
|
||
def json_dumps(payload: Any) -> str:
|
||
return json.dumps(payload, ensure_ascii=False, indent=2, sort_keys=True)
|