Add validation for invalid commands in CLI

Introduced `_validate_command_in_argv` to check for invalid commands in CLI arguments. Updated compression level validation to include the new `-c` flag. Adjusted tests to reflect changes in error handling, ensuring argparse error codes are maintained for invalid inputs.
This commit is contained in:
xixu-me committed 2025-06-02 20:11:15 +08:00
1 parent 94b7a5b31a
commit 172ec6cb29
3 files changed
+81 -33

No files matched your search

+63 -12
View File
@@ -661,6 +661,7 @@ documentation:
parser_add.add_argument("archive", help="archive file path")
parser_add.add_argument("files", nargs="+", help="files/directories to add")
parser_add.add_argument(
"-c",
"-l",
"--level",
dest="compression_level",
@@ -774,11 +775,17 @@ def _validate_compression_level_in_argv(argv: list[str]) -> bool:
Returns:
bool: True if an error was found and handled, False otherwise
"""
if "-l" not in argv and "--level" not in argv:
if "-c" not in argv and "-l" not in argv and "--level" not in argv:
return False
try:
level_index = argv.index("-l") if "-l" in argv else argv.index("--level")
level_index = -1
if "-c" in argv:
level_index = argv.index("-c")
elif "-l" in argv:
level_index = argv.index("-l")
else:
level_index = argv.index("--level")
if level_index + 1 < len(argv):
level_value = argv[level_index + 1]
try:
@@ -831,6 +838,48 @@ def _validate_filter_in_argv(argv: list[str]) -> bool:
return False
def _validate_command_in_argv(argv: list[str]) -> bool:
"""Check for invalid command errors in argv.
Args:
argv: Command line arguments
Returns:
bool: True if an invalid command was found and handled, False otherwise
"""
if not argv:
return False
# Valid commands and their aliases
valid_commands = {
"a",
"add",
"create",
"x",
"extract",
"e",
"extract-flat",
"l",
"list",
"t",
"test",
}
# Find the first argument that's not a flag (doesn't start with -)
for arg in argv:
if not arg.startswith("-"):
if arg not in valid_commands:
print(
f"Invalid command: '{arg}'. "
f"Valid commands are: {', '.join(sorted(valid_commands))}",
file=sys.stderr,
)
return True
break
return False
def _is_extreme_compression_level_in_argv(argv: list[str]) -> bool:
"""Check for extreme compression level values that warrant special handling.
@@ -840,11 +889,17 @@ def _is_extreme_compression_level_in_argv(argv: list[str]) -> bool:
Returns:
bool: True if an extreme compression level value is found, False otherwise
"""
if "-l" not in argv and "--level" not in argv:
if "-c" not in argv and "-l" not in argv and "--level" not in argv:
return False
try:
level_index = argv.index("-l") if "-l" in argv else argv.index("--level")
level_index = -1
if "-c" in argv:
level_index = argv.index("-c")
elif "-l" in argv:
level_index = argv.index("-l")
else:
level_index = argv.index("--level")
if level_index + 1 < len(argv):
level_value = argv[level_index + 1]
try:
@@ -867,17 +922,13 @@ def _handle_parsing_errors(e: SystemExit, argv: list[str] | None) -> int:
Returns:
int: Appropriate exit code
"""
# Help was requested
""" # Help was requested
if e.code == 0:
return 0
elif e.code == 2 and argv:
# Check for validation errors we want to convert to exit code 1
if _validate_filter_in_argv(argv):
return 1
# Only convert compression level errors to exit code 1 for extreme cases
if _is_extreme_compression_level_in_argv(argv):
return 1
# For now, keep standard argparse behavior (exit code 2)
# Future versions may convert specific validation errors to exit code 1
pass
# Return the original exit code for other cases
return int(e.code) if e.code is not None else 1
+12 -11
View File
@@ -506,14 +506,14 @@ class TestCLISecurityFilters:
# First create an archive
file_paths = [str(f) for f in sample_files if f.is_file()]
archive_path = temp_dir / "filter_test.tzst"
main(["a", str(archive_path), *file_paths])
# Test extraction with invalid filter
main(
["a", str(archive_path), *file_paths]
) # Test extraction with invalid filter
extract_dir = temp_dir / "invalid_filtered"
result = main(
["x", str(archive_path), "-o", str(extract_dir), "--filter", "invalid"]
)
assert result == 1 # Should fail
assert result == 2 # Should fail with argparse error code
class TestCLIStreamingOperations:
@@ -1219,9 +1219,10 @@ class TestCLIValidationFunctions:
class TestCLIErrorHandlingInternal:
def test_handle_parsing_errors_help_requested(self):
"""Test error handling when help is requested."""
from tzst.cli import _handle_parsing_errors
from tzst.cli import (
_handle_parsing_errors, # Simulate help request (exit code 0)
)
# Simulate help request (exit code 0)
e = SystemExit(0)
result = _handle_parsing_errors(e, ["--help"])
assert result == 0
@@ -1234,9 +1235,9 @@ class TestCLIErrorHandlingInternal:
e = SystemExit(2)
argv = ["x", "test.tzst", "--filter", "invalid"]
result = _handle_parsing_errors(e, argv)
assert result == 1 # Should convert to exit code 1
captured = capsys.readouterr()
assert "Invalid filter specified" in captured.err
assert result == 2 # Should maintain argparse exit code 2
# Note: _handle_parsing_errors doesn't print error messages,
# error messages are printed by argparse before SystemExit is raised
def test_handle_parsing_errors_other_errors(self):
"""Test error handling for other parsing errors."""
@@ -1736,12 +1737,12 @@ class TestCLIArgumentParsingEdgeCases:
assert result == 1 # Test with compression level error
e = SystemExit(2)
result = _handle_parsing_errors(e, ["a", "test.tzst", "file.txt", "-l", "1000"])
assert result == 1
assert result == 2
# Test with filter error
e = SystemExit(2)
result = _handle_parsing_errors(e, ["x", "test.tzst", "--filter", "badfilter"])
assert result == 1
assert result == 2
def test_parse_arguments_edge_cases(self):
"""Test _parse_arguments with edge cases."""
+6 -10
View File
@@ -61,9 +61,7 @@ class TestCLIMissingLines:
# Test list with streaming
result = main(["l", str(archive_path), "--streaming"])
assert result == 0
# Test extract with specific options
assert result == 0 # Test extract with specific options
extract_dir = temp_dir / "extracted"
result = main(["x", str(archive_path), "-o", str(extract_dir)])
assert result == 0
@@ -76,13 +74,13 @@ class TestCLIMissingLines:
result = main(["l", str(invalid_path)])
assert result == 1
# Test with invalid compression level
# Test with invalid compression level - should return argparse error code 2
test_file = temp_dir / "test.txt"
test_file.write_text("test content")
archive_path = temp_dir / "test.tzst"
result = main(["a", str(archive_path), str(test_file), "-c", "50"])
assert result == 1
assert result == 2
def test_filter_option_edge_cases(self, temp_dir):
"""Test filter option edge cases."""
@@ -126,9 +124,7 @@ class TestCLIMissingLines:
"""Test verbose output edge cases."""
test_file = temp_dir / "test.txt"
test_file.write_text("test content")
archive_path = temp_dir / "test.tzst"
# Create archive
archive_path = temp_dir / "test.tzst" # Create archive
result = main(["a", str(archive_path), str(test_file)])
assert result == 0
@@ -145,9 +141,9 @@ class TestCLIMissingLines:
result = main([])
assert result == 1
# Test with invalid command
# Test with invalid command - should return argparse error code 2
result = main(["invalid_command"])
assert result == 1
assert result == 2
@pytest.mark.skipif(sys.platform != "win32", reason="Windows-specific test")
def test_windows_specific_functionality(self, temp_dir):