diff --git a/.pre-commit-hooks.yaml b/.pre-commit-hooks.yaml index 4854796..d001081 100644 --- a/.pre-commit-hooks.yaml +++ b/.pre-commit-hooks.yaml @@ -6,7 +6,6 @@ entry: verify-alpha-spec language: python files: (^|/)dependencies[.]yaml$ - args: [--fix] additional_dependencies: - --extra-index-url=https://pypi.anaconda.org/rapidsai-wheels-nightly/simple - .[alpha-spec] @@ -18,7 +17,6 @@ files: (?x) ^[.]github/CODEOWNERS$ - args: [--fix] - id: verify-copyright name: verify-copyright description: make sure copyright headers are up to date @@ -32,14 +30,12 @@ (^|/)meta[.]yaml$| (^|/)pyproject[.]toml$| (^|/)setup[.]cfg$ - args: [--fix] - id: verify-dependencies name: verify-dependencies description: make sure dependencies.yaml follows RAPIDS conventions entry: verify-dependencies language: python files: (^|/)dependencies[.]yaml$ - args: [--fix] additional_dependencies: - --extra-index-url=https://pypi.anaconda.org/rapidsai-wheels-nightly/simple - .[dependencies] @@ -48,7 +44,6 @@ description: make sure RAPIDS version is not hard-coded in files entry: verify-hardcoded-version language: python - args: [--fix] exclude: | (?x) (^|/)devcontainer[.]json$| @@ -68,4 +63,3 @@ files: | (?x) (^|/)pyproject[.]toml$ - args: [--fix] diff --git a/src/rapids_pre_commit_hooks/copyright.py b/src/rapids_pre_commit_hooks/copyright.py index 5acb25d..022b44a 100644 --- a/src/rapids_pre_commit_hooks/copyright.py +++ b/src/rapids_pre_commit_hooks/copyright.py @@ -1,4 +1,4 @@ -# SPDX-FileCopyrightText: Copyright (c) 2024-2026, NVIDIA CORPORATION. +# SPDX-FileCopyrightText: Copyright (c) 2024-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 import dataclasses @@ -12,7 +12,6 @@ from typing import TYPE_CHECKING import Levenshtein -import git from .lint import Lines, Linter, LintMain @@ -21,6 +20,8 @@ from collections.abc import Callable, Generator, Iterable from typing import Optional + import git + from .lint import LintWarning, Span @@ -838,7 +839,7 @@ def get_target_branch(repo: "git.Repo", args: "argparse.Namespace") -> str: def get_target_branch_upstream_commit( repo: "git.Repo", args: "argparse.Namespace" -) -> git.Commit | None: +) -> "Optional[git.Commit]": # If no target branch can be determined, use HEAD if it exists target_branch_name = get_target_branch(repo, args) if target_branch_name is None: @@ -974,6 +975,8 @@ def find_blob( def check_copyright( args: "argparse.Namespace", ) -> "Callable[[Linter, argparse.Namespace], None]": + import git + try: repo = git.Repo() except git.InvalidGitRepositoryError: diff --git a/src/rapids_pre_commit_hooks/lint.py b/src/rapids_pre_commit_hooks/lint.py index 602b1b6..93aeaea 100644 --- a/src/rapids_pre_commit_hooks/lint.py +++ b/src/rapids_pre_commit_hooks/lint.py @@ -1,4 +1,4 @@ -# SPDX-FileCopyrightText: Copyright (c) 2024-2026, NVIDIA CORPORATION. +# SPDX-FileCopyrightText: Copyright (c) 2024-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 import argparse @@ -440,7 +440,11 @@ def __init__(self, warning_name: str) -> None: self.warning_name: str = warning_name self.argparser: argparse.ArgumentParser = argparse.ArgumentParser() self.argparser.add_argument( - "--fix", action="store_true", help="automatically fix warnings" + "--fix", + action=argparse.BooleanOptionalAction, + default=True, + help="whether or not to automatically fix warnings " + "(yes by default)", ) self.argparser.add_argument("files", nargs="+", metavar="file") diff --git a/tests/rapids_pre_commit_hooks/test_lint.py b/tests/rapids_pre_commit_hooks/test_lint.py index 9ec1ff8..0ac669c 100644 --- a/tests/rapids_pre_commit_hooks/test_lint.py +++ b/tests/rapids_pre_commit_hooks/test_lint.py @@ -751,7 +751,7 @@ def test_no_warnings_no_fix(self, hello_world_file): with ( patch( "sys.argv", - ["check-test", "--check-test", file.name], + ["check-test", "--check-test", "--no-fix", file.name], ), self.mock_console() as console, ): @@ -765,12 +765,13 @@ def test_no_warnings_no_fix(self, hello_world_file): call(highlight=False), ] - def test_no_warnings_fix(self, hello_world_file): + @pytest.mark.parametrize("fix_args", [[], ["--fix"]]) + def test_no_warnings_fix(self, fix_args, hello_world_file): file, _, _ = hello_world_file with ( patch( "sys.argv", - ["check-test", "--check-test", "--fix", file.name], + ["check-test", "--check-test", *fix_args, file.name], ), self.mock_console() as console, ): @@ -789,7 +790,7 @@ def test_warnings_no_fix(self, hello_world_file, hello_check): with ( patch( "sys.argv", - ["check-test", "--check-test", file.name], + ["check-test", "--check-test", "--no-fix", file.name], ), self.mock_console() as console, pytest.raises(SystemExit, match=r"^1$"), @@ -838,12 +839,13 @@ def test_warnings_no_fix(self, hello_world_file, hello_check): call().print(), ] - def test_warnings_fix(self, hello_world_file, hello_check): + @pytest.mark.parametrize("fix_args", [[], ["--fix"]]) + def test_warnings_fix(self, fix_args, hello_world_file, hello_check): file, content, spans = hello_world_file with ( patch( "sys.argv", - ["check-test", "--check-test", "--fix", file.name], + ["check-test", "--check-test", *fix_args, file.name], ), self.mock_console() as console, pytest.raises(SystemExit, match=r"^1$"), @@ -899,6 +901,7 @@ def test_warnings_note(self, hello_world_file, hello_check): "sys.argv", [ "check-test", + "--no-fix", "--check-test", "--check-test-note", file.name, @@ -1102,6 +1105,7 @@ def test_long_file(self, long_file, long_file_check, long_fix_check): "sys.argv", [ "check-test", + "--no-fix", file.name, ], ), @@ -1157,6 +1161,7 @@ def test_long_file_delete(self, long_file, long_delete_fix_check): "sys.argv", [ "check-test", + "--no-fix", file.name, ], ), @@ -1198,14 +1203,17 @@ def test_long_file_delete(self, long_file, long_delete_fix_check): call().print(), ] - def test_long_file_fix(self, long_file, long_file_check, long_fix_check): + @pytest.mark.parametrize("fix_args", [[], ["--fix"]]) + def test_long_file_fix( + self, fix_args, long_file, long_file_check, long_fix_check + ): file, content, spans = long_file with ( patch( "sys.argv", [ "check-test", - "--fix", + *fix_args, file.name, ], ), @@ -1255,14 +1263,17 @@ def test_long_file_fix(self, long_file, long_file_check, long_fix_check): call().print(), ] - def test_long_file_delete_fix(self, long_file, long_delete_fix_check): + @pytest.mark.parametrize("fix_args", [[], ["--fix"]]) + def test_long_file_delete_fix( + self, fix_args, long_file, long_delete_fix_check + ): file, content, spans = long_file with ( patch( "sys.argv", [ "check-test", - "--fix", + *fix_args, file.name, ], ),