Add label master skill

This commit is contained in:
wuyang6
2026-05-09 17:49:53 +08:00
parent 84b715a293
commit 010dc1a88d
166 changed files with 13099 additions and 3 deletions
+336
View File
@@ -0,0 +1,336 @@
#!/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())