diff --git a/scripts/pre_commit_lint.sh b/scripts/pre_commit_lint.sh index 02fd78dfbb3..d2cf7cd307f 100755 --- a/scripts/pre_commit_lint.sh +++ b/scripts/pre_commit_lint.sh @@ -113,15 +113,19 @@ python_checks() { on_interrupt() { trap - INT TERM rm -f "${python_log:-}" "${dash_log:-}" "${gen_log:-}" - kill 0 2>/dev/null + for job_pid in ${python_pid:-} ${dash_pid:-} ${gen_pid:-}; do + kill -- "-$job_pid" 2>/dev/null || true + done exit 130 } trap on_interrupt INT TERM if [ -n "$litellm_py_files" ]; then python_log=$(mktemp) + set -m python_checks > "$python_log" 2>&1 & python_pid=$! + set +m fi if [ -n "$e2e_py_files" ] && [ -z "$litellm_py_files" ]; then @@ -147,8 +151,10 @@ dashboard_checks() { if [ -n "$ui_prettier_files" ] || [ -n "$ui_eslint_files" ]; then dash_log=$(mktemp) + set -m dashboard_checks > "$dash_log" 2>&1 & dash_pid=$! + set +m fi genapi_checks() { @@ -183,8 +189,10 @@ genapi_checks() { if [ -n "$spec_files" ]; then gen_log=$(mktemp) + set -m genapi_checks > "$gen_log" 2>&1 & gen_pid=$! + set +m fi if [ -n "${python_pid:-}" ]; then diff --git a/tests/test_litellm/test_pre_commit_lint.py b/tests/test_litellm/test_pre_commit_lint.py index f9dbf17b8b2..92de89a3cd5 100644 --- a/tests/test_litellm/test_pre_commit_lint.py +++ b/tests/test_litellm/test_pre_commit_lint.py @@ -192,6 +192,30 @@ def test_interrupt_kills_background_jobs_and_removes_logs(tmp_path: Path) -> Non os.killpg(proc.pid, signal.SIGTERM) +def test_interrupt_spares_the_invoking_process(tmp_path: Path) -> None: + repo, bin_dir = _sandbox(tmp_path) + hang_dir = tmp_path / "hang" + hang_dir.mkdir() + marker = tmp_path / "invoker_survived" + proc = subprocess.Popen( + ["bash", "-c", 'trap : INT; "$1"; echo "$?" > "$2"', "bash", str(SCRIPT), str(marker)], + cwd=repo, + env=_env(repo, bin_dir, {"STUB_HANG_DIR": str(hang_dir)}), + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + start_new_session=True, + ) + try: + assert _wait_until((hang_dir / "make.started").exists, 10) + os.killpg(proc.pid, signal.SIGINT) + assert proc.wait(timeout=10) == 0 + assert _wait_until(marker.exists, 5) + assert marker.read_text().strip() == "130" + finally: + with suppress(ProcessLookupError, PermissionError): + os.killpg(proc.pid, signal.SIGTERM) + + @pytest.mark.parametrize( ("fail", "message"), [