Skip to content
Open
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
7 changes: 4 additions & 3 deletions scripts/check_headers.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,8 +52,8 @@

def get_file_header(filename):
"""Read the file header from the file."""
with open(filename, 'rt') as f:
# Only read the first 2048 bytes to avoid loading too much and the
with open(filename, 'rt', encoding='utf-8') as f:
# Only read the first 2048 characters to avoid loading too much and the
# license information should be in the first part anyway.
return f.read(2048)

Expand All @@ -74,7 +74,8 @@ def get_file_header(filename):
for filename in sorted(files_to_check):
contents = get_file_header(filename)
m = identification_re.match(contents)
if not bool(m) or m.group('filename') not in (filename, os.path.basename(filename)):
if not bool(m) or m.group('filename') not in (
filename, filename.replace(os.sep, '/'), os.path.basename(filename)):
print('%s: Incorrect file identification' % filename)
fail = True
if not bool(license_re.search(contents)):
Expand Down
108 changes: 108 additions & 0 deletions tests/test_check_headers.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,108 @@
# test_check_headers.py - tests for the source-header checker
#
# Copyright (C) 2026 Ryan Duguid
#
# This library is free software; you can redistribute it and/or
# modify it under the terms of the GNU Lesser General Public
# License as published by the Free Software Foundation; either
# version 2.1 of the License, or (at your option) any later version.
#
# This library is distributed in the hope that it will be useful,
# but WITHOUT ANY WARRANTY; without even the implied warranty of
# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the GNU
# Lesser General Public License for more details.
#
# You should have received a copy of the GNU Lesser General Public
# License along with this library; if not, see <https://www.gnu.org/licenses/>.

"""Tests for the standalone source-header checker."""

from __future__ import annotations

import os
import subprocess
import sys
import tempfile
import unittest
from pathlib import Path


class TestHeaders(unittest.TestCase):
"""Exercise header validation through the script's command-line interface."""

def setUp(self) -> None:
"""Read the standard licence header without importing the script."""
self.checker = Path(__file__).resolve().parents[1] / 'scripts' / 'check_headers.py'
self.header = self.checker.read_text(encoding='utf-8').split('"""', 1)[0]
self.assertIn('# check_headers.py - ', self.header)

def _check(self, filename: str, contents: str) -> subprocess.CompletedProcess[str]:
"""Run the checker against one fixture in an isolated directory."""
# The checker uses only the standard library, so skip site startup hooks.
command = [sys.executable, '-S']
if os.name == 'nt':
# Disable UTF-8 mode so it cannot hide the locale-decoding failure.
command.extend(['-X', 'utf8=0'])
command.append(str(self.checker))
with tempfile.TemporaryDirectory() as directory:
fixture = Path(directory) / filename
fixture.parent.mkdir(parents=True, exist_ok=True)
fixture.write_text(contents, encoding='utf-8')
return subprocess.run(command, cwd=directory, capture_output=True, text=True, encoding='utf-8', check=False)

def test_utf8_header(self) -> None:
"""Read a UTF-8 quote whose final byte is undefined in cp1252."""
header = self.header.replace('# check_headers.py - ', '# check_headers.py - \u201d ', 1)
result = self._check('check_headers.py', header)
self.assertEqual((result.returncode, result.stdout, result.stderr), (0, '', ''))

def test_relative_path(self) -> None:
"""Accept the portable relative filename used by updater headings."""
header = self.header.replace('# check_headers.py - ', '# update/check_headers.py - ', 1)
result = self._check('update/check_headers.py', header)
self.assertEqual((result.returncode, result.stdout, result.stderr), (0, '', ''))

def test_native_path(self) -> None:
"""Preserve acceptance of an exact native relative filename."""
native = str(Path('update') / 'check_headers.py')
header = self.header.replace('# check_headers.py - ', '# %s - ' % native, 1)
result = self._check('update/check_headers.py', header)
self.assertEqual((result.returncode, result.stdout, result.stderr), (0, '', ''))

def test_basename(self) -> None:
"""Preserve acceptance of a basename in a nested file."""
result = self._check('update/check_headers.py', self.header)
self.assertEqual((result.returncode, result.stdout, result.stderr), (0, '', ''))

def test_wrong_identification(self) -> None:
"""Reject a wrong filename and a wrong directory with the same basename."""
for name in ('other.py', 'other/check_headers.py'):
with self.subTest(name=name):
header = self.header.replace('# check_headers.py - ', '# %s - ' % name, 1)
result = self._check('update/check_headers.py', header)
expected = '%s: Incorrect file identification\n' % str(Path('update') / 'check_headers.py')
self.assertEqual((result.returncode, result.stdout, result.stderr), (1, expected, ''))

def test_missing_identification(self) -> None:
"""Reject a missing identification without raising an exception."""
header = self.header.replace('# check_headers.py - ', '# ', 1)
result = self._check('check_headers.py', header)
self.assertEqual(
(result.returncode, result.stdout, result.stderr),
(1, 'check_headers.py: Incorrect file identification\n', ''))

def test_wrong_licence(self) -> None:
"""Reject missing and materially altered licence text."""
for header in (self.header.split('# This library', 1)[0], self.header.replace('version 2.1', 'version 3.0')):
with self.subTest(header=header):
result = self._check('check_headers.py', header)
self.assertEqual(
(result.returncode, result.stdout, result.stderr),
(1, 'check_headers.py: Incorrect license text\n', ''))

def test_both_errors(self) -> None:
"""Report identification and licence errors in their existing order."""
result = self._check('other.py', self.header.replace('version 2.1', 'version 3.0'))
self.assertEqual(
(result.returncode, result.stdout, result.stderr),
(1, 'other.py: Incorrect file identification\nother.py: Incorrect license text\n', ''))
Loading