From 1ae798bcaa0b1f08d6bca3f4376f6874e1b432b8 Mon Sep 17 00:00:00 2001 From: csh Date: Thu, 30 Jul 2026 13:44:50 +0800 Subject: [PATCH] :sparkles: feat(tsl-codegen): support file and directory inputs --- tools/tsl-codegen/README.md | 17 +- tools/tsl-codegen/scripts/generate.py | 451 ++++++++++++++----- tools/tsl-codegen/tests/test_generate.py | 550 ++++++++++++++++------- tools/tsl-codegen/tests/test_pipeline.py | 85 ++-- 4 files changed, 791 insertions(+), 312 deletions(-) diff --git a/tools/tsl-codegen/README.md b/tools/tsl-codegen/README.md index 8394b07b..8c9bef32 100644 --- a/tools/tsl-codegen/README.md +++ b/tools/tsl-codegen/README.md @@ -250,15 +250,26 @@ skills/tsl-api-reference/references/codegen/project/base/array.md 目标文件不存在时,可以直接生成到 skill。例如: ```bash -python tools/tsl-codegen/scripts/generate.py tmp/my-api.yaml +python tools/tsl-codegen/scripts/generate.py --file tmp/my-api.yaml ``` json 使用相同命令: ```bash -python tools/tsl-codegen/scripts/generate.py tmp/my-api.json +python tools/tsl-codegen/scripts/generate.py --file tmp/my-api.json ``` +需要批量生成时,将录入文件放在同一目录并使用 `--dir`: + +```bash +python tools/tsl-codegen/scripts/generate.py --dir tmp/api-recordings +``` + +目录模式只处理目录中的直属文件,不进入子目录。默认处理 `.json`、`.yaml` 和 +`.yml`;`--format json` 只处理 `.json`,`--format yaml` 只处理 `.yaml` 和 +`.yml`。所有录入文件会先完成解析、校验和格式化;预检发现任何错误或输出目标冲突 +时,生成器会列出对应输入文件和具体原因,并且不写入任何 markdown。 + 生成器读取录入文件中的 `path`,默认写入 `skills/tsl-api-reference/references/codegen/project/.md`。不指定 `--scope` 时,scope 就是 `project` @@ -269,7 +280,7 @@ builtin 页面保持一致。未安装 Prettier 或格式化失败时,生成 需要使用自定义 scope 时,通过 `--scope` 指定单级目录名: ```bash -python tools/tsl-codegen/scripts/generate.py tmp/my-api.json --scope my-project +python tools/tsl-codegen/scripts/generate.py --file tmp/my-api.json --scope my-project ``` #### 修改现有叶子页 diff --git a/tools/tsl-codegen/scripts/generate.py b/tools/tsl-codegen/scripts/generate.py index 10297eb8..c0f5ddd9 100644 --- a/tools/tsl-codegen/scripts/generate.py +++ b/tools/tsl-codegen/scripts/generate.py @@ -26,9 +26,10 @@ Each param: name/type/desc required; optional (bool) -> `可选。` prefix; values (list of {value, desc}) -> a `**name 取值**` enum section. Usage (run from repo root): - python tools/tsl-codegen/scripts/generate.py entry.yml - python tools/tsl-codegen/scripts/generate.py entry.json \ + python tools/tsl-codegen/scripts/generate.py --file entry.yml + python tools/tsl-codegen/scripts/generate.py --file entry.json \ --scope my-project + python tools/tsl-codegen/scripts/generate.py --dir recordings """ import argparse @@ -38,15 +39,42 @@ import shutil import subprocess import sys import tempfile -from pathlib import Path +from pathlib import Path, PureWindowsPath REPO_ROOT = Path(__file__).resolve().parents[3] PRETTIER_CONFIG = REPO_ROOT / ".prettierrc.json" +FORMAT_SUFFIXES = { + None: (".json", ".yaml", ".yml"), + "json": (".json",), + "yaml": (".yaml", ".yml"), +} + + +class ChineseArgumentParser(argparse.ArgumentParser): + def error(self, message): + replacements = ( + ("unrecognized arguments:", "无法识别的参数:"), + ("the following arguments are required:", "缺少必需参数:"), + ("expected one argument", "需要一个参数值"), + ("invalid choice: ", "取值无效:"), + ) + if message.startswith("argument "): + message = "参数 " + message[len("argument ") :] + for source, target in replacements: + message = message.replace(source, target) + choice_marker = " (choose from " + if choice_marker in message and message.endswith(")"): + message = message[:-1].replace(choice_marker, "(可选值:", 1) + ")" + self.print_usage(sys.stderr) + self.exit(2, f"{self.prog}: 错误:{message}\n") + + +class GenerationError(Exception): + pass def die(msg): - print(f"ERROR: {msg}", file=sys.stderr) - raise SystemExit(1) + raise GenerationError(msg) def scope_name(value): @@ -65,7 +93,29 @@ def resolve_format(path, fmt): return "json" if suffix in (".yml", ".yaml"): return "yaml" - die(f"cannot infer format from extension '{suffix}'; " f"pass --format json|yaml") + die( + f"无法根据扩展名 {suffix or '(无扩展名)'} 判断录入格式;" + "请使用 --format json 或 --format yaml" + ) + + +def describe_json_error(exc): + translations = ( + ( + "Expecting property name enclosed in double quotes", + "对象属性名必须使用双引号", + ), + ("Expecting value", "此处缺少有效值"), + ("Expecting ',' delimiter", "此处缺少逗号分隔符"), + ("Extra data", "根值之后存在多余内容"), + ("Unterminated string", "字符串未结束"), + ("Invalid \\escape", "字符串中包含无效转义"), + ("Invalid control character", "字符串中包含无效控制字符"), + ) + for prefix, message in translations: + if exc.msg.startswith(prefix): + return message + return "语法无效" def load_entries(path, fmt=None): @@ -73,26 +123,40 @@ def load_entries(path, fmt=None): Format is chosen by --format when given, else by file extension. """ - text = path.read_text(encoding="utf-8") + try: + text = path.read_text(encoding="utf-8") + except UnicodeDecodeError: + die("文件不是有效的 UTF-8 编码") + except OSError as exc: + die(f"读取文件失败:{exc}") fmt = resolve_format(path, fmt) if fmt == "json": try: return json.loads(text) except json.JSONDecodeError as exc: - die(f"invalid JSON in {path}: {exc}") + die( + f"JSON 格式错误:第 {exc.lineno} 行,第 {exc.colno} 列:" + f"{describe_json_error(exc)}" + ) if fmt == "yaml": try: import yaml except ImportError: die( - "pyyaml is not installed; either `pip install pyyaml` or " - "convert the input to .json (json parses with the stdlib)" + "未安装 pyyaml;请运行 `python -m pip install pyyaml`," + "或改用无需额外依赖的 JSON 录入文件" ) try: return yaml.safe_load(text) except yaml.YAMLError as exc: - die(f"invalid YAML in {path}: {exc}") - die(f"unknown format '{fmt}'; use json or yaml") + mark = getattr(exc, "problem_mark", None) + if mark is None: + die("YAML 格式错误:语法无效") + die( + f"YAML 格式错误:第 {mark.line + 1} 行,第 {mark.column + 1} 列:" + "语法无效" + ) + die(f"不支持的录入格式:{fmt};只能使用 json 或 yaml") def escape_cell(text): @@ -106,16 +170,16 @@ def require(cond, msg): def require_mapping(value, where): - require(isinstance(value, dict), f"{where}: must be a mapping") + require(isinstance(value, dict), f"{where}:必须是对象") def reject_unknown(mapping, allowed, where): unknown = sorted(set(mapping) - set(allowed)) - require(not unknown, f"{where}: unknown field(s): {', '.join(unknown)}") + require(not unknown, f"{where}:存在未知字段:{', '.join(unknown)}") def non_empty_string(value, where): - require(isinstance(value, str) and value.strip(), f"{where}: must be non-empty") + require(isinstance(value, str) and value.strip(), f"{where}:不能为空") def optional_draft_string(value, where): @@ -127,7 +191,7 @@ def optional_draft_string(value, where): def validate_tags(tags, where): if tags is None: return - require(isinstance(tags, list) and tags, f"{where}: tags must be a non-empty list") + require(isinstance(tags, list) and tags, f"{where}:tags 必须是非空列表") for index, tag in enumerate(tags): non_empty_string(tag, f"{where}: tags[{index}]") @@ -136,57 +200,67 @@ def signature_names(signature, where): non_empty_string(signature, f"{where}: signature") left = signature.find("(") right = signature.rfind(")") - require(left > 0 and right == len(signature) - 1, f"{where}: invalid signature") + require(left > 0 and right == len(signature) - 1, f"{where}:signature 格式无效") name = signature[:left] non_empty_string(name, f"{where}: signature name") raw = signature[left + 1 : right].strip() if not raw: return name, [] names = [item.strip() for item in raw.split(",")] - require(all(names), f"{where}: signature contains an empty parameter") - require(len({item.casefold() for item in names}) == len(names), f"{where}: duplicate parameter name") + require(all(names), f"{where}:signature 中存在空参数名") + require( + len({item.casefold() for item in names}) == len(names), + f"{where}:signature 中存在重复参数名", + ) return name, names def validate_values(values, where): - require(isinstance(values, list) and values, f"{where}: values must be a non-empty list") + require(isinstance(values, list) and values, f"{where}:values 必须是非空列表") for index, item in enumerate(values): item_where = f"{where}[{index}]" require_mapping(item, item_where) reject_unknown(item, {"value", "desc"}, item_where) - require("value" in item, f"{item_where}: missing 'value'") + require("value" in item, f"{item_where}:缺少 value") non_empty_string(item.get("desc"), f"{item_where}: desc") def validate_params(params, expected_names, where): if not expected_names: - require(not params, f"{where}: nullary signature must not have params") + require(not params, f"{where}:无参数 signature 不能包含 params") return - require(isinstance(params, list), f"{where}: params must be a list") - require(len(params) == len(expected_names), f"{where}: params do not match signature") + require(isinstance(params, list), f"{where}:params 必须是列表") + require(len(params) == len(expected_names), f"{where}:params 与 signature 不匹配") actual_names = [] for index, param in enumerate(params): param_where = f"{where}: params[{index}]" require_mapping(param, param_where) - reject_unknown(param, {"name", "type", "desc", "optional", "values"}, param_where) + reject_unknown( + param, {"name", "type", "desc", "optional", "values"}, param_where + ) name = param.get("name") non_empty_string(name, f"{param_where}: name") non_empty_string(param.get("type"), f"{param_where}: type") non_empty_string(param.get("desc"), f"{param_where}: desc") if "optional" in param: - require(isinstance(param["optional"], bool), f"{param_where}: optional must be boolean") + require( + isinstance(param["optional"], bool), + f"{param_where}:optional 必须是布尔值", + ) if "values" in param: validate_values(param["values"], f"{param_where}: values") actual_names.append(name) require( [name.casefold() for name in actual_names] == [name.casefold() for name in expected_names], - f"{where}: params must follow signature order", + f"{where}:params 必须与 signature 中的参数顺序一致", ) def validate_examples(examples, where): - require(isinstance(examples, list) and examples, f"{where}: examples must be a non-empty list") + require( + isinstance(examples, list) and examples, f"{where}:examples 必须是非空列表" + ) for index, example in enumerate(examples): example_where = f"{where}: examples[{index}]" require_mapping(example, example_where) @@ -215,10 +289,13 @@ def validate_function(fn, where, *, returns_required, extra_fields=()): validate_tags(fn.get("tags"), where) validate_params(fn.get("params"), names, where) if returns_required: - non_empty_string(fn.get("returns"), f"{where}: missing 'returns'") + non_empty_string(fn.get("returns"), f"{where}:缺少 returns") elif "returns" in fn: optional_draft_string(fn["returns"], f"{where}: returns") - require(not ("example" in fn and "examples" in fn), f"{where}: use example or examples, not both") + require( + not ("example" in fn and "examples" in fn), + f"{where}:example 和 examples 不能同时存在", + ) if "example" in fn: non_empty_string(fn["example"], f"{where}: example") if "examples" in fn: @@ -229,10 +306,16 @@ def validate_function(fn, where, *, returns_required, extra_fields=()): def validate_class_member(member, where): require_mapping(member, where) kind = member.get("kind") - require(kind in {"method", "property", "field", "constant"}, f"{where}: unknown kind '{kind}'") + require( + kind in {"method", "property", "field", "constant"}, + f"{where}:未知 kind:{kind}", + ) non_empty_string(member.get("name"), f"{where}: name") visibility = member.get("visibility") - require(visibility in {"public", "protected"}, f"{where}: visibility must be public or protected") + require( + visibility in {"public", "protected"}, + f"{where}:visibility 只能是 public 或 protected", + ) non_empty_string(member.get("desc"), f"{where}: desc") validate_tags(member.get("tags"), where) @@ -240,25 +323,46 @@ def validate_class_member(member, where): reject_unknown( member, { - "kind", "name", "visibility", "binding", "signature", "desc", - "tags", "params", "returns", "modifiers", "example", "examples", + "kind", + "name", + "visibility", + "binding", + "signature", + "desc", + "tags", + "params", + "returns", + "modifiers", + "example", + "examples", }, where, ) - require(member.get("binding") in {"instance", "class"}, f"{where}: invalid binding") + require( + member.get("binding") in {"instance", "class"}, + f"{where}:binding 无效", + ) parsed_name = validate_function( member, where, returns_required=False, extra_fields={"kind", "name", "visibility", "binding", "modifiers"}, ) - require(parsed_name.casefold() == member["name"].casefold(), f"{where}: name and signature differ") + require( + parsed_name.casefold() == member["name"].casefold(), + f"{where}:name 与 signature 中的名称不一致", + ) if "modifiers" in member: modifiers = member["modifiers"] - require(isinstance(modifiers, list), f"{where}: modifiers must be a list") + require(isinstance(modifiers, list), f"{where}:modifiers 必须是列表") allowed = {"overload", "virtual", "override"} - require(all(item in allowed for item in modifiers), f"{where}: invalid modifier") - require(len(set(modifiers)) == len(modifiers), f"{where}: duplicate modifier") + require( + all(item in allowed for item in modifiers), + f"{where}:存在无效 modifier", + ) + require( + len(set(modifiers)) == len(modifiers), f"{where}:存在重复 modifier" + ) return common = {"kind", "name", "visibility", "desc", "tags"} @@ -266,7 +370,10 @@ def validate_class_member(member, where): reject_unknown(member, common | {"type", "params", "access"}, where) if "type" in member: optional_draft_string(member["type"], f"{where}: type") - require(member.get("access") in {"read", "write", "readwrite"}, f"{where}: invalid access") + require( + member.get("access") in {"read", "write", "readwrite"}, + f"{where}:access 无效", + ) params = member.get("params") if params: expected = [param.get("name") for param in params] @@ -276,30 +383,30 @@ def validate_class_member(member, where): reject_unknown(member, common | {"type", "static"}, where) non_empty_string(member.get("type"), f"{where}: type") if "static" in member: - require(isinstance(member["static"], bool), f"{where}: static must be boolean") + require(isinstance(member["static"], bool), f"{where}:static 必须是布尔值") return reject_unknown(member, common | {"type", "value", "static"}, where) - require("value" in member and member["value"] is not None, f"{where}: missing 'value'") + require("value" in member and member["value"] is not None, f"{where}:缺少 value") if isinstance(member["value"], str): non_empty_string(member["value"], f"{where}: value") if "type" in member: optional_draft_string(member["type"], f"{where}: type") if "static" in member: - require(isinstance(member["static"], bool), f"{where}: static must be boolean") + require(isinstance(member["static"], bool), f"{where}:static 必须是布尔值") def validate_class(cls, where): require_mapping(cls, where) reject_unknown(cls, {"kind", "name", "desc", "tags", "bases", "members"}, where) - require(cls.get("kind") == "class", f"{where}: kind must be class") + require(cls.get("kind") == "class", f"{where}:kind 必须是 class") non_empty_string(cls.get("name"), f"{where}: name") non_empty_string(cls.get("desc"), f"{where}: desc") validate_tags(cls.get("tags"), where) if "bases" in cls: - require(isinstance(cls["bases"], list), f"{where}: bases must be a list") + require(isinstance(cls["bases"], list), f"{where}:bases 必须是列表") for index, base in enumerate(cls["bases"]): non_empty_string(base, f"{where}: bases[{index}]") - require(isinstance(cls.get("members"), list), f"{where}: members must be a list") + require(isinstance(cls.get("members"), list), f"{where}:members 必须是列表") for index, member in enumerate(cls["members"]): validate_class_member(member, f"{where}: members[{index}]") @@ -307,7 +414,10 @@ def validate_class(cls, where): def validate_unit_member(member, where): require_mapping(member, where) kind = member.get("kind") - require(kind in {"function", "variable", "constant", "class"}, f"{where}: unknown kind '{kind}'") + require( + kind in {"function", "variable", "constant", "class"}, + f"{where}:未知 kind:{kind}", + ) if kind == "class": validate_class(member, where) return @@ -321,7 +431,10 @@ def validate_unit_member(member, where): returns_required=True, extra_fields={"kind", "name"}, ) - require(parsed_name.casefold() == member["name"].casefold(), f"{where}: name and signature differ") + require( + parsed_name.casefold() == member["name"].casefold(), + f"{where}:name 与 signature 中的名称不一致", + ) return common = {"kind", "name", "desc", "tags", "type"} if kind == "variable": @@ -329,7 +442,7 @@ def validate_unit_member(member, where): non_empty_string(member.get("type"), f"{where}: type") return reject_unknown(member, common | {"value"}, where) - require("value" in member and member["value"] is not None, f"{where}: missing 'value'") + require("value" in member and member["value"] is not None, f"{where}:缺少 value") if isinstance(member["value"], str): non_empty_string(member["value"], f"{where}: value") if "type" in member: @@ -339,13 +452,13 @@ def validate_unit_member(member, where): def validate_unit(unit, where): require_mapping(unit, where) reject_unknown(unit, {"kind", "name", "desc", "tags", "members"}, where) - require(unit.get("kind") == "unit", f"{where}: kind must be unit") + require(unit.get("kind") == "unit", f"{where}:kind 必须是 unit") non_empty_string(unit.get("name"), f"{where}: name") non_empty_string(unit.get("desc"), f"{where}: desc") validate_tags(unit.get("tags"), where) require( isinstance(unit.get("members"), list), - f"{where}: members must be a list", + f"{where}:members 必须是列表", ) for index, member in enumerate(unit["members"]): validate_unit_member(member, f"{where}: members[{index}]") @@ -354,7 +467,7 @@ def validate_unit(unit, where): def validate_top_level_function(declaration, where): require( declaration.get("kind") == "function", - f"{where}: kind must be function", + f"{where}:kind 必须是 function", ) non_empty_string(declaration.get("name"), f"{where}: name") parsed_name = validate_function( @@ -365,7 +478,7 @@ def validate_top_level_function(declaration, where): ) require( parsed_name.casefold() == declaration["name"].casefold(), - f"{where}: name and signature differ", + f"{where}:name 与 signature 中的名称不一致", ) @@ -374,7 +487,7 @@ def validate_declaration(declaration, where): kind = declaration.get("kind") require( kind in {"function", "class", "unit"}, - f"{where}: unknown kind '{kind}'", + f"{where}:未知 kind:{kind}", ) if kind == "function": validate_top_level_function(declaration, where) @@ -395,25 +508,25 @@ def validate_declaration_uniqueness(declarations): key = declaration["signature"].casefold() require( key not in function_signatures, - f"{where}: duplicate function signature", + f"{where}:function signature 重复", ) function_signatures.add(key) continue names = class_names if kind == "class" else unit_names key = declaration["name"].casefold() - require(key not in names, f"{where}: duplicate {kind} name") + require(key not in names, f"{where}:{kind} 名称重复") names.add(key) def validate_page(data): - require_mapping(data, "input root") - reject_unknown(data, {"module", "path", "declarations"}, "input root") - non_empty_string(data.get("module"), "input root: module") - non_empty_string(data.get("path"), "input root: path") + require_mapping(data, "录入根") + reject_unknown(data, {"module", "path", "declarations"}, "录入根") + non_empty_string(data.get("module"), "录入根:module") + non_empty_string(data.get("path"), "录入根:path") declarations = data.get("declarations") require( isinstance(declarations, list) and declarations, - "input root: declarations must be a non-empty list", + "录入根:declarations 必须是非空列表", ) for index, declaration in enumerate(declarations): validate_declaration(declaration, f"declarations[{index}]") @@ -424,7 +537,7 @@ def validate_page(data): def param_desc(param, where): """Description column text: prepend `可选。` for optional params.""" desc = param.get("desc") - require(desc, f"{where}: param '{param.get('name', '?')}' missing desc") + require(desc, f"{where}:参数 {param.get('name', '?')} 缺少 desc") if param.get("optional") and not desc.startswith("可选。"): return "可选。" + desc return desc @@ -436,8 +549,8 @@ def render_param_table(params, where): for param in params: name = param.get("name") ptype = param.get("type") - require(name, f"{where}: a param is missing 'name'") - require(ptype, f"{where}: param '{name}' missing 'type'") + require(name, f"{where}:存在缺少 name 的参数") + require(ptype, f"{where}:参数 {name} 缺少 type") desc = param_desc(param, where) lines.append( f"| `{escape_cell(name)}` | {escape_cell(ptype)} | {escape_cell(desc)} |" @@ -479,7 +592,9 @@ def render_examples(fn, heading_level): lines.append(f"// 输出:{output_lines[0]}") else: lines.append("// 输出:") - lines.extend(f"// {line}" if line else "//" for line in output_lines) + lines.extend( + f"// {line}" if line else "//" for line in output_lines + ) lines.append("```") lines.append("") while lines and not lines[-1]: @@ -520,9 +635,7 @@ def render_callable( returns_required=returns_required, extra_fields={"kind", "name", "visibility", "binding", "modifiers"}, ) - lines = render_intro( - f"{'#' * level} `{signature}`", declaration, fn - ) + lines = render_intro(f"{'#' * level} `{signature}`", declaration, fn) if show_visibility: lines.extend([f"可见性:`{fn['visibility']}`", ""]) if fn.get("modifiers"): @@ -563,11 +676,7 @@ def render_top_level_function(fn, index): def render_class_member(member, level, where): kind = member["kind"] if kind == "method": - declaration = ( - "class function" - if member["binding"] == "class" - else "function" - ) + declaration = "class function" if member["binding"] == "class" else "function" return render_callable( member, member["signature"], @@ -586,14 +695,14 @@ def render_class_member(member, level, where): signature = member["name"] if kind == "property" and member.get("params"): signature += "(" + ", ".join(param["name"] for param in member["params"]) + ")" - lines = render_intro( - f"{'#' * level} `{signature}`", declaration, member - ) + lines = render_intro(f"{'#' * level} `{signature}`", declaration, member) lines.extend([f"可见性:`{member['visibility']}`", ""]) if kind == "property": if member.get("type"): lines.extend([f"类型:{member['type']}", ""]) - access = {"read": "read", "write": "write", "readwrite": "read / write"}[member["access"]] + access = {"read": "read", "write": "write", "readwrite": "read / write"}[ + member["access"] + ] lines.append(f"访问:{access}") params = member.get("params") or [] if params: @@ -609,13 +718,13 @@ def render_class_member(member, level, where): def render_class(cls, level, where, *, page_root): - lines = render_intro( - f"{'#' * level} `{cls['name']}`", "class", cls - ) + lines = render_intro(f"{'#' * level} `{cls['name']}`", "class", cls) if cls.get("bases"): lines.extend(["父类:" + "、".join(f"`{base}`" for base in cls["bases"]), ""]) for index, member in enumerate(cls["members"]): - lines.extend(render_class_member(member, level + 1, f"{where}: members[{index}]")) + lines.extend( + render_class_member(member, level + 1, f"{where}: members[{index}]") + ) lines.append("") while lines and not lines[-1]: lines.pop() @@ -638,9 +747,7 @@ def render_unit_member(member, index, unit_where): if kind == "class": return render_class(member, 3, where, page_root=False) declaration = "var" if kind == "variable" else "const" - lines = render_intro( - f"### `{member['name']}`", declaration, member - ) + lines = render_intro(f"### `{member['name']}`", declaration, member) if kind == "variable": lines.append(f"类型:{member['type']}") else: @@ -718,13 +825,34 @@ def format_markdown(text): def output_path(data, scope): """Build the leaf-page destination from the recording file's relative path.""" relative = data.get("path") - require(relative, "input missing 'path'") - require(isinstance(relative, str), "'path' must be a string") - require("\\" not in relative, "'path' must use '/' as the separator") - relative_path = Path(relative) - require(not relative_path.is_absolute(), "'path' must be relative") - require(".." not in relative_path.parts, "'path' must not contain '..'") - require(relative_path.suffix == "", "'path' must not include a file extension") + require(relative, "录入数据缺少 path") + require(isinstance(relative, str), "path 必须是字符串") + recorded_path = PureWindowsPath(relative) + require( + not recorded_path.drive.startswith("\\\\"), + "path 必须是相对路径,不允许 UNC 路径", + ) + require( + not recorded_path.drive, + "path 必须是相对路径,不允许 Windows 盘符路径", + ) + require( + not recorded_path.root, + "path 必须是相对路径,不能以斜杠或反斜杠开头", + ) + require( + bool(recorded_path.parts), + "path 必须指向具体文档,不能只表示当前目录", + ) + require( + ".." not in recorded_path.parts, + "path 不能包含 '..',不允许跳转到父目录", + ) + require( + recorded_path.suffix == "", + "path 不能包含 .md 等文件扩展名", + ) + relative_path = Path(*recorded_path.parts) return ( Path("skills/tsl-api-reference/references/codegen") / scope @@ -750,11 +878,75 @@ def atomic_write(path, text): temporary_path.unlink() +def gather_directory_inputs(directory, fmt): + try: + if not directory.is_dir(): + die(f"输入目录不存在或不是目录:{directory}") + suffixes = FORMAT_SUFFIXES[fmt] + inputs = sorted( + ( + path + for path in directory.iterdir() + if path.is_file() and path.suffix.lower() in suffixes + ), + key=lambda path: (path.name.casefold(), path.name), + ) + except OSError as exc: + die(f"读取输入目录失败:{directory}:{exc}") + if not inputs: + displayed_suffixes = "、".join(suffixes) + die( + f"目录 {directory} 中未找到符合条件的直属录入文件" + f"({displayed_suffixes})" + ) + return inputs + + +def prepare_input(in_path, input_format, scope): + data = load_entries(in_path, input_format) + try: + rendered = render_page(data) + except GenerationError as exc: + raise GenerationError(f"录入数据校验失败:{exc}") from exc + try: + out_path = output_path(data, scope) + except GenerationError as exc: + raise GenerationError(f"输出路径无效:{exc}") from exc + try: + formatted = format_markdown(rendered) + except OSError as exc: + raise GenerationError(f"运行 Prettier 失败:{exc}") from exc + return in_path, out_path, formatted + + +def report_input_error(in_path, error): + print(f"错误:{in_path}:{error}", file=sys.stderr) + + +def find_output_collisions(prepared): + owners = {} + collisions = [] + for in_path, out_path, _ in prepared: + resolved_output = out_path.resolve() + first_input = owners.get(resolved_output) + if first_input is None: + owners[resolved_output] = in_path + continue + collisions.append( + f"输出目标冲突:{first_input} 和 {in_path} 都会写入 {resolved_output}" + ) + return collisions + + def main(argv=None): if hasattr(sys.stdout, "reconfigure"): sys.stdout.reconfigure(encoding="utf-8") - parser = argparse.ArgumentParser( + parser = ChineseArgumentParser( description="从 YAML/JSON 录入文件生成 TSL API 文档", + usage=( + "%(prog)s [--help] [--scope SCOPE] [--format {json,yaml}] " + "(--file INPUT_FILE | --dir INPUT_DIR | INPUT_FILE)" + ), add_help=False, allow_abbrev=False, ) @@ -764,9 +956,22 @@ def main(argv=None): help="显示本帮助并退出(不提供 -h 短选项)", ) parser.add_argument( - "input", + "legacy_input", + nargs="?", metavar="INPUT_FILE", - help="YAML/JSON 录入文件路径,例如 tmp/my-functions.yaml", + help="已废弃,请使用 --file;暂时兼容 YAML/JSON 录入文件路径", + ) + parser.add_argument( + "--file", + dest="input_file", + metavar="INPUT_FILE", + help="要生成的单个 YAML/JSON 录入文件", + ) + parser.add_argument( + "--dir", + dest="input_dir", + metavar="INPUT_DIR", + help="批量生成目录中的直属 YAML/JSON 录入文件", ) parser.add_argument( "--scope", @@ -781,18 +986,54 @@ def main(argv=None): ) args = parser.parse_args(argv) - in_path = Path(args.input) - if not in_path.is_file(): - die(f"input not found: {in_path}") - data = load_entries(in_path, args.format) + input_modes = (args.legacy_input, args.input_file, args.input_dir) + if sum(value is not None for value in input_modes) != 1: + parser.error("必须且只能指定一种输入方式:INPUT_FILE、--file 或 --dir") + is_batch = args.input_dir is not None + try: + if is_batch: + input_paths = gather_directory_inputs(Path(args.input_dir), args.format) + input_format = None + else: + in_path = Path(args.input_file or args.legacy_input) + if not in_path.is_file(): + die(f"输入文件不存在或不是普通文件:{in_path}") + input_paths = [in_path] + input_format = args.format + except GenerationError as exc: + print(f"错误:{exc}", file=sys.stderr) + return 1 - text = format_markdown(render_page(data)) - out_path = output_path(data, args.scope) - atomic_write(out_path, text) - print( - f"wrote {out_path}", - file=sys.stderr, - ) + prepared = [] + input_errors = [] + for in_path in input_paths: + try: + prepared.append(prepare_input(in_path, input_format, args.scope)) + except GenerationError as exc: + input_errors.append((in_path, exc)) + + collisions = find_output_collisions(prepared) + if input_errors or collisions: + for in_path, error in input_errors: + report_input_error(in_path, error) + for collision in collisions: + print(f"错误:{collision}", file=sys.stderr) + if is_batch: + issue_count = len(input_errors) + len(collisions) + print( + f"错误:批量生成已中止:发现 {issue_count} 个问题," + f"共检查 {len(input_paths)} 个文件;未写入任何 Markdown 文件", + file=sys.stderr, + ) + return 1 + + for in_path, out_path, text in prepared: + try: + atomic_write(out_path, text) + except OSError as exc: + report_input_error(in_path, f"写入 Markdown 失败:{out_path}:{exc}") + return 1 + print(f"已生成:{in_path} -> {out_path}", file=sys.stderr) return 0 diff --git a/tools/tsl-codegen/tests/test_generate.py b/tools/tsl-codegen/tests/test_generate.py index 7c76daed..dcec9be0 100644 --- a/tools/tsl-codegen/tests/test_generate.py +++ b/tools/tsl-codegen/tests/test_generate.py @@ -9,6 +9,11 @@ import unittest from pathlib import Path from unittest import mock +try: + import yaml +except ImportError: + yaml = None + SCRIPT = Path(__file__).parents[1] / "scripts" / "generate.py" REPO_ROOT = Path(__file__).resolve().parents[3] @@ -50,7 +55,17 @@ class DocGenCliTest(unittest.TestCase): def run_cli(self, *args, env=None): return subprocess.run( - [sys.executable, str(SCRIPT), str(self.input), *args], + [sys.executable, str(SCRIPT), "--file", str(self.input), *args], + capture_output=True, + text=True, + encoding="utf-8", + cwd=self.root, + env=env, + ) + + def run_raw_cli(self, *args, env=None): + return subprocess.run( + [sys.executable, str(SCRIPT), *map(str, args)], capture_output=True, text=True, encoding="utf-8", @@ -64,6 +79,29 @@ class DocGenCliTest(unittest.TestCase): encoding="utf-8", ) + def write_recording(self, path, relative, module="项目 / 批量"): + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text( + json.dumps( + { + "module": module, + "path": relative, + "declarations": [ + { + "kind": "function", + "name": "demo", + "signature": "demo()", + "desc": "示例函数。", + "returns": "nil", + } + ], + }, + ensure_ascii=False, + ), + encoding="utf-8", + ) + return path + def generated(self, relative, scope="project"): return ( self.root @@ -120,6 +158,232 @@ class DocGenCliTest(unittest.TestCase): output.read_text(encoding="utf-8").startswith("# 项目 / 示例\n") ) + def test_file_option_generates_one_recording(self): + result = self.run_raw_cli("--file", self.input) + + self.assertEqual(0, result.returncode, result.stderr) + self.assertTrue(self.generated("base/my_functions").is_file()) + + def test_legacy_positional_input_remains_supported(self): + result = self.run_raw_cli(self.input) + + self.assertEqual(0, result.returncode, result.stderr) + self.assertTrue(self.generated("base/my_functions").is_file()) + + def test_help_marks_legacy_input_as_deprecated(self): + result = self.run_raw_cli("--help") + + self.assertEqual(0, result.returncode, result.stderr) + self.assertIn("--file INPUT_FILE", result.stdout) + self.assertIn("--dir INPUT_DIR", result.stdout) + self.assertIn("已废弃,请使用 --file", result.stdout) + + def test_invalid_format_error_is_fully_chinese(self): + result = self.run_raw_cli("--file", self.input, "--format", "xml") + + self.assertEqual(2, result.returncode) + self.assertIn( + "参数 --format: 取值无效:'xml'(可选值:'json', 'yaml')", + result.stderr, + ) + self.assertNotIn("choose from", result.stderr) + + def test_exactly_one_input_mode_is_required(self): + cases = ( + (), + ("--file", self.input, "--dir", self.root), + (self.input, "--file", self.input), + ) + for args in cases: + with self.subTest(args=args): + result = self.run_raw_cli(*args) + + self.assertEqual(2, result.returncode) + self.assertIn( + "必须且只能指定一种输入方式:INPUT_FILE、--file 或 --dir", + result.stderr, + ) + + def test_dir_processes_only_direct_json_files_in_name_order(self): + input_dir = self.root / "recordings" + second = self.write_recording(input_dir / "b.json", "base/b") + first = self.write_recording(input_dir / "a.json", "base/a") + self.write_recording(input_dir / "nested" / "c.json", "base/c") + (input_dir / "notes.txt").write_text("ignore", encoding="utf-8") + + result = self.run_raw_cli("--dir", input_dir, "--format", "json") + + self.assertEqual(0, result.returncode, result.stderr) + self.assertTrue(self.generated("base/a").is_file()) + self.assertTrue(self.generated("base/b").is_file()) + self.assertFalse(self.generated("base/c").exists()) + self.assertLess( + result.stderr.index(str(first)), result.stderr.index(str(second)) + ) + + def test_dir_format_filters_out_other_supported_extensions(self): + input_dir = self.root / "recordings" + self.write_recording(input_dir / "only.json", "base/only") + (input_dir / "ignored.yaml").write_text(": invalid", encoding="utf-8") + + result = self.run_raw_cli("--dir", input_dir, "--format", "json") + + self.assertEqual(0, result.returncode, result.stderr) + self.assertTrue(self.generated("base/only").is_file()) + self.assertNotIn("ignored.yaml", result.stderr) + + @unittest.skipUnless(yaml is not None, "pyyaml is not installed") + def test_dir_without_format_processes_json_yaml_and_yml(self): + input_dir = self.root / "recordings" + self.write_recording(input_dir / "first.json", "base/first") + for name, relative in ( + ("second.yaml", "base/second"), + ("third.yml", "base/third"), + ): + data = { + "module": "项目 / 批量", + "path": relative, + "declarations": [ + { + "kind": "function", + "name": "demo", + "signature": "demo()", + "desc": "示例函数。", + "returns": "nil", + } + ], + } + (input_dir / name).write_text( + yaml.safe_dump(data, allow_unicode=True, sort_keys=False), + encoding="utf-8", + ) + + result = self.run_raw_cli("--dir", input_dir) + + self.assertEqual(0, result.returncode, result.stderr) + for relative in ("base/first", "base/second", "base/third"): + self.assertTrue(self.generated(relative).is_file()) + + def test_dir_without_matching_files_reports_directory_and_extensions(self): + input_dir = self.root / "recordings" + input_dir.mkdir() + (input_dir / "notes.txt").write_text("ignore", encoding="utf-8") + self.write_recording(input_dir / "nested" / "hidden.json", "base/hidden") + + result = self.run_raw_cli("--dir", input_dir) + + self.assertEqual(1, result.returncode) + self.assertIn(str(input_dir), result.stderr) + self.assertIn("未找到符合条件的直属录入文件", result.stderr) + self.assertIn(".json、.yaml、.yml", result.stderr) + + def test_dir_scan_error_reports_directory_and_specific_reason(self): + module = load_script() + input_dir = self.root / "recordings" + input_dir.mkdir() + + with mock.patch.object( + module.Path, "iterdir", side_effect=PermissionError("拒绝访问") + ): + with self.assertRaises(module.GenerationError) as caught: + module.gather_directory_inputs(input_dir, None) + + message = str(caught.exception) + self.assertIn(f"读取输入目录失败:{input_dir}", message) + self.assertIn("拒绝访问", message) + + def test_dir_collects_all_file_errors_before_writing_any_markdown(self): + input_dir = self.root / "recordings" + self.write_recording(input_dir / "a-valid.json", "base/valid") + (input_dir / "b-invalid-json.json").write_text("{", encoding="utf-8") + (input_dir / "c-invalid-data.json").write_text( + json.dumps( + { + "module": "项目 / 错误", + "path": "base/invalid_data", + "declarations": [], + }, + ensure_ascii=False, + ), + encoding="utf-8", + ) + existing = self.generated("base/valid") + existing.parent.mkdir(parents=True, exist_ok=True) + existing.write_text("原内容\n", encoding="utf-8") + + result = self.run_raw_cli("--dir", input_dir, "--format", "json") + + self.assertEqual(1, result.returncode) + self.assertIn(str(input_dir / "b-invalid-json.json"), result.stderr) + self.assertIn("JSON 格式错误", result.stderr) + self.assertIn(str(input_dir / "c-invalid-data.json"), result.stderr) + self.assertIn("录入数据校验失败", result.stderr) + self.assertIn("declarations 必须是非空列表", result.stderr) + self.assertIn("批量生成已中止", result.stderr) + self.assertIn("未写入任何 Markdown 文件", result.stderr) + self.assertEqual("原内容\n", existing.read_text(encoding="utf-8")) + + def test_dir_rejects_output_collisions_before_writing(self): + input_dir = self.root / "recordings" + first = self.write_recording(input_dir / "a.json", "base/collision") + second = self.write_recording(input_dir / "b.json", "base/collision") + output = self.generated("base/collision") + + result = self.run_raw_cli("--dir", input_dir, "--format", "json") + + self.assertEqual(1, result.returncode) + self.assertIn("输出目标冲突", result.stderr) + self.assertIn(str(first), result.stderr) + self.assertIn(str(second), result.stderr) + self.assertIn(str(output), result.stderr) + self.assertIn("未写入任何 Markdown 文件", result.stderr) + self.assertFalse(output.exists()) + + def test_single_file_json_error_reports_file_and_position_in_chinese(self): + self.input.write_text("{", encoding="utf-8") + + result = self.run_raw_cli("--file", self.input) + + self.assertEqual(1, result.returncode) + self.assertIn(f"错误:{self.input}:JSON 格式错误", result.stderr) + self.assertIn("第 1 行", result.stderr) + self.assertIn("第 2 列", result.stderr) + + def test_windows_path_separator_writes_nested_markdown(self): + data = json.loads(self.input.read_text(encoding="utf-8")) + data["path"] = r"base\windows_example" + self.write_input(data) + + result = self.run_raw_cli("--file", self.input) + + self.assertEqual(0, result.returncode, result.stderr) + self.assertTrue(self.generated("base/windows_example").is_file()) + + def test_invalid_paths_report_input_file_and_specific_reason(self): + cases = ( + (r"C:\base\example", "不允许 Windows 盘符路径"), + (r"\\server\share\example", "不允许 UNC 路径"), + (r"\base\example", "不能以斜杠或反斜杠开头"), + (r"base\..\example", "不允许跳转到父目录"), + (r"base\example.md", "不能包含 .md 等文件扩展名"), + (".", "必须指向具体文档"), + ) + for path, reason in cases: + with self.subTest(path=path): + data = json.loads(self.input.read_text(encoding="utf-8")) + data["path"] = path + self.write_input(data) + + result = self.run_raw_cli("--file", self.input) + + self.assertEqual(1, result.returncode, result.stderr) + self.assertIn( + f"错误:{self.input}:输出路径无效:", + result.stderr, + ) + self.assertIn(reason, result.stderr) + self.assertNotIn("Traceback", result.stderr) + def test_custom_scope_changes_first_destination_directory(self): result = self.run_cli("--scope", "my-project") output = ( @@ -138,7 +402,7 @@ class DocGenCliTest(unittest.TestCase): def test_output_option_is_rejected(self): result = self.run_cli("--output", str(self.root / "out.md")) self.assertNotEqual(result.returncode, 0) - self.assertIn("unrecognized arguments: --output", result.stderr) + self.assertIn("无法识别的参数: --output", result.stderr) def test_generated_markdown_is_formatted_by_repo_prettier(self): self.input.write_text( @@ -304,54 +568,54 @@ class DocGenCliTest(unittest.TestCase): "tags": ["组件", "示例"], "bases": ["BaseWidget"], "members": [ - { - "kind": "method", - "name": "Close", - "visibility": "public", - "binding": "instance", - "signature": "Close()", - "desc": "关闭组件。", - }, - { - "kind": "method", - "name": "Create", - "visibility": "protected", - "binding": "class", - "signature": "Create(name)", - "desc": "创建组件。", - "params": [ - { - "name": "name", - "type": "string", - "desc": "组件名称", - } - ], - "returns": "Widget", - "modifiers": ["overload"], - }, - { - "kind": "property", - "name": "Title", - "visibility": "public", - "desc": "组件标题。", - "type": "string", - "access": "readwrite", - }, - { - "kind": "field", - "name": "Count", - "visibility": "protected", - "desc": "组件数量。", - "type": "integer", - "static": True, - }, - { - "kind": "constant", - "name": "DefaultName", - "visibility": "public", - "desc": "默认名称。", - "value": "'widget'", - }, + { + "kind": "method", + "name": "Close", + "visibility": "public", + "binding": "instance", + "signature": "Close()", + "desc": "关闭组件。", + }, + { + "kind": "method", + "name": "Create", + "visibility": "protected", + "binding": "class", + "signature": "Create(name)", + "desc": "创建组件。", + "params": [ + { + "name": "name", + "type": "string", + "desc": "组件名称", + } + ], + "returns": "Widget", + "modifiers": ["overload"], + }, + { + "kind": "property", + "name": "Title", + "visibility": "public", + "desc": "组件标题。", + "type": "string", + "access": "readwrite", + }, + { + "kind": "field", + "name": "Count", + "visibility": "protected", + "desc": "组件数量。", + "type": "integer", + "static": True, + }, + { + "kind": "constant", + "name": "DefaultName", + "visibility": "public", + "desc": "默认名称。", + "value": "'widget'", + }, ], } ], @@ -428,49 +692,49 @@ class DocGenCliTest(unittest.TestCase): "name": "DemoUnit", "desc": "提供文档能力。", "members": [ - { - "kind": "constant", - "name": "DefaultSize", - "desc": "默认大小。", - "value": 100, - }, - { - "kind": "variable", - "name": "CurrentDocument", - "desc": "当前文档。", - "type": "Document", - }, - { - "kind": "function", - "name": "OpenDocument", - "signature": "OpenDocument(path)", - "desc": "打开文档。", - "params": [ - { - "name": "path", - "type": "string", - "desc": "文档路径", - } - ], - "returns": "Document", - }, - { - "kind": "class", - "name": "Document", - "desc": "文档对象。", - "bases": ["BaseDocument"], - "members": [ - { - "kind": "method", - "name": "Save", - "visibility": "public", - "binding": "instance", - "signature": "Save()", - "desc": "保存文档。", - "returns": "boolean", - } - ], - }, + { + "kind": "constant", + "name": "DefaultSize", + "desc": "默认大小。", + "value": 100, + }, + { + "kind": "variable", + "name": "CurrentDocument", + "desc": "当前文档。", + "type": "Document", + }, + { + "kind": "function", + "name": "OpenDocument", + "signature": "OpenDocument(path)", + "desc": "打开文档。", + "params": [ + { + "name": "path", + "type": "string", + "desc": "文档路径", + } + ], + "returns": "Document", + }, + { + "kind": "class", + "name": "Document", + "desc": "文档对象。", + "bases": ["BaseDocument"], + "members": [ + { + "kind": "method", + "name": "Save", + "visibility": "public", + "binding": "instance", + "signature": "Save()", + "desc": "保存文档。", + "returns": "boolean", + } + ], + }, ], } ], @@ -579,7 +843,7 @@ class DocGenCliTest(unittest.TestCase): result = self.run_cli() self.assertEqual(1, result.returncode) - self.assertIn("missing 'returns'", result.stderr) + self.assertIn("缺少 returns", result.stderr) def test_examples_list_renders_independent_fences_and_output_comments(self): self.write_input( @@ -625,7 +889,9 @@ class DocGenCliTest(unittest.TestCase): output = self.root / "existing.md" output.write_text("原内容\n", encoding="utf-8") - with mock.patch.object(module.os, "replace", side_effect=OSError("replace failed")): + with mock.patch.object( + module.os, "replace", side_effect=OSError("replace failed") + ): with self.assertRaisesRegex(OSError, "replace failed"): module.atomic_write(output, "新内容\n") @@ -634,9 +900,7 @@ class DocGenCliTest(unittest.TestCase): def test_legacy_root_keys_are_rejected_without_overwrite(self): legacy_values = { - "functions": [ - {"signature": "Old()", "desc": "旧函数。", "returns": "nil"} - ], + "functions": [{"signature": "Old()", "desc": "旧函数。", "returns": "nil"}], "class": {"name": "Old", "desc": "旧类。", "members": []}, "unit": {"name": "Old", "desc": "旧接口。", "members": []}, } @@ -651,9 +915,7 @@ class DocGenCliTest(unittest.TestCase): } ) - self.assert_rejected_without_overwrite( - relative, f"unknown field(s): {key}" - ) + self.assert_rejected_without_overwrite(relative, f"存在未知字段:{key}") def test_unknown_class_member_kind_is_rejected_without_overwrite(self): self.write_class_member_input( @@ -666,9 +928,7 @@ class DocGenCliTest(unittest.TestCase): "base/unknown_kind", ) - self.assert_rejected_without_overwrite( - "base/unknown_kind", "unknown kind 'event'" - ) + self.assert_rejected_without_overwrite("base/unknown_kind", "未知 kind:event") def test_method_static_field_is_rejected_without_overwrite(self): self.write_class_member_input( @@ -685,7 +945,7 @@ class DocGenCliTest(unittest.TestCase): ) self.assert_rejected_without_overwrite( - "base/method_static", "unknown field(s): static" + "base/method_static", "存在未知字段:static" ) def test_private_class_member_is_rejected_without_overwrite(self): @@ -701,7 +961,7 @@ class DocGenCliTest(unittest.TestCase): ) self.assert_rejected_without_overwrite( - "base/private_member", "visibility must be public or protected" + "base/private_member", "visibility 只能是 public 或 protected" ) def test_property_empty_type_is_treated_as_omitted(self): @@ -757,7 +1017,7 @@ class DocGenCliTest(unittest.TestCase): ) self.assert_rejected_without_overwrite( - "base/field_type", "members[0]: type: must be non-empty" + "base/field_type", "members[0]: type:不能为空" ) def test_unit_variable_missing_type_is_rejected_without_overwrite(self): @@ -783,7 +1043,7 @@ class DocGenCliTest(unittest.TestCase): ) self.assert_rejected_without_overwrite( - "base/variable_type", "members[0]: type: must be non-empty" + "base/variable_type", "members[0]: type:不能为空" ) def test_invalid_method_binding_is_rejected_without_overwrite(self): @@ -799,9 +1059,7 @@ class DocGenCliTest(unittest.TestCase): "base/invalid_binding", ) - self.assert_rejected_without_overwrite( - "base/invalid_binding", "invalid binding" - ) + self.assert_rejected_without_overwrite("base/invalid_binding", "binding 无效") def test_invalid_property_access_is_rejected_without_overwrite(self): self.write_class_member_input( @@ -816,9 +1074,7 @@ class DocGenCliTest(unittest.TestCase): "base/invalid_access", ) - self.assert_rejected_without_overwrite( - "base/invalid_access", "invalid access" - ) + self.assert_rejected_without_overwrite("base/invalid_access", "access 无效") def test_invalid_method_modifier_is_rejected_without_overwrite(self): self.write_class_member_input( @@ -834,9 +1090,7 @@ class DocGenCliTest(unittest.TestCase): "base/invalid_modifier", ) - self.assert_rejected_without_overwrite( - "base/invalid_modifier", "invalid modifier" - ) + self.assert_rejected_without_overwrite("base/invalid_modifier", "无效 modifier") def test_property_examples_are_rejected_without_overwrite(self): self.write_class_member_input( @@ -858,7 +1112,7 @@ class DocGenCliTest(unittest.TestCase): ) self.assert_rejected_without_overwrite( - "base/property_examples", "unknown field(s): examples" + "base/property_examples", "存在未知字段:examples" ) def test_unit_rejects_implementation_data_without_overwrite(self): @@ -879,7 +1133,7 @@ class DocGenCliTest(unittest.TestCase): ) self.assert_rejected_without_overwrite( - "base/unit_implementation", "unknown field(s): implementation" + "base/unit_implementation", "存在未知字段:implementation" ) def test_empty_declarations_are_rejected_without_overwrite(self): @@ -892,7 +1146,7 @@ class DocGenCliTest(unittest.TestCase): ) self.assert_rejected_without_overwrite( - "base/missing_branch", "declarations must be a non-empty list" + "base/missing_branch", "declarations 必须是非空列表" ) def test_unknown_declaration_kind_is_rejected_without_overwrite(self): @@ -905,7 +1159,7 @@ class DocGenCliTest(unittest.TestCase): ) self.assert_rejected_without_overwrite( - "base/unknown_declaration", "unknown kind 'procedure'" + "base/unknown_declaration", "未知 kind:procedure" ) def test_top_level_function_requires_name(self): @@ -946,7 +1200,7 @@ class DocGenCliTest(unittest.TestCase): ) self.assert_rejected_without_overwrite( - "base/function_name_mismatch", "name and signature differ" + "base/function_name_mismatch", "name 与 signature 中的名称不一致" ) def test_duplicate_function_signature_is_rejected(self): @@ -974,7 +1228,7 @@ class DocGenCliTest(unittest.TestCase): ) self.assert_rejected_without_overwrite( - "base/duplicate_function", "duplicate function signature" + "base/duplicate_function", "function signature 重复" ) def test_duplicate_class_and_unit_names_are_rejected(self): @@ -1019,9 +1273,7 @@ class DocGenCliTest(unittest.TestCase): } ) - self.assert_rejected_without_overwrite( - relative, f"duplicate {kind} name" - ) + self.assert_rejected_without_overwrite(relative, f"{kind} 名称重复") def test_function_overloads_and_cross_kind_same_name_are_allowed(self): self.write_input( @@ -1034,9 +1286,7 @@ class DocGenCliTest(unittest.TestCase): "name": "Open", "signature": "Open(path)", "desc": "按路径打开。", - "params": [ - {"name": "path", "type": "string", "desc": "路径"} - ], + "params": [{"name": "path", "type": "string", "desc": "路径"}], "returns": "nil", }, { @@ -1044,9 +1294,7 @@ class DocGenCliTest(unittest.TestCase): "name": "Open", "signature": "Open(mode)", "desc": "按模式打开。", - "params": [ - {"name": "mode", "type": "integer", "desc": "模式"} - ], + "params": [{"name": "mode", "type": "integer", "desc": "模式"}], "returns": "nil", }, { @@ -1097,9 +1345,7 @@ class DocGenCliTest(unittest.TestCase): "base/missing_value", ) - self.assert_rejected_without_overwrite( - "base/missing_value", "missing 'value'" - ) + self.assert_rejected_without_overwrite("base/missing_value", "缺少 value") def test_class_missing_description_is_rejected_without_overwrite(self): self.write_input( @@ -1119,7 +1365,7 @@ class DocGenCliTest(unittest.TestCase): self.assert_rejected_without_overwrite( "base/missing_class_desc", - "declarations[0]: desc: must be non-empty", + "declarations[0]: desc:不能为空", ) def test_method_name_and_parameter_order_must_match_signature(self): @@ -1150,7 +1396,7 @@ class DocGenCliTest(unittest.TestCase): relative = f"base/mismatch_{name}" self.write_class_member_input(member, relative) - expected = "differ" if name == "name" else "signature order" + expected = "名称不一致" if name == "name" else "参数顺序一致" self.assert_rejected_without_overwrite(relative, expected) def test_instance_field_and_static_constant_render_distinct_headings(self): @@ -1164,21 +1410,21 @@ class DocGenCliTest(unittest.TestCase): "name": "Bindings", "desc": "绑定示例。", "members": [ - { - "kind": "field", - "name": "Name", - "visibility": "public", - "desc": "名称。", - "type": "string", - }, - { - "kind": "constant", - "name": "Maximum", - "visibility": "protected", - "desc": "最大值。", - "value": 10, - "static": True, - }, + { + "kind": "field", + "name": "Name", + "visibility": "public", + "desc": "名称。", + "type": "string", + }, + { + "kind": "constant", + "name": "Maximum", + "visibility": "protected", + "desc": "最大值。", + "value": 10, + "static": True, + }, ], } ], diff --git a/tools/tsl-codegen/tests/test_pipeline.py b/tools/tsl-codegen/tests/test_pipeline.py index 06d42aa7..243e61c4 100644 --- a/tools/tsl-codegen/tests/test_pipeline.py +++ b/tools/tsl-codegen/tests/test_pipeline.py @@ -6,7 +6,6 @@ import textwrap import unittest from pathlib import Path - try: import yaml except ImportError: @@ -50,9 +49,7 @@ class UnifiedPipelineTest(unittest.TestCase): tsf_paths = [] for name, source in sources: tsf_path = self.root / f"{name}.tsf" - tsf_path.write_text( - textwrap.dedent(source).lstrip(), encoding="utf-8" - ) + tsf_path.write_text(textwrap.dedent(source).lstrip(), encoding="utf-8") tsf_paths.append(tsf_path) page_path = "base/mixed" @@ -76,38 +73,26 @@ class UnifiedPipelineTest(unittest.TestCase): recording_paths[output_format] = recording text = recording.read_text(encoding="utf-8") recordings[output_format] = ( - json.loads(text) - if output_format == "json" - else yaml.safe_load(text) + json.loads(text) if output_format == "json" else yaml.safe_load(text) ) self.assertEqual(recordings["json"], recordings["yaml"]) markdown = ( - self.skill_dir - / "references" - / "codegen" - / "project" - / f"{page_path}.md" + self.skill_dir / "references" / "codegen" / "project" / f"{page_path}.md" ) generated_markdown = {} for output_format, recording in recording_paths.items(): - generate = self.run_command(GENERATE, recording, cwd=self.root) + generate = self.run_command(GENERATE, "--file", recording, cwd=self.root) self.assertEqual(0, generate.returncode, generate.stderr) self.assertTrue(markdown.is_file()) - generated_markdown[output_format] = markdown.read_text( - encoding="utf-8" - ) + generated_markdown[output_format] = markdown.read_text(encoding="utf-8") lint = self.run_command(LINT, "--file", markdown, "--strict") self.assertEqual(0, lint.returncode, lint.stdout + lint.stderr) - self.assertEqual( - generated_markdown["json"], generated_markdown["yaml"] - ) + self.assertEqual(generated_markdown["json"], generated_markdown["yaml"]) build = self.run_command(BUILD_INDEX, "--skill-dir", self.skill_dir) self.assertEqual(0, build.returncode, build.stderr) - check = self.run_command( - BUILD_INDEX, "--skill-dir", self.skill_dir, "--check" - ) + check = self.run_command(BUILD_INDEX, "--skill-dir", self.skill_dir, "--check") self.assertEqual(0, check.returncode, check.stderr) lookup_outputs = {} @@ -206,26 +191,18 @@ class UnifiedPipelineTest(unittest.TestCase): [item["kind"] for item in recording["declarations"]], ) widget_members = recording["declarations"][0]["members"] - raw_title = next( - item for item in widget_members if item["name"] == "RawTitle" - ) + raw_title = next(item for item in widget_members if item["name"] == "RawTitle") self.assertEqual("", raw_title["type"]) self.assertEqual([], raw_title["params"]) self.assertIn("## `Widget`\n\n声明:class", markdown) self.assertIn("## `OpenWidget(path)`\n\n声明:function", markdown) self.assertIn("## `DemoUnit`\n\n声明:unit", markdown) - self.assertIn( - "### `Create()`\n\n声明:class function", lookups["Widget"] - ) - self.assertIn( - "### `Rename(name)`\n\n声明:function", lookups["Widget"] - ) + self.assertIn("### `Create()`\n\n声明:class function", lookups["Widget"]) + self.assertIn("### `Rename(name)`\n\n声明:function", lookups["Widget"]) self.assertNotIn("## `OpenWidget()`", lookups["Widget"]) self.assertIn("声明:function", lookups["OpenWidget"]) self.assertNotIn("## `DemoUnit`", lookups["OpenWidget"]) - self.assertIn( - "#### `Save()`\n\n声明:function", lookups["DemoUnit.Document"] - ) + self.assertIn("#### `Save()`\n\n声明:function", lookups["DemoUnit.Document"]) self.assertIn( "#### `Create()`\n\n声明:class function", lookups["DemoUnit.Document"], @@ -273,23 +250,35 @@ class UnifiedPipelineTest(unittest.TestCase): "property 类型可选", ): self.assertIn(fragment, readme) - self.assertIn( - "## `OpenXmlAttribute`\n\n声明:class", readme - ) + self.assertIn("## `OpenXmlAttribute`\n\n声明:class", readme) self.assertIn( "### `CreateVirtual(position, row_index, _story)`\n\n" "声明:class function", readme, ) self.assertNotIn("property/field/variable 缺少类型", readme) + self.assertIn( + "generate.py --file tmp/my-api.yaml", + readme, + ) + self.assertIn( + "generate.py --file tmp/my-api.json", + readme, + ) + self.assertIn( + "generate.py --dir tmp/api-recordings", + readme, + ) + self.assertIn("只处理目录中的直属文件", readme) + self.assertNotIn("generate.py tmp/my-api.yaml", readme) + self.assertNotIn("generate.py tmp/my-api.json", readme) + self.assertNotIn("已废弃,请使用 --file", readme) self.assertIn("混合页面", skill) self.assertIn("page#anchor", skill) def test_standard_class_methods_reuse_top_level_function_structure(self): standard = STANDARD.read_text(encoding="utf-8") - class_section = standard.split("### class\n", 1)[1].split( - "### unit\n", 1 - )[0] + class_section = standard.split("### class\n", 1)[1].split("### unit\n", 1)[0] normalized_class = " ".join(class_section.split()) normalized_standard = " ".join(standard.split()) @@ -321,9 +310,7 @@ class UnifiedPipelineTest(unittest.TestCase): self.assertEqual(expected_root, set(yaml_data)) self.assertEqual(json_data, yaml_data) - declarations = { - item["kind"]: item for item in json_data["declarations"] - } + declarations = {item["kind"]: item for item in json_data["declarations"]} self.assertEqual({"function", "class", "unit"}, set(declarations)) self.assertEqual( ["class", "function", "unit"], @@ -361,11 +348,7 @@ class UnifiedPipelineTest(unittest.TestCase): next(m for m in class_members if m["name"] == "Count")["static"] ) self.assertTrue( - next( - m - for m in class_members - if m["name"] == "MaximumAttributes" - )["static"] + next(m for m in class_members if m["name"] == "MaximumAttributes")["static"] ) unit_members = declarations["unit"]["members"] @@ -373,9 +356,7 @@ class UnifiedPipelineTest(unittest.TestCase): {"constant", "variable", "function", "class"}, {member["kind"] for member in unit_members}, ) - document = next( - member for member in unit_members if member["kind"] == "class" - ) + document = next(member for member in unit_members if member["kind"] == "class") self.assertEqual( {("method", "instance"), ("method", "class"), ("property", None)}, {(m["kind"], m.get("binding")) for m in document["members"]}, @@ -393,7 +374,7 @@ class UnifiedPipelineTest(unittest.TestCase): ("json", EXAMPLE_JSON), ("yaml", EXAMPLE_YAML), ): - generate = self.run_command(GENERATE, recording, cwd=self.root) + generate = self.run_command(GENERATE, "--file", recording, cwd=self.root) self.assertEqual(0, generate.returncode, generate.stderr) markdown_by_format[output_format] = page.read_bytes() lint = self.run_command(LINT, "--file", page, "--strict")