Improve CLI error handling and add tests
Enhanced error handling in `cli.py` for invalid compression levels, filters, and file path issues. Added `test_fixes.py` to verify CLI behavior, including tests for help output, invalid arguments, and edge cases.
This commit is contained in:
1 parent
6693713143
commit
dc12d122c4
2 files changed
+141
-13
No files matched your search
+73
-13
@@ -57,7 +57,9 @@ def validate_compression_level(value: str) -> int:
|
||||
value: String value from command line
|
||||
|
||||
Returns:
|
||||
int: Valid compression level (1-22) Raises:
|
||||
int: Valid compression level (1-22)
|
||||
|
||||
Raises:
|
||||
argparse.ArgumentTypeError: If value is not a valid compression level
|
||||
"""
|
||||
try:
|
||||
@@ -96,18 +98,30 @@ def cmd_add(args) -> int:
|
||||
Note:
|
||||
This function uses atomic file operations by default, creating the
|
||||
archive in a temporary file first, then atomically moving it to the
|
||||
final location to prevent incomplete archives.
|
||||
|
||||
See Also:
|
||||
final location to prevent incomplete archives. See Also:
|
||||
:func:`tzst.create_archive`: The underlying function for archive creation
|
||||
:meth:`TzstArchive.add`: The core method for adding files to archives
|
||||
"""
|
||||
try:
|
||||
archive_path = Path(args.archive)
|
||||
files: list[Path] = [Path(f) for f in args.files]
|
||||
files: list[Path] = []
|
||||
|
||||
# Process files with proper path handling for special characters
|
||||
for file_arg in args.files:
|
||||
file_path = Path(file_arg).resolve()
|
||||
files.append(file_path)
|
||||
|
||||
# Check if files exist with better error reporting
|
||||
missing_files = []
|
||||
for f in files:
|
||||
try:
|
||||
if not f.exists():
|
||||
missing_files.append(f)
|
||||
except OSError as e:
|
||||
# Handle path issues (invalid characters, permissions, etc.)
|
||||
print(f"Error: Cannot access file '{f}' - {e}", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
# Check if files exist
|
||||
missing_files = [f for f in files if not f.exists()]
|
||||
if missing_files:
|
||||
print(
|
||||
f"Error: Files not found - {', '.join(map(str, missing_files))}",
|
||||
@@ -511,12 +525,11 @@ Security Note:
|
||||
Documentation:
|
||||
https://github.com/xixu-me/tzst#readme
|
||||
"""
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
prog="tzst",
|
||||
epilog=epilog,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
add_help=False,
|
||||
add_help=True,
|
||||
)
|
||||
|
||||
# Add global arguments
|
||||
@@ -649,8 +662,8 @@ def main(argv: list[str] | None = None) -> int:
|
||||
Returns:
|
||||
int: Exit code for the program
|
||||
- 0: Success
|
||||
- 1: No command specified (help displayed)
|
||||
- 2: Argument parsing error (invalid arguments)
|
||||
- 1: Invalid compression level, filter, or command error
|
||||
- 2: Argument parsing error (help, unknown options)
|
||||
- Other codes: Specific to individual command handlers
|
||||
|
||||
Note:
|
||||
@@ -669,8 +682,55 @@ def main(argv: list[str] | None = None) -> int:
|
||||
try:
|
||||
args = parser.parse_args(argv)
|
||||
except SystemExit as e:
|
||||
# Handle argparse errors (invalid arguments, help, etc.)
|
||||
return e.code if e.code is not None else 1
|
||||
# Handle different types of argparse errors
|
||||
if e.code == 0: # Help was requested
|
||||
return 0
|
||||
elif e.code == 2 and argv:
|
||||
# Check for validation errors we want to convert to exit code 1
|
||||
# Check for compression level errors
|
||||
if "-l" in argv or "--level" in argv:
|
||||
try:
|
||||
level_index = (
|
||||
argv.index("-l") if "-l" in argv else argv.index("--level")
|
||||
)
|
||||
if level_index + 1 < len(argv):
|
||||
level_value = argv[level_index + 1]
|
||||
try:
|
||||
level = int(level_value)
|
||||
if not 1 <= level <= 22:
|
||||
print(
|
||||
f"Invalid compression level: {level}. "
|
||||
f"Must be between 1 and 22.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return 1
|
||||
except ValueError:
|
||||
print(
|
||||
f"Invalid compression level: '{level_value}'. "
|
||||
f"Must be an integer between 1 and 22.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return 1
|
||||
except (ValueError, IndexError):
|
||||
pass # Check for invalid filter errors
|
||||
if "--filter" in argv:
|
||||
try:
|
||||
filter_index = argv.index("--filter")
|
||||
if filter_index + 1 < len(argv):
|
||||
filter_value = argv[filter_index + 1]
|
||||
valid_filters = ["data", "tar", "fully_trusted"]
|
||||
if filter_value not in valid_filters:
|
||||
print(
|
||||
f"Invalid filter specified: {filter_value}. "
|
||||
f"Must be one of: {', '.join(valid_filters)}",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return 1
|
||||
except (ValueError, IndexError):
|
||||
pass
|
||||
|
||||
# Return the original exit code for other cases
|
||||
return int(e.code) if e.code is not None else 1
|
||||
|
||||
if not hasattr(args, "func"):
|
||||
parser.print_help()
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Test script to verify the main CLI fixes."""
|
||||
|
||||
from src.tzst.cli import main
|
||||
|
||||
|
||||
def test_help():
|
||||
"""Test that --help works and returns exit code 0."""
|
||||
result = main(["--help"])
|
||||
print(f"Help test: {'PASS' if result == 0 else 'FAIL'} (exit code: {result})")
|
||||
return result == 0
|
||||
|
||||
|
||||
def test_invalid_compression_level():
|
||||
"""Test that invalid compression level returns exit code 1."""
|
||||
result = main(["a", "test.tzst", "-l", "25", "file.txt"])
|
||||
print(
|
||||
f"Invalid compression level test: {'PASS' if result == 1 else 'FAIL'} (exit code: {result})"
|
||||
)
|
||||
return result == 1
|
||||
|
||||
|
||||
def test_invalid_filter():
|
||||
"""Test that invalid filter returns exit code 1."""
|
||||
result = main(["x", "test.tzst", "--filter", "invalid"])
|
||||
print(
|
||||
f"Invalid filter test: {'PASS' if result == 1 else 'FAIL'} (exit code: {result})"
|
||||
)
|
||||
return result == 1
|
||||
|
||||
|
||||
def test_non_integer_compression_level():
|
||||
"""Test that non-integer compression level returns exit code 1."""
|
||||
result = main(["a", "test.tzst", "-l", "abc", "file.txt"])
|
||||
print(
|
||||
f"Non-integer compression level test: {'PASS' if result == 1 else 'FAIL'} (exit code: {result})"
|
||||
)
|
||||
return result == 1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
print("Testing main CLI fixes:")
|
||||
print("=" * 50)
|
||||
|
||||
tests = [
|
||||
test_help,
|
||||
test_invalid_compression_level,
|
||||
test_invalid_filter,
|
||||
test_non_integer_compression_level,
|
||||
]
|
||||
|
||||
passed = 0
|
||||
total = len(tests)
|
||||
|
||||
for test in tests:
|
||||
try:
|
||||
if test():
|
||||
passed += 1
|
||||
except Exception as e:
|
||||
print(f"Test {test.__name__} failed with exception: {e}")
|
||||
|
||||
print("=" * 50)
|
||||
print(f"Results: {passed}/{total} tests passed")
|
||||
|
||||
if passed == total:
|
||||
print("✅ All main fixes are working correctly!")
|
||||
else:
|
||||
print("❌ Some fixes still need attention")
|
||||
Reference in new issue
Block a user