diff --git a/cycode/cli/apps/ai_guardrails/consts.py b/cycode/cli/apps/ai_guardrails/consts.py index 0898761d..7f75a436 100644 --- a/cycode/cli/apps/ai_guardrails/consts.py +++ b/cycode/cli/apps/ai_guardrails/consts.py @@ -34,6 +34,13 @@ class GuardrailCellMode(str, Enum): BLOCK = GuardrailsMode.BLOCK.value +class McpServerEnforceOn(str, Enum): + """Which MCP servers the unauthorized MCP server guardrail enforces on (its `enforce_on` setting).""" + + UNAUTHORIZED = 'unauthorized' # only servers explicitly marked Unauthorized + NOT_AUTHORIZED = 'not_authorized' # strict: anything that isn't Authorized, servers ASM hasn't seen included + + # Base CLI commands invoked from installed hooks. IDE classes append --ide flags # (and any other suffix) on top of these. CYCODE_SCAN_PROMPT_COMMAND = 'cycode ai-guardrails scan' diff --git a/cycode/cli/apps/ai_guardrails/ides/cursor.py b/cycode/cli/apps/ai_guardrails/ides/cursor.py index 1f41a483..24b3dff9 100644 --- a/cycode/cli/apps/ai_guardrails/ides/cursor.py +++ b/cycode/cli/apps/ai_guardrails/ides/cursor.py @@ -64,6 +64,50 @@ def _load_cursor_mcp_config(config_path: Optional[Path] = None) -> Optional[dict return None +def _cursor_mcp_config_paths(workspace_roots: object) -> list[Path]: + """The MCP configs a Cursor session may load servers from: each project's, then the user's.""" + roots = workspace_roots if isinstance(workspace_roots, list) else [] + paths = [Path(root) / _REPO_SUBDIR / _MCP_CONFIG_FILENAME for root in roots if isinstance(root, str) and root] + paths.append(_cursor_mcp_config_path()) + return paths + + +def _server_command_line(server: dict) -> Optional[str]: + command = server.get('command') + if not isinstance(command, str) or not command: + return None + args = server.get('args') if isinstance(server.get('args'), list) else [] + return ' '.join([command, *(str(arg) for arg in args)]) + + +def _resolve_mcp_server_name(raw_payload: dict) -> Optional[str]: + """The ``mcp.json`` entry name of the server a beforeMCPExecution payload called. + + Cursor identifies the server by its ``url`` (remote servers) or ``command`` (stdio servers) + rather than by the entry name the platform stores servers under, so look the entry up in the + MCP configs: a url matches exactly; a command matches the entry's command line (command + args), + or the entry name itself. Falls back to the raw ``command`` when no entry matches. + """ + url = raw_payload.get('url') + command = raw_payload.get('command') + if not url and not command: + return None + + for config_path in _cursor_mcp_config_paths(raw_payload.get('workspace_roots')): + servers = (_load_cursor_mcp_config(config_path) or {}).get('mcpServers') + if not isinstance(servers, dict): + continue + for name, server in servers.items(): + if not isinstance(server, dict): + continue + if url and server.get('url') == url: + return name + if command and (command == _server_command_line(server) or command.lower() == str(name).lower()): + return name + + return command + + class Cursor(IDE): name: ClassVar[str] = 'cursor' display_name: ClassVar[str] = 'Cursor' @@ -95,7 +139,9 @@ def parse_hook_payload(self, raw_payload: dict) -> AIHookPayload: ide_version=raw_payload.get('cursor_version'), prompt=raw_payload.get('prompt', ''), file_path=raw_payload.get('file_path') or raw_payload.get('path'), - mcp_server_name=raw_payload.get('command'), + mcp_server_name=( + _resolve_mcp_server_name(raw_payload) if canonical_event == AiHookEventType.MCP_EXECUTION else None + ), mcp_tool_name=raw_payload.get('tool_name') or raw_payload.get('tool'), mcp_arguments=(raw_payload.get('arguments') or raw_payload.get('tool_input') or raw_payload.get('input')), ) diff --git a/cycode/cli/apps/ai_guardrails/scan/guardrail_config.py b/cycode/cli/apps/ai_guardrails/scan/guardrail_config.py index 416262b8..695c1f56 100644 --- a/cycode/cli/apps/ai_guardrails/scan/guardrail_config.py +++ b/cycode/cli/apps/ai_guardrails/scan/guardrail_config.py @@ -3,7 +3,7 @@ session-start fetches the tenant's resolved guardrail config from the platform and writes it here; scans only read. Per-agent modes and sensitive-path globs are platform-owned - local policy files never carry them. An absent or corrupt cache means built-in defaults (Report -everywhere + the default globs), always synchronous. +everywhere + the default globs, with the unauthorized MCP server guardrail Off), always synchronous. """ import json @@ -12,7 +12,7 @@ from pathlib import Path from typing import Optional -from cycode.cli.apps.ai_guardrails.consts import GuardrailCellMode, PolicyMode +from cycode.cli.apps.ai_guardrails.consts import GuardrailCellMode, McpServerEnforceOn, PolicyMode from cycode.cli.apps.ai_guardrails.scan.consts import DEFAULT_SENSITIVE_PATH_GLOBS from cycode.cli.apps.ai_guardrails.scan.types import BlockReason from cycode.cli.consts import CYCODE_CONFIGURATION_DIRECTORY @@ -23,7 +23,7 @@ GUARDRAILS_CONFIG_FILE_NAME = 'ai-guardrails-config.json' -_DEFAULT_TTL_SECONDS = 900 +DEFAULT_TTL_SECONDS = 900 # Guardrail keys are the CLI's block-reason vocabulary. Anything else in the payload (a future # guardrail this CLI doesn't implement) is ignored - unknown config must never fail closed. @@ -34,14 +34,26 @@ BlockReason.SECRETS_IN_FILE, BlockReason.SENSITIVE_PATH, BlockReason.SECRETS_IN_MCP_ARGS, + BlockReason.UNAUTHORIZED_MCP_SERVER, ) ) +# Guardrails that are Off unless the platform switches them on for an agent. The rest default to +# Report; these would otherwise start enforcing on a tenant that never looked at them. +_DEFAULT_OFF_GUARDRAIL_KEYS = frozenset((BlockReason.UNAUTHORIZED_MCP_SERVER.value,)) + def get_config_cache_path() -> Path: return Path.home() / CYCODE_CONFIGURATION_DIRECTORY / GUARDRAILS_CONFIG_FILE_NAME +def default_mode_for(guardrail_key: str) -> str: + """A guardrail's mode for an agent the platform sent no cell for (or with no cache at all).""" + if guardrail_key in _DEFAULT_OFF_GUARDRAIL_KEYS: + return GuardrailCellMode.OFF.value + return GuardrailCellMode.REPORT.value + + def _default_sensitive_globs() -> list: return list(DEFAULT_SENSITIVE_PATH_GLOBS) @@ -63,7 +75,12 @@ def __post_init__(self) -> None: def mode_for(self, guardrail_key: str, ide_name: Optional[str]) -> str: """The platform keys the cells by our --ide names, so the lookup is direct.""" agents = (self._guardrails.get(guardrail_key) or {}).get('agents') or {} - return str(agents.get((ide_name or '').lower(), GuardrailCellMode.REPORT.value)).lower() + return str(agents.get((ide_name or '').lower(), default_mode_for(guardrail_key))).lower() + + def is_off_for_every_agent(self, guardrail_key: str) -> bool: + """No agent has this guardrail on - nothing it needs has to be fetched.""" + agents = (self._guardrails.get(guardrail_key) or {}).get('agents') or {} + return all(str(mode).lower() == GuardrailCellMode.OFF for mode in agents.values()) def _modes_for_event(self, event_name: str, ide_name: Optional[str]) -> list: return [ @@ -86,8 +103,16 @@ def sensitive_globs(self) -> list: globs = settings.get('globs') return globs if isinstance(globs, list) and globs else _default_sensitive_globs() + def mcp_server_enforce_on(self) -> str: + """Unknown values read as the default, which enforces on fewer servers - config must never fail closed.""" + settings = (self._guardrails.get(BlockReason.UNAUTHORIZED_MCP_SERVER) or {}).get('settings') or {} + enforce_on = str(settings.get('enforce_on') or '').lower() + if enforce_on == McpServerEnforceOn.NOT_AUTHORIZED: + return McpServerEnforceOn.NOT_AUTHORIZED.value + return McpServerEnforceOn.UNAUTHORIZED.value + def is_expired(self) -> bool: - ttl = self.payload.get('ttl_seconds') or _DEFAULT_TTL_SECONDS + ttl = self.payload.get('ttl_seconds') or DEFAULT_TTL_SECONDS return time.time() - self.fetched_at > ttl def needs_refresh(self, tenant_id: Optional[str]) -> bool: @@ -99,14 +124,15 @@ def apply_platform_config(policy: dict, config: Optional[GuardrailConfig], ide_n """Overlay the platform-owned enforcement config onto the local knobs-only policy. The platform is the only mode source: no cache (cold start) means the built-in defaults - - Report everywhere with the default globs - which equal an unconfigured tenant's platform - config, so behaviour is uniform either way. Each matrix cell lands on its own per-feature - action, so the two FileRead guardrails (content scan vs. sensitive path) keep independent modes. + Report everywhere with the default globs, the unauthorized MCP server guardrail Off - which + equal an unconfigured tenant's platform config, so behaviour is uniform either way. Each matrix + cell lands on its own per-feature action, so the guardrails sharing an event (FileRead: content + scan vs. sensitive path; McpExecution: argument scan vs. server authorization) keep independent modes. An all-Off event never reaches here at all: scan_command skips it. """ def cell(guardrail_key: str) -> str: - return config.mode_for(guardrail_key, ide_name) if config is not None else GuardrailCellMode.REPORT.value + return config.mode_for(guardrail_key, ide_name) if config is not None else default_mode_for(guardrail_key) def action(guardrail_key: str) -> str: return PolicyMode.BLOCK.value if cell(guardrail_key) == GuardrailCellMode.BLOCK else PolicyMode.WARN.value @@ -123,7 +149,14 @@ def action(guardrail_key: str) -> str: ) file_read['path_action'] = action(BlockReason.SENSITIVE_PATH) - policy.setdefault('mcp', {})['action'] = action(BlockReason.SECRETS_IN_MCP_ARGS) + mcp = policy.setdefault('mcp', {}) + mcp['scan_args'] = cell(BlockReason.SECRETS_IN_MCP_ARGS) != GuardrailCellMode.OFF + mcp['action'] = action(BlockReason.SECRETS_IN_MCP_ARGS) + mcp['check_server'] = cell(BlockReason.UNAUTHORIZED_MCP_SERVER) != GuardrailCellMode.OFF + mcp['server_action'] = action(BlockReason.UNAUTHORIZED_MCP_SERVER) + mcp['server_enforce_on'] = ( + config.mcp_server_enforce_on() if config is not None else McpServerEnforceOn.UNAUTHORIZED.value + ) def save_guardrail_config(payload: dict, tenant_id: Optional[str]) -> None: diff --git a/cycode/cli/apps/ai_guardrails/scan/handlers.py b/cycode/cli/apps/ai_guardrails/scan/handlers.py index 54cd92a0..4caeefb6 100644 --- a/cycode/cli/apps/ai_guardrails/scan/handlers.py +++ b/cycode/cli/apps/ai_guardrails/scan/handlers.py @@ -20,8 +20,9 @@ if TYPE_CHECKING: from cycode.cli.apps.ai_guardrails.scan.guardrail_config import GuardrailConfig -from cycode.cli.apps.ai_guardrails.consts import GuardrailsMode, PolicyMode +from cycode.cli.apps.ai_guardrails.consts import GuardrailsMode, McpServerEnforceOn, PolicyMode from cycode.cli.apps.ai_guardrails.ides.base import HookDecision +from cycode.cli.apps.ai_guardrails.scan.mcp_server_status import is_enforced, load_mcp_server_statuses from cycode.cli.apps.ai_guardrails.scan.payload import AIHookPayload from cycode.cli.apps.ai_guardrails.scan.policy import get_policy_value from cycode.cli.apps.ai_guardrails.scan.types import ( @@ -218,6 +219,16 @@ class _ArgScanFeature: deny_agent_message: str ask_message: Callable[[str], str] ask_agent_message: str + scan_enabled: bool = True + + +class _PreScanFinding(NamedTuple): + """A guardrail that fired on the call itself, before its text is scanned (e.g. an unauthorized MCP server).""" + + block_reason: BlockReason + mode: GuardrailsMode + deny_message: str + deny_agent_message: str def _handle_arg_scan( @@ -226,8 +237,14 @@ def _handle_arg_scan( policy: dict, feature: _ArgScanFeature, scan_text: str, + pre_scan_finding: Optional[_PreScanFinding] = None, ) -> HookDecision: - """Shared scan + decision flow for MCP_EXECUTION and COMMAND_EXEC events.""" + """Shared scan + decision flow for MCP_EXECUTION and COMMAND_EXEC events. + + A pre-scan finding in Block mode denies without scanning. In Report mode it marks the event + warned and the scan still runs; a secret the scan then finds takes over the response and the + event's block reason and outcome, since that is what the user is shown. + """ ai_client = ctx.obj['ai_security_client'] max_bytes = get_policy_value(policy, 'secrets', 'max_bytes', default=200000) @@ -240,6 +257,18 @@ def _handle_arg_scan( error_message = None try: + if pre_scan_finding is not None: + block_reason = pre_scan_finding.block_reason + if pre_scan_finding.mode == GuardrailsMode.BLOCK: + outcome = AIHookOutcome.BLOCKED + return HookDecision.deny( + feature.event_type, pre_scan_finding.deny_message, pre_scan_finding.deny_agent_message + ) + outcome = AIHookOutcome.WARNED + + if not feature.scan_enabled: + return HookDecision.allow(feature.event_type) + scan_outcome = _scan_text_for_secrets( ctx, clipped, @@ -283,8 +312,56 @@ def _handle_arg_scan( ) +def _check_mcp_server_authorization(payload: AIHookPayload, policy: dict) -> Optional[_PreScanFinding]: + """The unauthorized MCP server guardrail: a finding when the called server is enforced, else None. + + Fails open - no server name, or no cached statuses, means the server is let through. + """ + mcp_config = get_policy_value(policy, 'mcp', default={}) + if not get_policy_value(mcp_config, 'check_server', default=False): + return None + + alias = payload.mcp_server_name + if not alias: + logger.debug('No MCP server name in the payload; skipping the server authorization check') + return None + + statuses = load_mcp_server_statuses() + if statuses is None: + logger.debug('No cached MCP server statuses; skipping the server authorization check') + return None + + server = statuses.match(alias) + if server is not None and server.alias.lower() != alias.lower(): + # A normalized or plugin-namespaced name: report the alias the platform resolves servers by. + payload.mcp_server_name = server.alias + + enforce_on = get_policy_value(mcp_config, 'server_enforce_on', default=McpServerEnforceOn.UNAUTHORIZED.value) + status = server.status if server is not None else None + if not is_enforced(status, enforce_on): + return None + + logger.debug( + 'MCP server is not authorized, %s', + {'mcp_server_name': payload.mcp_server_name, 'status': status, 'enforce_on': enforce_on}, + ) + server_name = payload.mcp_server_name + return _PreScanFinding( + block_reason=BlockReason.UNAUTHORIZED_MCP_SERVER, + mode=get_effective_mode(mcp_config, action_key='server_action'), + deny_message=( + f"Cycode blocked MCP server '{server_name}': it is not authorized in your organization. " + 'Contact your admin to authorize it.' + ), + deny_agent_message=( + f"The MCP server '{server_name}' is not authorized in this organization. " + 'Do not retry its tools or reach it another way.' + ), + ) + + def handle_before_mcp_execution(ctx: typer.Context, payload: AIHookPayload, policy: dict) -> HookDecision: - """Scan MCP tool arguments for secrets before execution.""" + """Check the MCP server is authorized, then scan the tool arguments for secrets before execution.""" tool = payload.mcp_tool_name or 'unknown' args = payload.mcp_arguments or {} args_text = args if isinstance(args, str) else json.dumps(args) @@ -298,8 +375,10 @@ def handle_before_mcp_execution(ctx: typer.Context, payload: AIHookPayload, poli deny_agent_message='Do not pass secrets to tools. Use secret references (name/id) instead.', ask_message=lambda v: f'Allow MCP tool call "{tool}"? {v}', ask_agent_message='Possible secrets detected in tool arguments; proceed with caution.', + scan_enabled=get_policy_value(policy, 'mcp', 'scan_args', default=True), ), scan_text=args_text, + pre_scan_finding=_check_mcp_server_authorization(payload, policy), ) diff --git a/cycode/cli/apps/ai_guardrails/scan/mcp_server_status.py b/cycode/cli/apps/ai_guardrails/scan/mcp_server_status.py new file mode 100644 index 00000000..b3a5d693 --- /dev/null +++ b/cycode/cli/apps/ai_guardrails/scan/mcp_server_status.py @@ -0,0 +1,157 @@ +"""MCP server authorization status cache, for the unauthorized MCP server guardrail. + +session-start fetches the authorization status of every MCP server the platform ingested from +this user's devices (Inventory -> AI Governance) and writes it here; the pre-MCP-execution hook +only reads it. An absent or corrupt cache means no status is known, so the guardrail fails open. +""" + +import json +import re +import time +from dataclasses import dataclass, field +from enum import Enum +from pathlib import Path +from typing import NamedTuple, Optional + +from cycode.cli.apps.ai_guardrails.consts import McpServerEnforceOn +from cycode.cli.apps.ai_guardrails.scan.guardrail_config import DEFAULT_TTL_SECONDS +from cycode.cli.consts import CYCODE_CONFIGURATION_DIRECTORY +from cycode.cli.utils.path_utils import atomic_write_text, quarantine_corrupt_file +from cycode.logger import get_logger + +logger = get_logger('AI Guardrails') + +MCP_SERVER_STATUSES_FILE_NAME = 'ai-guardrails-mcp-servers.json' + +# Claude Code namespaces a plugin's servers in tool names as `plugin__`, while the +# platform stores the server under its config key - the name the session sweep reported. +_PLUGIN_ALIAS_PREFIX = 'plugin_' + + +class McpServerAuthorizationStatus(str, Enum): + AUTHORIZED = 'Authorized' + UNREVIEWED = 'Unreviewed' + UNAUTHORIZED = 'Unauthorized' + + +# When one alias maps to several servers (e.g. the same name configured differently on two +# devices), the most restrictive status wins. +_RESTRICTIVENESS = { + McpServerAuthorizationStatus.AUTHORIZED: 0, + McpServerAuthorizationStatus.UNREVIEWED: 1, + McpServerAuthorizationStatus.UNAUTHORIZED: 2, +} + + +def get_mcp_server_statuses_cache_path() -> Path: + return Path.home() / CYCODE_CONFIGURATION_DIRECTORY / MCP_SERVER_STATUSES_FILE_NAME + + +def parse_status(raw_status: object) -> McpServerAuthorizationStatus: + """Case-insensitive; an unknown status reads as Unreviewed (no decision was made on the server).""" + lowered = str(raw_status or '').lower() + for status in McpServerAuthorizationStatus: + if status.value.lower() == lowered: + return status + return McpServerAuthorizationStatus.UNREVIEWED + + +def _normalize_alias(alias: str) -> str: + """The form an IDE may turn a config name into inside a tool name (Claude Code keeps [A-Za-z0-9_-]).""" + return re.sub(r'[^a-z0-9_-]', '_', alias.lower()) + + +def is_enforced(status: Optional[McpServerAuthorizationStatus], enforce_on: str) -> bool: + """Whether a server with this status (None: the platform never saw it) is enforced under `enforce_on`.""" + if status == McpServerAuthorizationStatus.UNAUTHORIZED: + return True + return enforce_on == McpServerEnforceOn.NOT_AUTHORIZED and status != McpServerAuthorizationStatus.AUTHORIZED + + +class McpServerMatch(NamedTuple): + alias: str # as the platform stores it + status: McpServerAuthorizationStatus + + +@dataclass +class McpServerStatuses: + servers: list + fetched_at: float + tenant_id: Optional[str] = None + ttl_seconds: float = DEFAULT_TTL_SECONDS + _by_alias: dict = field(init=False, repr=False) + + def __post_init__(self) -> None: + # Lowered alias -> (stored alias, most restrictive status). The platform matches aliases + # case-insensitively, so the CLI does too. + self._by_alias = {} + for server in self.servers: + if not isinstance(server, dict) or not server.get('alias'): + continue + alias = str(server['alias']) + status = parse_status(server.get('status')) + known = self._by_alias.get(alias.lower()) + if known is None or _RESTRICTIVENESS[status] > _RESTRICTIVENESS[known.status]: + self._by_alias[alias.lower()] = McpServerMatch(alias, status) + + def match(self, alias: str) -> Optional[McpServerMatch]: + """The platform's entry for the alias a hook reported, or None when the platform has none.""" + exact = self._by_alias.get(alias.lower()) + if exact is not None: + return exact + + wanted = _normalize_alias(alias) + normalized = [known for key, known in self._by_alias.items() if _normalize_alias(key) == wanted] + if not normalized and wanted.startswith(_PLUGIN_ALIAS_PREFIX): + # The plugin name may itself contain '_', so the server is the longest known suffix. + suffixes = [ + (len(_normalize_alias(key)), known) + for key, known in self._by_alias.items() + if wanted.endswith(f'_{_normalize_alias(key)}') + ] + longest = max((length for length, _ in suffixes), default=None) + normalized = [known for length, known in suffixes if length == longest] + + return max(normalized, key=lambda known: _RESTRICTIVENESS[known.status], default=None) + + def is_expired(self) -> bool: + return time.time() - self.fetched_at > self.ttl_seconds + + def needs_refresh(self, tenant_id: Optional[str]) -> bool: + """Expired, or fetched for another tenant (the user switched tenants since).""" + return self.is_expired() or self.tenant_id != tenant_id + + +def save_mcp_server_statuses(servers: list, tenant_id: Optional[str], ttl_seconds: float = DEFAULT_TTL_SECONDS) -> None: + """Persist fetched statuses; a failed write just leaves the previous cache in place.""" + path = get_mcp_server_statuses_cache_path() + content = {'fetched_at': time.time(), 'tenant_id': tenant_id, 'ttl_seconds': ttl_seconds, 'servers': servers} + try: + path.parent.mkdir(parents=True, exist_ok=True) + atomic_write_text(str(path), json.dumps(content)) + except Exception as e: + logger.debug('Failed to save MCP server statuses cache', exc_info=e) + + +def load_mcp_server_statuses() -> Optional[McpServerStatuses]: + """The cached statuses, or None when the cache is absent or corrupt (quarantined).""" + path = get_mcp_server_statuses_cache_path() + if not path.exists(): + return None + + try: + with open(path, encoding='UTF-8') as file: + content = json.load(file) + servers = content['servers'] + if not isinstance(servers, list): + raise ValueError('servers is not a list') + return McpServerStatuses( + servers=servers, + fetched_at=float(content['fetched_at']), + tenant_id=content.get('tenant_id'), + ttl_seconds=float(content.get('ttl_seconds') or DEFAULT_TTL_SECONDS), + ) + except Exception as e: + logger.warning('MCP server statuses cache is corrupt and will be moved aside', exc_info=e) + quarantine_corrupt_file(str(path)) + return None diff --git a/cycode/cli/apps/ai_guardrails/scan/types.py b/cycode/cli/apps/ai_guardrails/scan/types.py index 5d18e07d..3ffb6280 100644 --- a/cycode/cli/apps/ai_guardrails/scan/types.py +++ b/cycode/cli/apps/ai_guardrails/scan/types.py @@ -40,6 +40,7 @@ class BlockReason(StrEnum): SECRETS_IN_FILE = 'secrets_in_file' SECRETS_IN_MCP_ARGS = 'secrets_in_mcp_args' SENSITIVE_PATH = 'sensitive_path' + UNAUTHORIZED_MCP_SERVER = 'unauthorized_mcp_server' SCAN_FAILURE = 'scan_failure' diff --git a/cycode/cli/apps/ai_guardrails/session_start_command.py b/cycode/cli/apps/ai_guardrails/session_start_command.py index 282fab78..1f1e9b2a 100644 --- a/cycode/cli/apps/ai_guardrails/session_start_command.py +++ b/cycode/cli/apps/ai_guardrails/session_start_command.py @@ -15,7 +15,16 @@ collect_all_skills, get_ide, ) -from cycode.cli.apps.ai_guardrails.scan.guardrail_config import load_guardrail_config, save_guardrail_config +from cycode.cli.apps.ai_guardrails.scan.guardrail_config import ( + DEFAULT_TTL_SECONDS, + load_guardrail_config, + save_guardrail_config, +) +from cycode.cli.apps.ai_guardrails.scan.mcp_server_status import ( + load_mcp_server_statuses, + save_mcp_server_statuses, +) +from cycode.cli.apps.ai_guardrails.scan.types import BlockReason from cycode.cli.apps.ai_guardrails.scan.utils import read_stdin_text, safe_json_parse from cycode.cli.apps.auth.auth_common import get_authorization_info from cycode.cli.apps.auth.auth_manager import AuthManager @@ -78,12 +87,12 @@ def _report_session_context( ai_client: 'AISecurityManagerClient', user_email: Optional[str], tenant_id: Optional[str], -) -> None: +) -> bool: """Report the device + cross-IDE session context to the AI security manager. Never raises. The device context is always reported. MCP configs and skills are collected from every registered IDE, not just the triggering one. Unchanged payloads are skipped via a hash cache - until the TTL expires. + until the TTL expires. Returns whether a report was accepted in this call. """ try: config_files_by_ide, enabled_plugins = collect_all_session_contexts() @@ -105,12 +114,14 @@ def _report_session_context( digest = _session_context_digest(report) if _should_skip_report(digest, tenant_id): logger.debug('Session context unchanged; skipping report') - return + return False if ai_client.report_session_context(**report): _save_report_cache(digest, tenant_id) + return True except Exception as e: logger.debug('Failed to report session context', exc_info=e) + return False def session_start_command( @@ -167,11 +178,14 @@ def session_start_command( logger.debug('Failed to create conversation during session start', exc_info=e) # Report session context (device + cross-IDE MCP servers and plugins) - _report_session_context(ai_client, session_payload.ide_user_email, auth_info.tenant_id) + context_reported = _report_session_context(ai_client, session_payload.ide_user_email, auth_info.tenant_id) # SessionStart precedes the first prompt hook in every IDE, so scans normally find a cache. _sync_guardrail_config(ai_client, auth_info.tenant_id) + # After the report, so the statuses cover the MCP configs the platform just ingested. + _sync_mcp_server_statuses(ai_client, auth_info.tenant_id, force=context_reported) + def _sync_guardrail_config(ai_client: 'AISecurityManagerClient', tenant_id: Optional[str]) -> None: """Refresh the guardrail config cache when it is expired or belongs to another tenant. @@ -188,3 +202,27 @@ def _sync_guardrail_config(ai_client: 'AISecurityManagerClient', tenant_id: Opti if resolved: save_guardrail_config(resolved, tenant_id) logger.debug('Guardrail config cache updated') + + +def _sync_mcp_server_statuses(ai_client: 'AISecurityManagerClient', tenant_id: Optional[str], force: bool) -> None: + """Refresh the MCP server authorization statuses the unauthorized MCP server guardrail reads. + + Fetched only while the guardrail is on for some agent. Refreshed when the cache is expired or + belongs to another tenant, and whenever a changed session context was just reported (``force``): + that report may have brought servers the cached statuses don't cover yet. A failed fetch keeps + the previous cache. + """ + config = load_guardrail_config() + if config is None or config.is_off_for_every_agent(BlockReason.UNAUTHORIZED_MCP_SERVER): + logger.debug('Unauthorized MCP server guardrail is off, skipping MCP server statuses fetch') + return + + cached = load_mcp_server_statuses() + if not force and cached is not None and not cached.needs_refresh(tenant_id): + logger.debug('MCP server statuses cache is fresh, skipping fetch') + return + + servers = ai_client.get_mcp_server_statuses() + if servers is not None: + save_mcp_server_statuses(servers, tenant_id, config.payload.get('ttl_seconds') or DEFAULT_TTL_SECONDS) + logger.debug('MCP server statuses cache updated') diff --git a/cycode/cyclient/ai_security_manager_client.py b/cycode/cyclient/ai_security_manager_client.py index c3a9fb95..8ffea4ef 100644 --- a/cycode/cyclient/ai_security_manager_client.py +++ b/cycode/cyclient/ai_security_manager_client.py @@ -19,6 +19,7 @@ class AISecurityManagerClient: _EVENTS_PATH = 'v4/ai-security/interactions/events' _SESSION_CONTEXT_PATH = 'v4/ai-security/interactions/session-context' _RESOLVED_GUARDRAILS_PATH = 'v4/ai-security/guardrails/resolved' + _MCP_SERVER_STATUSES_PATH = 'v4/ai-security/authorization/mcp/servers' def __init__(self, client: CycodeClientBase, service_config: 'AISecurityManagerServiceConfigBase') -> None: self.client = client @@ -103,6 +104,21 @@ def get_resolved_guardrails(self) -> Optional[dict]: logger.debug('Failed to fetch resolved guardrail config', exc_info=e) return None + def get_mcp_server_statuses(self) -> Optional[list]: + """Fetch the authorization status of every MCP server on the caller's devices. + + Returns the ``[{alias, normalized_id, status}]`` rows, or None when the fetch failed. + """ + try: + response = self.client.get(self._build_endpoint_path(self._MCP_SERVER_STATUSES_PATH)) + servers = response.json().get('servers') + if not isinstance(servers, list): + raise ValueError('servers is not a list') + return servers + except Exception as e: + logger.debug('Failed to fetch MCP server statuses', exc_info=e) + return None + def report_session_context( self, hostname: Optional[str] = None, diff --git a/tests/cli/commands/ai_guardrails/ides/test_cursor.py b/tests/cli/commands/ai_guardrails/ides/test_cursor.py index 0547019e..8ed0f7d5 100644 --- a/tests/cli/commands/ai_guardrails/ides/test_cursor.py +++ b/tests/cli/commands/ai_guardrails/ides/test_cursor.py @@ -49,7 +49,8 @@ def test_parse_file_read_payload() -> None: assert unified.file_path == '/path/to/secret.env' -def test_parse_mcp_execution_payload() -> None: +def test_parse_mcp_execution_payload(fs: FakeFilesystem) -> None: + # No MCP config on disk: the raw command is kept as the server name. args: dict[str, Any] = {'resource_type': 'merge_request', 'parent_id': 'org/repo', 'resource_id': '4'} unified = Cursor().parse_hook_payload( { @@ -66,6 +67,72 @@ def test_parse_mcp_execution_payload() -> None: assert unified.mcp_arguments == args +def _write_mcp_config(fs: FakeFilesystem, path: Path, servers: dict) -> None: + fs.create_file(str(path), contents=json.dumps({'mcpServers': servers})) + + +def _parse_mcp(**fields: Any) -> Any: + return Cursor().parse_hook_payload({'hook_event_name': 'beforeMCPExecution', 'tool_name': 't', **fields}) + + +_USER_MCP_CONFIG = Path.home() / '.cursor' / 'mcp.json' + + +def test_mcp_server_name_resolved_by_url(fs: FakeFilesystem) -> None: + _write_mcp_config( + fs, + _USER_MCP_CONFIG, + {'notion': {'url': 'https://mcp.notion.com/mcp'}, 'other': {'url': 'https://example.com/mcp'}}, + ) + + assert _parse_mcp(url='https://mcp.notion.com/mcp').mcp_server_name == 'notion' + + +def test_mcp_server_name_resolved_by_command_line(fs: FakeFilesystem) -> None: + _write_mcp_config( + fs, + _USER_MCP_CONFIG, + { + 'filesystem': {'command': 'npx', 'args': ['-y', '@modelcontextprotocol/server-filesystem']}, + 'github': {'command': 'npx', 'args': ['-y', '@modelcontextprotocol/server-github']}, + 'local-db': {'command': '/usr/local/bin/db-mcp'}, + }, + ) + + assert _parse_mcp(command='npx -y @modelcontextprotocol/server-github').mcp_server_name == 'github' + assert _parse_mcp(command='/usr/local/bin/db-mcp').mcp_server_name == 'local-db' + + +def test_mcp_server_name_resolved_by_entry_name(fs: FakeFilesystem) -> None: + _write_mcp_config(fs, _USER_MCP_CONFIG, {'gitlab': {'command': '/opt/homebrew/bin/gitlab-mcp'}}) + + assert _parse_mcp(command='GitLab').mcp_server_name == 'gitlab' + + +def test_mcp_server_name_prefers_the_project_config(fs: FakeFilesystem) -> None: + _write_mcp_config(fs, _USER_MCP_CONFIG, {'user-notion': {'url': 'https://mcp.notion.com/mcp'}}) + _write_mcp_config( + fs, Path('/work/repo/.cursor/mcp.json'), {'project-notion': {'url': 'https://mcp.notion.com/mcp'}} + ) + + unified = _parse_mcp(url='https://mcp.notion.com/mcp', workspace_roots=['/work/repo']) + + assert unified.mcp_server_name == 'project-notion' + + +def test_mcp_server_name_without_a_matching_entry_falls_back(fs: FakeFilesystem) -> None: + _write_mcp_config(fs, _USER_MCP_CONFIG, {'notion': {'url': 'https://mcp.notion.com/mcp'}}) + + assert _parse_mcp(command='npx some-unknown-server').mcp_server_name == 'npx some-unknown-server' + assert _parse_mcp(url='https://unknown.example.com/mcp').mcp_server_name is None + + +def test_mcp_server_name_with_a_corrupt_config_falls_back(fs: FakeFilesystem) -> None: + fs.create_file(str(_USER_MCP_CONFIG), contents='not json {') + + assert _parse_mcp(command='GitLab').mcp_server_name == 'GitLab' + + def test_parse_alternative_field_names() -> None: """Cursor's payload has alternative names for some fields.""" fr = Cursor().parse_hook_payload({'hook_event_name': 'beforeReadFile', 'path': '/alt/path.txt'}) diff --git a/tests/cli/commands/ai_guardrails/scan/conftest.py b/tests/cli/commands/ai_guardrails/scan/conftest.py index 86c1b2ae..0c9f5948 100644 --- a/tests/cli/commands/ai_guardrails/scan/conftest.py +++ b/tests/cli/commands/ai_guardrails/scan/conftest.py @@ -10,9 +10,14 @@ def resolved_guardrails_payload( sensitive_path: str = 'Report', mcp: str = 'Report', globs: Optional[list] = None, + mcp_server: Optional[str] = None, + enforce_on: str = 'unauthorized', ) -> dict: - """A platform resolved-config payload with the given per-guardrail modes for the cursor agent.""" - return { + """A platform resolved-config payload with the given per-guardrail modes for the cursor agent. + + The unauthorized MCP server guardrail is only in the payload when ``mcp_server`` is given. + """ + payload = { 'ttl_seconds': 900, 'guardrails': [ {'key': 'secrets_in_prompt', 'event_type': 'Prompt', 'agents': {'cursor': prompt, 'claude-code': 'Block'}}, @@ -26,9 +31,20 @@ def resolved_guardrails_payload( {'key': 'secrets_in_mcp_args', 'event_type': 'McpExecution', 'agents': {'cursor': mcp}}, ], } + if mcp_server is not None: + payload['guardrails'].append( + { + 'key': 'unauthorized_mcp_server', + 'policy_type': 'UnauthorizedAiTools', + 'event_type': 'McpExecution', + 'agents': {'cursor': mcp_server}, + 'settings': {'enforce_on': enforce_on}, + } + ) + return payload -def platform_config(fetched_at: Optional[float] = None, **modes: str) -> GuardrailConfig: +def platform_config(fetched_at: Optional[float] = None, **modes: Optional[str]) -> GuardrailConfig: """A cached platform config; keyword args are the per-guardrail modes (see resolved_guardrails_payload).""" return GuardrailConfig( payload=resolved_guardrails_payload(**modes), diff --git a/tests/cli/commands/ai_guardrails/scan/test_guardrail_config.py b/tests/cli/commands/ai_guardrails/scan/test_guardrail_config.py index efde1d26..e1f0e22a 100644 --- a/tests/cli/commands/ai_guardrails/scan/test_guardrail_config.py +++ b/tests/cli/commands/ai_guardrails/scan/test_guardrail_config.py @@ -82,7 +82,7 @@ def test_unknown_guardrail_keys_are_ignored() -> None: payload = _payload() # A future guardrail this CLI doesn't implement must never affect decisions (fail-open). payload['guardrails'].append( - {'key': 'unauthorized_mcp_server', 'event_type': 'McpExecution', 'agents': {'cursor': 'Block'}} + {'key': 'future_guardrail', 'event_type': 'McpExecution', 'agents': {'cursor': 'Block'}} ) config = GuardrailConfig(payload=payload, fetched_at=time.time()) @@ -151,3 +151,71 @@ def test_apply_platform_config_sensitive_path_off_clears_globs() -> None: assert policy['file_read']['deny_globs'] == [] assert policy['file_read']['scan_content'] is True + + +# --- unauthorized MCP server guardrail --- + + +def test_unauthorized_mcp_server_defaults_to_off() -> None: + # Unlike the other guardrails: a missing entry, or a missing agent cell, must not enable it. + missing_entry = _config() + assert missing_entry.mode_for('unauthorized_mcp_server', 'cursor') == 'off' + assert missing_entry.is_off_for_every_agent('unauthorized_mcp_server') is True + + missing_cell = _config(mcp_server='Block') + assert missing_cell.mode_for('unauthorized_mcp_server', 'codex') == 'off' + assert missing_cell.can_event_block('McpExecution', 'codex') is False + + +def test_unauthorized_mcp_server_block_makes_the_event_blockable() -> None: + config = _config(mcp_server='Block') + + assert config.can_event_block('McpExecution', 'cursor') is True + assert config.is_off_for_every_agent('unauthorized_mcp_server') is False + + +def test_mcp_event_is_off_only_when_both_mcp_guardrails_are_off() -> None: + assert _config(mcp='Off').is_event_off('McpExecution', 'cursor') is True + assert _config(mcp='Off', mcp_server='Off').is_event_off('McpExecution', 'cursor') is True + assert _config(mcp='Off', mcp_server='Report').is_event_off('McpExecution', 'cursor') is False + + +@pytest.mark.parametrize( + ('enforce_on', 'expected'), + [ + ('unauthorized', 'unauthorized'), + ('not_authorized', 'not_authorized'), + ('Not_Authorized', 'not_authorized'), + ('something_new', 'unauthorized'), + ('', 'unauthorized'), + ], +) +def test_mcp_server_enforce_on(enforce_on: str, expected: str) -> None: + assert _config(mcp_server='Report', enforce_on=enforce_on).mcp_server_enforce_on() == expected + + +def test_apply_platform_config_without_cache_leaves_the_mcp_server_check_off() -> None: + policy: dict = {} + + apply_platform_config(policy, None, 'cursor') + + assert policy['mcp']['check_server'] is False + assert policy['mcp']['scan_args'] is True + assert policy['mcp']['server_enforce_on'] == 'unauthorized' + + +def test_apply_platform_config_mcp_server_cells() -> None: + policy: dict = {} + apply_platform_config(policy, _config(mcp='Off', mcp_server='Block', enforce_on='not_authorized'), 'cursor') + + assert policy['mcp']['scan_args'] is False + assert policy['mcp']['check_server'] is True + assert policy['mcp']['server_action'] == 'block' + assert policy['mcp']['server_enforce_on'] == 'not_authorized' + + policy = {} + apply_platform_config(policy, _config(mcp_server='Report'), 'cursor') + + assert policy['mcp']['check_server'] is True + assert policy['mcp']['server_action'] == 'warn' + assert policy['mcp']['action'] == 'warn' diff --git a/tests/cli/commands/ai_guardrails/scan/test_handlers.py b/tests/cli/commands/ai_guardrails/scan/test_handlers.py index 0e4acd8a..3b19ba7d 100644 --- a/tests/cli/commands/ai_guardrails/scan/test_handlers.py +++ b/tests/cli/commands/ai_guardrails/scan/test_handlers.py @@ -1,7 +1,8 @@ """Tests for AI guardrails handlers.""" import os -from typing import Any +import time +from typing import Any, Optional from unittest.mock import MagicMock, patch import pytest @@ -20,6 +21,7 @@ handle_before_read_file, handle_before_submit_prompt, ) +from cycode.cli.apps.ai_guardrails.scan.mcp_server_status import McpServerStatuses from cycode.cli.apps.ai_guardrails.scan.payload import AIHookPayload from cycode.cli.apps.ai_guardrails.scan.types import AiHookEventType, AIHookOutcome, BlockReason from cycode.cli.apps.ai_guardrails.scan.utils import MAX_VIOLATION_DETAIL_LINES, build_violation_summary @@ -544,6 +546,178 @@ def test_handle_before_mcp_execution_with_secrets_warned( assert call_args.args[2] == AIHookOutcome.WARNED +# Tests for the unauthorized MCP server guardrail + + +def _server_check_policy( + default_policy: dict[str, Any], server_action: str = 'block', enforce_on: str = 'unauthorized' +) -> dict[str, Any]: + default_policy['mcp'].update(check_server=True, server_action=server_action, server_enforce_on=enforce_on) + return default_policy + + +def _mcp_payload(server: Optional[str] = 'github') -> AIHookPayload: + return AIHookPayload( + event_name='McpExecution', + conversation_id='conv-1', + ide_provider='claude-code', + mcp_server_name=server, + mcp_tool_name='create_issue', + mcp_arguments={'title': 'hello'}, + ) + + +def _cached_statuses(*rows: tuple[str, str]) -> McpServerStatuses: + return McpServerStatuses(servers=[{'alias': a, 'status': s} for a, s in rows], fetched_at=time.time()) + + +def _reported_event(mock_ctx: MagicMock) -> tuple[AIHookOutcome, Optional[BlockReason]]: + call_args = mock_ctx.obj['ai_security_client'].create_event.call_args + return call_args.args[2], call_args.kwargs['block_reason'] + + +@patch('cycode.cli.apps.ai_guardrails.scan.handlers._scan_text_for_secrets') +@patch('cycode.cli.apps.ai_guardrails.scan.handlers.load_mcp_server_statuses') +def test_unauthorized_mcp_server_block_denies_without_scanning( + mock_statuses: MagicMock, mock_scan: MagicMock, mock_ctx: MagicMock, default_policy: dict[str, Any] +) -> None: + mock_statuses.return_value = _cached_statuses(('github', 'Unauthorized')) + + result = handle_before_mcp_execution(mock_ctx, _mcp_payload(), _server_check_policy(default_policy)) + + assert result.action == DecisionAction.DENY + assert result.user_message == ( + "Cycode blocked MCP server 'github': it is not authorized in your organization. " + 'Contact your admin to authorize it.' + ) + assert "'github'" in result.agent_message + mock_scan.assert_not_called() + assert _reported_event(mock_ctx) == (AIHookOutcome.BLOCKED, BlockReason.UNAUTHORIZED_MCP_SERVER) + + +@patch('cycode.cli.apps.ai_guardrails.scan.handlers._scan_text_for_secrets') +@patch('cycode.cli.apps.ai_guardrails.scan.handlers.load_mcp_server_statuses') +def test_unauthorized_mcp_server_report_allows_warns_and_still_scans( + mock_statuses: MagicMock, mock_scan: MagicMock, mock_ctx: MagicMock, default_policy: dict[str, Any] +) -> None: + mock_statuses.return_value = _cached_statuses(('github', 'Unauthorized')) + mock_scan.return_value = ScanOutcome(scan_id='scan-1') + + result = handle_before_mcp_execution(mock_ctx, _mcp_payload(), _server_check_policy(default_policy, 'warn')) + + assert result == HookDecision.allow(AiHookEventType.MCP_EXECUTION) + mock_scan.assert_called_once() + assert _reported_event(mock_ctx) == (AIHookOutcome.WARNED, BlockReason.UNAUTHORIZED_MCP_SERVER) + + +@patch('cycode.cli.apps.ai_guardrails.scan.handlers._scan_text_for_secrets') +@patch('cycode.cli.apps.ai_guardrails.scan.handlers.load_mcp_server_statuses') +def test_unauthorized_mcp_server_report_then_secret_block_takes_over( + mock_statuses: MagicMock, mock_scan: MagicMock, mock_ctx: MagicMock, default_policy: dict[str, Any] +) -> None: + mock_statuses.return_value = _cached_statuses(('github', 'Unauthorized')) + mock_scan.return_value = ScanOutcome('Found 1 secret: token', 'scan-1', GuardrailsMode.BLOCK) + + result = handle_before_mcp_execution(mock_ctx, _mcp_payload(), _server_check_policy(default_policy, 'warn')) + + assert result.action == DecisionAction.DENY + assert 'Found 1 secret: token' in result.user_message + assert _reported_event(mock_ctx) == (AIHookOutcome.BLOCKED, BlockReason.SECRETS_IN_MCP_ARGS) + + +@pytest.mark.parametrize( + ('rows', 'enforce_on', 'server', 'expected_action'), + [ + # Default enforce_on: only explicitly Unauthorized servers. + ((('github', 'Unreviewed'),), 'unauthorized', 'github', DecisionAction.ALLOW), + ((('github', 'Authorized'),), 'unauthorized', 'github', DecisionAction.ALLOW), + ((), 'unauthorized', 'github', DecisionAction.ALLOW), + # Strict: anything not Authorized, a server the platform never saw included. + ((('github', 'Unreviewed'),), 'not_authorized', 'github', DecisionAction.DENY), + ((), 'not_authorized', 'github', DecisionAction.DENY), + ((('github', 'Authorized'),), 'not_authorized', 'github', DecisionAction.ALLOW), + # One alias, several servers: the most restrictive status wins. + ((('github', 'Authorized'), ('GitHub', 'Unauthorized')), 'unauthorized', 'github', DecisionAction.DENY), + # No server name to check: fail open, even in strict mode. + ((), 'not_authorized', None, DecisionAction.ALLOW), + ], +) +@patch('cycode.cli.apps.ai_guardrails.scan.handlers._scan_text_for_secrets', return_value=ScanOutcome()) +@patch('cycode.cli.apps.ai_guardrails.scan.handlers.load_mcp_server_statuses') +def test_unauthorized_mcp_server_enforcement( + mock_statuses: MagicMock, + mock_scan: MagicMock, + mock_ctx: MagicMock, + default_policy: dict[str, Any], + rows: tuple, + enforce_on: str, + server: Optional[str], + expected_action: DecisionAction, +) -> None: + mock_statuses.return_value = _cached_statuses(*rows) + policy = _server_check_policy(default_policy, enforce_on=enforce_on) + + result = handle_before_mcp_execution(mock_ctx, _mcp_payload(server), policy) + + assert result.action == expected_action + + +@patch('cycode.cli.apps.ai_guardrails.scan.handlers._scan_text_for_secrets', return_value=ScanOutcome()) +@patch('cycode.cli.apps.ai_guardrails.scan.handlers.load_mcp_server_statuses', return_value=None) +def test_unauthorized_mcp_server_without_cached_statuses_fails_open( + mock_statuses: MagicMock, mock_scan: MagicMock, mock_ctx: MagicMock, default_policy: dict[str, Any] +) -> None: + policy = _server_check_policy(default_policy, enforce_on='not_authorized') + + result = handle_before_mcp_execution(mock_ctx, _mcp_payload(), policy) + + assert result == HookDecision.allow(AiHookEventType.MCP_EXECUTION) + assert _reported_event(mock_ctx) == (AIHookOutcome.ALLOWED, None) + + +@patch('cycode.cli.apps.ai_guardrails.scan.handlers._scan_text_for_secrets', return_value=ScanOutcome()) +@patch('cycode.cli.apps.ai_guardrails.scan.handlers.load_mcp_server_statuses') +def test_unauthorized_mcp_server_off_never_reads_statuses( + mock_statuses: MagicMock, mock_scan: MagicMock, mock_ctx: MagicMock, default_policy: dict[str, Any] +) -> None: + # The default policy carries no platform overlay: the check is off. + result = handle_before_mcp_execution(mock_ctx, _mcp_payload(), default_policy) + + assert result == HookDecision.allow(AiHookEventType.MCP_EXECUTION) + mock_statuses.assert_not_called() + + +@patch('cycode.cli.apps.ai_guardrails.scan.handlers._scan_text_for_secrets') +@patch('cycode.cli.apps.ai_guardrails.scan.handlers.load_mcp_server_statuses') +def test_plugin_namespaced_server_is_reported_under_the_platform_alias( + mock_statuses: MagicMock, mock_scan: MagicMock, mock_ctx: MagicMock, default_policy: dict[str, Any] +) -> None: + mock_statuses.return_value = _cached_statuses(('sentry', 'Unauthorized')) + payload = _mcp_payload('plugin_cycode-dev_sentry') + + result = handle_before_mcp_execution(mock_ctx, payload, _server_check_policy(default_policy)) + + assert result.action == DecisionAction.DENY + assert "'sentry'" in result.user_message + assert mock_ctx.obj['ai_security_client'].create_event.call_args.args[0].mcp_server_name == 'sentry' + + +@patch('cycode.cli.apps.ai_guardrails.scan.handlers._scan_text_for_secrets') +@patch('cycode.cli.apps.ai_guardrails.scan.handlers.load_mcp_server_statuses') +def test_args_scan_off_skips_the_scan_when_only_the_server_check_is_on( + mock_statuses: MagicMock, mock_scan: MagicMock, mock_ctx: MagicMock, default_policy: dict[str, Any] +) -> None: + mock_statuses.return_value = _cached_statuses(('github', 'Authorized')) + policy = _server_check_policy(default_policy) + policy['mcp']['scan_args'] = False + + result = handle_before_mcp_execution(mock_ctx, _mcp_payload(), policy) + + assert result == HookDecision.allow(AiHookEventType.MCP_EXECUTION) + mock_scan.assert_not_called() + assert _reported_event(mock_ctx) == (AIHookOutcome.ALLOWED, None) + + def test_get_effective_mode_reads_the_guardrails_action() -> None: assert get_effective_mode({'action': 'block'}) == GuardrailsMode.BLOCK assert get_effective_mode({'action': 'warn'}) == GuardrailsMode.REPORT diff --git a/tests/cli/commands/ai_guardrails/scan/test_mcp_server_status.py b/tests/cli/commands/ai_guardrails/scan/test_mcp_server_status.py new file mode 100644 index 00000000..21437cd3 --- /dev/null +++ b/tests/cli/commands/ai_guardrails/scan/test_mcp_server_status.py @@ -0,0 +1,149 @@ +"""Tests for the MCP server authorization status cache.""" + +import time +from pathlib import Path +from typing import Optional + +import pytest +from pyfakefs.fake_filesystem import FakeFilesystem + +from cycode.cli.apps.ai_guardrails.scan.mcp_server_status import ( + McpServerAuthorizationStatus, + McpServerStatuses, + get_mcp_server_statuses_cache_path, + is_enforced, + load_mcp_server_statuses, + parse_status, + save_mcp_server_statuses, +) + +_AUTHORIZED = McpServerAuthorizationStatus.AUTHORIZED +_UNREVIEWED = McpServerAuthorizationStatus.UNREVIEWED +_UNAUTHORIZED = McpServerAuthorizationStatus.UNAUTHORIZED + + +@pytest.fixture(autouse=True) +def _fake_home(fs: FakeFilesystem, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv('HOME', '/home/testuser') + fs.create_dir('/home/testuser') + + +def _statuses(*rows: tuple[str, str]) -> McpServerStatuses: + return McpServerStatuses( + servers=[{'alias': alias, 'normalized_id': f'id:{alias}', 'status': status} for alias, status in rows], + fetched_at=time.time(), + ) + + +def test_cache_file_sits_next_to_the_guardrail_config() -> None: + assert get_mcp_server_statuses_cache_path() == Path.home() / '.cycode' / 'ai-guardrails-mcp-servers.json' + + +@pytest.mark.parametrize( + ('raw', 'expected'), + [ + ('Authorized', _AUTHORIZED), + ('unauthorized', _UNAUTHORIZED), + ('UNREVIEWED', _UNREVIEWED), + ('Pending', _UNREVIEWED), + (None, _UNREVIEWED), + ], +) +def test_parse_status_is_case_insensitive_and_unknown_reads_unreviewed( + raw: Optional[str], expected: McpServerAuthorizationStatus +) -> None: + assert parse_status(raw) == expected + + +@pytest.mark.parametrize( + ('status', 'enforce_on', 'expected'), + [ + (_UNAUTHORIZED, 'unauthorized', True), + (_UNREVIEWED, 'unauthorized', False), + (_AUTHORIZED, 'unauthorized', False), + (None, 'unauthorized', False), + (_UNAUTHORIZED, 'not_authorized', True), + (_UNREVIEWED, 'not_authorized', True), + (None, 'not_authorized', True), + (_AUTHORIZED, 'not_authorized', False), + ], +) +def test_is_enforced(status: Optional[McpServerAuthorizationStatus], enforce_on: str, expected: bool) -> None: + assert is_enforced(status, enforce_on) is expected + + +def test_match_is_case_insensitive_and_returns_the_stored_alias() -> None: + match = _statuses(('GitHub', 'Unauthorized')).match('github') + + assert match is not None + assert match.alias == 'GitHub' + assert match.status == _UNAUTHORIZED + + +def test_match_unknown_alias_returns_none() -> None: + assert _statuses(('github', 'Authorized')).match('notion') is None + + +def test_the_most_restrictive_status_wins_for_one_alias() -> None: + statuses = _statuses(('github', 'Authorized'), ('github', 'Unauthorized'), ('github', 'Unreviewed')) + assert statuses.match('github').status == _UNAUTHORIZED + + statuses = _statuses(('notion', 'Authorized'), ('Notion', 'Unreviewed')) + assert statuses.match('notion').status == _UNREVIEWED + + +def test_match_normalized_name() -> None: + # Claude Code turns characters outside [A-Za-z0-9_-] into '_' in tool names. + match = _statuses(('my.server', 'Unauthorized')).match('my_server') + + assert match is not None + assert match.alias == 'my.server' + + +def test_match_plugin_namespaced_name_by_longest_server_suffix() -> None: + statuses = _statuses(('sentry', 'Authorized'), ('dev_sentry', 'Unauthorized'), ('other', 'Unauthorized')) + + match = statuses.match('plugin_cycode-dev_sentry') + assert match is not None + assert match.alias == 'sentry' + + match = statuses.match('plugin_cycode_dev_sentry') + assert match is not None + assert match.alias == 'dev_sentry' + + assert statuses.match('plugin_cycode-dev_unknown') is None + + +def test_rows_without_an_alias_are_ignored() -> None: + statuses = McpServerStatuses(servers=[{'status': 'Unauthorized'}, 'garbage', {'alias': ''}], fetched_at=time.time()) + assert statuses.match('') is None + + +def test_save_and_load_round_trip() -> None: + save_mcp_server_statuses([{'alias': 'github', 'status': 'Unauthorized'}], 'tenant-a', ttl_seconds=60) + + statuses = load_mcp_server_statuses() + + assert statuses is not None + assert statuses.match('github').status == _UNAUTHORIZED + assert statuses.ttl_seconds == 60 + assert statuses.needs_refresh('tenant-a') is False + assert statuses.needs_refresh('tenant-b') is True + + +def test_expired_cache_needs_refresh() -> None: + statuses = McpServerStatuses(servers=[], fetched_at=time.time() - 10_000, tenant_id='tenant-a') + assert statuses.needs_refresh('tenant-a') is True + + +def test_load_missing_cache_returns_none() -> None: + assert load_mcp_server_statuses() is None + + +def test_corrupt_cache_is_quarantined(fs: FakeFilesystem) -> None: + path = get_mcp_server_statuses_cache_path() + fs.create_file(str(path), contents='{"servers": {}}') + + assert load_mcp_server_statuses() is None + assert not path.exists() + assert Path(f'{path}.corrupt').exists() diff --git a/tests/cli/commands/ai_guardrails/scan/test_scan_command.py b/tests/cli/commands/ai_guardrails/scan/test_scan_command.py index e60e6047..08c641ad 100644 --- a/tests/cli/commands/ai_guardrails/scan/test_scan_command.py +++ b/tests/cli/commands/ai_guardrails/scan/test_scan_command.py @@ -449,6 +449,55 @@ def test_detach_decision_is_per_event( scan_command(mock_ctx, ide='cursor') mock_respawn.assert_called_once() + @pytest.mark.parametrize( + ('mcp_server', 'expect_detach'), + [(None, True), ('off', True), ('report', True), ('block', False)], + ) + def test_unauthorized_mcp_server_block_keeps_the_mcp_event_synchronous( + self, + mock_ctx: MagicMock, + mocker: MockerFixture, + mock_respawn: MagicMock, + not_detached: None, + mcp_server: Optional[str], + expect_detach: bool, + ) -> None: + config = platform_config(mcp='report', mcp_server=mcp_server) + mocker.patch('cycode.cli.apps.ai_guardrails.scan.scan_command._initialize_clients') + handler = MagicMock(return_value=HookDecision.allow(AiHookEventType.MCP_EXECUTION)) + mocker.patch('cycode.cli.apps.ai_guardrails.scan.scan_command.get_handler_for_event', return_value=handler) + mocker.patch('cycode.cli.apps.ai_guardrails.scan.scan_command.load_guardrail_config', return_value=config) + mocker.patch('cycode.cli.apps.ai_guardrails.scan.scan_command.load_policy', return_value={'fail_open': True}) + + mocker.patch('sys.stdin', StringIO(json.dumps({'hook_event_name': 'beforeMCPExecution', 'tool_name': 't'}))) + scan_command(mock_ctx, ide='cursor') + + assert mock_respawn.called is expect_detach + assert handler.called is not expect_detach + + def test_mcp_event_runs_when_only_the_server_check_is_on( + self, + mock_ctx: MagicMock, + mocker: MockerFixture, + mock_respawn: MagicMock, + not_detached: None, + ) -> None: + config = platform_config(mcp='off', mcp_server='block') + mocker.patch('cycode.cli.apps.ai_guardrails.scan.scan_command._initialize_clients') + handler = MagicMock(return_value=HookDecision.allow(AiHookEventType.MCP_EXECUTION)) + mocker.patch('cycode.cli.apps.ai_guardrails.scan.scan_command.get_handler_for_event', return_value=handler) + mocker.patch('cycode.cli.apps.ai_guardrails.scan.scan_command.load_guardrail_config', return_value=config) + mocker.patch('cycode.cli.apps.ai_guardrails.scan.scan_command.load_policy', return_value={'fail_open': True}) + + mocker.patch('sys.stdin', StringIO(json.dumps({'hook_event_name': 'beforeMCPExecution', 'tool_name': 't'}))) + scan_command(mock_ctx, ide='cursor') + + handler.assert_called_once() + policy = handler.call_args.args[2] + assert policy['mcp']['scan_args'] is False + assert policy['mcp']['check_server'] is True + assert policy['mcp']['server_action'] == 'block' + def test_failed_respawn_falls_back_to_synchronous_scan( self, mock_ctx: MagicMock, diff --git a/tests/cli/commands/ai_guardrails/test_session_start_command.py b/tests/cli/commands/ai_guardrails/test_session_start_command.py index 2ace737a..6eb5b931 100644 --- a/tests/cli/commands/ai_guardrails/test_session_start_command.py +++ b/tests/cli/commands/ai_guardrails/test_session_start_command.py @@ -16,6 +16,7 @@ from cycode.cli.apps.ai_guardrails.ides import copilot as _copilot_mod from cycode.cli.apps.ai_guardrails.ides import cursor as _cursor_mod from cycode.cli.apps.ai_guardrails.scan.guardrail_config import GuardrailConfig +from cycode.cli.apps.ai_guardrails.scan.mcp_server_status import McpServerStatuses from cycode.cli.apps.ai_guardrails.session_start_command import session_start_command @@ -36,6 +37,15 @@ def mock_save_guardrail_config(monkeypatch: pytest.MonkeyPatch) -> MagicMock: return save_mock +@pytest.fixture(autouse=True) +def mock_save_mcp_server_statuses(monkeypatch: pytest.MonkeyPatch) -> MagicMock: + """Keep tests hermetic: never read or write the real MCP server statuses cache.""" + save_mock = MagicMock() + monkeypatch.setattr(_session_start_mod, 'save_mcp_server_statuses', save_mock) + monkeypatch.setattr(_session_start_mod, 'load_mcp_server_statuses', MagicMock(return_value=None)) + return save_mock + + @pytest.fixture(autouse=True) def _isolated_session_context_cache(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: """Keep the dedup cache away from the real ~/.cycode in every test.""" @@ -702,6 +712,133 @@ def test_fresh_guardrail_config_cache_skips_fetch( mock_save_guardrail_config.assert_called_once_with(mock_ai_client.get_resolved_guardrails.return_value, 'tenant-b') +# MCP server statuses (unauthorized MCP server guardrail) + + +def _guardrail_config_with_mcp_server(agents: dict) -> GuardrailConfig: + payload = { + 'ttl_seconds': 600, + 'guardrails': [ + { + 'key': 'unauthorized_mcp_server', + 'event_type': 'McpExecution', + 'agents': agents, + 'settings': {'enforce_on': 'unauthorized'}, + } + ], + } + return GuardrailConfig(payload=payload, fetched_at=time.time(), tenant_id='tenant-a') + + +_SERVERS = [{'alias': 'github', 'normalized_id': 'pkg:gh', 'status': 'Unauthorized'}] + + +@pytest.mark.parametrize( + 'config', + [ + None, + GuardrailConfig(payload={'guardrails': []}, fetched_at=time.time(), tenant_id='tenant-a'), + _guardrail_config_with_mcp_server({'cursor': 'Off', 'claude-code': 'Off'}), + ], +) +def test_mcp_server_statuses_not_fetched_while_the_guardrail_is_off( + config: object, monkeypatch: pytest.MonkeyPatch, mock_save_mcp_server_statuses: MagicMock +) -> None: + monkeypatch.setattr(_session_start_mod, 'load_guardrail_config', MagicMock(return_value=config)) + ai_client = MagicMock() + + _session_start_mod._sync_mcp_server_statuses(ai_client, 'tenant-a', force=True) + + ai_client.get_mcp_server_statuses.assert_not_called() + mock_save_mcp_server_statuses.assert_not_called() + + +def test_mcp_server_statuses_fetched_when_on_for_some_agent( + monkeypatch: pytest.MonkeyPatch, mock_save_mcp_server_statuses: MagicMock +) -> None: + config = _guardrail_config_with_mcp_server({'cursor': 'Off', 'claude-code': 'Report'}) + monkeypatch.setattr(_session_start_mod, 'load_guardrail_config', MagicMock(return_value=config)) + ai_client = MagicMock() + ai_client.get_mcp_server_statuses.return_value = _SERVERS + + _session_start_mod._sync_mcp_server_statuses(ai_client, 'tenant-a', force=False) + + # Cached with the guardrail config's TTL. + mock_save_mcp_server_statuses.assert_called_once_with(_SERVERS, 'tenant-a', 600) + + +def test_fresh_mcp_server_statuses_skip_the_fetch_unless_forced( + monkeypatch: pytest.MonkeyPatch, mock_save_mcp_server_statuses: MagicMock +) -> None: + config = _guardrail_config_with_mcp_server({'cursor': 'Block'}) + cached = McpServerStatuses(servers=[], fetched_at=time.time(), tenant_id='tenant-a') + monkeypatch.setattr(_session_start_mod, 'load_guardrail_config', MagicMock(return_value=config)) + monkeypatch.setattr(_session_start_mod, 'load_mcp_server_statuses', MagicMock(return_value=cached)) + ai_client = MagicMock() + ai_client.get_mcp_server_statuses.return_value = _SERVERS + + _session_start_mod._sync_mcp_server_statuses(ai_client, 'tenant-a', force=False) + ai_client.get_mcp_server_statuses.assert_not_called() + + # Another tenant's cache is refetched. + _session_start_mod._sync_mcp_server_statuses(ai_client, 'tenant-b', force=False) + assert ai_client.get_mcp_server_statuses.call_count == 1 + + # A just-reported session context may have brought new servers: refetch even when fresh. + _session_start_mod._sync_mcp_server_statuses(ai_client, 'tenant-a', force=True) + assert ai_client.get_mcp_server_statuses.call_count == 2 + + +def test_failed_mcp_server_statuses_fetch_keeps_the_cache( + monkeypatch: pytest.MonkeyPatch, mock_save_mcp_server_statuses: MagicMock +) -> None: + config = _guardrail_config_with_mcp_server({'cursor': 'Block'}) + monkeypatch.setattr(_session_start_mod, 'load_guardrail_config', MagicMock(return_value=config)) + ai_client = MagicMock() + ai_client.get_mcp_server_statuses.return_value = None + + _session_start_mod._sync_mcp_server_statuses(ai_client, 'tenant-a', force=True) + + mock_save_mcp_server_statuses.assert_not_called() + + +@patch.object(_session_start_mod, 'collect_all_session_contexts') +@patch.object(_session_start_mod, 'get_ai_security_manager_client') +@patch.object(_session_start_mod, 'get_authorization_info') +def test_session_start_fetches_mcp_server_statuses_after_reporting_the_context( + mock_get_auth: MagicMock, + mock_get_client: MagicMock, + mock_collect: MagicMock, + mock_ctx: MagicMock, + monkeypatch: pytest.MonkeyPatch, + mock_save_mcp_server_statuses: MagicMock, +) -> None: + mock_get_auth.return_value = MagicMock(tenant_id='tenant-a') + mock_collect.return_value = ({'cursor': {'path': '/p', 'content': 'c'}}, {}) + config = _guardrail_config_with_mcp_server({'cursor': 'Block'}) + monkeypatch.setattr(_session_start_mod, 'load_guardrail_config', MagicMock(return_value=config)) + # A fresh statuses cache: only the context report forces the refetch. + cached = McpServerStatuses(servers=[], fetched_at=time.time(), tenant_id='tenant-a') + monkeypatch.setattr(_session_start_mod, 'load_mcp_server_statuses', MagicMock(return_value=cached)) + calls: list = [] + ai_client = MagicMock() + ai_client.report_session_context.side_effect = lambda **_: calls.append('report') or True + ai_client.get_mcp_server_statuses.side_effect = lambda: calls.append('statuses') or _SERVERS + mock_get_client.return_value = ai_client + + with patch('sys.stdin', new=StringIO(json.dumps({'conversation_id': 'conv-1'}))): + session_start_command(mock_ctx, ide='cursor') + + assert calls == ['report', 'statuses'] + mock_save_mcp_server_statuses.assert_called_once_with(_SERVERS, 'tenant-a', 600) + + # The unchanged context is not re-reported, so the fresh cache is kept. + with patch('sys.stdin', new=StringIO(json.dumps({'conversation_id': 'conv-2'}))): + session_start_command(mock_ctx, ide='cursor') + + assert calls == ['report', 'statuses'] + + # Skills reporting diff --git a/tests/cyclient/test_ai_security_manager_client.py b/tests/cyclient/test_ai_security_manager_client.py index 4d37e020..3847960e 100644 --- a/tests/cyclient/test_ai_security_manager_client.py +++ b/tests/cyclient/test_ai_security_manager_client.py @@ -74,3 +74,26 @@ def test_create_event_without_a_conversation_posts_nothing() -> None: client.create_event(AIHookPayload(event_name='Prompt'), AiHookEventType.PROMPT, AIHookOutcome.ALLOWED) http_client.post.assert_not_called() + + +def test_get_mcp_server_statuses_returns_the_server_rows() -> None: + client, http_client = _build_client() + servers = [{'alias': 'github', 'normalized_id': 'pkg:gh', 'status': 'Unauthorized'}] + http_client.get.return_value.json.return_value = {'servers': servers} + + assert client.get_mcp_server_statuses() == servers + http_client.get.assert_called_once_with('v4/ai-security/authorization/mcp/servers') + + +def test_get_mcp_server_statuses_failure_returns_none() -> None: + client, http_client = _build_client() + http_client.get.side_effect = RuntimeError('boom') + + assert client.get_mcp_server_statuses() is None + + +def test_get_mcp_server_statuses_malformed_response_returns_none() -> None: + client, http_client = _build_client() + http_client.get.return_value.json.return_value = {'servers': {'github': 'Unauthorized'}} + + assert client.get_mcp_server_statuses() is None