Files
zk-data-agent/skills/intervention-data/scripts/intervention_common.py
2026-05-19 11:59:49 +08:00

331 lines
11 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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)