Files
zk-data-agent/skills/label-master/scripts/validate_label_output.py
T
2026-05-09 17:49:53 +08:00

337 lines
14 KiB
Python
Executable File
Raw 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
"""校验标签输出格式是否合法。
这个脚本只做“格式、存在性、引用关系”的确定性校验,不判断 query 语义。
适合在 Agent 给出 target 后、写入训练/评测数据前调用。
"""
from __future__ import annotations
import argparse
import ast
import json
import re
import sys
from pathlib import Path
from typing import Any
try:
from build_label_manifest import DEFAULT_OUTPUT, SKILL_ROOT, build_manifest
except ImportError: # pragma: no cover - 兼容从其他 cwd 直接执行
sys.path.insert(0, str(Path(__file__).resolve().parent))
from build_label_manifest import DEFAULT_OUTPUT, SKILL_ROOT, build_manifest
AGENT_RE = re.compile(r'^Agent\(\s*tag\s*=\s*["\'](.+?)["\']\s*\)$')
ASSIGNMENT_RE = re.compile(r"^(x\d+)\s*=\s*(.+)$")
def normalize_quotes(text: str) -> str:
return (
text.replace("", '"')
.replace("", '"')
.replace("", "'")
.replace("", "'")
.strip()
)
def load_manifest(path: Path) -> dict[str, Any]:
if path.exists():
return json.loads(path.read_text(encoding="utf-8"))
return build_manifest()
def names_by_key(items: list[dict[str, Any]]) -> dict[str, dict[str, Any]]:
return {item["name"]: item for item in items if item.get("name")}
class LabelOutputValidator:
def __init__(self, manifest: dict[str, Any]):
self.manifest = manifest
self.agent_tags = set(manifest.get("agent_tags", []))
self.functions = names_by_key(manifest.get("functions", []))
self.objects = names_by_key(manifest.get("objects", []))
self.intents = names_by_key(manifest.get("intents", []))
def validate(self, raw_target: str) -> dict[str, Any]:
target = normalize_quotes(raw_target)
result: dict[str, Any] = {
"valid": True,
"detected_type": "unknown",
"normalized_output": target,
"errors": [],
"warnings": [],
"references": [],
}
if not target:
self.error(result, "输出为空")
return result
dimensions, target_body = self.split_dimension_prefix(target)
if dimensions:
result["dimensions"] = dimensions
target_body = target_body.strip()
if target_body.startswith("["):
self.validate_group_json(target_body, result)
elif AGENT_RE.match(target_body):
self.validate_agent_target(target_body, result)
elif self.looks_like_call_program(target_body):
self.validate_call_program(target_body, result)
elif target_body in self.intents:
result["detected_type"] = "intent"
result["references"].append({"type": "intent", "name": target_body})
elif target_body in self.agent_tags:
result["detected_type"] = "bare_label"
result["warnings"].append("只输出了标签名;如果当前任务需要可消费 target,建议明确 intent、function 或 Agent 包装形式")
result["references"].append({"type": "label", "name": target_body})
else:
self.error(result, f"无法识别的输出形式:{target_body}")
if dimensions:
result["normalized_output"] = "\n".join([*dimensions, result["normalized_output"]])
result["valid"] = not result["errors"]
return result
def split_dimension_prefix(self, target: str) -> tuple[list[str], str]:
dimensions: list[str] = []
lines = [line.strip() for line in target.splitlines() if line.strip()]
while lines and re.fullmatch(r"(complex|multi_instruction|auto_task)\s*=\s*(true|false)", lines[0], re.I):
key, value = [part.strip().lower() for part in lines.pop(0).split("=", 1)]
dimensions.append(f"{key}={value}")
return dimensions, "\n".join(lines)
def validate_agent_target(self, target: str, result: dict[str, Any]) -> None:
result["detected_type"] = "agent"
match = AGENT_RE.match(target)
assert match is not None
tag = match.group(1)
result["normalized_output"] = f'Agent(tag="{tag}")'
if tag not in self.agent_tags:
self.error(result, f'Agent tag 不在标签知识库中:{tag}')
return
result["references"].append({"type": "agent_tag", "name": tag})
def looks_like_call_program(self, target: str) -> bool:
lines = [line.strip() for line in target.splitlines() if line.strip()]
return bool(lines) and all("(" in line and line.endswith(")") for line in lines)
def validate_call_program(self, target: str, result: dict[str, Any]) -> None:
result["detected_type"] = "function_program"
variables: dict[str, str] = {}
lines = [line.strip() for line in target.splitlines() if line.strip()]
for index, line in enumerate(lines):
assignment = ASSIGNMENT_RE.match(line)
if assignment:
var_name, expr = assignment.group(1), assignment.group(2).strip()
call_info = self.parse_call(expr, result, line)
if not call_info:
continue
object_name = call_info["name"]
if object_name not in self.objects:
self.error(result, f"变量 {var_name} 只能绑定 object,但 {object_name} 不在对象目录中")
continue
self.validate_call_kwargs(call_info, self.objects[object_name], result)
variables[var_name] = object_name
result["references"].append({"type": "object", "name": object_name, "variable": var_name})
continue
call_info = self.parse_call(line, result, line)
if not call_info:
continue
function_name = call_info["name"]
if function_name not in self.functions:
self.error(result, f"最终调用必须是已定义 function,但 {function_name} 不在函数目录中")
continue
self.validate_call_kwargs(call_info, self.functions[function_name], result)
self.validate_variable_refs(call_info, variables, result)
result["references"].append({"type": "function", "name": function_name})
if index != len(lines) - 1:
self.error(result, f"非最后一行不能直接调用 function:{line}")
def parse_call(self, expr: str, result: dict[str, Any], raw_line: str) -> dict[str, Any] | None:
try:
node = ast.parse(expr, mode="eval").body
except SyntaxError as exc:
self.error(result, f"调用语法无法解析:{raw_line} ({exc.msg})")
return None
if not isinstance(node, ast.Call) or not isinstance(node.func, ast.Name):
self.error(result, f"不是合法的函数调用:{raw_line}")
return None
return {
"name": node.func.id,
"keywords": node.keywords,
"positional_count": len(node.args),
"raw": raw_line,
}
def validate_call_kwargs(self, call_info: dict[str, Any], spec: dict[str, Any], result: dict[str, Any]) -> None:
if call_info["positional_count"]:
self.error(result, f"暂不支持位置参数,请使用命名参数:{call_info['raw']}")
params = spec.get("params", {})
allow_unknown = bool(spec.get("allow_unknown_params"))
for keyword in call_info["keywords"]:
if keyword.arg is None:
self.error(result, f"暂不支持 **kwargs{call_info['raw']}")
continue
if not allow_unknown and keyword.arg not in params:
self.error(result, f"{call_info['name']} 未定义参数:{keyword.arg}")
continue
enum_values = params.get(keyword.arg, {}).get("enum_values", [])
if enum_values:
self.validate_enum(keyword, enum_values, call_info, result)
def validate_enum(self, keyword: ast.keyword, enum_values: list[str], call_info: dict[str, Any], result: dict[str, Any]) -> None:
values = self.literal_string_values(keyword.value)
for value in values:
if value not in enum_values:
self.error(result, f"{call_info['name']}.{keyword.arg} 枚举值不合法:{value},可选:{', '.join(enum_values)}")
def literal_string_values(self, node: ast.AST) -> list[str]:
if isinstance(node, ast.Constant) and isinstance(node.value, str):
return [node.value]
if isinstance(node, ast.List):
values: list[str] = []
for item in node.elts:
values.extend(self.literal_string_values(item))
return values
return []
def validate_variable_refs(self, call_info: dict[str, Any], variables: dict[str, str], result: dict[str, Any]) -> None:
referenced = sorted({node.id for keyword in call_info["keywords"] for node in ast.walk(keyword.value) if isinstance(node, ast.Name)})
for name in referenced:
if name not in variables:
self.error(result, f"引用了未定义 object 变量:{name}")
def validate_group_json(self, target: str, result: dict[str, Any]) -> None:
result["detected_type"] = "group_json"
try:
groups = json.loads(target)
except json.JSONDecodeError as exc:
self.error(result, f"多指令/自动任务 JSON 解析失败:{exc.msg}")
return
if not isinstance(groups, list):
self.error(result, "多指令/自动任务输出必须是 JSON list")
return
for group_index, group in enumerate(groups):
if not isinstance(group, dict):
self.error(result, f"{group_index} 个 group 不是 object")
continue
if "condition" not in group:
self.error(result, f"{group_index} 个 group 缺少 condition")
querys = group.get("querys")
if not isinstance(querys, list):
self.error(result, f"{group_index} 个 group 的 querys 必须是 list")
continue
for query_index, item in enumerate(querys):
self.validate_group_item(item, group_index, query_index, result)
def validate_group_item(self, item: Any, group_index: int, query_index: int, result: dict[str, Any]) -> None:
prefix = f"{group_index} 个 group 的第 {query_index} 条 query"
if not isinstance(item, dict):
self.error(result, f"{prefix} 不是 object")
return
if not item.get("subquery"):
self.error(result, f"{prefix} 缺少 subquery")
value_fields = [field for field in ("function", "intent", "target", "label") if item.get(field)]
if not value_fields:
self.error(result, f"{prefix} 缺少 function、intent、target 或 label")
return
if len(value_fields) > 1:
self.error(result, f"{prefix} 同时包含多个输出字段:{', '.join(value_fields)}")
return
field = value_fields[0]
value = str(item[field])
if field == "function":
nested = self.validate(value)
self.merge_nested(result, nested, prefix)
elif field == "intent":
if value not in self.intents and value not in self.agent_tags:
self.error(result, f"{prefix} 的 intent 不在知识库中:{value}")
else:
result["references"].append({"type": "intent", "name": value})
else:
nested = self.validate(value)
self.merge_nested(result, nested, prefix)
def merge_nested(self, result: dict[str, Any], nested: dict[str, Any], prefix: str) -> None:
for error in nested.get("errors", []):
self.error(result, f"{prefix}: {error}")
for warning in nested.get("warnings", []):
result["warnings"].append(f"{prefix}: {warning}")
result["references"].extend(nested.get("references", []))
def error(self, result: dict[str, Any], message: str) -> None:
result["errors"].append(message)
def load_targets_from_file(path: Path, field: str) -> list[str]:
text = path.read_text(encoding="utf-8").strip()
if not text:
return []
if path.suffix == ".jsonl":
targets: list[str] = []
for line_no, line in enumerate(text.splitlines(), start=1):
if not line.strip():
continue
item = json.loads(line)
if field not in item:
raise ValueError(f"{path}:{line_no} 缺少字段 {field}")
targets.append(str(item[field]))
return targets
data = json.loads(text)
if isinstance(data, dict):
if field in data:
return [str(data[field])]
if "records" in data and isinstance(data["records"], list):
return [str(item[field]) for item in data["records"] if isinstance(item, dict) and field in item]
if isinstance(data, list):
return [str(item[field]) for item in data if isinstance(item, dict) and field in item]
raise ValueError(f"无法从 {path} 读取字段 {field}")
def main(argv: list[str] | None = None) -> int:
parser = argparse.ArgumentParser(description="校验标签输出格式")
parser.add_argument("--target", help="直接校验一个 target 字符串")
parser.add_argument("--file", help="从 JSON/JSONL 文件中批量读取 target")
parser.add_argument("--field", default="target", help="批量文件中的字段名,默认 target")
parser.add_argument("--manifest", default=str(DEFAULT_OUTPUT), help="manifest JSON 路径")
args = parser.parse_args(argv)
if not args.target and not args.file:
parser.error("必须提供 --target 或 --file")
manifest_path = Path(args.manifest)
if not manifest_path.is_absolute():
manifest_path = SKILL_ROOT / manifest_path
validator = LabelOutputValidator(load_manifest(manifest_path))
targets = [args.target] if args.target else load_targets_from_file(Path(args.file), args.field)
results = [validator.validate(target or "") for target in targets]
payload: dict[str, Any]
if len(results) == 1:
payload = results[0]
else:
payload = {
"valid": all(item["valid"] for item in results),
"total": len(results),
"invalid": sum(1 for item in results if not item["valid"]),
"results": results,
}
print(json.dumps(payload, ensure_ascii=False, indent=2, sort_keys=True))
return 0 if payload["valid"] else 1
if __name__ == "__main__":
raise SystemExit(main())