diff --git a/cycode/cli/apps/ai_guardrails/scan/handlers.py b/cycode/cli/apps/ai_guardrails/scan/handlers.py index 54cd92a0..c4a37b35 100644 --- a/cycode/cli/apps/ai_guardrails/scan/handlers.py +++ b/cycode/cli/apps/ai_guardrails/scan/handlers.py @@ -10,10 +10,9 @@ import json import os +import threading from dataclasses import dataclass -from multiprocessing.pool import ThreadPool -from multiprocessing.pool import TimeoutError as PoolTimeoutError -from typing import TYPE_CHECKING, Callable, NamedTuple, Optional +from typing import TYPE_CHECKING, Any, Callable, NamedTuple, Optional import typer @@ -370,6 +369,13 @@ def _setup_scan_context(ctx: typer.Context) -> typer.Context: return ctx +def _run_scan(scan_func: Callable, documents: list[Document], scan_result: dict[str, Any]) -> None: + try: + scan_result['value'] = scan_func(documents) + except BaseException as e: + scan_result['error'] = e + + def _perform_scan( ctx: typer.Context, documents: list[Document], scan_parameters: dict, timeout_seconds: float ) -> ScanOutcome: @@ -384,15 +390,17 @@ def _perform_scan( ctx, is_git_diff=False, is_commit_range=False, scan_parameters=scan_parameters ) - # Use ThreadPool.apply_async with timeout to abort if scan takes too long - # This uses the same ThreadPool mechanism as run_parallel_batched_scan but with timeout support - with ThreadPool(processes=1) as pool: - result = pool.apply_async(scan_batch_thread_func, (documents,)) - try: - _, error, local_scan_result = result.get(timeout=timeout_seconds) - except PoolTimeoutError: - logger.debug('Scan timed out after %s seconds', timeout_seconds) - raise RuntimeError(f'Scan timed out after {timeout_seconds} seconds') from None + scan_result: dict[str, Any] = {} + scan_thread = threading.Thread(target=_run_scan, args=(scan_batch_thread_func, documents, scan_result), daemon=True) + scan_thread.start() + scan_thread.join(timeout_seconds) + if scan_thread.is_alive(): + logger.debug('Scan timed out after %s seconds', timeout_seconds) + raise RuntimeError(f'Scan timed out after {timeout_seconds} seconds') + if 'error' in scan_result: + raise scan_result['error'] + + _, error, local_scan_result = scan_result['value'] # Check if scan failed - raise exception to trigger fail_open policy if error: diff --git a/cycode/cli/utils/git_proxy.py b/cycode/cli/utils/git_proxy.py index beaafdd0..0aa9c78e 100644 --- a/cycode/cli/utils/git_proxy.py +++ b/cycode/cli/utils/git_proxy.py @@ -1,5 +1,6 @@ import types from abc import ABC, abstractmethod +from functools import cache from typing import TYPE_CHECKING, Optional _GIT_ERROR_MESSAGE = """ @@ -10,11 +11,6 @@ by setting the GIT_PYTHON_GIT_EXECUTABLE= environment variable. """.strip().replace('\n', ' ') -try: - import git -except ImportError: - git = None - if TYPE_CHECKING: from git import PathLike, Repo @@ -23,6 +19,16 @@ class GitProxyError(Exception): pass +# GitPython runs `git version` on import, so it is imported on first use rather than at CLI startup +@cache +def _import_git() -> Optional[types.ModuleType]: + try: + import git + except ImportError: + return None + return git + + class _AbstractGitProxy(ABC): @abstractmethod def get_repo(self, path: Optional['PathLike'] = None, *args, **kwargs) -> 'Repo': ... @@ -52,46 +58,54 @@ def get_git_command_error(self) -> type[BaseException]: class _GitProxy(_AbstractGitProxy): + def __init__(self, git_module: types.ModuleType) -> None: + self._git = git_module + def get_repo(self, path: Optional['PathLike'] = None, *args, **kwargs) -> 'Repo': - return git.Repo(path, *args, **kwargs) + return self._git.Repo(path, *args, **kwargs) def get_null_tree(self) -> object: - return git.NULL_TREE + return self._git.NULL_TREE def get_invalid_git_repository_error(self) -> type[BaseException]: - return git.InvalidGitRepositoryError + return self._git.InvalidGitRepositoryError def get_git_command_error(self) -> type[BaseException]: - return git.GitCommandError + return self._git.GitCommandError def get_git_proxy(git_module: Optional[types.ModuleType]) -> _AbstractGitProxy: - return _GitProxy() if git_module else _DummyGitProxy() + return _GitProxy(git_module) if git_module else _DummyGitProxy() class GitProxyManager(_AbstractGitProxy): """We are using this manager for easy unit testing and mocking of the git module.""" def __init__(self) -> None: - self._git_proxy = get_git_proxy(git) + self._git_proxy: Optional[_AbstractGitProxy] = None + + def _get_git_proxy(self) -> _AbstractGitProxy: + if self._git_proxy is None: + self._git_proxy = get_git_proxy(_import_git()) + return self._git_proxy def _set_dummy_git_proxy(self) -> None: self._git_proxy = _DummyGitProxy() def _set_git_proxy(self) -> None: - self._git_proxy = _GitProxy() + self._git_proxy = _GitProxy(_import_git()) def get_repo(self, path: Optional['PathLike'] = None, *args, **kwargs) -> 'Repo': - return self._git_proxy.get_repo(path, *args, **kwargs) + return self._get_git_proxy().get_repo(path, *args, **kwargs) def get_null_tree(self) -> object: - return self._git_proxy.get_null_tree() + return self._get_git_proxy().get_null_tree() def get_invalid_git_repository_error(self) -> type[BaseException]: - return self._git_proxy.get_invalid_git_repository_error() + return self._get_git_proxy().get_invalid_git_repository_error() def get_git_command_error(self) -> type[BaseException]: - return self._git_proxy.get_git_command_error() + return self._get_git_proxy().get_git_command_error() git_proxy = GitProxyManager() diff --git a/tests/cli/commands/ai_guardrails/scan/test_handlers.py b/tests/cli/commands/ai_guardrails/scan/test_handlers.py index 0e4acd8a..5235f1a3 100644 --- a/tests/cli/commands/ai_guardrails/scan/test_handlers.py +++ b/tests/cli/commands/ai_guardrails/scan/test_handlers.py @@ -1,6 +1,8 @@ """Tests for AI guardrails handlers.""" import os +import threading +from multiprocessing import synchronize from typing import Any from unittest.mock import MagicMock, patch @@ -410,15 +412,51 @@ def test_perform_scan_no_violation_when_all_detections_excluded(mock_ctx: MagicM ) document = Document(path='prompt-content.txt', content='some content', is_git_diff_format=False) - with patch( - 'cycode.cli.apps.ai_guardrails.scan.handlers._get_scan_documents_thread_func', - return_value=lambda batch: ('scan-id-123', None, local_scan_result), + with ( + patch( + 'cycode.cli.apps.ai_guardrails.scan.handlers._get_scan_documents_thread_func', + return_value=lambda batch: ('scan-id-123', None, local_scan_result), + ), + patch.object( + synchronize.SemLock, '__init__', autospec=True, side_effect=synchronize.SemLock.__init__ + ) as mock_semlock_init, ): scan_outcome = _perform_scan(mock_ctx, [document], {}, timeout_seconds=5.0) assert scan_outcome.violation_summary is None assert scan_outcome.scan_id == 'scan-id-123' assert scan_outcome.verdict == GuardrailsMode.BLOCK + mock_semlock_init.assert_not_called() + + +def test_perform_scan_raises_on_timeout_and_on_scan_exception(mock_ctx: MagicMock) -> None: + document = Document(path='prompt-content.txt', content='some content', is_git_diff_format=False) + release_hung_scan = threading.Event() + + def hung_scan(_: list[Document]) -> None: + release_hung_scan.wait() + + def crashing_scan(_: list[Document]) -> None: + raise ValueError('boom') + + try: + with ( + patch( + 'cycode.cli.apps.ai_guardrails.scan.handlers._get_scan_documents_thread_func', return_value=hung_scan + ), + pytest.raises(RuntimeError, match='Scan timed out'), + ): + _perform_scan(mock_ctx, [document], {}, timeout_seconds=0.05) + finally: + release_hung_scan.set() + + with ( + patch( + 'cycode.cli.apps.ai_guardrails.scan.handlers._get_scan_documents_thread_func', return_value=crashing_scan + ), + pytest.raises(ValueError, match='boom'), + ): + _perform_scan(mock_ctx, [document], {}, timeout_seconds=5.0) def _local_scan_result_with_detections(*shas: str) -> LocalScanResult: diff --git a/tests/utils/test_git_proxy.py b/tests/utils/test_git_proxy.py index 62416361..ae31a43a 100644 --- a/tests/utils/test_git_proxy.py +++ b/tests/utils/test_git_proxy.py @@ -1,10 +1,18 @@ import os import tempfile +from unittest.mock import patch import git as real_git import pytest -from cycode.cli.utils.git_proxy import _GIT_ERROR_MESSAGE, GitProxyError, _DummyGitProxy, _GitProxy, get_git_proxy +from cycode.cli.utils.git_proxy import ( + _GIT_ERROR_MESSAGE, + GitProxyError, + GitProxyManager, + _DummyGitProxy, + _GitProxy, + get_git_proxy, +) def test_get_git_proxy() -> None: @@ -15,6 +23,20 @@ def test_get_git_proxy() -> None: assert isinstance(proxy2, _GitProxy) +def test_git_proxy_manager_imports_git_on_first_use_only() -> None: + with patch('cycode.cli.utils.git_proxy._import_git', return_value=real_git) as mock_import_git: + manager = GitProxyManager() + # Importing GitPython runs `git version`, which commands that never touch git shouldn't pay for + mock_import_git.assert_not_called() + + assert manager.get_null_tree() is real_git.NULL_TREE + assert manager.get_git_command_error() is real_git.GitCommandError + mock_import_git.assert_called_once() + + with patch('cycode.cli.utils.git_proxy._import_git', return_value=None): + assert GitProxyManager().get_git_command_error() is GitProxyError + + def test_dummy_git_proxy() -> None: proxy = _DummyGitProxy() @@ -31,7 +53,7 @@ def test_dummy_git_proxy() -> None: def test_git_proxy() -> None: - proxy = _GitProxy() + proxy = _GitProxy(real_git) repo = proxy.get_repo(os.getcwd(), search_parent_directories=True) assert isinstance(repo, real_git.Repo)