Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 20 additions & 12 deletions cycode/cli/apps/ai_guardrails/scan/handlers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand Down
46 changes: 30 additions & 16 deletions cycode/cli/utils/git_proxy.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import types
from abc import ABC, abstractmethod
from functools import cache
from typing import TYPE_CHECKING, Optional

_GIT_ERROR_MESSAGE = """
Expand All @@ -10,11 +11,6 @@
by setting the GIT_PYTHON_GIT_EXECUTABLE=<path/to/git> environment variable.
""".strip().replace('\n', ' ')

try:
import git
except ImportError:
git = None

if TYPE_CHECKING:
from git import PathLike, Repo

Expand All @@ -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': ...
Expand Down Expand Up @@ -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()
44 changes: 41 additions & 3 deletions tests/cli/commands/ai_guardrails/scan/test_handlers.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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:
Expand Down
26 changes: 24 additions & 2 deletions tests/utils/test_git_proxy.py
Original file line number Diff line number Diff line change
@@ -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:
Expand All @@ -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()

Expand All @@ -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)
Expand Down
Loading