diff --git a/src/tzst/cli.py b/src/tzst/cli.py index 537c75f..5b6ca94 100644 --- a/src/tzst/cli.py +++ b/src/tzst/cli.py @@ -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() diff --git a/test_fixes.py b/test_fixes.py new file mode 100644 index 0000000..96af581 --- /dev/null +++ b/test_fixes.py @@ -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")