Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 0 additions & 6 deletions .pre-commit-hooks.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand All @@ -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
Expand All @@ -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]
Expand All @@ -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$|
Expand All @@ -68,4 +63,3 @@
files: |
(?x)
(^|/)pyproject[.]toml$
args: [--fix]
9 changes: 6 additions & 3 deletions src/rapids_pre_commit_hooks/copyright.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -12,7 +12,6 @@
from typing import TYPE_CHECKING

import Levenshtein
import git

from .lint import Lines, Linter, LintMain

Expand All @@ -21,6 +20,8 @@
from collections.abc import Callable, Generator, Iterable
from typing import Optional

import git

from .lint import LintWarning, Span


Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
8 changes: 6 additions & 2 deletions src/rapids_pre_commit_hooks/lint.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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")

Expand Down
31 changes: 21 additions & 10 deletions tests/rapids_pre_commit_hooks/test_lint.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
):
Expand All @@ -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,
):
Expand All @@ -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$"),
Expand Down Expand Up @@ -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$"),
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
],
),
Expand Down Expand Up @@ -1157,6 +1161,7 @@ def test_long_file_delete(self, long_file, long_delete_fix_check):
"sys.argv",
[
"check-test",
"--no-fix",
file.name,
],
),
Expand Down Expand Up @@ -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,
],
),
Expand Down Expand Up @@ -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,
],
),
Expand Down
Loading