diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 8781056..60f9de6 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -29,6 +29,12 @@ repos: language: python pass_filenames: false always_run: true + - id: regen-openapi-spec + name: regenerate OpenAPI spec (staged) + entry: python -B scripts/regenerate_openapi_if_needed.py --staged + language: python + pass_filenames: false + always_run: true - id: guard-openapi-sync name: guard generated OpenAPI sync (staged) entry: python -B scripts/check_openapi_sync.py --staged diff --git a/scripts/regenerate_openapi_if_needed.py b/scripts/regenerate_openapi_if_needed.py new file mode 100644 index 0000000..2087f8e --- /dev/null +++ b/scripts/regenerate_openapi_if_needed.py @@ -0,0 +1,84 @@ +from __future__ import annotations + +import argparse +import importlib.util +import sys +from pathlib import Path + + +def _repo_root() -> Path: + return Path(__file__).resolve().parents[1] + + +def _load_sync_guard(): + module_path = _repo_root() / "scripts" / "check_openapi_sync.py" + spec = importlib.util.spec_from_file_location("openapi_sync_guard_mod", module_path) + if spec is None or spec.loader is None: + raise RuntimeError("Failed to load check_openapi_sync.py") + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def regenerate_openapi_if_needed( + *, + openapi_path: str | Path | None = None, + contract_path: str | Path | None = None, + guard_module=None, + write_openapi_yaml=None, +) -> tuple[bool, str]: + guard = guard_module or _load_sync_guard() + root = _repo_root() + openapi_file = ( + Path(openapi_path) if openapi_path else root / "docs" / "openapi.yaml" + ) + contract_file = ( + Path(contract_path) + if contract_path + else root / "docs" / "release" / "api_contract.md" + ) + ok, _ = guard.validate_openapi_sync( + openapi_path=openapi_file, + contract_path=contract_file, + ) + if ok: + return False, "" + + writer = write_openapi_yaml + if writer is None: + _ensure_repo_on_path() + from services.openapi_generation import write_openapi_yaml as writer + + output = writer(openapi_file, contract_path=contract_file) + return True, f"[OpenClaw] Regenerated generated spec: {output}" + + +def _ensure_repo_on_path() -> None: + root = str(_repo_root()) + if root not in sys.path: + sys.path.insert(0, root) + + +def main(argv: list[str] | None = None) -> int: + parser = argparse.ArgumentParser( + description="Regenerate docs/openapi.yaml when generator inputs changed." + ) + parser.add_argument( + "--staged", + action="store_true", + help="Only regenerate when staged changes touch OpenAPI generator/spec sources.", + ) + args = parser.parse_args(argv) + + guard = _load_sync_guard() + if args.staged and not guard.should_validate_openapi(guard._get_staged_paths()): + return 0 + + changed, message = regenerate_openapi_if_needed() + if changed and message: + print(message) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/test_openapi_regen_hook.py b/tests/test_openapi_regen_hook.py new file mode 100644 index 0000000..6a8df25 --- /dev/null +++ b/tests/test_openapi_regen_hook.py @@ -0,0 +1,71 @@ +import importlib.util +import tempfile +import unittest +from pathlib import Path +from types import SimpleNamespace + + +def _load_module(): + root = Path(__file__).resolve().parents[1] + module_path = root / "scripts" / "regenerate_openapi_if_needed.py" + spec = importlib.util.spec_from_file_location("openapi_regen_hook_mod", module_path) + if spec is None or spec.loader is None: + raise RuntimeError("Failed to load regenerate_openapi_if_needed.py") + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +class TestOpenApiRegenHook(unittest.TestCase): + def setUp(self): + self.mod = _load_module() + + def test_regenerate_openapi_if_needed_noop_when_already_synced(self): + guard = SimpleNamespace( + validate_openapi_sync=lambda **_: (True, ""), + ) + with tempfile.TemporaryDirectory() as td: + openapi_path = Path(td) / "openapi.yaml" + contract_path = Path(td) / "api_contract.md" + openapi_path.write_text("expected\n", encoding="utf-8") + contract_path.write_text("dummy", encoding="utf-8") + changed, message = self.mod.regenerate_openapi_if_needed( + openapi_path=openapi_path, + contract_path=contract_path, + guard_module=guard, + write_openapi_yaml=lambda *_args, **_kwargs: self.fail( + "writer should not run when spec is already synced" + ), + ) + self.assertFalse(changed) + self.assertEqual(message, "") + self.assertEqual(openapi_path.read_text(encoding="utf-8"), "expected\n") + + def test_regenerate_openapi_if_needed_rewrites_drifted_spec(self): + guard = SimpleNamespace( + validate_openapi_sync=lambda **_: (False, "drift"), + ) + + def writer(out_path, *, contract_path): + Path(out_path).write_text("generated\n", encoding="utf-8") + self.assertTrue(Path(contract_path).exists()) + return Path(out_path) + + with tempfile.TemporaryDirectory() as td: + openapi_path = Path(td) / "openapi.yaml" + contract_path = Path(td) / "api_contract.md" + openapi_path.write_text("hand-edited\n", encoding="utf-8") + contract_path.write_text("dummy", encoding="utf-8") + changed, message = self.mod.regenerate_openapi_if_needed( + openapi_path=openapi_path, + contract_path=contract_path, + guard_module=guard, + write_openapi_yaml=writer, + ) + self.assertTrue(changed) + self.assertIn("Regenerated generated spec", message) + self.assertEqual(openapi_path.read_text(encoding="utf-8"), "generated\n") + + +if __name__ == "__main__": + unittest.main()