Files
hse-2026/homework/cli.py
T

751 lines
29 KiB
Python

"""Portable student CLI for registration and solution submission."""
from __future__ import annotations
import argparse
import getpass
import os
import queue
import re
import stat
import subprocess
import sys
import tempfile
import threading
import unicodedata
import uuid
import zipfile
from collections.abc import Callable, Sequence
from datetime import datetime
from pathlib import Path
from typing import Any, cast
from urllib.parse import urlsplit, urlunsplit
import httpx
import yaml
DEFAULT_CONFIG = Path(__file__).resolve().with_name("cli.yaml")
TOTAL_REQUEST_TIMEOUT_SECONDS = 120.0
SUCCESS_FIELDS = frozenset({"username", "password", "repo", "token"})
ERROR_FIELDS = frozenset({"code", "message", "request_id"})
REGISTER_ERROR_CODES = {
400: frozenset({"INVALID_REQUEST", "STUDENT_NOT_FOUND"}),
401: frozenset({"INVALID_REGISTRATION_TOKEN"}),
409: frozenset({"STUDENT_ALREADY_REGISTERED", "EMAIL_ALREADY_REGISTERED"}),
413: frozenset({"PAYLOAD_TOO_LARGE"}),
415: frozenset({"UNSUPPORTED_MEDIA_TYPE"}),
500: frozenset({"INTERNAL_ERROR"}),
502: frozenset({"GITEA_OPERATION_FAILED"}),
503: frozenset({"CONFIG_INVALID"}),
}
SUBMIT_ERROR_CODES = {
400: frozenset({"INVALID_REQUEST", "INVALID_ARCHIVE"}),
401: frozenset({"ACCESS_TOKEN_REQUIRED", "INVALID_ACCESS_TOKEN"}),
403: frozenset({"SUBMISSION_DEADLINE_EXPIRED"}),
404: frozenset({"ASSIGNMENT_NOT_FOUND"}),
409: frozenset({"SUBMISSION_IN_PROGRESS", "SOLUTION_UNCHANGED"}),
413: frozenset({"PAYLOAD_TOO_LARGE"}),
415: frozenset({"UNSUPPORTED_MEDIA_TYPE"}),
500: frozenset({"INTERNAL_ERROR"}),
502: frozenset({"GITEA_OPERATION_FAILED"}),
503: frozenset({"CONFIG_INVALID"}),
}
SUBMIT_SUCCESS_FIELDS = frozenset(
{"submission_id", "accepted_at", "late_submission", "workflow_run"}
)
ASSIGNMENT_ID_RE = re.compile(r"[a-z0-9][a-z0-9-]{0,63}\Z")
_DRIVE_PATH_RE = re.compile(r"[A-Za-z]:")
_UUID_RE = re.compile(
r"[0-9A-Fa-f]{8}-[0-9A-Fa-f]{4}-[0-9A-Fa-f]{4}-"
r"[0-9A-Fa-f]{4}-[0-9A-Fa-f]{12}\Z"
)
_RFC3339_RE = re.compile(
r"\d{4}-\d{2}-\d{2}[Tt]\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:[Zz]|[+-]\d{2}:\d{2})\Z"
)
_COPY_CHUNK_SIZE = 64 * 1024
MAX_PATH_SEGMENT_UTF8_BYTES = 255
MAX_SOLUTION_PATH_UTF8_BYTES = 1024
class CLIError(RuntimeError):
"""A safe command failure with its stable process exit code."""
def __init__(self, message: str, exit_code: int) -> None:
super().__init__(message)
self.exit_code = exit_code
class BackendRejected(CLIError):
"""A successfully validated backend error envelope."""
def __init__(self, message: str, exit_code: int, backend_code: str) -> None:
super().__init__(message, exit_code)
self.backend_code = backend_code
ClientFactory = Callable[[], httpx.Client]
def normalize_server_url(value: str) -> str:
"""Normalize an absolute root HTTP(S) backend URL without exposing user info."""
if (
not value
or any(char.isspace() or ord(char) < 0x20 or ord(char) == 0x7F for char in value)
or "?" in value
or "#" in value
):
raise CLIError("server URL is invalid", 2)
try:
parsed = urlsplit(value)
port = parsed.port
except ValueError as exc:
raise CLIError("server URL is invalid", 2) from exc
if (
parsed.scheme.lower() not in {"http", "https"}
or not parsed.hostname
or parsed.username is not None
or parsed.password is not None
or parsed.path not in {"", "/"}
or parsed.query
or parsed.fragment
):
raise CLIError("server URL is invalid", 2)
host = parsed.hostname.lower()
if ":" in host:
host = f"[{host}]"
netloc = host if port is None else f"{host}:{port}"
normalized = urlunsplit((parsed.scheme.lower(), netloc, "", "", ""))
try:
httpx.URL(f"{normalized}/register")
except httpx.InvalidURL as exc:
raise CLIError("server URL is invalid", 2) from exc
return normalized
def _default_client() -> httpx.Client:
return httpx.Client(
timeout=httpx.Timeout(120.0, connect=10.0),
follow_redirects=False,
verify=True,
)
def _post_registration_with_deadline(
client_factory: ClientFactory,
url: str,
payload: dict[str, str],
total_timeout_seconds: float,
) -> httpx.Response:
outcomes: queue.SimpleQueue[httpx.Response | BaseException] = queue.SimpleQueue()
def send() -> None:
try:
with client_factory() as client:
outcomes.put(client.post(url, json=payload))
except BaseException as exc:
outcomes.put(exc)
worker = threading.Thread(target=send, name="distsys-cli-request", daemon=True)
worker.start()
worker.join(total_timeout_seconds)
if worker.is_alive():
raise CLIError("network request failed", 4)
outcome = outcomes.get()
if isinstance(outcome, BaseException):
raise outcome
return outcome
def _json_object(response: httpx.Response) -> dict[str, Any]:
content_type = response.headers.get("content-type", "").split(";", 1)[0].strip().lower()
if content_type != "application/json":
raise CLIError("backend returned an invalid response", 4)
try:
value = response.json()
except ValueError as exc:
raise CLIError("backend returned an invalid response", 4) from exc
if not isinstance(value, dict):
raise CLIError("backend returned an invalid response", 4)
return value
def _request_id(response: httpx.Response) -> str:
request_id = response.headers.get("X-Request-ID")
if not request_id:
raise CLIError("backend returned an invalid response", 4)
return str(request_id)
def _validate_success(response: httpx.Response) -> tuple[dict[str, str], str]:
request_id = _request_id(response)
if response.headers.get("Cache-Control") != "no-store":
raise CLIError("backend returned an invalid response", 4)
body = _json_object(response)
if set(body) != SUCCESS_FIELDS or any(
not isinstance(body.get(field), str) or not body[field] for field in SUCCESS_FIELDS
):
raise CLIError("backend returned an invalid response", 4)
repository = str(body["repo"])
try:
parsed_repository = urlsplit(repository)
except ValueError as exc:
raise CLIError("backend returned an invalid response", 4) from exc
if (
any(char.isspace() for char in repository)
or not parsed_repository.scheme
or not parsed_repository.netloc
):
raise CLIError("backend returned an invalid response", 4)
return {field: str(body[field]) for field in SUCCESS_FIELDS}, request_id
def _raise_backend_error(
response: httpx.Response,
allowed_errors: dict[int, frozenset[str]] = REGISTER_ERROR_CODES,
) -> None:
request_id = _request_id(response)
body = _json_object(response)
if (
set(body) != ERROR_FIELDS
or not all(isinstance(body.get(field), str) and body[field] for field in ERROR_FIELDS)
or body["request_id"] != request_id
):
raise CLIError("backend returned an invalid response", 4)
allowed_codes = allowed_errors.get(response.status_code)
if allowed_codes is None or body["code"] not in allowed_codes:
raise CLIError("backend returned an invalid response", 4)
status_class = response.status_code // 100
exit_code = 3 if status_class == 4 else 4
raise BackendRejected(
f"{body['message']} (request ID: {request_id})",
exit_code,
str(body["code"]),
)
class _StrictLoader(yaml.SafeLoader):
"""YAML loader that rejects duplicate mapping keys."""
def _strict_mapping(loader: _StrictLoader, node: yaml.MappingNode, deep: bool = False) -> Any:
loader.flatten_mapping(node)
result: dict[Any, Any] = {}
for key_node, value_node in node.value:
key = loader.construct_object(key_node, deep=deep)
try:
duplicate = key in result
except TypeError as exc:
raise yaml.constructor.ConstructorError(
None, None, "mapping key is not hashable", key_node.start_mark
) from exc
if duplicate:
raise yaml.constructor.ConstructorError(
None, None, "duplicate mapping key", key_node.start_mark
)
result[key] = loader.construct_object(value_node, deep=deep)
return result
_StrictLoader.add_constructor(yaml.resolver.BaseResolver.DEFAULT_MAPPING_TAG, _strict_mapping)
def _load_config(config_path: Path | None) -> tuple[str, str]:
path = DEFAULT_CONFIG if config_path is None else config_path
if config_path is not None and not path.is_absolute():
path = Path.cwd() / path
try:
documents = list(yaml.load_all(path.read_text(encoding="utf-8"), Loader=_StrictLoader))
except (OSError, UnicodeError, yaml.YAMLError) as exc:
raise CLIError("configuration is invalid", 2) from exc
if len(documents) != 1 or not isinstance(documents[0], dict):
raise CLIError("configuration is invalid", 2)
value = documents[0]
if set(value) - {"server_url", "access_token"} or "server_url" not in value:
raise CLIError("configuration is invalid", 2)
server_url = value.get("server_url")
access_token = value.get("access_token")
if not isinstance(server_url, str) or not isinstance(access_token, str) or not access_token:
raise CLIError("configuration is invalid or has no access token", 2)
return normalize_server_url(server_url), access_token
def _normalized_segment(value: str) -> str:
if (
not value
or value in {".", ".."}
or "\\" in value
or _DRIVE_PATH_RE.match(value)
or any(unicodedata.category(character) == "Cc" for character in value)
):
raise CLIError("solution contains an unsupported path", 2)
normalized = unicodedata.normalize("NFC", value)
if (
normalized in {"", ".", "..", ".git"}
or len(normalized.encode("utf-8")) > MAX_PATH_SEGMENT_UTF8_BYTES
):
raise CLIError("solution contains an unsupported path", 2)
return normalized
def _validate_solution_path(path: tuple[str, ...]) -> None:
if len("/".join(path).encode("utf-8")) > MAX_SOLUTION_PATH_UTF8_BYTES:
raise CLIError("solution contains an unsupported path", 2)
def _collect_solution_entries(
source: Path,
) -> list[tuple[Path, tuple[str, ...], bool, os.stat_result]]:
entries: list[tuple[Path, tuple[str, ...], bool, os.stat_result]] = []
normalized_paths: dict[tuple[str, ...], bool] = {}
def visit(directory: Path, relative: tuple[str, ...]) -> None:
try:
children = sorted(os.scandir(directory), key=lambda entry: entry.name)
except OSError as exc:
raise CLIError("solution directory cannot be read", 2) from exc
for child in children:
normalized = (*relative, _normalized_segment(child.name))
_validate_solution_path(normalized)
try:
metadata = child.stat(follow_symlinks=False)
except OSError as exc:
raise CLIError("solution entry cannot be read", 2) from exc
mode = metadata.st_mode
if stat.S_ISLNK(mode):
raise CLIError("solution contains a symbolic link", 2)
directory_entry = stat.S_ISDIR(mode)
regular = stat.S_ISREG(mode)
if not directory_entry and not regular:
raise CLIError("solution contains a non-regular entry", 2)
if regular and getattr(metadata, "st_nlink", 1) > 1:
raise CLIError("solution contains a hard-linked file", 2)
if normalized in normalized_paths:
raise CLIError("solution contains duplicate normalized paths", 2)
for length in range(1, len(normalized)):
ancestor = normalized[:length]
if ancestor in normalized_paths and not normalized_paths[ancestor]:
raise CLIError("solution contains a file/directory conflict", 2)
normalized_paths[normalized] = directory_entry
entries.append((Path(child.path), normalized, directory_entry, metadata))
if directory_entry:
visit(Path(child.path), normalized)
visit(source, ())
if not any(not directory for _, _, directory, _ in entries):
raise CLIError("solution directory contains no regular files", 2)
return entries
def _zip_info(name: str, mode: int, *, directory: bool) -> zipfile.ZipInfo:
info = zipfile.ZipInfo(name + ("/" if directory else ""))
info.create_system = 3
file_type = stat.S_IFDIR if directory else stat.S_IFREG
info.external_attr = (file_type | mode) << 16
info.compress_type = zipfile.ZIP_DEFLATED
return info
def _source_is_executable(mode: int, platform_name: str = os.name) -> bool:
"""Return whether this platform exposes any executable bit for a source file."""
return platform_name == "posix" and bool(mode & 0o111)
def _resolve_solution_directory(solution_dir: Path) -> Path:
try:
source = solution_dir.resolve(strict=True)
except OSError as exc:
raise CLIError("solution directory does not exist", 2) from exc
if not source.is_dir():
raise CLIError("solution path is not a directory", 2)
return source
def build_solution_archive(solution_dir: Path, destination: Path) -> None:
"""Build one normalized, link-free archive below a synthetic solution root."""
source = _resolve_solution_directory(solution_dir)
entries = _collect_solution_entries(source)
try:
with zipfile.ZipFile(destination, "w", allowZip64=True) as archive:
archive.writestr(_zip_info("solution", 0o755, directory=True), b"")
for path, relative, directory, selected_metadata in entries:
member_name = "solution/" + "/".join(relative)
if directory:
archive.writestr(_zip_info(member_name, 0o755, directory=True), b"")
continue
executable = _source_is_executable(selected_metadata.st_mode)
permissions = 0o755 if executable else 0o644
flags = os.O_RDONLY | getattr(os, "O_BINARY", 0) | getattr(os, "O_NOFOLLOW", 0)
descriptor = os.open(path, flags)
try:
current = os.fstat(descriptor)
if getattr(current, "st_nlink", 1) > 1:
raise CLIError("solution contains a hard-linked file", 2)
identity_changed = os.name == "posix" and (
current.st_dev != selected_metadata.st_dev
or current.st_ino != selected_metadata.st_ino
)
metadata_changed = (
current.st_size != selected_metadata.st_size
or current.st_mtime_ns != selected_metadata.st_mtime_ns
)
if not stat.S_ISREG(current.st_mode) or identity_changed or metadata_changed:
raise CLIError("solution changed while it was archived", 2)
with os.fdopen(descriptor, "rb", closefd=True) as source_file:
descriptor = -1
info = _zip_info(member_name, permissions, directory=False)
with archive.open(info, "w") as output:
while chunk := source_file.read(_COPY_CHUNK_SIZE):
output.write(chunk)
final = os.fstat(source_file.fileno())
if (
final.st_size != current.st_size
or final.st_mtime_ns != current.st_mtime_ns
):
raise CLIError("solution changed while it was archived", 2)
finally:
if descriptor >= 0:
os.close(descriptor)
except CLIError:
raise
except (OSError, ValueError, zipfile.BadZipFile) as exc:
raise CLIError("solution archive could not be created", 1) from exc
def _post_archive_with_deadline(
client_factory: ClientFactory,
url: str,
archive_path: Path,
access_token: str,
total_timeout_seconds: float,
) -> httpx.Response:
outcomes: queue.SimpleQueue[httpx.Response | BaseException] = queue.SimpleQueue()
# Windows does not permit unlinking a file held open by the request worker.
# Detach the request body from the temporary path before enforcing the outer
# wall-clock deadline so timeout cleanup remains immediate.
windows_content = archive_path.read_bytes() if os.name == "nt" else None
headers = {
"Authorization": f"token {access_token}",
"Content-Type": "application/zip",
}
def send() -> None:
try:
with client_factory() as client:
if windows_content is not None:
outcomes.put(client.post(url, content=windows_content, headers=headers))
else:
with archive_path.open("rb") as content:
outcomes.put(client.post(url, content=content, headers=headers))
except BaseException as exc:
outcomes.put(exc)
finally:
_remove_temporary_archive(archive_path)
worker = threading.Thread(target=send, name="distsys-cli-request", daemon=True)
worker.start()
worker.join(total_timeout_seconds)
if worker.is_alive():
raise CLIError("network request failed", 4)
outcome = outcomes.get()
if isinstance(outcome, BaseException):
raise outcome
return outcome
def _remove_temporary_archive(path: Path) -> None:
try:
path.unlink(missing_ok=True)
except OSError:
pass
def _validate_submit_success(response: httpx.Response) -> dict[str, str | bool]:
_request_id(response)
body = _json_object(response)
if set(body) != SUBMIT_SUCCESS_FIELDS:
raise CLIError("backend returned an invalid response", 4)
submission_id = body.get("submission_id")
accepted_at = body.get("accepted_at")
late = body.get("late_submission")
workflow_run = body.get("workflow_run")
if not isinstance(submission_id, str) or not submission_id:
raise CLIError("backend returned an invalid response", 4)
if not isinstance(accepted_at, str) or not accepted_at:
raise CLIError("backend returned an invalid response", 4)
if not isinstance(workflow_run, str) or not workflow_run:
raise CLIError("backend returned an invalid response", 4)
if not isinstance(late, bool):
raise CLIError("backend returned an invalid response", 4)
if _UUID_RE.fullmatch(submission_id) is None or _RFC3339_RE.fullmatch(accepted_at) is None:
raise CLIError("backend returned an invalid response", 4)
try:
uuid.UUID(submission_id)
normalized_time = accepted_at[:-1] + "+00:00" if accepted_at[-1] in "Zz" else accepted_at
parsed_time = datetime.fromisoformat(normalized_time)
parsed_url = urlsplit(workflow_run)
workflow_port = parsed_url.port
parsed_http_url = httpx.URL(workflow_run)
except (UnicodeError, ValueError, TypeError, httpx.InvalidURL) as exc:
raise CLIError("backend returned an invalid response", 4) from exc
if (
parsed_time.tzinfo is None
or accepted_at.endswith("-00:00")
or parsed_url.scheme not in {"http", "https"}
or not parsed_url.netloc
or parsed_url.hostname is None
or parsed_url.username is not None
or parsed_url.password is not None
or workflow_port == 0
or not parsed_http_url.host
or any(
character.isspace() or ord(character) < 0x20 or ord(character) == 0x7F
for character in workflow_run
)
):
raise CLIError("backend returned an invalid response", 4)
return {
"submission_id": submission_id,
"accepted_at": accepted_at,
"late_submission": late,
"workflow_run": workflow_run,
}
def _format_accepted_at(accepted_at: str) -> str:
normalized_time = accepted_at[:-1] + "+00:00" if accepted_at[-1] in "Zz" else accepted_at
return datetime.fromisoformat(normalized_time).strftime("%Y-%m-%d %H:%M:%S")
def submit(
assignment_id: str,
solution_dir: Path,
config_path: Path | None,
*,
client_factory: ClientFactory = _default_client,
total_timeout_seconds: float = TOTAL_REQUEST_TIMEOUT_SECONDS,
) -> None:
if ASSIGNMENT_ID_RE.fullmatch(assignment_id) is None:
raise CLIError("assignment ID is invalid", 2)
server_url, access_token = _load_config(config_path)
resolved_solution = _resolve_solution_directory(solution_dir)
descriptor, archive_name = tempfile.mkstemp(prefix="distsys-solution-", suffix=".zip")
os.close(descriptor)
archive_path = Path(archive_name)
try:
try:
archive_path.resolve().relative_to(resolved_solution)
except ValueError:
pass
else:
raise CLIError("temporary archive location overlaps the solution directory", 2)
build_solution_archive(resolved_solution, archive_path)
try:
response = _post_archive_with_deadline(
client_factory,
f"{server_url}/submit/{assignment_id}",
archive_path,
access_token,
total_timeout_seconds,
)
except (httpx.HTTPError, httpx.InvalidURL, OSError) as exc:
raise CLIError("network request failed", 4) from exc
if response.status_code != 201:
try:
_raise_backend_error(response, SUBMIT_ERROR_CODES)
except BackendRejected as exc:
if exc.backend_code == "SOLUTION_UNCHANGED":
raise CLIError(f"{exc}; no grading run was created", 3) from None
raise
result = _validate_submit_success(response)
finally:
_remove_temporary_archive(archive_path)
print(f"Submission ID: {result['submission_id']}")
print(f"Accepted at: {_format_accepted_at(cast(str, result['accepted_at']))}")
print(f"Late submission: {str(result['late_submission']).lower()}")
print(f"Workflow run: {result['workflow_run']}")
def _restrict_windows(path: Path) -> bool:
creation_flags = getattr(subprocess, "CREATE_NO_WINDOW", 0)
try:
identity = subprocess.run(
["whoami"],
check=False,
stdout=subprocess.PIPE,
stderr=subprocess.DEVNULL,
text=True,
creationflags=creation_flags,
)
principal = identity.stdout.strip()
if identity.returncode != 0 or not principal:
return False
completed = subprocess.run(
[
"icacls",
str(path),
"/inheritance:r",
"/grant:r",
f"{principal}:(F)",
],
check=False,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
creationflags=creation_flags,
)
except OSError:
return False
return completed.returncode == 0
def _write_config(destination: Path, server_url: str, access_token: str) -> bool:
"""Synchronize a private sibling temporary file, then atomically replace."""
payload = yaml.safe_dump(
{"server_url": server_url, "access_token": access_token},
allow_unicode=True,
sort_keys=False,
).encode("utf-8")
descriptor = -1
temporary: Path | None = None
permission_warning = False
try:
descriptor, temporary_name = tempfile.mkstemp(
prefix=f".{destination.name}.", dir=destination.parent
)
temporary = Path(temporary_name)
with os.fdopen(descriptor, "wb", closefd=True) as stream:
descriptor = -1
stream.write(payload)
stream.flush()
os.fsync(stream.fileno())
if os.name == "posix":
temporary.chmod(stat.S_IRUSR | stat.S_IWUSR)
if os.name == "nt" and not _restrict_windows(temporary):
permission_warning = True
os.replace(temporary, destination)
temporary = None
finally:
if descriptor >= 0:
os.close(descriptor)
if temporary is not None:
try:
temporary.unlink()
except OSError:
pass
return permission_warning
def _resolve_destination(value: Path | None) -> Path:
destination = DEFAULT_CONFIG if value is None else value
if value is not None and not destination.is_absolute():
destination = Path.cwd() / destination
try:
destination = destination.resolve(strict=False)
destination.parent.mkdir(parents=True, exist_ok=True)
except OSError as exc:
raise CLIError("configuration destination is not usable", 2) from exc
if destination.exists() and destination.is_dir():
raise CLIError("configuration destination is not a file", 2)
return destination
def _masked_email(email: str) -> str:
local, separator, domain = email.rpartition("@")
if not separator or not local or not domain:
return "***"
return f"{local[0]}***@{domain}"
def register(
server_url: str,
config_path: Path | None,
*,
input_fn: Callable[[str], str] = input,
password_fn: Callable[[str], str] = getpass.getpass,
client_factory: ClientFactory = _default_client,
total_timeout_seconds: float = TOTAL_REQUEST_TIMEOUT_SECONDS,
) -> None:
normalized_url = normalize_server_url(server_url)
destination = _resolve_destination(config_path)
registration_token = password_fn("Registration token: ")
name = input_fn("Full name: ")
email = input_fn("Email: ")
try:
response = _post_registration_with_deadline(
client_factory,
f"{normalized_url}/register",
{"name": name, "email": email, "reg_token": registration_token},
total_timeout_seconds,
)
except (httpx.HTTPError, httpx.InvalidURL) as exc:
raise CLIError("network request failed", 4) from exc
if response.status_code != 201:
_raise_backend_error(response)
result, request_id = _validate_success(response)
try:
permission_warning = _write_config(destination, normalized_url, result["token"])
except OSError as exc:
raise CLIError(
"Registration succeeded, but credentials could not be stored. "
f"Contact course staff with request ID: {request_id}",
1,
) from exc
if permission_warning:
print(
"warning: could not restrict configuration permissions to the current user",
file=sys.stderr,
)
print(f"Username: {result['username']}")
print(f"Temporary password: {result['password']}")
print(f"Repository: {result['repo']}")
print(f"Registration email: Check {_masked_email(email)} for your Gitea account details.")
print("If it does not arrive within 5 minutes, check spam and contact course staff.")
print("Action required: Sign in to Gitea and change the temporary password before submitting.")
def parse_args(arguments: Sequence[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--config", type=Path, default=None)
subparsers = parser.add_subparsers(dest="command", required=True)
register_parser = subparsers.add_parser("register", help="create a student account")
register_parser.add_argument("--config", type=Path, default=argparse.SUPPRESS)
register_parser.add_argument("server_url")
submit_parser = subparsers.add_parser("submit", help="submit a solution for grading")
submit_parser.add_argument("--config", type=Path, default=argparse.SUPPRESS)
submit_parser.add_argument("--assignment-id")
submit_parser.add_argument("--solution-dir", type=Path)
return parser.parse_args(arguments)
def _resolve_submit_options(
assignment_id: str | None, solution_dir: Path | None
) -> tuple[str, Path]:
if assignment_id is None and solution_dir is None:
working_directory = Path.cwd()
default_solution = working_directory / "solution"
if not default_solution.is_dir():
raise CLIError("submit without options requires a solution subdirectory", 2)
return working_directory.name, default_solution
if assignment_id is None or solution_dir is None:
raise CLIError("--assignment-id and --solution-dir must be provided together", 2)
return assignment_id, solution_dir
def main(arguments: Sequence[str] | None = None) -> int:
args = parse_args(arguments)
try:
if args.command == "register":
register(args.server_url, args.config)
elif args.command == "submit":
assignment_id, solution_dir = _resolve_submit_options(
args.assignment_id, args.solution_dir
)
submit(assignment_id, solution_dir, args.config)
except CLIError as exc:
print(str(exc), file=sys.stderr)
return exc.exit_code
return 0
if __name__ == "__main__":
raise SystemExit(main())