Merge pull request #1558 from meeb/tcely-shasum-updates

fix(shasum): version 1.1.4
This commit is contained in:
tcely
2026-08-20 08:32:25 -04:00
committed by GitHub
2 changed files with 68 additions and 60 deletions

View File

@@ -19,8 +19,8 @@ CHUNK_SIZE = (1024) * 32 # KiB
PROG_NAME = Path(__file__).stem PROG_NAME = Path(__file__).stem
# Versioning # Versioning
VERSION = (1, 1, 3) VERSION = (1, 1, 4)
VERSION_STR = "v" + ".".join(map(str, VERSION)) VERSION_STR = 'v' + '.'.join(map(str, VERSION))
def _std_base(*args, **kwargs): def _std_base(*args, **kwargs):
try: try:
@@ -44,9 +44,9 @@ def parse_args():
"""Configures and returns command line arguments.""" """Configures and returns command line arguments."""
parser = argparse.ArgumentParser( parser = argparse.ArgumentParser(
prog=PROG_NAME, prog=PROG_NAME,
description="Verify file checksums, enforcing strict line formatting and skipping missing files.", description='Verify file checksums, enforcing strict line formatting and skipping missing files.',
formatter_class=argparse.RawDescriptionHelpFormatter, formatter_class=argparse.RawDescriptionHelpFormatter,
epilog=f""" epilog=f'''
Examples: Examples:
Standard (sha256sum style): Standard (sha256sum style):
python3 {PROG_NAME}.py sums.txt python3 {PROG_NAME}.py sums.txt
@@ -61,14 +61,14 @@ Notes:
- Missing files are skipped (non-fatal). - Missing files are skipped (non-fatal).
- Uses all available CPU cores for hashing. - Uses all available CPU cores for hashing.
- Rejects UTF-16 manifests to ensure audit integrity. - Rejects UTF-16 manifests to ensure audit integrity.
""" '''
) )
parser.add_argument("-a", "--algorithm", default="sha256", parser.add_argument('-a', '--algorithm', default='sha256',
help="checksum algorithm to use (default: sha256)") help='checksum algorithm to use (default: sha256)')
parser.add_argument("file", nargs="?", default="-", parser.add_argument('file', nargs='?', default='-',
help="checksum file to read (default: '-' for stdin)") help='checksum file to read (default: "-" for stdin)')
parser.add_argument("-v", "--version", action="version", parser.add_argument('-v', '--version', action='version',
version=f"%(prog)s {VERSION_STR} (Python {platform.python_version()})") version=f'%(prog)s {VERSION_STR} (Python {platform.python_version()})')
return parser.parse_args() return parser.parse_args()
def get_algo_suggestion(word, available): def get_algo_suggestion(word, available):
@@ -133,9 +133,9 @@ def validate_algo(algo_name):
if algo_lower not in available: if algo_lower not in available:
suggestion = get_algo_suggestion(algo_lower, available) suggestion = get_algo_suggestion(algo_lower, available)
error_msg = f"{PROG_NAME}: Error: Unsupported algorithm '{algo_name}'" error_msg = f'{PROG_NAME}: Error: Unsupported algorithm "{algo_name}"'
if suggestion: if suggestion:
error_msg += f". Did you mean '{suggestion}?'" error_msg += f'. Did you mean "{suggestion}?"'
stderr(error_msg) stderr(error_msg)
sys.exit(1) sys.exit(1)
@@ -144,33 +144,33 @@ def validate_algo(algo_name):
def get_input_and_format(file_arg): def get_input_and_format(file_arg):
"""Verifies argument is a file, handles stdin, and determines format.""" """Verifies argument is a file, handles stdin, and determines format."""
def error_exit(message, /, label): def error_exit(message, /, label):
stderr(f"{PROG_NAME}: error: {label}: {message}") stderr(f'{PROG_NAME}: error: {label}: {message}')
sys.exit(1) sys.exit(1)
if "-" == file_arg: if '-' == file_arg:
label = "stdin" label = 'stdin'
raw_data = sys.stdin.buffer.read() raw_data = sys.stdin.buffer.read()
if raw_data.startswith(b'\xef\xbb\xbf'): if raw_data.startswith(b'\xef\xbb\xbf'):
error_exit("UTF-8 BOM detected; please provide a clean stream", label) error_exit('UTF-8 BOM detected; please provide a clean stream', label)
if raw_data.startswith((b'\xff\xfe', b'\xfe\xff')): if raw_data.startswith((b'\xff\xfe', b'\xfe\xff')):
error_exit("UTF-16 BOM detected; only UTF-8 (without BOM) is supported", label) error_exit('UTF-16 BOM detected; only UTF-8 (without BOM) is supported', label)
try: try:
lines = raw_data.decode('utf-8').splitlines() lines = raw_data.decode('utf-8').splitlines()
except UnicodeDecodeError: except UnicodeDecodeError:
error_exit("invalid UTF-8 encoding", label) error_exit('invalid UTF-8 encoding', label)
else: else:
label = file_arg label = file_arg
p = Path(file_arg) p = Path(file_arg)
if not p.exists(): if not p.exists():
error_exit("No such file or directory") error_exit('No such file or directory')
if not p.is_file(): if not p.is_file():
error_exit("Is not a regular file") error_exit('Is not a regular file')
with p.open('rb') as f: with p.open('rb') as f:
raw_data = f.read(2) raw_data = f.read(2)
if raw_data.startswith((b'\xff\xfe', b'\xfe\xff')): if raw_data.startswith((b'\xff\xfe', b'\xfe\xff')):
error_exit("UTF-16 manifest detected; only UTF-8 (with/without BOM) is supported", label) error_exit('UTF-16 manifest detected; only UTF-8 (with/without BOM) is supported', label)
lines = p.read_text(encoding='utf-8-sig').splitlines() lines = p.read_text(encoding='utf-8-sig').splitlines()
@@ -213,7 +213,7 @@ def path_resolve(path_str, *, strict=False):
resolved.resolve(strict=True) resolved.resolve(strict=True)
except OSError as e: except OSError as e:
if e.errno == errno.ELOOP: if e.errno == errno.ELOOP:
raise RuntimeError(f"Symlink loop detected: {path_str}") from e raise RuntimeError(f'Symlink loop detected: {path_str}') from e
raise raise
return resolved return resolved
@@ -227,7 +227,6 @@ def verify_checksums(line_data, is_tag, label, algorithm):
file_buffer_pool = queue.Queue() file_buffer_pool = queue.Queue()
files_verified = 0 files_verified = 0
format_errors = 0 format_errors = 0
hexdigest_args = []
is_windows = "Windows" == platform.system() is_windows = "Windows" == platform.system()
max_pending_tasks = 50_000 max_pending_tasks = 50_000
semaphore = threading.Semaphore(max_pending_tasks) semaphore = threading.Semaphore(max_pending_tasks)
@@ -250,7 +249,7 @@ def verify_checksums(line_data, is_tag, label, algorithm):
with target_path.open('rb') as f: with target_path.open('rb') as f:
actual_read = f.readinto(buf) actual_read = f.readinto(buf)
return buf, actual_read return buf, actual_read
except (queue.Empty, IOError, OSError): except (queue.Empty, OSError):
return None, 0 return None, 0
def return_buffer(buffer, /, pool = file_buffer_pool): def return_buffer(buffer, /, pool = file_buffer_pool):
@@ -270,11 +269,11 @@ def verify_checksums(line_data, is_tag, label, algorithm):
match = pattern.match(line) match = pattern.match(line)
if not match: if not match:
msg = f"{label}:{line_no}: WARNING: improperly formatted line" msg = f'{label}:{line_no}: WARNING: improperly formatted line'
if is_tag: if is_tag:
found_match = algo_extractor.match(line) found_match = algo_extractor.match(line)
if found_match: if found_match:
msg += f" (found {found_match.group(1)})" msg += f' (found {found_match.group(1)})'
stderr(msg) stderr(msg)
format_errors += 1 format_errors += 1
exit_code = 1 exit_code = 1
@@ -285,14 +284,14 @@ def verify_checksums(line_data, is_tag, label, algorithm):
else: else:
expected_hash, mode_char, filename_str = match.groups() expected_hash, mode_char, filename_str = match.groups()
if is_windows and ' ' == mode_char: if is_windows and ' ' == mode_char:
msg = f"{label}:{line_no}: {filename_str}: WARNING: text conversion is not supported; hashing as binary" msg = f'{label}:{line_no}: {filename_str}: WARNING: text conversion is not supported; hashing as binary'
stderr(msg) stderr(msg)
if '#' in filename_str: if '#' in filename_str:
msg = ( msg = (
f"{label}:{line_no}: WARNING: filename contained a " f'{label}:{line_no}: WARNING: filename contained a '
"'#' character; inline comments are not supported and " '"#" character; inline comments are not supported and '
"this will be treated as part of the literal filename." 'this will be treated as part of the literal filename.'
) )
stderr(msg) stderr(msg)
@@ -301,15 +300,15 @@ def verify_checksums(line_data, is_tag, label, algorithm):
# Security: Prevent Path Traversal # Security: Prevent Path Traversal
# Skip files outside the directory to prevent traversal attacks # Skip files outside the directory to prevent traversal attacks
# Resolve to absolute path and check if it's within CWD # Resolve to absolute path and check if it's within CWD
msg = f"{label}:{line_no}: WARNING: " msg = f'{label}:{line_no}: WARNING: '
try: try:
abs_target = path_resolve(target_path, strict=False) abs_target = path_resolve(target_path, strict=False)
if abs_cwd not in abs_target.parents and abs_target != abs_cwd: if abs_cwd not in abs_target.parents and abs_target != abs_cwd:
msg += "skipping path that is outside the current directory" msg += 'skipping path that is outside the current directory'
stderr(msg) stderr(msg)
continue continue
except (OSError, RuntimeError) as e: except (OSError, RuntimeError) as e:
msg += f"skipping path that could not be resolved: {e}" msg += f'skipping path that could not be resolved: {e}'
stderr(msg) stderr(msg)
continue continue
@@ -336,40 +335,47 @@ def verify_checksums(line_data, is_tag, label, algorithm):
current.st_ino != original_stat.st_ino or current.st_ino != original_stat.st_ino or
current.st_ctime != original_stat.st_ctime current.st_ctime != original_stat.st_ctime
) )
return "File modified during processing" if changed else None return 'File modified during processing' if changed else None
except (OSError, RuntimeError): except (OSError, RuntimeError):
return "Metadata access failed" return 'Metadata access failed'
def update_chunks(hasher, byte_array, total_length, chunk_size):
with memoryview(byte_array) as view:
for begin in range(0, total_length, chunk_size):
end = min(total_length, chunk_size + begin)
hasher.update(view[begin:end])
try: try:
# Pre-hash integrity check # Pre-hash integrity check
if err := check_file_integrity(target_path, stat): if err := check_file_integrity(target_path, stat):
if 'File modified' in err: if 'File modified' in err:
err = "File modified since scan" err = 'File modified since scan'
return target_path, False, err return target_path, False, err
hasher = hashlib.new(algorithm) hasher = hashlib.new(algorithm)
if buffer is not None: if buffer:
hasher.update(buffer[:actual_len]) update_chunks(hasher, buffer, actual_len, CHUNK_SIZE)
return_buffer(buffer) return_buffer(buffer)
buffer = None buffer = None
else: else:
if stat is None:
stat = target_path.stat()
data = bytearray(file_buffer_size)
with target_path.open('rb') as fb: with target_path.open('rb') as fb:
if sys.version_info >= (3, 11): while actual_read := fb.readinto(data):
hasher = hashlib.file_digest(fb, algorithm) update_chunks(hasher, data, actual_read, CHUNK_SIZE)
else: data = None
for chunk in iter(lambda: fb.read(CHUNK_SIZE), b""):
hasher.update(chunk)
# Post-hash integrity check # Post-hash integrity check
if err := check_file_integrity(target_path, stat): if err := check_file_integrity(target_path, stat):
return target_path, False, err return target_path, False, err
hexdigest_args.clear() hexdigest_args = []
if algorithm.startswith('shake'): if algorithm.startswith('shake'):
hexdigest_args.append(len(expected_hash) // 2) hexdigest_args.append(len(expected_hash) // 2)
is_ok = hasher.hexdigest(*hexdigest_args) == expected_hash.lower() is_ok = hasher.hexdigest(*hexdigest_args) == expected_hash.lower()
return target_path, is_ok, None return target_path, is_ok, None
except (IOError, OSError) as e: except OSError as e:
return target_path, False, str(e) return target_path, False, str(e)
finally: finally:
if buffer is not None: if buffer is not None:
@@ -381,21 +387,22 @@ def verify_checksums(line_data, is_tag, label, algorithm):
semaphore.release() semaphore.release()
try: try:
_, ok, err = future.result() _, ok, err = future.result()
msg = f"{path}: " msg = f'{path}: '
if ok: if ok:
msg += "OK" msg += 'OK'
files_verified += 1 files_verified += 1
else: else:
msg += "FAILED" msg += 'FAILED'
if err: if err:
msg += f" (Error: {err})" msg += f' (Error: {err})'
checksum_failures += 1 checksum_failures += 1
exit_code = 1 exit_code = 1
stdout(msg) stdout(msg)
# ruff: ignore[BLE001]
except Exception as e: except Exception as e:
checksum_failures += 1 checksum_failures += 1
exit_code = 1 exit_code = 1
stdout(f"{path}: FAILED (Unexpected Error: {e!r})") stdout(f'{path}: FAILED (Unexpected Error: {e!r})')
tasks.sort(key=lambda x: x[0], reverse=True) tasks.sort(key=lambda x: x[0], reverse=True)
with ThreadPoolExecutor() as executor: with ThreadPoolExecutor() as executor:
@@ -404,11 +411,12 @@ def verify_checksums(line_data, is_tag, label, algorithm):
_, target_path, stat, expected_hash, buffer, blen = tasks.pop() _, target_path, stat, expected_hash, buffer, blen = tasks.pop()
if buffer is None: if buffer is None:
blen = 0
assigned_buffer, actual_read = fill_buffer(target_path, stat=stat) assigned_buffer, actual_read = fill_buffer(target_path, stat=stat)
if assigned_buffer and actual_read: if assigned_buffer and 0 <= actual_read:
buffer = assigned_buffer buffer = assigned_buffer
blen = actual_read blen = actual_read
elif assigned_buffer: elif assigned_buffer is not None:
return_buffer(assigned_buffer) return_buffer(assigned_buffer)
future = executor.submit( future = executor.submit(
@@ -425,19 +433,19 @@ def verify_checksums(line_data, is_tag, label, algorithm):
def warning(msg, singular, plural, /, count): def warning(msg, singular, plural, /, count):
if 0 < count: if 0 < count:
alt = plural if 1 < count else singular alt = singular if 1 == count else plural
prefix = f"{PROG_NAME}: WARNING: {count} " prefix = f'{PROG_NAME}: WARNING: {count} '
stderr(prefix + msg.format(alt)) stderr(prefix + msg.format(alt))
warning("line{0} improperly formatted", " is", "s are", count=format_errors) warning('line{0} improperly formatted', ' is', 's are', count=format_errors)
warning("computed checksum{0} did NOT match", "", "s", count=checksum_failures) warning('computed checksum{0} did NOT match', '', 's', count=checksum_failures)
if 0 == files_verified: if 0 == files_verified:
exit_code = 1 exit_code = 1
msg = f"{PROG_NAME}: WARNING: {label}: no file was verified" msg = f'{PROG_NAME}: WARNING: {label}: no file was verified'
stderr(msg) stderr(msg)
sys.exit(exit_code) sys.exit(exit_code)
if __name__ == "__main__": if '__main__' == __name__:
args = parse_args() args = parse_args()
algo = validate_algo(args.algorithm) algo = validate_algo(args.algorithm)
content, tag_mode, file_label = get_input_and_format(args.file) content, tag_mode, file_label = get_input_and_format(args.file)

View File

@@ -412,7 +412,7 @@ class TestShasum(unittest.TestCase):
exit_code, out, err = self.run_verify(content, False, "dirty.txt", "sha256") exit_code, out, err = self.run_verify(content, False, "dirty.txt", "sha256")
# 1. Should warn about the '#' in the filename # 1. Should warn about the '#' in the filename
self.assertIn("contained a '#' character", err) self.assertIn('contained a "#" character', err)
# 2. Should eventually report no files verified (since 'target.txt #...' doesn't exist) # 2. Should eventually report no files verified (since 'target.txt #...' doesn't exist)
self.assertIn("no file was verified", err) self.assertIn("no file was verified", err)
self.assertEqual(exit_code, 1) self.assertEqual(exit_code, 1)