diff --git a/src/skillspector/cleanup.py b/src/skillspector/cleanup.py index bb5ff4a1f..10e0f3318 100644 --- a/src/skillspector/cleanup.py +++ b/src/skillspector/cleanup.py @@ -6,8 +6,12 @@ import os import shutil import stat -from collections.abc import Callable +from collections.abc import Callable, Iterator, Mapping +from contextlib import contextmanager from pathlib import Path +from typing import Any + +from langchain_core.callbacks import BaseCallbackHandler from skillspector.python_ast import clear_python_ast_cache @@ -43,6 +47,66 @@ def remove_temp_tree(path: str | Path) -> None: shutil.rmtree(path, onexc=_retry_writable) +SCAN_TEMP_DIR_PREFIX = "skillspector_" + + +def _is_scan_temp_dir(value: object) -> bool: + """Return whether ``value`` names an existing, non-symlink scan temp directory.""" + if not isinstance(value, str) or not value: + return False + path = Path(value) + if not path.name.startswith(SCAN_TEMP_DIR_PREFIX): + return False + try: + return path.is_dir() and not path.is_symlink() and not os.path.isjunction(path) + except OSError: + return False + + +class TempDirTracker(BaseCallbackHandler): + """Remember the temp directory a graph run materializes, as soon as it exists. + + ``resolve_input`` reports the directory in its output as ``temp_dir_for_cleanup``, + but a run that raises or is interrupted returns no result for + :func:`cleanup_result`. Pass the tracker in the run's ``callbacks`` and use + :meth:`removing_on_error` around the run. + """ + + run_inline = True + + def __init__(self) -> None: + super().__init__() + self.temp_dir: str | None = None + + def on_chain_end(self, outputs: Any, **kwargs: Any) -> None: + """Record ``temp_dir_for_cleanup`` from a node or graph output. + + The recorded path is later removed recursively, so only a value that + looks like a scan temp dir is accepted: an existing directory, not a + symlink, whose name carries the ``skillspector_`` prefix that + ``InputHandler`` gives ``mkdtemp``. Anything else is ignored and never + replaces a path already recorded. + """ + if isinstance(outputs, Mapping): + temp_dir = outputs.get("temp_dir_for_cleanup") + if _is_scan_temp_dir(temp_dir): + self.temp_dir = temp_dir + + def remove(self) -> None: + """Remove the recorded temp directory, if any.""" + if self.temp_dir: + remove_temp_tree(self.temp_dir) + + @contextmanager + def removing_on_error(self) -> Iterator[None]: + """Remove the recorded temp directory if the wrapped run raises or is interrupted.""" + try: + yield + except BaseException: + self.remove() + raise + + def cleanup_result(result: dict[str, object]) -> None: """Release scan-local resources and remove a temp dir if set.""" python_ast_cache_key = result.get("python_ast_cache_key") diff --git a/src/skillspector/cli.py b/src/skillspector/cli.py index 32560d70a..0b8ad5cbb 100644 --- a/src/skillspector/cli.py +++ b/src/skillspector/cli.py @@ -44,7 +44,7 @@ from rich.tree import Tree from skillspector import __version__, transitive -from skillspector.cleanup import cleanup_result +from skillspector.cleanup import TempDirTracker, cleanup_result from skillspector.constants import RISK_THRESHOLD from skillspector.graph_proxy import graph from skillspector.input_handler import validate_local_input_path @@ -1512,8 +1512,13 @@ def _run_graph_scan( if initial_inspection_ledger: state["inspection_ledger"] = initial_inspection_ledger trace_config = _build_trace_config(input_path, format, no_llm) + # A scan that raises or is interrupted returns no result for the caller to + # clean up, so remove the temp directory resolve_input made here instead. + temp_dir_tracker = TempDirTracker() + trace_config["callbacks"] = [temp_dir_tracker] if not stream_progress: - return cast(dict[str, object], graph.invoke(state, config=trace_config)) + with temp_dir_tracker.removing_on_error(): + return cast(dict[str, object], graph.invoke(state, config=trace_config)) analyzer_node_ids = _wired_analyzer_node_ids() total_analyzers = len(analyzer_node_ids) @@ -1531,6 +1536,7 @@ def _run_graph_scan( console=err_console, transient=True, ) as progress, + temp_dir_tracker.removing_on_error(), ): warnings.filterwarnings("ignore", category=UserWarning, module="pydantic") task_id = progress.add_task("Resolving input...", total=total_steps) @@ -2527,21 +2533,29 @@ def _scan_skill( active_visited.add(transitive.canonicalize_source_identity(input_path)) except ValueError: pass - return _scan_transitive( - initial_result=result, - format=format, - no_llm=no_llm, - max_depth=transitive_depth, - transitive_allow_prefix=transitive_allow_prefix, - transitive_deny_prefix=transitive_deny_prefix, - baseline=baseline, - show_suppressed=show_suppressed, - visited=active_visited, - scan_cache=transitive_cache, - yara_dir=yara_dir, - traversal=transitive_traversal, - source_local_only=source_local_only, - ) + # The root graph has returned, so its tracker no longer guards the root's + # temp dir. If the transitive phase is interrupted or raises, nothing is + # returned for the caller's cleanup_result, so remove it here. On success + # the merged result carries the same temp_dir_for_cleanup for the caller. + try: + return _scan_transitive( + initial_result=result, + format=format, + no_llm=no_llm, + max_depth=transitive_depth, + transitive_allow_prefix=transitive_allow_prefix, + transitive_deny_prefix=transitive_deny_prefix, + baseline=baseline, + show_suppressed=show_suppressed, + visited=active_visited, + scan_cache=transitive_cache, + yara_dir=yara_dir, + traversal=transitive_traversal, + source_local_only=source_local_only, + ) + except BaseException: + cleanup_result(result) + raise def _multi_skill_public_record_count(result: dict[str, object]) -> int: diff --git a/src/skillspector/mcp_server.py b/src/skillspector/mcp_server.py index 57699867b..4933e0ae4 100644 --- a/src/skillspector/mcp_server.py +++ b/src/skillspector/mcp_server.py @@ -32,7 +32,7 @@ from typing import TYPE_CHECKING, Any from skillspector import __version__ -from skillspector.cleanup import cleanup_result +from skillspector.cleanup import TempDirTracker, cleanup_result from skillspector.constants import RISK_THRESHOLD from skillspector.graph import graph from skillspector.graph_proxy import restore_package_graph_export @@ -153,10 +153,14 @@ async def run_scan( ) result: dict[str, Any] | None = None + # A cancelled or failed scan returns no result to clean up; the tracker still + # knows the temp directory resolve_input made. + temp_dir_tracker = TempDirTracker() try: result = await graph.ainvoke( state, config={ + "callbacks": [temp_dir_tracker], "run_name": "skillspector-mcp-scan", "tags": ["skillspector", "mcp"], "metadata": { @@ -228,6 +232,8 @@ async def run_scan( finally: if result is not None: cleanup_result(result) + else: + temp_dir_tracker.remove() def build_server(name: str = "skillspector", *, allow_local_targets: bool = False) -> FastMCP: diff --git a/tests/unit/test_cleanup.py b/tests/unit/test_cleanup.py index 83b7e4fbe..3611435c3 100644 --- a/tests/unit/test_cleanup.py +++ b/tests/unit/test_cleanup.py @@ -11,7 +11,7 @@ import pytest -from skillspector.cleanup import _retry_writable, cleanup_result +from skillspector.cleanup import TempDirTracker, _retry_writable, cleanup_result from skillspector.input_handler import InputHandler @@ -163,3 +163,111 @@ def test_input_handler_cleanup_removes_read_only_git_objects( assert not temp_dir.exists() assert handler.temp_dir_for_cleanup() is None + + +def _scan_temp_dir(root: Path) -> Path: + """Create a directory named like the ones ``InputHandler`` materializes.""" + path = root / "skillspector_abc123" + path.mkdir() + (path / "SKILL.md").write_text("# Skill\n", encoding="utf-8") + return path + + +def test_tracker_records_a_scan_temp_dir(tmp_path: Path) -> None: + """A node output naming an existing skillspector_ directory is recorded.""" + temp_dir = _scan_temp_dir(tmp_path) + tracker = TempDirTracker() + + tracker.on_chain_end({"temp_dir_for_cleanup": str(temp_dir)}) + + assert tracker.temp_dir == str(temp_dir) + + +@pytest.mark.parametrize("later", [None, "", 42, ["x"]]) +def test_tracker_keeps_recorded_path_when_a_later_output_has_none( + tmp_path: Path, later: object +) -> None: + """A later output without a usable path does not erase the recorded one.""" + temp_dir = _scan_temp_dir(tmp_path) + tracker = TempDirTracker() + tracker.on_chain_end({"temp_dir_for_cleanup": str(temp_dir)}) + + tracker.on_chain_end({"temp_dir_for_cleanup": later}) + tracker.on_chain_end({"other": "value"}) + tracker.on_chain_end("not a mapping") + + assert tracker.temp_dir == str(temp_dir) + + +def test_tracker_ignores_paths_that_are_not_scan_temp_dirs(tmp_path: Path) -> None: + """Only an existing, non-symlink skillspector_ directory can become a deletion target.""" + user_dir = tmp_path / "my-skill" + user_dir.mkdir() + missing = tmp_path / "skillspector_missing" + a_file = tmp_path / "skillspector_file" + a_file.write_text("x", encoding="utf-8") + tracker = TempDirTracker() + + for candidate in (user_dir, missing, a_file): + tracker.on_chain_end({"temp_dir_for_cleanup": str(candidate)}) + + assert tracker.temp_dir is None + assert user_dir.exists() + + +@pytest.mark.skipif(not hasattr(os, "symlink"), reason="symlinks unavailable") +def test_tracker_ignores_a_symlink_named_like_a_scan_temp_dir(tmp_path: Path) -> None: + """A symlink is never recorded, even with the scan prefix and a directory target.""" + target = tmp_path / "keep" + target.mkdir() + link = tmp_path / "skillspector_link" + try: + link.symlink_to(target, target_is_directory=True) + except OSError: + pytest.skip("cannot create symlinks here") + tracker = TempDirTracker() + + tracker.on_chain_end({"temp_dir_for_cleanup": str(link)}) + + assert tracker.temp_dir is None + + +def test_tracker_remove_twice_is_a_no_op(tmp_path: Path) -> None: + """Removing an already-removed directory does not raise.""" + temp_dir = _scan_temp_dir(tmp_path) + tracker = TempDirTracker() + tracker.on_chain_end({"temp_dir_for_cleanup": str(temp_dir)}) + + tracker.remove() + tracker.remove() + + assert not temp_dir.exists() + + +@pytest.mark.parametrize("error", [KeyboardInterrupt, RuntimeError]) +def test_removing_on_error_reraises_the_original_exception( + tmp_path: Path, error: type[BaseException] +) -> None: + """The wrapped run's exception propagates unchanged after the directory is removed.""" + temp_dir = _scan_temp_dir(tmp_path) + tracker = TempDirTracker() + original = error("stopped") + + with pytest.raises(error) as raised: + with tracker.removing_on_error(): + tracker.on_chain_end({"temp_dir_for_cleanup": str(temp_dir)}) + raise original + + assert raised.value is original + assert not temp_dir.exists() + + +def test_removing_on_error_leaves_the_directory_on_success(tmp_path: Path) -> None: + """A run that completes leaves removal to the caller's cleanup_result.""" + temp_dir = _scan_temp_dir(tmp_path) + tracker = TempDirTracker() + + with tracker.removing_on_error(): + tracker.on_chain_end({"temp_dir_for_cleanup": str(temp_dir)}) + + assert temp_dir.exists() diff --git a/tests/unit/test_cli.py b/tests/unit/test_cli.py index 465e130a8..a559fd4dc 100644 --- a/tests/unit/test_cli.py +++ b/tests/unit/test_cli.py @@ -22,6 +22,8 @@ import shutil import subprocess import sys +import tempfile +import zipfile from collections.abc import Callable, Iterator from contextlib import AbstractContextManager, ExitStack, contextmanager, nullcontext from importlib import import_module @@ -173,6 +175,118 @@ def fake_stream( assert result == final_state +def _stop_scan_after_input_resolution( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, error: type[BaseException] +) -> tuple[Path, list[Path]]: + """Build a zipped skill, record scan temp dirs, and make the report step raise.""" + archive = tmp_path / "skill.zip" + with zipfile.ZipFile(archive, "w") as bundle: + bundle.writestr("SKILL.md", "---\nname: demo\ndescription: d\n---\n# Demo\n") + created: list[Path] = [] + real_mkdtemp = tempfile.mkdtemp + + def recording_mkdtemp(*args: Any, **kwargs: Any) -> str: + """Create scan temp dirs under tmp_path and remember them.""" + if kwargs.get("prefix") != "skillspector_": + return real_mkdtemp(*args, **kwargs) + path = real_mkdtemp(*args, **{**kwargs, "dir": tmp_path}) + created.append(Path(path)) + return path + + def stop(*args: Any, **kwargs: Any) -> str: + """Stand in for an interrupt or a failure after the input is materialized.""" + raise error("scan stopped") + + monkeypatch.setattr("skillspector.input_handler.tempfile.mkdtemp", recording_mkdtemp) + monkeypatch.setattr("skillspector.nodes.report._format_json", stop) + return archive, created + + +@pytest.mark.parametrize("stream_progress", [False, True]) +@pytest.mark.parametrize("error", [KeyboardInterrupt, RuntimeError]) +def test_scan_that_stops_early_removes_its_temp_dir( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, + stream_progress: bool, + error: type[BaseException], +) -> None: + """A scan interrupted or failing after input resolution removes its temp dir.""" + archive, created = _stop_scan_after_input_resolution(monkeypatch, tmp_path, error) + + with pytest.raises(error): + cli._run_graph_scan( + input_path=str(archive), + format=FormatChoice.json, + no_llm=True, + stream_progress=stream_progress, + ) + + assert created + assert not any(path.exists() for path in created) + + +@pytest.mark.parametrize("error", [KeyboardInterrupt, RuntimeError]) +def test_transitive_scan_stopped_in_a_child_removes_the_root_temp_dir( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, error: type[BaseException] +) -> None: + """A --transitive scan of a zipped root that stops while a child is scanning removes the root's temp dir.""" + archive = tmp_path / "skill.zip" + with zipfile.ZipFile(archive, "w") as bundle: + bundle.writestr( + "SKILL.md", + "---\nname: demo\ndescription: d\n---\n# Demo\n\n" + "Install the helper skill from https://github.com/org/dep.git first.\n", + ) + created: list[Path] = [] + real_mkdtemp = tempfile.mkdtemp + + def recording_mkdtemp(*args: Any, **kwargs: Any) -> str: + """Create scan temp dirs under tmp_path and remember them.""" + if kwargs.get("prefix") != "skillspector_": + return real_mkdtemp(*args, **kwargs) + path = real_mkdtemp(*args, **{**kwargs, "dir": tmp_path}) + created.append(Path(path)) + return path + + real_scan_for_source = cli._run_graph_scan_for_source + child_inputs: list[str] = [] + + def scan_for_source(**kwargs: Any) -> dict[str, object]: + """Run the root scan for real and stop the first transitive child.""" + if kwargs["input_path"] == str(archive): + return real_scan_for_source(**kwargs) + child_inputs.append(kwargs["input_path"]) + raise error("child scan stopped") + + monkeypatch.setattr("skillspector.input_handler.tempfile.mkdtemp", recording_mkdtemp) + monkeypatch.setattr(cli, "_run_graph_scan_for_source", scan_for_source) + if error is RuntimeError: + # A child failure is recorded as a warning; make the merge step raise instead. + def stop_merge(*args: Any, **kwargs: Any) -> None: + raise error("merge stopped") + + monkeypatch.setattr(cli, "_ensure_required_failure_events", stop_merge) + + with pytest.raises(error): + cli._scan_skill( + input_path=str(archive), + format=FormatChoice.json, + no_llm=True, + baseline=None, + yara_rules_dir=None, + verbose=False, + show_suppressed=False, + transitive_enabled=True, + transitive_depth=1, + transitive_allow_prefix=None, + transitive_deny_prefix=None, + ) + + assert child_inputs, "the root should have reached a transitive child scan" + assert created + assert not any(path.exists() for path in created) + + @pytest.mark.parametrize( ("terminal", "verbose", "expected"), [(True, False, True), (False, False, False), (True, True, False)], diff --git a/tests/unit/test_mcp_server.py b/tests/unit/test_mcp_server.py index 23cf6733f..1f1b45b1c 100644 --- a/tests/unit/test_mcp_server.py +++ b/tests/unit/test_mcp_server.py @@ -19,8 +19,12 @@ import json import os import sys +import tempfile +import threading +import zipfile from pathlib import Path from types import SimpleNamespace +from typing import Any from unittest.mock import AsyncMock, MagicMock import pytest @@ -59,6 +63,78 @@ async def test_run_scan_returns_structured_verdict( assert result["report"] # non-empty rendered report +def _record_scan_temp_dirs( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> tuple[Path, list[Path]]: + """Build a zipped skill and create the scan temp dirs under tmp_path, recording them.""" + archive = tmp_path / "skill.zip" + with zipfile.ZipFile(archive, "w") as bundle: + bundle.writestr("SKILL.md", "---\nname: demo\ndescription: d\n---\n# Demo\n") + created: list[Path] = [] + real_mkdtemp = tempfile.mkdtemp + + def recording_mkdtemp(*args: Any, **kwargs: Any) -> str: + """Create scan temp dirs under tmp_path and remember them.""" + if kwargs.get("prefix") != "skillspector_": + return real_mkdtemp(*args, **kwargs) + path = real_mkdtemp(*args, **{**kwargs, "dir": tmp_path}) + created.append(Path(path)) + return path + + monkeypatch.setattr("skillspector.input_handler.tempfile.mkdtemp", recording_mkdtemp) + return archive, created + + +async def test_run_scan_that_fails_removes_its_temp_dir( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """A scan failing after input resolution removes the temp dir it made.""" + monkeypatch.setattr(mcp_server, "is_llm_available", lambda: (False, "no llm")) + archive, created = _record_scan_temp_dirs(monkeypatch, tmp_path) + + def fail(*args: Any, **kwargs: Any) -> str: + """Stand in for a failure after the input is materialized.""" + raise RuntimeError("report failed") + + monkeypatch.setattr("skillspector.nodes.report._format_json", fail) + + with pytest.raises(RuntimeError, match="report failed"): + await run_scan(str(archive), use_llm=False, output_format="json") + + assert created + assert not any(path.exists() for path in created) + + +async def test_run_scan_that_is_cancelled_removes_its_temp_dir( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """Cancelling the tool call after input resolution removes the temp dir.""" + monkeypatch.setattr(mcp_server, "is_llm_available", lambda: (False, "no llm")) + archive, created = _record_scan_temp_dirs(monkeypatch, tmp_path) + reporting = threading.Event() + release = threading.Event() + + def block(*args: Any, **kwargs: Any) -> str: + """Hold the report step so the caller can cancel the running scan.""" + reporting.set() + release.wait(timeout=10) + raise RuntimeError("released after cancellation") + + monkeypatch.setattr("skillspector.nodes.report._format_json", block) + + scan = asyncio.create_task(run_scan(str(archive), use_llm=False, output_format="json")) + try: + assert await asyncio.to_thread(reporting.wait, 10) + scan.cancel() + with pytest.raises(asyncio.CancelledError): + await scan + finally: + release.set() + + assert created + assert not any(path.exists() for path in created) + + @pytest.mark.parametrize( ("body", "expect_p6"), [