diff --git a/mcp_server/changelog.d/mcp-app-tool-failures.added.md b/mcp_server/changelog.d/mcp-app-tool-failures.added.md new file mode 100644 index 0000000000..c51daa6c86 --- /dev/null +++ b/mcp_server/changelog.d/mcp-app-tool-failures.added.md @@ -0,0 +1 @@ +Prowler App tools now report a failure as an MCP tool execution error (`isError: true`, explanation in `content`) instead of as a successful result carrying an `{"error": ...}` object, which clients and models read as a success diff --git a/mcp_server/changelog.d/mcp-compliance-scan-or-provider.changed.md b/mcp_server/changelog.d/mcp-compliance-scan-or-provider.changed.md new file mode 100644 index 0000000000..04eebef825 --- /dev/null +++ b/mcp_server/changelog.d/mcp-compliance-scan-or-provider.changed.md @@ -0,0 +1 @@ +`prowler_get_compliance_framework_state_details` now rejects a call that passes both `scan_id` and `provider_id` instead of silently ignoring the provider, which could report on a scan belonging to a different provider than the one that was asked about diff --git a/mcp_server/prowler_mcp_server/lib/errors.py b/mcp_server/prowler_mcp_server/lib/errors.py index 29e2db6481..9491bf87eb 100644 --- a/mcp_server/prowler_mcp_server/lib/errors.py +++ b/mcp_server/prowler_mcp_server/lib/errors.py @@ -19,10 +19,17 @@ class ProwlerAPIError(Exception): Attributes: status_code: HTTP status the API answered with detail: JSON:API `errors[0].detail`, None when there is none to trust + payload: Parsed JSON body, for a tool that has to read the answer rather + than only report it """ def __init__( - self, message: str, status_code: int, *, detail: str | None = None + self, + message: str, + status_code: int, + *, + detail: str | None = None, + payload: dict[str, Any] | None = None, ) -> None: super().__init__(message) self.status_code: int = status_code @@ -32,6 +39,12 @@ class ProwlerAPIError(Exception): # that must never be repeated to a model -- and None for a 5xx, see # `jsonapi_detail`. self.detail: str | None = detail + # Not every error status means the request failed: Prowler answers 404 + # with the result itself when a query ran and matched nothing. A tool + # reads this to tell such an answer apart from a real failure. It is the + # upstream body, so it is read structurally and never relayed as text -- + # `detail` above is the only part of it that may be repeated to a model. + self.payload: dict[str, Any] | None = payload class ProwlerAPIUnreachable(Exception): @@ -71,6 +84,10 @@ class InvalidArgument(ValueError): """An argument this server rejected before any request went out.""" +class CredentialError(Exception): + """The credential the caller sent is missing, malformed or expired.""" + + # ------------------------------------------------------------------- messages @@ -154,6 +171,16 @@ def _describe_failure(exc: BaseException) -> str | None: "current state before sending it again." ) + if isinstance(exc, CredentialError): + # Not an argument problem, so it is worth saying that plainly: the + # answer is a credential the user has to fix, not another attempt. + return ( + f"This request carried no usable credential: {exc}. Retrying or " + "changing the arguments will not help -- the client has to send an " + "'Authorization: Bearer ' header holding a valid Prowler API " + "key or an unexpired JWT." + ) + if isinstance(exc, ProwlerAPIUnreachable): # The only failure a model can turn into a duplicate write by repeating. return ( diff --git a/mcp_server/prowler_mcp_server/lib/types.py b/mcp_server/prowler_mcp_server/lib/types.py new file mode 100644 index 0000000000..5381b21e8a --- /dev/null +++ b/mcp_server/prowler_mcp_server/lib/types.py @@ -0,0 +1,18 @@ +"""Argument types shared by every tool in this server.""" + +from typing import Annotated + +from pydantic import StringConstraints + +# The identifiers tools take -- a scan UUID, a query id, a Jira project key -- +# are required because there is nothing sensible to do without them. A model +# that does not have one to hand tends to send an empty string rather than omit +# the argument, and an empty string is not caught by "required": it travels into +# a URL path or a request body and comes back as a 404 or an opaque API error +# ("This field may not be blank") that says nothing about which argument was at +# fault. Rejecting it here names the argument instead, and `minLength` puts the +# constraint in the tool schema so a client can see it before calling. +# +# Whitespace is stripped first, so " abc " is accepted as "abc" and " " is +# rejected like "". +NonBlankStr = Annotated[str, StringConstraints(strip_whitespace=True, min_length=1)] diff --git a/mcp_server/prowler_mcp_server/prowler_app/models/attack_paths.py b/mcp_server/prowler_mcp_server/prowler_app/models/attack_paths.py index bbe2eb7401..acaacac540 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/models/attack_paths.py +++ b/mcp_server/prowler_mcp_server/prowler_app/models/attack_paths.py @@ -354,6 +354,14 @@ class AttackPathQueryResult(MinimalSerializerMixin, BaseModel): relationships: list[AttackPathsGraphRelationship] = Field( default_factory=list, description="Relationships connecting the nodes" ) + # A graph with nothing in it serializes to `{}`, since the mixin drops empty + # lists. That reads as an answer that went missing rather than as the finding + # it is -- the query ran and this account has no such attack path -- so the + # empty case carries a sentence saying so. + message: str | None = Field( + default=None, + description="Present only when the query matched nothing, to say the query ran and found no attack path rather than leaving an empty result to interpret", + ) @classmethod def from_api_response( @@ -368,7 +376,15 @@ class AttackPathQueryResult(MinimalSerializerMixin, BaseModel): Returns: AttackPathQueryResult with parsed data and summary """ - attributes = response.get("data", {}).get("attributes") + data = response.get("data") + attributes = data.get("attributes") if data is not None else None + # Prowler spells a graph with nothing in it either as empty lists or as + # a null `attributes`. Both say the same thing -- the query ran and + # matched nothing -- so the null reads as the empty graph it stands for + # instead of crashing the parse. + if attributes is None: + attributes = {} + nodes_data = attributes.get("nodes", []) relationships_data = attributes.get("relationships", []) diff --git a/mcp_server/prowler_mcp_server/prowler_app/models/providers.py b/mcp_server/prowler_mcp_server/prowler_app/models/providers.py index af9509a963..322d85186c 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/models/providers.py +++ b/mcp_server/prowler_mcp_server/prowler_app/models/providers.py @@ -2,7 +2,7 @@ from typing import Any, Literal -from pydantic import BaseModel +from pydantic import BaseModel, ConfigDict, Field from prowler_mcp_server.prowler_app.models.base import MinimalSerializerMixin @@ -104,6 +104,29 @@ class ProvidersListResponse(BaseModel): ) +class ProviderDeletionResult(MinimalSerializerMixin, BaseModel): + """Outcome of a provider deletion. + + Prowler deletes a provider in a background task, so the answer is not always + a finished deletion. A deletion that never started is raised as an error + instead of being reported here: this model only describes a deletion Prowler + accepted and began. + """ + + model_config = ConfigDict(frozen=True) + + status: Literal["deleted", "in_progress"] = Field( + description="Outcome of the deletion: 'deleted' when Prowler finished removing the provider, 'in_progress' when the background task was accepted and is still running, which is normal for a provider with many scans and findings" + ) + task_id: str | None = Field( + default=None, + description="UUIDv4 of the background deletion task, present when the deletion did not finish within the polling window so its state can be checked later", + ) + message: str = Field( + description="Human-readable description of what happened and what to do next" + ) + + class ProviderConnectionStatus(MinimalSerializerMixin, BaseModel): """Result of provider connection operation.""" diff --git a/mcp_server/prowler_mcp_server/prowler_app/models/scans.py b/mcp_server/prowler_mcp_server/prowler_app/models/scans.py index f8eef988ce..3fc095eaa8 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/models/scans.py +++ b/mcp_server/prowler_mcp_server/prowler_app/models/scans.py @@ -191,18 +191,18 @@ class ScansListResponse(BaseModel): class ScanCreationResult(MinimalSerializerMixin, BaseModel): - """Result of scan creation operation. + """Result of a scan creation that succeeded. - Used by trigger_scan() to communicate the outcome of scan creation. - Status indicates whether scan was created successfully or failed. + Used by trigger_scan(). A scan that was not created leaves the tool as an + error instead of being reported here, so this model only ever describes a + scan that exists -- which is why it carries no success flag: a field with + one reachable value says nothing, and inviting a reader to branch on it + suggests there is a failure shape to look for here. There is not; the + failure is the error. """ - scan: DetailedScan | None = Field( - default=None, - description="Detailed scan information if creation succeeded, None otherwise", - ) - status: Literal["success", "failed"] = Field( - description="Outcome of scan creation: success (scan created successfully) or failed (error)" + scan: DetailedScan = Field( + description="Detailed information about the scan that was created" ) message: str = Field( description="Human-readable message describing the scan creation result" @@ -210,13 +210,26 @@ class ScanCreationResult(MinimalSerializerMixin, BaseModel): class ScheduleCreationResult(MinimalSerializerMixin, BaseModel): - """Result of async schedule creation operation. + """Result of a daily schedule creation that succeeded. - Used by schedule_daily_scan() to communicate scheduling outcome. + Used by schedule_daily_scan(). Prowler commits the schedule inside the + request that creates it, so an answer means it exists; a provider that + already has one is refused with a 409 and leaves the tool as an error. That + leaves nothing for a success flag to distinguish, so there is none. """ - scheduled: bool = Field( - description="Whether the daily scan schedule was created successfully" + first_run_state: ( + Literal[ + "available", "scheduled", "executing", "completed", "failed", "cancelled" + ] + | None + ) = Field( + default=None, + description=( + "State of the first scan Prowler starts immediately alongside the schedule. " + "This describes that one run, not the recurring schedule, which stands " + "regardless of it" + ), ) message: str = Field( description="Human-readable message describing the scheduling result" diff --git a/mcp_server/prowler_mcp_server/prowler_app/server.py b/mcp_server/prowler_mcp_server/prowler_app/server.py index e8e854144f..3129a50d09 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/server.py +++ b/mcp_server/prowler_mcp_server/prowler_app/server.py @@ -3,7 +3,7 @@ from fastmcp import FastMCP from prowler_mcp_server.prowler_app.utils.tool_loader import load_all_tools # Initialize MCP server -app_mcp_server = FastMCP("prowler-app") +app_mcp_server = FastMCP("prowler-app", mask_error_details=True) # Auto-discover and load all tools from the tools package load_all_tools(app_mcp_server) diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/attack_paths.py b/mcp_server/prowler_mcp_server/prowler_app/tools/attack_paths.py index 5bd66760fa..b1914dc795 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/attack_paths.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/attack_paths.py @@ -7,8 +7,11 @@ through cloud infrastructure relationships. from typing import Any, Literal +from fastmcp.exceptions import ToolError from pydantic import Field +from prowler_mcp_server.lib.errors import ProwlerAPIError +from prowler_mcp_server.lib.types import NonBlankStr from prowler_mcp_server.prowler_app.models.attack_paths import ( AttackPathCartographySchema, AttackPathQuery, @@ -76,50 +79,47 @@ class AttackPathsTools(BaseTool): 2. Use prowler_list_attack_paths_queries to see available queries for a scan 3. Use prowler_run_attack_paths_query to execute analysis """ - try: - # Validate pagination - self.api_client.validate_page_size(page_size) + # Validate pagination + self.api_client.validate_page_size(page_size) - # Build query parameters - params: dict[str, Any] = { - "page[size]": page_size, - "page[number]": page_number, - } + # Build query parameters + params: dict[str, Any] = { + "page[size]": page_size, + "page[number]": page_number, + } - # Apply provider filters - if provider_id: - params["filter[provider__in]"] = provider_id - if provider_type: - params["filter[provider_type__in]"] = provider_type + # Apply provider filters + if provider_id: + params["filter[provider__in]"] = provider_id + if provider_type: + params["filter[provider_type__in]"] = provider_type - # Apply state filter - if state: - params["filter[state__in]"] = state + # Apply state filter + if state: + params["filter[state__in]"] = state - clean_params = self.api_client.build_filter_params(params) + clean_params = self.api_client.build_filter_params(params) - api_response = await self.api_client.get( - "/attack-paths-scans", params=clean_params - ) - simplified_response = AttackPathScansListResponse.from_api_response( - api_response - ) + api_response = await self.api_client.get( + "/attack-paths-scans", params=clean_params + ) + simplified_response = AttackPathScansListResponse.from_api_response( + api_response + ) - return simplified_response.model_dump() - except Exception as e: - self.logger.error(f"Failed to list attack paths scans: {e}") - return {"error": f"Failed to list attack paths scans: {str(e)}"} + return simplified_response.model_dump() async def list_attack_paths_queries( self, - scan_id: str = Field( - description="UUID of a COMPLETED attack paths scan. Use `prowler_list_attack_paths_scans` with state=['completed'] to find scan IDs" + scan_id: NonBlankStr = Field( + description="UUID of a COMPLETED attack paths scan, as returned by `prowler_list_attack_paths_scans` with state=['completed']. This is NOT a regular scan ID: an ID from `prowler_search_scans` or `prowler_get_scan` names a different resource and is rejected here" ), ) -> list[dict[str, Any]]: """Discover available Attack Paths queries for a completed scan. IMPORTANT: The scan must be in 'completed' state to list queries. - Queries are provider-specific + Attack Paths covers AWS providers only, so only an AWS provider has an + Attack Paths scan to name here, and every query is an AWS one. Each query includes: - id: Query identifier to use with run_attack_paths_query @@ -141,23 +141,32 @@ class AttackPathsTools(BaseTool): api_response = await self.api_client.get( f"/attack-paths-scans/{scan_id}/queries" ) + except ProwlerAPIError as e: + # A 404 here is Prowler failing to resolve `scan_id` to an Attack + # Paths scan, and its own reason for it -- a bare "Not found." -- + # does not say what kind of ID it was looking for. The mistake it + # stands for is a regular scan ID: an Attack Paths scan is a separate + # resource with IDs of its own, and Prowler only creates one for an + # AWS provider, so a scan of any other provider has none to pass. + # + # The endpoint answers 404 for a second thing -- a provider type with + # no query catalog -- but that one cannot happen: a scan only exists + # where Attack Paths runs, which is AWS, and AWS has a catalog. + if e.status_code == 404: + raise self._unknown_scan_error(scan_id) + raise - return [ - AttackPathQuery.from_api_response(query).model_dump() - for query in api_response.get("data", []) - ] - except Exception as e: - self.logger.error( - f"Failed to list attack paths queries for scan {scan_id}: {e}" - ) - return [{"error": f"Failed to list attack paths queries: {str(e)}"}] + return [ + AttackPathQuery.from_api_response(query).model_dump() + for query in api_response.get("data", []) + ] async def run_attack_paths_query( self, - scan_id: str = Field( + scan_id: NonBlankStr = Field( description="UUID of a COMPLETED attack paths scan. The scan must be in 'completed' state" ), - query_id: str = Field( + query_id: NonBlankStr = Field( description="Query ID to execute (e.g., 'aws-internet-exposed-ec2-sensitive-s3-access'). Use `prowler_list_attack_paths_queries` to discover available queries" ), parameters: dict[str, str] = Field( @@ -198,39 +207,61 @@ class AttackPathsTools(BaseTool): 3. Execute this tool with appropriate parameters 4. Analyze the returned graph for security insights """ - try: - # Build the request payload following JSON:API format - request_data: dict[str, Any] = { - "data": { - "type": "attack-paths-query-run-requests", - "attributes": { - "id": query_id, - }, + # Build the request payload following JSON:API format + request_data: dict[str, Any] = { + "data": { + "type": "attack-paths-query-run-requests", + "attributes": { + "id": query_id, }, - } + }, + } - # Add parameters if provided - if parameters: - request_data["data"]["attributes"]["parameters"] = parameters + # Add parameters if provided + if parameters: + request_data["data"]["attributes"]["parameters"] = parameters + try: api_response = await self.api_client.post( f"/attack-paths-scans/{scan_id}/queries/run", json_data=request_data, ) + except ProwlerAPIError as e: + # Prowler answers a query that matched nothing with 404 and the empty + # result as the body. That is an answer -- this account has no such + # attack path, which is the good outcome -- so it is returned rather + # than raised: reporting it as a failure invites a retry of a call + # whose arguments were right, and hides a clean result. + if e.status_code == 404 and isinstance(e.payload, dict): + if "data" in e.payload: + api_response = e.payload + else: + # No result body, so `scan_id` did not resolve to an Attack + # Paths scan. An unknown query_id is a 400, not this. + raise self._unknown_scan_error(scan_id) + else: + raise - # Parse the response - query_result = AttackPathQueryResult.from_api_response(api_response) + # Parse the response + query_result = AttackPathQueryResult.from_api_response(api_response) - return query_result.model_dump() - except Exception as e: - self.logger.error( - f"Failed to run attack paths query '{query_id}' on scan {scan_id}: {e}" + if not query_result.nodes: + query_result = query_result.model_copy( + update={ + "message": ( + f"The query '{query_id}' ran against scan {scan_id} and matched " + "nothing, so this provider has no attack path of that shape. " + "The scan and the query ID were both valid; running it again " + "will return the same thing." + ) + } ) - return {"error": f"Failed to run attack paths query '{query_id}': {str(e)}"} + + return query_result.model_dump() async def get_attack_paths_cartography_schema( self, - scan_id: str = Field( + scan_id: NonBlankStr = Field( description="UUID of a COMPLETED attack paths scan. Use `prowler_list_attack_paths_scans` with state=['completed'] to find scan IDs" ), ) -> dict[str, Any]: @@ -262,18 +293,43 @@ class AttackPathsTools(BaseTool): api_response = await self.api_client.get( f"/attack-paths-scans/{scan_id}/schema" ) + except ProwlerAPIError as e: + # Two 404s again, told apart by whether Prowler wrote a JSON:API + # error. Absent means the scan resolved and its graph simply records + # no Cartography module, so the ID is not the thing to change. + if e.status_code == 404: + if e.detail is None: + raise ToolError( + f"Scan {scan_id} has no Cartography schema recorded, so there is " + "nothing to write custom queries against. Use " + "prowler_list_attack_paths_queries for the ready-made queries of " + "this scan, which do not need the schema." + ) + else: + raise self._unknown_scan_error(scan_id) + raise - schema = AttackPathCartographySchema.from_api_response(api_response) + schema = AttackPathCartographySchema.from_api_response(api_response) - schema_content = await self.api_client.fetch_external_url( - schema.raw_schema_url - ) + schema_content = await self.api_client.fetch_external_url(schema.raw_schema_url) - return schema.model_copy( - update={"schema_content": schema_content} - ).model_dump() - except Exception as e: - self.logger.error( - f"Failed to get cartography schema for scan {scan_id}: {e}" - ) - return {"error": f"Failed to get cartography schema: {str(e)}"} + return schema.model_copy(update={"schema_content": schema_content}).model_dump() + + # Private helper methods + + @staticmethod + def _unknown_scan_error(scan_id: str) -> ToolError: + """Describe a scan ID Prowler could not resolve to an Attack Paths scan. + + Returns: + The ``ToolError`` for the caller to raise. Built without a ``from`` + clause on purpose: the sentence is the final word, not a wrapper + around the API's. + """ + return ToolError( + f"Prowler has no Attack Paths scan with ID {scan_id}. These are a " + "different resource from regular scans and only exist for AWS " + "providers, so an ID from prowler_search_scans or prowler_get_scan " + "never resolves here. Use prowler_list_attack_paths_scans to get an " + "ID these tools take." + ) diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/compliance.py b/mcp_server/prowler_mcp_server/prowler_app/tools/compliance.py index 33cdd22a69..b7f06be9a8 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/compliance.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/compliance.py @@ -6,8 +6,11 @@ across all cloud providers. from typing import Any +from fastmcp.exceptions import ToolError from pydantic import Field +from prowler_mcp_server.lib.errors import InvalidArgument +from prowler_mcp_server.lib.types import NonBlankStr from prowler_mcp_server.prowler_app.models.compliance import ( ComplianceFrameworksListResponse, ComplianceRequirementAttributesListResponse, @@ -34,7 +37,7 @@ class ComplianceTools(BaseTool): The scan_id of the latest completed scan for the provider. Raises: - ValueError: If no completed scans are found for the provider. + ToolError: If no completed scans are found for the provider """ scan_params = { "filter[provider]": provider_id, @@ -48,7 +51,7 @@ class ComplianceTools(BaseTool): scans_data = scans_response.get("data", []) if not scans_data: - raise ValueError( + raise ToolError( f"No completed scans found for provider {provider_id}. " "Run a scan first using prowler_trigger_scan." ) @@ -93,18 +96,15 @@ class ComplianceTools(BaseTool): 2. Use prowler_get_compliance_framework_state_details with a specific compliance_id to see which requirements failed """ if not scan_id and not provider_id: - return { - "error": "Either scan_id or provider_id must be provided. Use prowler_search_providers to find provider IDs or prowler_list_scans to find scan IDs." - } + raise InvalidArgument( + "Either scan_id or provider_id must be provided. Use prowler_search_providers to find provider IDs or prowler_list_scans to find scan IDs." + ) elif scan_id and provider_id: - return { - "error": "Provide either scan_id or provider_id, not both. To get compliance data for a specific scan, use scan_id. To get data for the latest scan of a provider, use provider_id." - } + raise InvalidArgument( + "Provide either scan_id or provider_id, not both. To get compliance data for a specific scan, use scan_id. To get data for the latest scan of a provider, use provider_id." + ) elif not scan_id and provider_id: - try: - scan_id = await self._get_latest_scan_id_for_provider(provider_id) - except ValueError as e: - return {"error": str(e)} + scan_id = await self._get_latest_scan_id_for_provider(provider_id) params: dict[str, Any] = {"filter[scan_id]": scan_id} @@ -253,16 +253,16 @@ class ComplianceTools(BaseTool): async def get_compliance_framework_state_details( self, - compliance_id: str = Field( + compliance_id: NonBlankStr = Field( description="Compliance framework ID to get details for (e.g., 'cis_1.5_aws', 'pci_dss_v4.0_aws'). You can get compliance IDs from prowler_get_compliance_overview or consulting Prowler Hub/Prowler Documentation that you can also find in form of tools in this MCP Server", ), scan_id: str | None = Field( default=None, - description="UUID of a specific scan to get compliance data for. Required if provider_id is not specified.", + description="UUID of a specific scan to get compliance data for. Required if provider_id is not specified. Do not pass it together with provider_id.", ), provider_id: str | None = Field( default=None, - description="Prowler's internal UUID (v4) for a specific provider. If provided without scan_id, the tool will automatically find the latest completed scan for this provider. Use `prowler_search_providers` tool to find provider IDs.", + description="Prowler's internal UUID (v4) for a specific provider. The tool will automatically find the latest completed scan for this provider. Use `prowler_search_providers` tool to find provider IDs. Do not pass it together with scan_id.", ), ) -> dict[str, Any]: """Get detailed requirement-level breakdown for a specific compliance framework. @@ -283,8 +283,8 @@ class ComplianceTools(BaseTool): - Use prowler_get_finding_details with these finding IDs for more details and remediation guidance Default behavior: - - Requires either scan_id OR provider_id - - With provider_id (no scan_id): Automatically finds the latest completed scan for that provider + - Requires exactly one of scan_id OR provider_id; providing both is rejected + - With provider_id: Automatically finds the latest completed scan for that provider - With scan_id: Uses that specific scan's compliance data - Only shows failed requirements with their associated failed finding IDs @@ -293,21 +293,22 @@ class ComplianceTools(BaseTool): 2. Use this tool with the compliance_id to see failed requirements and their finding IDs 3. Use prowler_get_finding_details with the finding IDs to get remediation guidance """ - # Validate that either scan_id or provider_id is provided + # Exactly one of the two: taking scan_id and ignoring provider_id would + # answer for whatever provider that scan belongs to, which is not + # necessarily the one the caller named. if not scan_id and not provider_id: - return { - "error": "Either scan_id or provider_id must be provided. Use prowler_search_providers to find provider IDs or prowler_list_scans to find scan IDs." - } + raise InvalidArgument( + "Either scan_id or provider_id must be provided. Use prowler_search_providers to find provider IDs or prowler_list_scans to find scan IDs." + ) + elif scan_id and provider_id: + raise InvalidArgument( + "Provide either scan_id or provider_id, not both. To get compliance data for a specific scan, use scan_id. To get data for the latest scan of a provider, use provider_id." + ) # Resolve provider_id to latest scan_id if needed resolved_scan_id = scan_id if not scan_id and provider_id: - try: - resolved_scan_id = await self._get_latest_scan_id_for_provider( - provider_id - ) - except ValueError as e: - return {"error": str(e)} + resolved_scan_id = await self._get_latest_scan_id_for_provider(provider_id) # Build params for requirements endpoint params: dict[str, Any] = { diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/finding_groups.py b/mcp_server/prowler_mcp_server/prowler_app/tools/finding_groups.py index 05adf8db2b..54bc3bcbcd 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/finding_groups.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/finding_groups.py @@ -6,8 +6,10 @@ This module provides read-only tools for finding group triage and drill-downs. from typing import Any, Literal from urllib.parse import quote +from fastmcp.exceptions import ToolError from pydantic import Field +from prowler_mcp_server.lib.types import NonBlankStr from prowler_mcp_server.prowler_app.models.finding_groups import ( DetailedFindingGroup, FindingGroupResourcesListResponse, @@ -236,50 +238,46 @@ class FindingGroupsTools(BaseTool): prowler_get_finding_group_details for complete counters or prowler_list_finding_group_resources to drill into affected resources. """ - try: - self.api_client.validate_page_size(page_size) - date_range, params = self._base_date_params(date_from, date_to) - endpoint = self._group_endpoint(date_range) + self.api_client.validate_page_size(page_size) + date_range, params = self._base_date_params(date_from, date_to) + endpoint = self._group_endpoint(date_range) - self._apply_common_filters( - params, - provider, - provider_type, - provider_uid, - provider_alias, - region, - service, - resource_type, - resource_name, - resource_uid, - resource_group, - category, - check_id, - check_title, - severity, - status, - muted, - delta, - ) + self._apply_common_filters( + params, + provider, + provider_type, + provider_uid, + provider_alias, + region, + service, + resource_type, + resource_name, + resource_uid, + resource_group, + category, + check_id, + check_title, + severity, + status, + muted, + delta, + ) - params["filter[include_muted]"] = self._bool_value(include_muted) - params["page[size]"] = page_size - params["page[number]"] = page_number - params["fields[finding-groups]"] = GROUP_LIST_FIELDS - if sort: - params["sort"] = sort + params["filter[include_muted]"] = self._bool_value(include_muted) + params["page[size]"] = page_size + params["page[number]"] = page_number + params["fields[finding-groups]"] = GROUP_LIST_FIELDS + if sort: + params["sort"] = sort - clean_params = self.api_client.build_filter_params(params) - api_response = await self.api_client.get(endpoint, params=clean_params) - response = FindingGroupsListResponse.from_api_response(api_response) - return response.model_dump() - except Exception as e: - self.logger.error(f"Error listing finding groups: {e}") - return {"error": str(e), "status": "failed"} + clean_params = self.api_client.build_filter_params(params) + api_response = await self.api_client.get(endpoint, params=clean_params) + response = FindingGroupsListResponse.from_api_response(api_response) + return response.model_dump() async def get_finding_group_details( self, - check_id: str = Field( + check_id: NonBlankStr = Field( description="Public check ID that identifies the finding group. This is not a UUID." ), date_from: str | None = Field( @@ -297,39 +295,37 @@ class FindingGroupsTools(BaseTool): or historical data when dates are provided. Fully muted groups are included by default so accepted risk does not look like a missing group. """ - try: - date_range, params = self._base_date_params(date_from, date_to) - endpoint = self._group_endpoint(date_range) + date_range, params = self._base_date_params(date_from, date_to) + endpoint = self._group_endpoint(date_range) - params.update( - { - "filter[check_id]": check_id, - "filter[include_muted]": True, - "page[size]": 1, - "page[number]": 1, - "fields[finding-groups]": GROUP_DETAIL_FIELDS, - } + params.update( + { + "filter[check_id]": check_id, + "filter[include_muted]": True, + "page[size]": 1, + "page[number]": 1, + "fields[finding-groups]": GROUP_DETAIL_FIELDS, + } + ) + + clean_params = self.api_client.build_filter_params(params) + api_response = await self.api_client.get(endpoint, params=clean_params) + data = api_response.get("data", []) + + if not data: + # No `from`: this names the check and the tool that lists valid ones, + # neither of which the shared classifier can know. + raise ToolError( + f"No finding group exists for check '{check_id}' in this scan. Use " + "prowler_list_finding_groups to see the checks that have findings." ) - clean_params = self.api_client.build_filter_params(params) - api_response = await self.api_client.get(endpoint, params=clean_params) - data = api_response.get("data", []) - - if not data: - return { - "error": f"Finding group '{check_id}' not found.", - "status": "not_found", - } - - group = DetailedFindingGroup.from_api_response(data[0]) - return group.model_dump() - except Exception as e: - self.logger.error(f"Error getting finding group details: {e}") - return {"error": str(e), "status": "failed"} + group = DetailedFindingGroup.from_api_response(data[0]) + return group.model_dump() async def list_finding_group_resources( self, - check_id: str = Field( + check_id: NonBlankStr = Field( description="Public check ID that identifies the finding group. This is not a UUID." ), provider: list[str] = Field( @@ -426,45 +422,41 @@ class FindingGroupsTools(BaseTool): `finding_id`. Use `prowler_get_finding_details(finding_id)` to retrieve complete remediation guidance for a specific resource finding. """ - try: - self.api_client.validate_page_size(page_size) - date_range, params = self._base_date_params(date_from, date_to) - endpoint = self._resource_endpoint(check_id, date_range) + self.api_client.validate_page_size(page_size) + date_range, params = self._base_date_params(date_from, date_to) + endpoint = self._resource_endpoint(check_id, date_range) - if muted is None and not self._bool_value(include_muted): - muted = False + if muted is None and not self._bool_value(include_muted): + muted = False - self._apply_common_filters( - params, - provider, - provider_type, - provider_uid, - provider_alias, - region, - service, - resource_type, - resource_name, - resource_uid, - resource_group, - category, - [], - None, - severity, - status, - muted, - delta, - ) + self._apply_common_filters( + params, + provider, + provider_type, + provider_uid, + provider_alias, + region, + service, + resource_type, + resource_name, + resource_uid, + resource_group, + category, + [], + None, + severity, + status, + muted, + delta, + ) - params["page[size]"] = page_size - params["page[number]"] = page_number - params["fields[finding-group-resources]"] = RESOURCE_FIELDS - if sort: - params["sort"] = sort + params["page[size]"] = page_size + params["page[number]"] = page_number + params["fields[finding-group-resources]"] = RESOURCE_FIELDS + if sort: + params["sort"] = sort - clean_params = self.api_client.build_filter_params(params) - api_response = await self.api_client.get(endpoint, params=clean_params) - response = FindingGroupResourcesListResponse.from_api_response(api_response) - return response.model_dump() - except Exception as e: - self.logger.error(f"Error listing finding group resources: {e}") - return {"error": str(e), "status": "failed"} + clean_params = self.api_client.build_filter_params(params) + api_response = await self.api_client.get(endpoint, params=clean_params) + response = FindingGroupResourcesListResponse.from_api_response(api_response) + return response.model_dump() diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/findings.py b/mcp_server/prowler_mcp_server/prowler_app/tools/findings.py index b556101cab..860dfa806f 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/findings.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/findings.py @@ -8,6 +8,7 @@ from typing import Any, Literal from pydantic import Field +from prowler_mcp_server.lib.types import NonBlankStr from prowler_mcp_server.prowler_app.models.findings import ( DetailedFinding, FindingsListResponse, @@ -180,7 +181,7 @@ class FindingsTools(BaseTool): async def get_finding_details( self, - finding_id: str = Field( + finding_id: NonBlankStr = Field( description="UUID of the finding to retrieve (must be a valid UUID format, e.g., '019ac0d6-90d5-73e9-9acf-c22e256f1bac'). Returns an error if the finding ID is invalid or not found." ), ) -> dict[str, Any]: diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/integrations.py b/mcp_server/prowler_mcp_server/prowler_app/tools/integrations.py index d0aac0182a..f458a2ee5f 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/integrations.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/integrations.py @@ -9,8 +9,11 @@ This module provides tools for managing where Prowler sends its results, includi import json from typing import Any +from fastmcp.exceptions import ToolError from pydantic import Field +from prowler_mcp_server.lib.errors import CredentialError, InvalidArgument +from prowler_mcp_server.lib.types import NonBlankStr from prowler_mcp_server.prowler_app.models.integrations import ( DetailedIntegration, IntegrationConnectionStatus, @@ -126,7 +129,7 @@ class IntegrationsTools(BaseTool): async def get_integration( self, - integration_id: str = Field( + integration_id: NonBlankStr = Field( description="UUID of the integration to retrieve. Must be a valid UUID format (e.g., '019ac0d6-90d5-73e9-9acf-c22e256f1bac'). Use prowler_list_integrations to find it." ), ) -> dict[str, Any]: @@ -157,7 +160,7 @@ class IntegrationsTools(BaseTool): async def create_amazon_s3_integration( self, - bucket_name: str = Field( + bucket_name: NonBlankStr = Field( description="Name of the S3 bucket where Prowler will upload the scan outputs (CSV, HTML, OCSF JSON and compliance reports)." ), output_directory: str = Field( @@ -168,15 +171,15 @@ class IntegrationsTools(BaseTool): default=[], description="Prowler UUIDs of the providers whose scan outputs are exported to this bucket. Use prowler_search_providers to find them. Leave empty to attach no provider yet.", ), - role_arn: str | None = Field( + role_arn: NonBlankStr | None = Field( default=None, description="ARN of the IAM role Prowler assumes to write to the bucket (e.g. 'arn:aws:iam::123456789012:role/ProwlerS3Integration'). Recommended over static keys.", ), - external_id: str | None = Field( + external_id: NonBlankStr | None = Field( default=None, description="External ID required by the trust policy of the assumed role. In Prowler Cloud this is the tenant ID.", ), - role_session_name: str | None = Field( + role_session_name: NonBlankStr | None = Field( default=None, description="Identifier for the role session, useful to track it in AWS logs. Only letters, digits and the characters =,.@_- are allowed.", ), @@ -184,15 +187,15 @@ class IntegrationsTools(BaseTool): default=3600, description="Duration of the assumed role session in seconds. Must be between 900 and 43200. Defaults to 3600 when omitted.", ), - aws_access_key_id: str | None = Field( + aws_access_key_id: NonBlankStr | None = Field( default=None, description="AWS access key ID. Only needed when the Prowler deployment has no ambient AWS credentials.", ), - aws_secret_access_key: str | None = Field( + aws_secret_access_key: NonBlankStr | None = Field( default=None, description="AWS secret access key. Required when 'aws_access_key_id' is provided.", ), - aws_session_token: str | None = Field( + aws_session_token: NonBlankStr | None = Field( default=None, description="AWS session token, only for temporary credentials.", ), @@ -244,34 +247,30 @@ class IntegrationsTools(BaseTool): """ self.logger.info(f"Creating Amazon S3 integration for bucket {bucket_name}...") - try: - credentials = self._build_aws_credentials( - role_arn=role_arn, - external_id=external_id, - role_session_name=role_session_name, - session_duration=session_duration, - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, - ) + credentials = self._build_aws_credentials( + role_arn=role_arn, + external_id=external_id, + role_session_name=role_session_name, + session_duration=session_duration, + aws_access_key_id=aws_access_key_id, + aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, + ) - return await self._create_integration( - integration_type="amazon_s3", - configuration={ - "bucket_name": bucket_name, - "output_directory": output_directory, - }, - credentials=credentials, - provider_ids=provider_ids, - enabled=enabled, - ) - except Exception as e: - self.logger.error(f"Amazon S3 integration creation failed: {e}") - return {"error": str(e), "status": "failed"} + return await self._create_integration( + integration_type="amazon_s3", + configuration={ + "bucket_name": bucket_name, + "output_directory": output_directory, + }, + credentials=credentials, + provider_ids=provider_ids, + enabled=enabled, + ) async def create_aws_security_hub_integration( self, - provider_id: str = Field( + provider_id: NonBlankStr = Field( description="Prowler UUID of the AWS provider whose findings are sent to Security Hub. It must be an AWS provider, and it can only have one Security Hub integration. Use prowler_search_providers with provider_type=['aws'] to find it." ), send_only_fails: bool = Field( @@ -282,15 +281,15 @@ class IntegrationsTools(BaseTool): default=False, description="When true, findings that are no longer present in the latest scan are archived in Security Hub.", ), - role_arn: str | None = Field( + role_arn: NonBlankStr | None = Field( default=None, description="ARN of a dedicated IAM role Prowler assumes to write to Security Hub. Leave every credential parameter empty to reuse the credentials already stored for the provider, which is the recommended setup.", ), - external_id: str | None = Field( + external_id: NonBlankStr | None = Field( default=None, description="External ID required by the trust policy of the assumed role.", ), - role_session_name: str | None = Field( + role_session_name: NonBlankStr | None = Field( default=None, description="Identifier for the role session, useful to track it in AWS logs. Only letters, digits and the characters =,.@_- are allowed.", ), @@ -298,14 +297,14 @@ class IntegrationsTools(BaseTool): default=None, description="Duration of the assumed role session in seconds. Must be between 900 and 43200. Defaults to 3600 when omitted.", ), - aws_access_key_id: str | None = Field( + aws_access_key_id: NonBlankStr | None = Field( default=None, description="AWS access key ID for dedicated credentials." ), - aws_secret_access_key: str | None = Field( + aws_secret_access_key: NonBlankStr | None = Field( default=None, description="AWS secret access key. Required when 'aws_access_key_id' is provided.", ), - aws_session_token: str | None = Field( + aws_session_token: NonBlankStr | None = Field( default=None, description="AWS session token, only for temporary credentials.", ), @@ -344,40 +343,36 @@ class IntegrationsTools(BaseTool): f"Creating AWS Security Hub integration for provider {provider_id}..." ) - try: - credentials = self._build_aws_credentials( - role_arn=role_arn, - external_id=external_id, - role_session_name=role_session_name, - session_duration=session_duration, - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, - ) + credentials = self._build_aws_credentials( + role_arn=role_arn, + external_id=external_id, + role_session_name=role_session_name, + session_duration=session_duration, + aws_access_key_id=aws_access_key_id, + aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, + ) - return await self._create_integration( - integration_type="aws_security_hub", - configuration={ - "send_only_fails": send_only_fails, - "archive_previous_findings": archive_previous_findings, - }, - credentials=credentials, - provider_ids=[provider_id], - enabled=enabled, - ) - except Exception as e: - self.logger.error(f"AWS Security Hub integration creation failed: {e}") - return {"error": str(e), "status": "failed"} + return await self._create_integration( + integration_type="aws_security_hub", + configuration={ + "send_only_fails": send_only_fails, + "archive_previous_findings": archive_previous_findings, + }, + credentials=credentials, + provider_ids=[provider_id], + enabled=enabled, + ) async def create_jira_integration( self, - domain: str = Field( + domain: NonBlankStr = Field( description="Atlassian site name, without the '.atlassian.net' suffix. For the site 'https://acme.atlassian.net' the value is 'acme'. Full URLs are accepted and normalized automatically." ), - user_mail: str = Field( + user_mail: NonBlankStr = Field( description="Email address of the Atlassian account that owns the API token." ), - api_token: str = Field( + api_token: NonBlankStr = Field( description="Atlassian API token, created from the account settings. It needs the 'read:jira-user', 'read:jira-work' and 'write:jira-work' scopes." ), enabled: bool = Field( @@ -416,31 +411,25 @@ class IntegrationsTools(BaseTool): 3. Use prowler_get_jira_issue_types with that project key to pick an issue type 4. Use prowler_send_findings_to_jira to create the work items """ - try: - normalized_domain = self._normalize_atlassian_domain(domain) - self.logger.info( - f"Creating Jira integration for domain {normalized_domain}..." - ) + normalized_domain = self._normalize_atlassian_domain(domain) + self.logger.info(f"Creating Jira integration for domain {normalized_domain}...") - return await self._create_integration( - integration_type="jira", - # Jira rejects any configuration in the payload, the API generates it - configuration={}, - credentials={ - "domain": normalized_domain, - "user_mail": user_mail, - "api_token": api_token, - }, - provider_ids=[], - enabled=enabled, - ) - except Exception as e: - self.logger.error(f"Jira integration creation failed: {e}") - return {"error": str(e), "status": "failed"} + return await self._create_integration( + integration_type="jira", + # Jira rejects any configuration in the payload, the API generates it + configuration={}, + credentials={ + "domain": normalized_domain, + "user_mail": user_mail, + "api_token": api_token, + }, + provider_ids=[], + enabled=enabled, + ) async def update_integration( self, - integration_id: str = Field( + integration_id: NonBlankStr = Field( description="UUID of the integration to update. Use prowler_list_integrations to find it." ), enabled: bool | None = Field( @@ -494,96 +483,86 @@ class IntegrationsTools(BaseTool): """ self.logger.info(f"Updating integration {integration_id}...") - try: - current = DetailedIntegration.from_api_response( - await self._get_integration_raw(integration_id) - ) - integration_type = current.integration_type + current = DetailedIntegration.from_api_response( + await self._get_integration_raw(integration_id) + ) + integration_type = current.integration_type - if provider_ids is not None: - if integration_type == "jira": - raise ValueError( - "Jira integrations are tenant-wide and cannot be attached to providers." - ) - if integration_type == "aws_security_hub" and len(provider_ids) != 1: - raise ValueError( - "AWS Security Hub integrations must stay attached to exactly one AWS " - f"provider, got {len(provider_ids)}. Pass a single provider ID, or use " - "prowler_delete_integration to stop sending findings to Security Hub." - ) - - attributes: dict[str, Any] = {} - if enabled is not None: - attributes["enabled"] = enabled - - if credentials is not None: - attributes["credentials"] = self._validate_credentials( - integration_type, self._as_dict(credentials, "credentials") + if provider_ids is not None: + if integration_type == "jira": + raise InvalidArgument( + "Jira integrations are tenant-wide and cannot be attached to providers." + ) + if integration_type == "aws_security_hub" and len(provider_ids) != 1: + raise InvalidArgument( + "AWS Security Hub integrations must stay attached to exactly one AWS " + f"provider, got {len(provider_ids)}. Pass a single provider ID, or use " + "prowler_delete_integration to stop sending findings to Security Hub." ) - if configuration is not None: - if integration_type == "jira": - raise ValueError( - "Jira integrations do not accept a configuration: it is generated by Prowler. " - "Update the credentials instead, or run prowler_test_integration_connection to " - "refresh the available projects and issue types." - ) - merged = dict(current.configuration) - merged.update(self._as_dict(configuration, "configuration")) - # Server-owned, the API repopulates it from the connection check - merged.pop("regions", None) - merged.pop("enabled_regions", None) - attributes["configuration"] = merged + attributes: dict[str, Any] = {} + if enabled is not None: + attributes["enabled"] = enabled - if not attributes and provider_ids is None: - self.logger.info("No changes provided, returning the current state") - return current.model_dump() + if credentials is not None: + attributes["credentials"] = self._validate_credentials( + integration_type, self._as_dict(credentials, "credentials") + ) - update_body: dict[str, Any] = { - "data": { - "type": "integrations", - "id": integration_id, - "attributes": attributes, - } + if configuration is not None: + if integration_type == "jira": + raise InvalidArgument( + "Jira integrations do not accept a configuration: it is generated by Prowler. " + "Update the credentials instead, or run prowler_test_integration_connection to " + "refresh the available projects and issue types." + ) + merged = dict(current.configuration) + merged.update(self._as_dict(configuration, "configuration")) + # Server-owned, the API repopulates it from the connection check + merged.pop("regions", None) + merged.pop("enabled_regions", None) + attributes["configuration"] = merged + + if not attributes and provider_ids is None: + self.logger.info("No changes provided, returning the current state") + return current.model_dump() + + update_body: dict[str, Any] = { + "data": { + "type": "integrations", + "id": integration_id, + "attributes": attributes, } - if provider_ids is not None: - update_body["data"]["relationships"] = _providers_relationship( - provider_ids - ) + } + if provider_ids is not None: + update_body["data"]["relationships"] = _providers_relationship(provider_ids) - await self.api_client.patch( - f"/integrations/{integration_id}", json_data=update_body - ) + await self.api_client.patch( + f"/integrations/{integration_id}", json_data=update_body + ) - # A different provider means different effective credentials and different - # discovered configuration, so the stored connection state is stale too - providers_changed = provider_ids is not None and set(provider_ids) != set( - current.provider_ids - ) - recheck_connection = ( - credentials is not None - or configuration is not None - or providers_changed - ) - connection_status = ( - await self._test_connection(integration_id) - if recheck_connection - else None - ) + # A different provider means different effective credentials and different + # discovered configuration, so the stored connection state is stale too + providers_changed = provider_ids is not None and set(provider_ids) != set( + current.provider_ids + ) + recheck_connection = ( + credentials is not None or configuration is not None or providers_changed + ) + connection_status = ( + await self._test_connection(integration_id) if recheck_connection else None + ) - updated = await self._get_integration_raw(integration_id) - if connection_status is not None: - return IntegrationConnectionStatus.create( - updated, connection_status - ).model_dump() - return DetailedIntegration.from_api_response(updated).model_dump() - except Exception as e: - self.logger.error(f"Integration update failed: {e}") - return {"error": str(e), "status": "failed"} + updated = await self._get_integration_raw(integration_id) + if connection_status is not None: + return IntegrationConnectionStatus.create( + updated, connection_status + ).model_dump() + return DetailedIntegration.from_api_response(updated).model_dump() async def delete_integration( self, - integration_id: str = Field( + integration_id: NonBlankStr = Field( description="UUID of the integration to permanently remove. Use prowler_list_integrations to find it." ), ) -> dict[str, Any]: @@ -606,22 +585,15 @@ class IntegrationsTools(BaseTool): """ self.logger.info(f"Deleting integration {integration_id}...") - try: - await self.api_client.delete(f"/integrations/{integration_id}") - return { - "deleted": True, - "message": f"Integration {integration_id} deleted successfully", - } - except Exception as e: - self.logger.error(f"Integration deletion failed: {e}") - return { - "deleted": False, - "message": f"Integration {integration_id} deletion failed: {str(e)}", - } + await self.api_client.delete(f"/integrations/{integration_id}") + # No `deleted` flag: an integration that was not deleted leaves this tool + # as an error, so the flag could only ever be True and a reader branching + # on it would be looking for a shape that does not exist. + return {"message": f"Integration {integration_id} deleted successfully"} async def test_integration_connection( self, - integration_id: str = Field( + integration_id: NonBlankStr = Field( description="UUID of the integration to check. Use prowler_list_integrations to find it." ), ) -> dict[str, Any]: @@ -654,10 +626,10 @@ class IntegrationsTools(BaseTool): async def get_jira_issue_types( self, - integration_id: str = Field( + integration_id: NonBlankStr = Field( description="UUID of the Jira integration. Use prowler_list_integrations with integration_type=['jira'] to find it." ), - project_key: str = Field( + project_key: NonBlankStr = Field( description="Key of the Jira project to read the issue types from (e.g. 'PROJ'). It must be one of the keys in the 'projects' mapping of the integration configuration." ), ) -> dict[str, Any]: @@ -692,13 +664,13 @@ class IntegrationsTools(BaseTool): async def send_findings_to_jira( self, - integration_id: str = Field( + integration_id: NonBlankStr = Field( description="UUID of the Jira integration to send the findings through. It must be enabled." ), - project_key: str = Field( + project_key: NonBlankStr = Field( description="Key of the Jira project the work items are created in (e.g. 'PROJ'). It must be one of the keys in the 'projects' mapping of the integration configuration." ), - issue_type: str = Field( + issue_type: NonBlankStr = Field( description="Jira issue type for the created work items (e.g. 'Task', 'Bug', 'Story'). It must be one of the values returned by prowler_get_jira_issue_types for this project." ), finding_ids: list[str] = Field( @@ -783,20 +755,26 @@ class IntegrationsTools(BaseTool): return self._jira_dispatch_unknown( task_id=None, error=( - f"the request that starts the dispatch failed on the server: {e} " + "the request that starts the dispatch failed on Prowler's side. " "It may have been queued anyway." ), ) self.logger.error(f"Jira dispatch was rejected by Prowler: {e}") return self._jira_dispatch_rejected(str(e)) + except CredentialError: + # Authentication happens before the request goes out, so nothing was + # queued. It is raised rather than reported as a dispatch outcome: + # there is no partial state to describe, and the shared classifier + # says what has to be fixed, which no retry of this call can. + raise except Exception as e: # No answer came back, so the request may still have been accepted self.logger.error(f"Jira dispatch could not be started: {e}") return self._jira_dispatch_unknown( task_id=None, error=( - f"the request that starts the dispatch got no answer: {e} " + "the request that starts the dispatch got no answer. " "It may have been accepted anyway." ), ) @@ -866,7 +844,7 @@ class IntegrationsTools(BaseTool): normalized = normalized.removesuffix(".atlassian.net") if not normalized: - raise ValueError( + raise InvalidArgument( f"Invalid Jira domain: {domain}. Provide the Atlassian site name, for example " "'acme' for the site 'https://acme.atlassian.net'." ) @@ -890,7 +868,7 @@ class IntegrationsTools(BaseTool): if not isinstance(credentials.get(key), str) or not credentials[key].strip() ] if missing: - raise ValueError( + raise InvalidArgument( "Jira credentials are replaced as a whole, so 'domain', 'user_mail' and " f"'api_token' are all required. Missing or empty: {', '.join(missing)}. " "Sending an incomplete object would destroy the stored credentials and break " @@ -908,29 +886,33 @@ class IntegrationsTools(BaseTool): try: value = json.loads(value) except json.JSONDecodeError as e: - raise ValueError(f"Invalid JSON for {param_name}: {e}") + raise InvalidArgument(f"Invalid JSON for {param_name}: {e}") from e if not isinstance(value, dict): - raise ValueError(f"{param_name} must be a JSON object.") + raise InvalidArgument(f"{param_name} must be a JSON object.") return value async def _get_integration_raw(self, integration_id: str) -> dict[str, Any]: """Fetch the raw JSON:API resource of an integration. Raises: - ValueError: If the payload does not contain a usable integration resource + ToolError: If the payload does not contain a usable integration resource. + Raised without a ``from`` clause because these messages name the + integration and the tool that lists valid IDs, and the two cases + are reported differently: a missing resource is the caller's + mistake, a resource without attributes is the API's. """ response = await self.api_client.get(f"/integrations/{integration_id}") integration = response.get("data") if not isinstance(integration, dict) or not integration.get("id"): - raise ValueError( + raise ToolError( f"Integration {integration_id} was not found. Use prowler_list_integrations " "to get a valid integration ID." ) if not isinstance(integration.get("attributes"), dict): - raise ValueError( + raise ToolError( f"Prowler returned integration {integration_id} without its attributes, so " "its state cannot be read." ) @@ -970,7 +952,9 @@ class IntegrationsTools(BaseTool): integration_id = api_response.get("data", {}).get("id") if not integration_id: - raise ValueError( + # The integration may well exist, so this must not read as "nothing + # happened" and invite a duplicate. + raise ToolError( "Prowler accepted the integration creation but did not return its ID, so the " "connection could not be checked. Use prowler_list_integrations to see whether " "the integration exists before creating it again." @@ -981,11 +965,17 @@ class IntegrationsTools(BaseTool): try: integration = await self._get_integration_raw(integration_id) except Exception as e: - # The integration exists, so surface its ID instead of a plain read failure - raise ValueError( - f"Integration {integration_id} was created, but reading its state failed: {e} " + # The integration exists, so surface its ID instead of a plain read + # failure. No `from` clause: a cause would let the shared classifier + # replace this with a sentence that does not mention the ID. The + # failure text stays in the log, where the classifier would keep it. + self.logger.error( + f"Integration {integration_id} could not be read back: {e}" + ) + raise ToolError( + f"Integration {integration_id} was created, but reading its state failed. " "Use prowler_get_integration with that ID to check it." - ) from e + ) return IntegrationConnectionStatus.create( integration, connection_status @@ -1030,7 +1020,7 @@ class IntegrationsTools(BaseTool): return { "connected": None, "error": ( - f"The connection check could not be completed: {e} This says nothing " + "The connection check could not be completed. This says nothing " "about the stored credentials, run prowler_test_integration_connection " "to check them again." ), diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/muting.py b/mcp_server/prowler_mcp_server/prowler_app/tools/muting.py index 37e1504165..69a5f8884c 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/muting.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/muting.py @@ -8,8 +8,11 @@ This module provides tools for managing finding muting in Prowler, including: import json from typing import Any +from fastmcp.exceptions import ToolError from pydantic import Field +from prowler_mcp_server.lib.errors import InvalidArgument +from prowler_mcp_server.lib.types import NonBlankStr from prowler_mcp_server.prowler_app.models.muting import ( DetailedMuteRule, MutelistResponse, @@ -28,10 +31,31 @@ class MutingTools(BaseTool): # ===== MUTELIST TOOLS ===== + async def _get_mutelist_raw(self) -> dict[str, Any] | None: + """Return the tenant's mutelist, or None when it has none. + + Returns: + The mutelist configuration, or None when the tenant has none + """ + params = { + "filter[processor_type]": "mutelist", + "fields[processors]": "processor_type,configuration,inserted_at,updated_at", + } + + clean_params = self.api_client.build_filter_params(params) + api_response = await self.api_client.get("/processors", params=clean_params) + + data = api_response.get("data", []) + if not data: + return None + + # Only one mutelist can exist per tenant + return MutelistResponse.from_api_response(data[0]).model_dump() + async def get_mutelist(self) -> dict[str, Any]: """Retrieve the current mutelist configuration for the tenant. - IMPORTANT: Only one mutelist can exist per tenant. Returns an error message if no mutelist exists. + IMPORTANT: Only one mutelist can exist per tenant. Fails with a message saying so if no mutelist exists. For detailed information about mutelist structure and configuration, search Prowler documentation using prowler_docs_search tool available in this MCP Server. @@ -47,26 +71,15 @@ class MutingTools(BaseTool): """ self.logger.info("Retrieving mutelist configuration...") - # Query processors filtered by type=mutelist - params = { - "filter[processor_type]": "mutelist", - "fields[processors]": "processor_type,configuration,inserted_at,updated_at", - } - - clean_params = self.api_client.build_filter_params(params) - api_response = await self.api_client.get("/processors", params=clean_params) - - data = api_response.get("data", []) - - if len(data) == 0: - return { - "error": "No mutelist found", - "message": "No mutelist configuration exists for this tenant. Use prowler_set_mutelist to create one.", - } - - # Return the first (and only) mutelist - mutelist = MutelistResponse.from_api_response(data[0]) - return mutelist.model_dump() + mutelist = await self._get_mutelist_raw() + if mutelist is None: + # No `from`: this names the tool that creates one, which the shared + # classifier cannot know. + raise ToolError( + "No mutelist configuration exists for this tenant. Use " + "prowler_set_mutelist to create one." + ) + return mutelist async def set_mutelist( self, @@ -128,9 +141,9 @@ Structure: configuration = json.loads(configuration) # Check if mutelist already exists - existing_mutelist = await self.get_mutelist() + existing_mutelist = await self._get_mutelist_raw() - if "error" in existing_mutelist: + if existing_mutelist is None: # Create new mutelist self.logger.info("Creating new mutelist...") create_body = { @@ -183,21 +196,22 @@ Structure: self.logger.info("Deleting mutelist configuration...") # Get existing mutelist - existing_mutelist = await self.get_mutelist() + existing_mutelist = await self._get_mutelist_raw() - if "error" in existing_mutelist: - return { - "success": False, - "message": "No mutelist found to delete", - } + if existing_mutelist is None: + raise ToolError( + "There is no mutelist configuration to delete. Use " + "prowler_get_mutelist to confirm the current state." + ) # Delete the mutelist mutelist_id = existing_mutelist["id"] await self.api_client.delete(f"/processors/{mutelist_id}") + # No success flag: a deletion that did not happen leaves this tool as an + # error, so there is no second shape for one to distinguish. return { - "success": True, - "message": "Mutelist deleted successfully", + "message": "Mutelist deleted successfully. Findings it had muted stay muted." } # ===== MUTE RULES TOOLS ===== @@ -268,7 +282,7 @@ Structure: elif enabled.lower() == "false": params["filter[enabled]"] = False else: - raise ValueError( + raise InvalidArgument( f"Invalid enabled value: {enabled}. Valid values are True, False, 'true', 'false' or None." ) if search: @@ -282,7 +296,7 @@ Structure: async def get_mute_rule( self, - rule_id: str = Field( + rule_id: NonBlankStr = Field( description="UUID of the mute rule to retrieve. Must be a valid UUID format (e.g., '019ac0d6-90d5-73e9-9acf-c22e256f1bac')." ), ) -> dict[str, Any]: @@ -316,10 +330,10 @@ Structure: async def create_mute_rule( self, - name: str = Field( + name: NonBlankStr = Field( description="Name for the mute rule. Should be descriptive and meaningful (e.g., 'Dev S3 Public Access', 'Test Environment IMDSv1')." ), - reason: str = Field( + reason: NonBlankStr = Field( description="Reason for muting these findings. Document why this security issue is acceptable or intentional (e.g., 'Development environment with controlled access', 'Legacy application requires IMDSv1')." ), finding_ids: list[str] = Field( @@ -367,14 +381,14 @@ Structure: async def update_mute_rule( self, - rule_id: str = Field( + rule_id: NonBlankStr = Field( description="UUID of the mute rule to update. Must be a valid UUID format." ), - name: str | None = Field( + name: NonBlankStr | None = Field( default=None, description="New name for the rule. If not specified, name remains unchanged.", ), - reason: str | None = Field( + reason: NonBlankStr | None = Field( default=None, description="New reason for the rule. If not specified, reason remains unchanged.", ), @@ -435,7 +449,7 @@ Structure: async def delete_mute_rule( self, - rule_id: str = Field( + rule_id: NonBlankStr = Field( description="UUID of the mute rule to delete. Must be a valid UUID format." ), ) -> dict[str, Any]: @@ -457,15 +471,18 @@ Structure: """ self.logger.info(f"Deleting mute rule {rule_id}...") - result = await self.api_client.delete(f"/mute-rules/{rule_id}") + # A deletion that did not happen answers with an error status, which + # leaves this tool as an error. Reaching this line means Prowler accepted + # it, whether it answered 204 with no body or 200 with the deleted + # resource, so there is no second outcome to report: the previous + # "Failed to delete mute rule" fired on the shape of the answer rather + # than on anything having gone wrong, and said nothing a caller could act + # on. + await self.api_client.delete(f"/mute-rules/{rule_id}") - if result.get("success"): - return { - "success": True, - "message": "Mute rule deleted successfully", - } - else: - return { - "success": False, - "message": "Failed to delete mute rule", - } + return { + "message": ( + f"Mute rule {rule_id} deleted successfully. The findings it muted stay " + "muted." + ) + } diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/providers.py b/mcp_server/prowler_mcp_server/prowler_app/tools/providers.py index 3ba417d677..542289c652 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/providers.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/providers.py @@ -6,10 +6,14 @@ including searching, connecting, and deleting providers. from typing import Any +from fastmcp.exceptions import ToolError from pydantic import Field +from prowler_mcp_server.lib.errors import InvalidArgument +from prowler_mcp_server.lib.types import NonBlankStr from prowler_mcp_server.prowler_app.models.providers import ( ProviderConnectionStatus, + ProviderDeletionResult, ProvidersListResponse, ) from prowler_mcp_server.prowler_app.tools.base import BaseTool @@ -95,7 +99,7 @@ class ProvidersTools(BaseTool): elif connected.lower() == "false": params["filter[connected]"] = False else: - raise ValueError( + raise InvalidArgument( f"Invalid connected value: {connected}. Valid values are True, False, 'true', 'false' or None." ) @@ -128,13 +132,13 @@ class ProvidersTools(BaseTool): async def connect_provider( self, - provider_uid: str = Field( + provider_uid: NonBlankStr = Field( description="Provider's unique identifier. For supported UID provider formats, please refer to Prowler Hub/Prowler Documentation that you can also find in form of tools in this MCP Server" ), - provider_type: str = Field( + provider_type: NonBlankStr = Field( description="Type of provider to be scanned with Prowler. Valid values include: 'aws', 'azure', 'gcp', 'kubernetes'... For more valid values, please refer to Prowler Hub/Prowler Documentation that you can also find in form of tools in this MCP Server." ), - alias: str | None = Field( + alias: NonBlankStr | None = Field( default=None, description="Human-friendly name for this provider. Optional but recommended for easy identification. Use descriptive names to distinguish multiple accounts of the same type.", ), @@ -291,7 +295,7 @@ class ProvidersTools(BaseTool): async def delete_provider( self, - provider_id: str = Field( + provider_id: NonBlankStr = Field( description="Prowler's internal UUID (v4) for the provider to permanently remove, generated when the provider was registered in the system. Use `prowler_search_providers` tool to find the provider_id if you only know the alias or the provider's own identifier (provider_uid)" ), ) -> dict[str, Any]: @@ -300,33 +304,120 @@ class ProvidersTools(BaseTool): WARNING: This is a destructive operation that cannot be undone. The provider will need to be re-added with prowler_connect_provider if you want to scan it again. - The tool always returns the deletion status and message. + Prowler removes the provider and everything attached to it (its scans, findings and + resources) in a background task, so a large provider can take longer than the time + this tool waits for it. + + The result includes: + - status: 'deleted' when Prowler finished removing the provider, 'in_progress' when + the deletion was accepted and is still running + - task_id: the background task, present when the deletion was still running + + NEVER send the deletion again while status='in_progress'. Use prowler_search_providers + to check whether the provider is gone. """ self.logger.info(f"Deleting provider {provider_id}...") - try: - # Initiate the deletion task - task_response = await self.api_client.delete(f"/providers/{provider_id}") - task_id = task_response.get("data", {}).get("id") - # Poll until task completes (with 60 second timeout) + # A failure of the request itself is left to the shared classifier: the + # deletion never started, so there is no partial state to describe. + task_response = await self.api_client.delete(f"/providers/{provider_id}") + task_id = task_response.get("data", {}).get("id") + + if not task_id: + # The deletion may well be running, so this must not read as "nothing + # happened". No `from` clause: this names the provider and the tool + # that checks it, neither of which the shared classifier can know. + raise ToolError( + f"Prowler accepted the deletion of provider {provider_id} but did not " + "return the ID of the background task, so its outcome cannot be checked. " + "Use prowler_search_providers to see whether the provider is still there " + "before sending the deletion again." + ) + + try: await self.api_client.poll_task_until_complete( task_id=task_id, timeout=60, poll_interval=1.0 ) - - # If we reach here, the task completed successfully - return { - "deleted": True, - "message": f"Provider {provider_id} deleted successfully", - } except Exception as e: - self.logger.error(f"Provider deletion failed: {e}") - return { - "deleted": False, - "message": f"Provider {provider_id} deletion failed: {str(e)}", - } + self.logger.error(f"Provider deletion did not complete cleanly: {e}") + return await self._provider_deletion_fallback(provider_id, task_id) + + return ProviderDeletionResult( + status="deleted", + message=f"Provider {provider_id} deleted successfully", + ).model_dump() # Private helper methods + async def _provider_deletion_fallback( + self, provider_id: str, task_id: str + ) -> dict[str, Any]: + """Report a provider deletion whose polling did not end on a completed task. + + Running out of the polling window is not a failure: Prowler removes the + provider together with its scans, findings and resources, which outlives + 60 seconds on a large account. The deletion was accepted and is still + going, so calling it failed would be wrong twice over -- it is not, and + it invites a retry of a destructive call already in flight. + + The task is read once more here, because polling gives up on the clock + rather than on the task: a deletion that finished just after the last + poll is a finished deletion and is reported as one. + + Only a task that actually stopped is an error, and it is raised rather + than returned, because then the provider is still there. + + Raises: + ToolError: If the deletion task ended without deleting the provider. + Raised without a ``from`` clause because the message names what + was left behind, which the shared classifier cannot know. + """ + state = None + try: + task = await self.api_client.get(f"/tasks/{task_id}") + state = task.get("data", {}).get("attributes", {}).get("state") + except Exception as e: + self.logger.error(f"Could not read the state of task {task_id}: {e}") + + if state == "completed": + # The deletion outran the polling window by a moment, not by more. + return ProviderDeletionResult( + status="deleted", + message=f"Provider {provider_id} deleted successfully", + ).model_dump() + + if state in ("failed", "cancelled"): + # The failure that got us here is logged, not relayed: it carries + # upstream text, and the classifier masks exactly this kind of + # message when a tool does not write it itself. + raise ToolError( + f"The task deleting provider {provider_id} ended as '{state}', so the " + "provider was not deleted. Prowler removes a provider together with its " + "scans, findings and resources, so part of that may already be gone. Use " + "prowler_search_providers to check the current state." + ) + + if state is None: + message = ( + f"The deletion of provider {provider_id} was accepted, but its progress " + "could not be read, so whether it finished is unknown. Do not " + "send the deletion again. Use prowler_search_providers to check whether " + "the provider is gone." + ) + else: + message = ( + f"The deletion of provider {provider_id} was accepted and is still " + f"running (task state '{state}'), which is normal for a provider with " + "many scans and findings. Do not send the deletion again. Use " + "prowler_search_providers to check whether it is gone." + ) + + return ProviderDeletionResult( + status="in_progress", + task_id=task_id, + message=message, + ).model_dump() + async def _check_provider_exists(self, provider_uid: str) -> str | None: """Check if a provider already exists by its UID. @@ -357,7 +448,7 @@ class ProvidersTools(BaseTool): return prowler_provider_id else: # Multiple providers with the same UID is a data integrity issue - raise Exception( + raise ToolError( f"Data integrity error: Found {len(providers)} providers with UID '{provider_uid}'. " f"Each provider UID should be unique. Please contact support or manually clean up duplicate providers." ) @@ -392,7 +483,11 @@ class ProvidersTools(BaseTool): provider_id = await self._check_provider_exists(provider_uid) if provider_id is None: - raise Exception(f"Provider {provider_uid} creation failed") + raise ToolError( + f"Prowler accepted the creation of provider {provider_uid} but the " + "provider cannot be found afterwards. Use prowler_search_providers to " + "check whether it exists before creating it again." + ) return provider_id async def _update_provider_alias( @@ -418,7 +513,10 @@ class ProvidersTools(BaseTool): f"/providers/{prowler_provider_id}", json_data=update_body ) if result.get("data", {}).get("attributes", {}).get("alias") != alias: - raise Exception(f"Provider {prowler_provider_id} alias update failed") + raise ToolError( + f"Provider {prowler_provider_id} exists, but its alias was not updated. " + "Use prowler_search_providers to read its current alias." + ) def _determine_secret_type(self, credentials: dict[str, Any]) -> str: """Determine the secret type from credentials structure. @@ -443,29 +541,32 @@ class ProvidersTools(BaseTool): prowler_provider_id: The Prowler-generated provider ID Returns: - The secret ID if exists, None otherwise - """ - try: - response = await self.api_client.get( - "/providers/secrets", - params={"filter[provider]": prowler_provider_id}, - ) - secrets = response.get("data", []) + The secret ID if the provider has one, None if it has none - if len(secrets) > 0: - secret_id = secrets[0].get("id") - self.logger.info( - f"Found existing secret {secret_id} for provider {prowler_provider_id}" - ) - return secret_id - else: - self.logger.info( - f"No existing secret found for provider {prowler_provider_id}" - ) - return None - except Exception as e: - self.logger.error(f"Error checking for existing secret: {e}") - return None + Raises: + Exception: If the lookup itself failed, so that "no secret" is never + reported for a provider whose secret could not be read + """ + # A failure here is not swallowed into None. None means "this provider has + # no secret", which sends `_store_credentials` down the create branch, and + # a provider holds at most one secret: creating a second one is refused, + # and the caller would be told its credentials were rejected when all that + # actually failed was this read. + response = await self.api_client.get( + "/providers/secrets", + params={"filter[provider]": prowler_provider_id}, + ) + secrets = response.get("data", []) + + if len(secrets) > 0: + secret_id = secrets[0].get("id") + self.logger.info( + f"Found existing secret {secret_id} for provider {prowler_provider_id}" + ) + return secret_id + + self.logger.info(f"No existing secret found for provider {prowler_provider_id}") + return None async def _get_secret_type(self, secret_id: str) -> str | None: """Get the secret type for a given secret ID. @@ -573,13 +674,24 @@ class ProvidersTools(BaseTool): raise async def _test_connection(self, prowler_provider_id: str) -> dict[str, Any]: - """Test connection to a provider. + """Test connection to a provider and wait for the result. + + A test that could not be run is reported as 'connected: None', which + `ProviderConnectionStatus` renders as 'not_tested', rather than as a + failure. Credentials that do not work come back as a completed task + carrying 'connected: False', so an exception here never describes them: + it means this server could not get the test run at all -- an expired + Prowler credential, a rate limit, a test that outlived the timeout. + Reporting that as 'failed' would blame the provider's credentials for + something they did not cause, and send the caller off to fix a working + role. Args: prowler_provider_id: The Prowler-generated provider ID Returns: - Connection status dictionary with 'connected' boolean and optional 'error' message + Connection status dictionary with a 'connected' boolean or None, and + an optional 'error' message """ self.logger.info(f"Testing connection for provider {prowler_provider_id}...") try: @@ -589,6 +701,11 @@ class ProvidersTools(BaseTool): ) task_id = task_response.get("data", {}).get("id") + if not task_id: + raise ValueError( + "Prowler did not return the ID of the connection test task." + ) + # Poll until task completes (with 60 second timeout) completed_task = await self.api_client.poll_task_until_complete( task_id=task_id, timeout=60, poll_interval=1.0 @@ -596,13 +713,26 @@ class ProvidersTools(BaseTool): # Extract the result from the completed task task_result = ( - completed_task.get("data", {}).get("attributes", {}).get("result", {}) + completed_task.get("data", {}).get("attributes", {}).get("result") ) + if not isinstance(task_result, dict): + raise ValueError( + "The connection test task completed without reporting a result." + ) + return task_result except Exception as e: - self.logger.error(f"Connection test failed: {e}") - return {"connected": False, "error": str(e)} + self.logger.error(f"Connection test could not be completed: {e}") + return { + "connected": None, + "error": ( + "The connection test could not be completed. This says nothing " + "about the provider's credentials, they were never tested. Use " + "prowler_search_providers to read the connection state Prowler has " + "stored for this provider." + ), + } async def _get_final_provider_state( self, prowler_provider_id: str diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/resources.py b/mcp_server/prowler_mcp_server/prowler_app/tools/resources.py index 88fcca25ae..96a24fb4b2 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/resources.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/resources.py @@ -8,6 +8,7 @@ from typing import Any from pydantic import Field +from prowler_mcp_server.lib.types import NonBlankStr from prowler_mcp_server.prowler_app.models.resources import ( DetailedResource, ResourceEventsResponse, @@ -176,7 +177,7 @@ class ResourcesTools(BaseTool): async def get_resource( self, - resource_id: str = Field( + resource_id: NonBlankStr = Field( description="Prowler's internal UUID (v4) for the resource to retrieve, generated when the resource was discovered in the system. Use `prowler_list_resources` tool to find the right ID" ), ) -> dict[str, Any]: @@ -347,7 +348,7 @@ class ResourcesTools(BaseTool): async def get_resource_events( self, - resource_id: str = Field( + resource_id: NonBlankStr = Field( description="Prowler's internal UUID (v4) for the resource. Use `prowler_list_resources` to find the right ID, or get it from a finding's resource relationship via `prowler_get_finding_details`." ), lookback_days: int = Field( diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/roles.py b/mcp_server/prowler_mcp_server/prowler_app/tools/roles.py index 113694d8e8..8fc4daa20b 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/roles.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/roles.py @@ -11,8 +11,11 @@ adding to it. from typing import Any +from fastmcp.exceptions import ToolError from pydantic import Field +from prowler_mcp_server.lib.errors import ProwlerAPIError +from prowler_mcp_server.lib.types import NonBlankStr from prowler_mcp_server.prowler_app.models.roles import ( DetailedRole, RolesListResponse, @@ -70,7 +73,7 @@ class RolesTools(BaseTool): async def get_role( self, - role_id: str = Field( + role_id: NonBlankStr = Field( description="Prowler's internal UUID (v4) for the role to retrieve. Use `prowler_list_roles` to find role IDs if you only know a name." ), ) -> dict[str, Any]: @@ -98,7 +101,7 @@ class RolesTools(BaseTool): async def get_user_roles( self, - user_id: str = Field( + user_id: NonBlankStr = Field( description="Prowler's internal UUID (v4) for the user whose roles you want. Use `prowler_list_users` to find user IDs, or `prowler_get_current_user` for the caller." ), ) -> dict[str, Any]: @@ -124,10 +127,10 @@ class RolesTools(BaseTool): async def set_user_role( self, - user_id: str = Field( + user_id: NonBlankStr = Field( description="Prowler's internal UUID (v4) for the user whose role you want to set. Use `prowler_list_users` to find user IDs." ), - role_id: str = Field( + role_id: NonBlankStr = Field( description="Prowler's internal UUID (v4) for the role the user should hold. Use `prowler_list_roles` to find role IDs." ), ) -> dict[str, Any]: @@ -166,11 +169,20 @@ class RolesTools(BaseTool): # user with no role at all. Confirm the role exists before replacing. try: await self.api_client.get(f"/roles/{role_id}") - except Exception as e: - raise ValueError( - f"Role {role_id} could not be read ({e}), so user {user_id} was left " - f"unchanged. Use `prowler_list_roles` to find a valid role ID." - ) from e + except ProwlerAPIError as e: + if e.status_code != 404: + # Only a not-found says anything about the role ID. A permission + # error, a rate limit or a server error is about the request, so + # it goes to the shared classifier rather than being reported as + # an ID the caller should replace. + raise + # No `from` clause: this says what state the user was left in, which + # the shared classifier cannot know, and a cause would let it replace + # this message with its own. + raise ToolError( + f"Role {role_id} does not exist in this tenant, so user {user_id} was " + f"left unchanged. Use `prowler_list_roles` to find a valid role ID." + ) # PATCH replaces the user's whole role set with this single role, the # same call the Prowler UI makes when changing a user's role. diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/scans.py b/mcp_server/prowler_mcp_server/prowler_app/tools/scans.py index 21d1431b71..106fe69d0d 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/scans.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/scans.py @@ -5,8 +5,10 @@ This module provides tools for managing and monitoring Prowler security scans. from typing import Any, Literal +from fastmcp.exceptions import ToolError from pydantic import Field +from prowler_mcp_server.lib.types import NonBlankStr from prowler_mcp_server.prowler_app.models.scans import ( DetailedScan, ScanCreationResult, @@ -127,7 +129,7 @@ class ScansTools(BaseTool): async def get_scan( self, - scan_id: str = Field( + scan_id: NonBlankStr = Field( description="Prowler's internal UUID (v4) for the scan to retrieve, generated when the scan was created (e.g., '123e4567-e89b-12d3-a456-426614174000'). Use `prowler_list_scans` tool to find scan IDs" ), ) -> dict[str, Any]: @@ -171,10 +173,10 @@ class ScansTools(BaseTool): async def trigger_scan( self, - provider_id: str = Field( + provider_id: NonBlankStr = Field( description="Prowler's internal UUID (v4) for the provider to scan, generated when the provider was registered in the system (e.g., '4d0e2614-6385-4fa7-bf0b-c2e2f75c6877'). Use `prowler_search_providers` tool to find the provider ID" ), - name: str | None = Field( + name: NonBlankStr | None = Field( default=None, description="Optional human-friendly name for the scan. Use descriptive names to identify scan purpose or context, e.g., 'Weekly Production Security Audit', 'Pre-Deployment Validation', 'Compliance Check Q4 2025'", ), @@ -191,60 +193,70 @@ class ScansTools(BaseTool): 3. Use `prowler_get_scan` with the returned scan 'id' to monitor progress 4. Once completed, use `prowler_search_security_findings` to analyze results """ - try: - # Build request data - request_data: dict[str, Any] = { - "data": { - "type": "scans", - "attributes": {}, - "relationships": { - "provider": { - "data": { - "type": "providers", - "id": provider_id, - }, + # Build request data + request_data: dict[str, Any] = { + "data": { + "type": "scans", + "attributes": {}, + "relationships": { + "provider": { + "data": { + "type": "providers", + "id": provider_id, }, }, }, - } - if name: - request_data["data"]["attributes"]["name"] = name + }, + } + if name: + request_data["data"]["attributes"]["name"] = name - # Create scan (returns Task) - self.logger.info(f"Creating scan for provider {provider_id}") - task_response = await self.api_client.post("/scans", json_data=request_data) + # Create scan (returns Task) + self.logger.info(f"Creating scan for provider {provider_id}") + task_response = await self.api_client.post("/scans", json_data=request_data) - scan_id = ( - task_response.get("data", {}) - .get("attributes", {}) - .get("task_args", {}) - .get("scan_id", None) + scan_id = ( + task_response.get("data", {}) + .get("attributes", {}) + .get("task_args", {}) + .get("scan_id", None) + ) + + if not scan_id: + # The scan may well have been queued, so this must not read as + # "nothing happened" and invite a duplicate run. No `from` clause: + # this names the provider and the tool that checks for the scan, + # neither of which the shared classifier can know. + raise ToolError( + "Prowler accepted the scan but did not return its ID, so it " + "cannot be looked up. Use prowler_list_scans for provider " + f"{provider_id} to see whether a scan is already running before " + "triggering another one." ) - if not scan_id: - raise Exception("No scan_id returned from scan creation") - - self.logger.info(f"Scan created successfully: {scan_id}") + # The scan exists from here on, so a failure to read it back must name + # the ID rather than read as "the scan was not created". + try: scan_response = await self.api_client.get(f"/scans/{scan_id}") scan_info = DetailedScan.from_api_response(scan_response["data"]) - - return ScanCreationResult( - scan=scan_info, - status="success", - message=f"Scan {scan_id} created successfully. The scan may take some time to complete. Use prowler_get_scan tool with this ID to monitor progress.", - ).model_dump() - except Exception as e: - self.logger.error(f"Scan creation failed: {e}") - return ScanCreationResult( - scan=None, - status="failed", - message=f"Scan creation failed: {str(e)}", - ).model_dump() + # The failure itself is logged, not relayed: what it says is the + # shared classifier's to mask, and what the caller needs is the ID. + self.logger.error(f"Scan {scan_id} could not be read back: {e}") + raise ToolError( + f"Scan {scan_id} was created for provider {provider_id}, but reading " + "its state failed. Use prowler_get_scan with that ID to monitor " + "it. Do not trigger the scan again." + ) + + return ScanCreationResult( + scan=scan_info, + message=f"Scan {scan_id} created successfully. The scan may take some time to complete. Use prowler_get_scan tool with this ID to monitor progress.", + ).model_dump() async def schedule_daily_scan( self, - provider_id: str = Field( + provider_id: NonBlankStr = Field( description="Prowler's internal UUID (v4) for the provider to scan, generated when the provider was registered in the system (e.g., '4d0e2614-6385-4fa7-bf0b-c2e2f75c6877'). Use `prowler_search_providers` tool to find the provider ID" ), ) -> dict[str, Any]: @@ -280,26 +292,49 @@ class ScansTools(BaseTool): }, }, ) - task_state = ( + + # Reaching this line means the schedule exists. Prowler commits the + # recurring schedule and its first scan inside the transaction that + # serves this request, so an answer at all means it was created; a + # provider that already has one is refused with a 409 instead, which + # leaves this tool as an error. + # + # The task in the answer is the FIRST scan run, queued to start a few + # seconds later, not the schedule. Its state therefore says nothing + # about whether the schedule was created, and reporting it as the + # outcome would call a schedule that exists a failure and invite a + # retry that can only hit that 409. + first_run_state = ( task_response.get("data", {}).get("attributes", {}).get("state", None) ) - if task_state == "available": - return_message = "Daily schedule created successfully. The schedule is being set up in the background. Use prowler_list_scans with provider_id filter to view scheduled scans." - else: - return_message = "Daily schedule creation failed. Please try again later." + message = ( + f"Daily schedule created for provider {provider_id}. Prowler will scan it " + "every 24 hours until the provider is deleted. Use prowler_list_scans with " + "this provider_id and trigger='scheduled' to view its scheduled scans." + ) + + if first_run_state in ("failed", "cancelled"): + # Worth saying: the schedule stands, but the run that was supposed to + # start now will not produce findings, and only a manual scan fills + # the gap before tomorrow. + message = ( + f"{message} Note that the first scan, which Prowler starts immediately, " + f"ended as '{first_run_state}'. The daily schedule is unaffected, but " + "use prowler_trigger_scan if you need results before the next run." + ) return ScheduleCreationResult( - scheduled=(task_state == "available"), - message=return_message, + first_run_state=first_run_state, + message=message, ).model_dump() async def update_scan( self, - scan_id: str = Field( + scan_id: NonBlankStr = Field( description="Prowler's internal UUID (v4) for the scan to update, generated when the scan was created (e.g., '123e4567-e89b-12d3-a456-426614174000'). Use `prowler_list_scans` tool to find the scan ID if you only know the provider or scan name. Returns an error if the scan ID is invalid or not found." ), - name: str = Field( + name: NonBlankStr = Field( description="New human-friendly name for the scan (3-100 characters). Use descriptive names to improve organization and tracking, e.g., 'Production Security Audit - Q4 2025', 'Post-Deployment Compliance Check'. IMPORTANT: Only the scan name can be updated - other attributes (state, progress, duration) are read-only and managed by the system." ), ) -> dict[str, Any]: diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/users.py b/mcp_server/prowler_mcp_server/prowler_app/tools/users.py index a7e31b60ad..f9b71ccd84 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/users.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/users.py @@ -9,6 +9,7 @@ from typing import Any from pydantic import Field +from prowler_mcp_server.lib.types import NonBlankStr from prowler_mcp_server.prowler_app.models.users import ( DetailedUser, UsersListResponse, @@ -79,7 +80,7 @@ class UsersTools(BaseTool): async def get_user( self, - user_id: str = Field( + user_id: NonBlankStr = Field( description="Prowler's internal UUID (v4) for the user to retrieve. Use `prowler_list_users` to find user IDs if you only know a name or email." ), ) -> dict[str, Any]: diff --git a/mcp_server/prowler_mcp_server/prowler_app/utils/api_client.py b/mcp_server/prowler_mcp_server/prowler_app/utils/api_client.py index 217496a347..cb4fbd87e8 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/utils/api_client.py +++ b/mcp_server/prowler_mcp_server/prowler_app/utils/api_client.py @@ -118,7 +118,21 @@ class ProwlerAPIClient(metaclass=SingletonMeta): if detail: message = f"{message} - {detail}" - raise ProwlerAPIError(message, status, detail=detail) from e + # Carried on the exception, not into the message: a tool needs the + # body to tell an answer with an error status -- a 404 holding the + # empty result of a query that matched nothing -- apart from a + # request that actually failed. + try: + body = e.response.json() + except ValueError: + body = None + + raise ProwlerAPIError( + message, + status, + detail=detail, + payload=body if isinstance(body, dict) else None, + ) from e except httpx.RequestError as e: # No answer came back, so whether the request was applied is unknown. logger.error(f"Error during {method.value} {path}: {e}") diff --git a/mcp_server/prowler_mcp_server/prowler_app/utils/auth.py b/mcp_server/prowler_mcp_server/prowler_app/utils/auth.py index eff5d3a117..1af63920ec 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/utils/auth.py +++ b/mcp_server/prowler_mcp_server/prowler_app/utils/auth.py @@ -6,6 +6,7 @@ from datetime import datetime from fastmcp.server.dependencies import get_http_headers from prowler_mcp_server import __version__ +from prowler_mcp_server.lib.errors import CredentialError from prowler_mcp_server.lib.logger import logger @@ -64,7 +65,12 @@ class ProwlerAppAuth: # Decode and parse JSON decoded = base64.b64decode(base64_payload).decode("utf-8") - return json.loads(decoded) + payload = json.loads(decoded) + + # A JWT payload is a JSON object. A list or a scalar decodes just as + # cleanly, so the type is checked here rather than left to blow up as + # an AttributeError on the first claim read. + return payload if isinstance(payload, dict) else None except Exception as e: logger.warning(f"Failed to parse JWT token: {e}") return None @@ -76,14 +82,16 @@ class ProwlerAppAuth: authorization_header = headers.get("authorization", None) if not authorization_header: - raise ValueError("No authorization header provided") + raise CredentialError("No Authorization header was sent") - # Extract token from Bearer header - if authorization_header.startswith("Bearer "): - token = authorization_header.replace("Bearer ", "") - else: - raise ValueError( - "Invalid authorization header format. Expected 'Bearer '" + # Extract token from Bearer header. Authentication scheme names are + # case-insensitive (RFC 7235), and only the scheme prefix is removed: + # a token that happens to contain the word again keeps it. + scheme, _, credential = authorization_header.partition(" ") + token = credential.strip() + if scheme.lower() != "bearer" or not token: + raise CredentialError( + "The Authorization header is not in 'Bearer ' form" ) # Check if it's an API key or JWT token @@ -94,17 +102,29 @@ class ProwlerAppAuth: # JWT token - validate and check expiration payload = self._parse_jwt(token) if not payload: - raise ValueError("Invalid JWT token format") + raise CredentialError("The token is not a readable JWT") + + # Check if token is expired. `exp` is a numeric date in the + # spec, so a missing or non-numeric one makes the token + # unusable rather than merely stale -- comparing it would raise + # a TypeError and leave the failure masked as unclassified. + exp = payload.get("exp") + if isinstance(exp, bool) or not isinstance(exp, (int, float)): + raise CredentialError( + "The token carries no readable 'exp' expiration claim" + ) - # Check if token is expired now = int(datetime.now().timestamp()) - exp = payload.get("exp", 0) if exp <= now: - raise ValueError("Token has expired") + raise CredentialError("The token has expired") return token else: - raise ValueError(f"Invalid mode: {self.mode}") + # PROWLER_MCP_TRANSPORT_MODE holds something this server does not + # support. Nothing about a call caused it and nothing about a call + # can fix it, so it stays unclassified: masked for the model, logged + # for whoever runs the server. + raise RuntimeError(f"Invalid mode: {self.mode}") async def get_valid_token(self) -> str: """Get a valid token (API key or JWT token).""" diff --git a/mcp_server/prowler_mcp_server/prowler_documentation/server.py b/mcp_server/prowler_mcp_server/prowler_documentation/server.py index 9588302168..de0756a505 100644 --- a/mcp_server/prowler_mcp_server/prowler_documentation/server.py +++ b/mcp_server/prowler_mcp_server/prowler_documentation/server.py @@ -3,6 +3,7 @@ from typing import Any from fastmcp import FastMCP from pydantic import Field +from prowler_mcp_server.lib.types import NonBlankStr from prowler_mcp_server.prowler_documentation.search_engine import ( ProwlerDocsSearchEngine, ) @@ -14,7 +15,9 @@ prowler_docs_search_engine = ProwlerDocsSearchEngine() @docs_mcp_server.tool() def search( - term: str = Field(description="The term to search for in the documentation"), + term: NonBlankStr = Field( + description="The term to search for in the documentation" + ), page_size: int = Field( 5, description="Number of top results to return. It must be between 1 and 20.", @@ -39,7 +42,7 @@ def search( @docs_mcp_server.tool() def get_document( - doc_path: str = Field( + doc_path: NonBlankStr = Field( description="Path to the documentation file to retrieve. It is the same as the 'path' field of the search results. Use `prowler_docs_search` to find the path first." ), ) -> dict[str, str]: diff --git a/mcp_server/prowler_mcp_server/prowler_hub/server.py b/mcp_server/prowler_mcp_server/prowler_hub/server.py index 41e83eca90..8d18a2c240 100644 --- a/mcp_server/prowler_mcp_server/prowler_hub/server.py +++ b/mcp_server/prowler_mcp_server/prowler_hub/server.py @@ -9,6 +9,7 @@ from fastmcp import FastMCP from pydantic import Field from prowler_mcp_server import __version__ +from prowler_mcp_server.lib.types import NonBlankStr # Initialize FastMCP for Prowler Hub hub_mcp_server = FastMCP("prowler-hub") @@ -149,7 +150,7 @@ async def list_checks( @hub_mcp_server.tool() async def semantic_search_checks( - term: str = Field( + term: NonBlankStr = Field( description="Search term. Examples: 'public access', 'encryption', 'MFA', 'logging'.", ), ) -> dict: @@ -208,7 +209,7 @@ async def semantic_search_checks( @hub_mcp_server.tool() async def get_check_details( - check_id: str = Field( + check_id: NonBlankStr = Field( description="The check ID to retrieve details for. Example: 's3_bucket_level_public_access_block'" ), ) -> dict: @@ -346,10 +347,10 @@ async def get_check_details( @hub_mcp_server.tool() async def get_check_code( - provider_id: str = Field( + provider_id: NonBlankStr = Field( description="Prowler Provider ID. Example: 'aws', 'azure', 'gcp', 'kubernetes'. Use `prowler_hub_list_providers` to get available provider IDs.", ), - check_id: str = Field( + check_id: NonBlankStr = Field( description="The check ID. Example: 's3_bucket_public_access'. Get IDs from `prowler_hub_list_checks` or `prowler_hub_search_checks`.", ), ) -> dict: @@ -392,10 +393,10 @@ async def get_check_code( @hub_mcp_server.tool() async def get_check_fixer( - provider_id: str = Field( + provider_id: NonBlankStr = Field( description="Prowler Provider ID. Example: 'aws', 'azure', 'gcp', 'kubernetes'. Use `prowler_hub_list_providers` to get available provider IDs.", ), - check_id: str = Field( + check_id: NonBlankStr = Field( description="The check ID. Example: 's3_bucket_public_access'. Get IDs from `prowler_hub_list_checks` or `prowler_hub_search_checks`.", ), ) -> dict: @@ -517,7 +518,7 @@ async def list_compliances( @hub_mcp_server.tool() async def semantic_search_compliances( - term: str = Field( + term: NonBlankStr = Field( description="Search term. Examples: 'CIS', 'HIPAA', 'PCI', 'GDPR', 'SOC2', 'NIST'.", ), ) -> dict: @@ -568,7 +569,7 @@ async def semantic_search_compliances( @hub_mcp_server.tool() async def get_compliance_details( - compliance_id: str = Field( + compliance_id: NonBlankStr = Field( description="The compliance framework ID to retrieve details for. Example: 'cis_4.0_aws'. Use `prowler_hub_list_compliances` or `prowler_hub_semantic_search_compliances` to find available compliance IDs.", ), ) -> dict: @@ -708,7 +709,7 @@ async def list_providers() -> dict: @hub_mcp_server.tool() async def get_provider_services( - provider_id: str = Field( + provider_id: NonBlankStr = Field( description="The provider ID to get services for. Example: 'aws', 'azure', 'gcp', 'kubernetes'. Use `prowler_hub_list_providers` to get available provider IDs.", ), ) -> dict: diff --git a/mcp_server/tests/lib/test_errors.py b/mcp_server/tests/lib/test_errors.py index 5986d883c4..cd32752b42 100644 --- a/mcp_server/tests/lib/test_errors.py +++ b/mcp_server/tests/lib/test_errors.py @@ -11,7 +11,11 @@ import pytest from fastmcp import Client from pydantic import BaseModel, ValidationError -from prowler_mcp_server.lib.errors import InvalidArgument, _describe_failure +from prowler_mcp_server.lib.errors import ( + CredentialError, + InvalidArgument, + _describe_failure, +) from prowler_mcp_server.prowler_app.utils.api_client import ( ProwlerAPIError, ProwlerAPIInvalidResponse, @@ -99,6 +103,19 @@ def test_an_argument_this_server_rejected_is_repeated_verbatim(): assert message == "page_size must be between 1 and 1000." +def test_a_credential_caught_here_is_answered_like_the_401_it_would_have_got(): + """It is not an argument problem, and saying so stops a pointless retry.""" + message = _describe_failure(CredentialError("the token has expired")) + + assert "the token has expired" in message + assert "changing the arguments will not help" in message + + +def test_a_transport_this_server_cannot_serve_is_left_masked(): + """No call caused a bad PROWLER_MCP_TRANSPORT_MODE and no call can fix it.""" + assert _describe_failure(RuntimeError("Invalid mode: websocket")) is None + + def test_a_pydantic_rejection_names_the_field_without_echoing_the_value(): """Pydantic quotes `input_value` back, and these tools take credentials.""" @@ -195,3 +212,18 @@ async def test_an_unreadable_api_answer_does_not_reach_the_agent_as_a_bad_argume assert result.isError is True assert "gateway timeout" not in result.content[0].text assert "argument" not in result.content[0].text + + +async def test_a_tool_specific_message_survives_masking( + mcp_root_server, mock_api_client, mock_router +): + """A `ToolError` raised without a `from` clause is the final word.""" + mock_router.add("GET", "/api/v1/integrations/i1", json={"data": None}) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_get_integration", {"integration_id": "i1"} + ) + + assert result.isError is True + assert "prowler_list_integrations" in result.content[0].text diff --git a/mcp_server/tests/lib/test_types.py b/mcp_server/tests/lib/test_types.py new file mode 100644 index 0000000000..8af25a7aae --- /dev/null +++ b/mcp_server/tests/lib/test_types.py @@ -0,0 +1,125 @@ +"""Tests for the argument types every tool shares. + +The bug these pin: "required" alone does not stop a blank identifier. A model +that has no scan or query id to hand sends ``""`` rather than omitting the +argument, and an unguarded empty string travels into a URL path or a request +body -- where it comes back as a 404, or as an API rejection ("This field may +not be blank") that names no argument and leaves the model with nothing to fix. +""" + +import pytest +from fastmcp import Client + +from tests.helpers.jsonapi import jsonapi_collection, jsonapi_resource + +SCAN_ID = "019ac0d6-90d5-73e9-9acf-c22e256f1bac" +QUERIES = f"/api/v1/attack-paths-scans/{SCAN_ID}/queries" + + +@pytest.mark.parametrize("query_id", ["", " "], ids=["empty", "whitespace-only"]) +async def test_a_blank_identifier_is_rejected_before_any_request_goes_out( + mcp_root_server, mock_api_client, mock_router, query_id +): + """The reported failure: a blank `query_id` reached Prowler as a 400. + + The message has to name the argument. Prowler's own answer to the blank value + ("This field may not be blank") does not say which field, so the model had no + way to tell `scan_id` from `query_id` from the reply. + """ + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_run_attack_paths_query", + {"scan_id": SCAN_ID, "query_id": query_id}, + ) + + assert result.isError is True + assert "query_id" in result.content[0].text + assert mock_router.requests == [] + + +async def test_an_identifier_keeps_its_surrounding_whitespace_out_of_the_url( + mcp_root_server, mock_api_client, mock_router +): + """A padded id is the same id, and a raw one would build a URL-escaped path.""" + mock_router.add("GET", QUERIES, json=jsonapi_collection([])) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_list_attack_paths_queries", {"scan_id": f" {SCAN_ID} "} + ) + + assert result.isError is False + assert mock_router.paths() == [f"GET {QUERIES}"] + + +async def test_a_blank_optional_value_is_rejected_rather_than_written( + mcp_root_server, mock_api_client, mock_router +): + """An omitted optional means "leave it alone"; a blank one would blank the field. + + The API refuses it, so the only difference an unguarded blank makes is a + round trip and an error that names nothing. + """ + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_update_mute_rule", {"rule_id": SCAN_ID, "name": ""} + ) + + assert result.isError is True + assert "name" in result.content[0].text + assert mock_router.requests == [] + + +async def test_an_omitted_optional_string_is_still_omitted( + mcp_root_server, mock_api_client, mock_router +): + """`NonBlankStr | None` must not turn "not provided" into a rejection.""" + mock_router.add( + "GET", + f"/api/v1/mute-rules/{SCAN_ID}", + json={ + "data": jsonapi_resource( + "mute-rules", + SCAN_ID, + { + "name": "unchanged", + "reason": "already reviewed", + "enabled": True, + "finding_uids": [], + }, + ) + }, + ) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_update_mute_rule", {"rule_id": SCAN_ID} + ) + + assert result.isError is False + + +async def test_every_required_string_argument_is_guarded_against_a_blank( + mcp_root_server, +): + """A guard only one tool carries is one the next tool will be written without. + + Declared as `minLength` rather than checked inside each tool, so a client sees + the constraint in the schema before it calls. + """ + async with Client(mcp_root_server) as client: + tools = await client.list_tools() + + unguarded = [ + f"{tool.name}.{name}" + for tool in tools + for name, schema in tool.inputSchema.get("properties", {}).items() + # Plain required strings only. A union such as `dict | str` takes a JSON + # string, where a blank is a parse failure the classifier already + # explains, and a blank filter is a filter that matches everything. + if schema.get("type") == "string" + and name in tool.inputSchema.get("required", []) + and schema.get("minLength") != 1 + ] + + assert unguarded == [] diff --git a/mcp_server/tests/prowler_app/tools/test_attack_paths.py b/mcp_server/tests/prowler_app/tools/test_attack_paths.py new file mode 100644 index 0000000000..3f20549796 --- /dev/null +++ b/mcp_server/tests/prowler_app/tools/test_attack_paths.py @@ -0,0 +1,225 @@ +"""Tests for the Attack Paths tools. + +An Attack Paths scan is a separate resource from a regular scan, with IDs of its +own, and Prowler only creates one for an AWS provider. So the 404 these tools get +is almost always a regular scan ID passed where an Attack Paths one belongs -- +and Prowler's own reason for it, a bare "Not found.", names neither the resource +it looked in nor the tool that returns the right ID. +""" + +import pytest +from fastmcp import Client + +from tests.helpers.jsonapi import jsonapi_collection, jsonapi_resource + +QUERIES = "/api/v1/attack-paths-scans/s1/queries" + + +async def test_an_id_that_is_not_an_attack_paths_scan_says_which_tool_returns_one( + mcp_root_server, mock_api_client, mock_router +): + """Relaying "Not found." sends an agent to re-check an ID it cannot fix. + + The reply has to name the confusion it stands for: regular scan IDs do not + resolve here, and only AWS providers have an Attack Paths scan at all. + """ + mock_router.add( + "GET", + QUERIES, + status=404, + json={"errors": [{"status": "404", "detail": "Not found."}]}, + ) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="different resource from regular scans"): + await client.call_tool( + "prowler_list_attack_paths_queries", {"scan_id": "s1"} + ) + + +async def test_the_answer_names_the_tool_that_returns_a_usable_id( + mcp_root_server, mock_api_client, mock_router +): + """An explanation with no next step leaves the agent guessing IDs.""" + mock_router.add( + "GET", + QUERIES, + status=404, + json={"errors": [{"status": "404", "detail": "Not found."}]}, + ) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="prowler_list_attack_paths_scans"): + await client.call_tool( + "prowler_list_attack_paths_queries", {"scan_id": "s1"} + ) + + +async def test_a_failure_that_is_not_a_404_keeps_the_shared_message( + mcp_root_server, mock_api_client, mock_router +): + """Only the 404 means a bad ID. A 403 is a permission the ID cannot fix.""" + mock_router.add( + "GET", + QUERIES, + status=403, + json={"errors": [{"status": "403", "detail": "Denied."}]}, + ) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="prowler_get_current_user"): + await client.call_tool( + "prowler_list_attack_paths_queries", {"scan_id": "s1"} + ) + + +async def test_queries_come_back_as_a_list( + mcp_root_server, mock_api_client, mock_router +): + """The success path is unchanged.""" + mock_router.add( + "GET", + QUERIES, + json=jsonapi_collection( + [ + jsonapi_resource( + "attack-paths-queries", + "aws-ec2-instances-internet-exposed", + { + "name": "Internet exposed EC2", + "description": "Find internet-exposed EC2 instances", + "provider": "aws", + "parameters": [], + }, + ) + ] + ), + ) + + async with Client(mcp_root_server) as client: + result = await client.call_tool( + "prowler_list_attack_paths_queries", {"scan_id": "s1"} + ) + + assert result.data[0]["id"] == "aws-ec2-instances-internet-exposed" + + +# ------------------------------------------------------------- running a query + +RUN = "/api/v1/attack-paths-scans/s1/queries/run" +SCHEMA = "/api/v1/attack-paths-scans/s1/schema" + +EMPTY_RESULT = { + "data": { + "type": "attack-paths-query-results", + "id": "s1", + "attributes": {"nodes": [], "relationships": []}, + } +} + + +def _run_args(query_id: str = "aws-ec2-instances-internet-exposed") -> dict[str, str]: + """Arguments for a query run against the mocked scan.""" + return {"scan_id": "s1", "query_id": query_id} + + +async def test_a_query_that_matched_nothing_is_an_answer_not_a_failure( + mcp_root_server, mock_api_client, mock_router +): + """Prowler answers a query that matched nothing with 404 and the result body. + + Raising on the status called a clean account a failed call and sent the agent + off to re-check arguments that were right. + """ + mock_router.add("POST", RUN, status=404, json=EMPTY_RESULT) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_run_attack_paths_query", _run_args() + ) + + assert result.isError is False + assert "matched nothing" in result.structuredContent["message"] + + +async def test_an_empty_result_does_not_come_back_as_an_empty_object( + mcp_root_server, mock_api_client, mock_router +): + """The serializer drops empty lists, so `{}` is all that would be left.""" + mock_router.add("POST", RUN, status=404, json=EMPTY_RESULT) + + async with Client(mcp_root_server) as client: + result = await client.call_tool("prowler_run_attack_paths_query", _run_args()) + + assert result.data != {} + + +async def test_a_null_graph_is_read_as_an_empty_one( + mcp_root_server, mock_api_client, mock_router +): + """Prowler can spell the empty graph as `null` rather than as empty lists. + + Reading `null` as if it were a graph crashed the parse, turning the same + "nothing matched" answer into an error the agent could not act on. + """ + mock_router.add("POST", RUN, status=404, json={"data": {"attributes": None}}) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_run_attack_paths_query", _run_args() + ) + + assert result.isError is False + assert "matched nothing" in result.structuredContent["message"] + + +async def test_a_run_against_an_unknown_scan_still_names_the_confusion( + mcp_root_server, mock_api_client, mock_router +): + """A 404 with no result body is the ID being wrong, not an empty answer.""" + mock_router.add( + "POST", + RUN, + status=404, + json={"errors": [{"status": "404", "detail": "Not found."}]}, + ) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="different resource from regular scans"): + await client.call_tool("prowler_run_attack_paths_query", _run_args()) + + +async def test_a_scan_whose_graph_records_no_schema_says_so( + mcp_root_server, mock_api_client, mock_router +): + """This 404 is about the graph, not the ID, so it must not blame the ID.""" + mock_router.add( + "GET", + SCHEMA, + status=404, + json={"detail": "No cartography schema metadata found for this provider"}, + ) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="no Cartography schema recorded"): + await client.call_tool( + "prowler_get_attack_paths_cartography_schema", {"scan_id": "s1"} + ) + + +async def test_a_schema_request_for_an_unknown_scan_names_the_confusion( + mcp_root_server, mock_api_client, mock_router +): + """The other 404 here is the ID, and Prowler writes a JSON:API error for it.""" + mock_router.add( + "GET", + SCHEMA, + status=404, + json={"errors": [{"status": "404", "detail": "Not found."}]}, + ) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="different resource from regular scans"): + await client.call_tool( + "prowler_get_attack_paths_cartography_schema", {"scan_id": "s1"} + ) diff --git a/mcp_server/tests/prowler_app/tools/test_compliance.py b/mcp_server/tests/prowler_app/tools/test_compliance.py new file mode 100644 index 0000000000..373e1d7f6c --- /dev/null +++ b/mcp_server/tests/prowler_app/tools/test_compliance.py @@ -0,0 +1,102 @@ +"""Tests for the compliance tools. + +Both compliance tools answer for exactly one scan. ``scan_id`` names it +directly; ``provider_id`` names it indirectly, as "the latest completed scan of +this provider". Passing both is not a refinement of either -- the scan the +caller named may belong to a different provider entirely -- so it is rejected +rather than resolved by preferring one, which would answer confidently for a +provider nobody asked about. + +Tools are driven through an in-memory MCP client so FastMCP resolves the +pydantic ``Field`` defaults and a raised failure arrives the way a client sees +it. +""" + +import pytest +from fastmcp import Client + +from tests.helpers.jsonapi import jsonapi_collection, jsonapi_resource + +SCANS = "/api/v1/scans" +OVERVIEWS = "/api/v1/compliance-overviews" +REQUIREMENTS = f"{OVERVIEWS}/requirements" + +TOOLS = [ + "prowler_get_compliance_overview", + "prowler_get_compliance_framework_state_details", +] + + +def arguments(tool: str, **overrides) -> dict: + """Build the arguments for either tool, which differ only in compliance_id.""" + payload = dict(overrides) + if tool.endswith("framework_state_details"): + payload["compliance_id"] = "cis_1.5_aws" + return payload + + +@pytest.mark.parametrize("tool", TOOLS) +async def test_neither_a_scan_nor_a_provider_is_refused_before_any_request( + mcp_root_server, mock_api_client, mock_router, tool +): + """There is no scan to answer for, and no way to guess one.""" + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="must be provided"): + await client.call_tool(tool, arguments(tool)) + + assert mock_router.paths() == [] + + +@pytest.mark.parametrize("tool", TOOLS) +async def test_a_scan_and_a_provider_together_are_refused_rather_than_reconciled( + mcp_root_server, mock_api_client, mock_router, tool +): + """Silently keeping the scan would answer for whichever provider owns it. + + That report names a scan the caller did ask for, so nothing about it looks + wrong -- while the provider they also named went unread. + """ + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="not both"): + await client.call_tool( + tool, arguments(tool, scan_id="s1", provider_id="p1") + ) + + assert mock_router.paths() == [] + + +@pytest.mark.parametrize("tool", TOOLS) +async def test_a_provider_on_its_own_resolves_to_its_latest_completed_scan( + mcp_root_server, mock_api_client, mock_router, tool +): + """The indirection is the point of accepting a provider at all.""" + mock_router.add( + "GET", + SCANS, + json=jsonapi_collection( + [jsonapi_resource("scans", "s9", {"state": "completed"})] + ), + ) + mock_router.add("GET", OVERVIEWS, json=jsonapi_collection([])) + mock_router.add("GET", REQUIREMENTS, json=jsonapi_collection([])) + # Each tool reads the compliance state from its own endpoint; both filter it + # by the scan that had to be resolved first. + read = OVERVIEWS if tool.endswith("overview") else REQUIREMENTS + + async with Client(mcp_root_server) as client: + await client.call_tool(tool, arguments(tool, provider_id="p1")) + + assert mock_router.query_params("GET", SCANS)["filter[provider]"] == "p1" + assert mock_router.query_params("GET", read)["filter[scan_id]"] == "s9" + + +@pytest.mark.parametrize("tool", TOOLS) +async def test_a_provider_with_no_completed_scan_is_named_as_the_bad_argument( + mcp_root_server, mock_api_client, mock_router, tool +): + """Nothing has been scanned yet, so there is no compliance state to report.""" + mock_router.add("GET", SCANS, json=jsonapi_collection([])) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="No completed scans found for provider p1"): + await client.call_tool(tool, arguments(tool, provider_id="p1")) diff --git a/mcp_server/tests/prowler_app/tools/test_integrations.py b/mcp_server/tests/prowler_app/tools/test_integrations.py index 9d165a60c6..d94dd8a21d 100644 --- a/mcp_server/tests/prowler_app/tools/test_integrations.py +++ b/mcp_server/tests/prowler_app/tools/test_integrations.py @@ -312,17 +312,16 @@ async def test_creating_a_jira_integration_rejects_an_empty_domain( ): """A domain that normalizes to nothing is caught before the round trip.""" async with Client(mcp_root_server) as client: - result = await client.call_tool( - "prowler_create_jira_integration", - { - "domain": "https://", - "user_mail": "security@acme.com", - "api_token": "fake-atlassian-token-for-testing", - }, - ) + with pytest.raises(Exception, match="Invalid Jira domain"): + await client.call_tool( + "prowler_create_jira_integration", + { + "domain": "https://", + "user_mail": "security@acme.com", + "api_token": "fake-atlassian-token-for-testing", + }, + ) - assert result.data["status"] == "failed" - assert "Invalid Jira domain" in result.data["error"] assert mock_router.requests == [] @@ -334,14 +333,10 @@ async def test_creating_a_jira_integration_rejects_an_empty_domain( ], ids=["security-hub", "amazon-s3"], ) -async def test_a_rejected_creation_is_reported_rather_than_raised( +async def test_a_rejected_creation_fails_with_the_api_reason( mcp_root_server, mock_api_client, mock_router, tool, arguments ): - """Write tools answer with an error object so the agent can act on it. - - A raised exception reaches the model as a tool failure with no detail, and - the API's message is exactly what tells it what to do next. - """ + """A refused creation is a tool error, and it still carries the API's reason.""" mock_router.add( "POST", INTEGRATIONS, @@ -350,10 +345,8 @@ async def test_a_rejected_creation_is_reported_rather_than_raised( ) async with Client(mcp_root_server) as client: - result = await client.call_tool(tool, arguments) - - assert result.data["status"] == "failed" - assert "already has this integration" in result.data["error"] + with pytest.raises(Exception, match="already has this integration"): + await client.call_tool(tool, arguments) async def test_a_creation_with_no_id_back_warns_before_a_blind_retry( @@ -367,12 +360,11 @@ async def test_a_creation_with_no_id_back_warns_before_a_blind_retry( mock_router.add("POST", INTEGRATIONS, json={"data": {}}) async with Client(mcp_root_server) as client: - result = await client.call_tool( - "prowler_create_amazon_s3_integration", {"bucket_name": "my-reports"} - ) + with pytest.raises(Exception, match="did not return its ID"): + await client.call_tool( + "prowler_create_amazon_s3_integration", {"bucket_name": "my-reports"} + ) - assert result.data["status"] == "failed" - assert "did not return its ID" in result.data["error"] assert mock_router.paths() == [f"POST {INTEGRATIONS}"] @@ -395,12 +387,10 @@ async def test_a_creation_whose_read_back_fails_still_hands_over_the_id( ) async with Client(mcp_root_server) as client: - result = await client.call_tool( - "prowler_create_amazon_s3_integration", {"bucket_name": "my-reports"} - ) - - assert result.data["status"] == "failed" - assert "Integration i1 was created" in result.data["error"] + with pytest.raises(Exception, match="Integration i1 was created"): + await client.call_tool( + "prowler_create_amazon_s3_integration", {"bucket_name": "my-reports"} + ) async def test_a_connection_check_that_cannot_run_is_not_reported_as_a_failure( @@ -542,13 +532,12 @@ async def test_a_configuration_that_is_not_an_object_is_rejected_before_the_writ stub_integration(mock_router, S3_ATTRIBUTES) async with Client(mcp_root_server) as client: - result = await client.call_tool( - "prowler_update_integration", - {"integration_id": "i1", "configuration": configuration}, - ) + with pytest.raises(Exception, match=message): + await client.call_tool( + "prowler_update_integration", + {"integration_id": "i1", "configuration": configuration}, + ) - assert result.data["status"] == "failed" - assert message in result.data["error"] assert f"PATCH {INTEGRATION}" not in mock_router.paths() @@ -643,13 +632,12 @@ async def test_updating_a_jira_configuration_is_refused( stub_integration(mock_router, JIRA_ATTRIBUTES) async with Client(mcp_root_server) as client: - result = await client.call_tool( - "prowler_update_integration", - {"integration_id": "i1", "configuration": {"domain": "other"}}, - ) + with pytest.raises(Exception, match="do not accept a configuration"): + await client.call_tool( + "prowler_update_integration", + {"integration_id": "i1", "configuration": {"domain": "other"}}, + ) - assert result.data["status"] == "failed" - assert "do not accept a configuration" in result.data["error"] assert f"PATCH {INTEGRATION}" not in mock_router.paths() @@ -660,12 +648,12 @@ async def test_attaching_a_jira_integration_to_a_provider_is_refused( stub_integration(mock_router, JIRA_ATTRIBUTES) async with Client(mcp_root_server) as client: - result = await client.call_tool( - "prowler_update_integration", - {"integration_id": "i1", "provider_ids": ["p1"]}, - ) + with pytest.raises(Exception, match="tenant-wide"): + await client.call_tool( + "prowler_update_integration", + {"integration_id": "i1", "provider_ids": ["p1"]}, + ) - assert "tenant-wide" in result.data["error"] assert f"PATCH {INTEGRATION}" not in mock_router.paths() @@ -683,12 +671,12 @@ async def test_security_hub_must_keep_exactly_one_provider( stub_integration(mock_router, SECURITY_HUB_ATTRIBUTES, provider_ids=("p1",)) async with Client(mcp_root_server) as client: - result = await client.call_tool( - "prowler_update_integration", - {"integration_id": "i1", "provider_ids": provider_ids}, - ) + with pytest.raises(Exception, match="exactly one AWS provider"): + await client.call_tool( + "prowler_update_integration", + {"integration_id": "i1", "provider_ids": provider_ids}, + ) - assert "exactly one AWS provider" in result.data["error"] assert f"PATCH {INTEGRATION}" not in mock_router.paths() @@ -708,12 +696,12 @@ async def test_partial_jira_credentials_are_refused_to_protect_the_stored_ones( stub_integration(mock_router, JIRA_ATTRIBUTES) async with Client(mcp_root_server) as client: - result = await client.call_tool( - "prowler_update_integration", - {"integration_id": "i1", "credentials": credentials}, - ) + with pytest.raises(Exception, match="replaced as a whole"): + await client.call_tool( + "prowler_update_integration", + {"integration_id": "i1", "credentials": credentials}, + ) - assert "replaced as a whole" in result.data["error"] assert f"PATCH {INTEGRATION}" not in mock_router.paths() @@ -751,13 +739,14 @@ async def test_replacing_jira_credentials_normalizes_the_domain( # ------------------------------------------------- delete and connection tools -async def test_deleting_an_integration_reports_the_outcome_either_way( +async def test_deleting_an_integration_confirms_it_happened( mcp_root_server, mock_api_client, mock_router ): - """Deletion is irreversible, so both outcomes are stated explicitly. + """Deletion is irreversible, so a success says so rather than staying silent. - A bare exception would leave the agent unsure whether the credentials are - gone, and a retry of a delete that actually succeeded reads as a new failure. + It says so in the message and nowhere else: a `deleted: true` flag could only + ever be true, because an integration that was not deleted leaves the tool as + an error. """ mock_router.add("DELETE", INTEGRATION, status=204) @@ -766,24 +755,23 @@ async def test_deleting_an_integration_reports_the_outcome_either_way( "prowler_delete_integration", {"integration_id": "i1"} ) - assert result.data["deleted"] is True + assert "i1 deleted successfully" in result.data["message"] + assert "deleted" not in result.data -async def test_a_failed_deletion_says_it_did_not_happen( +async def test_a_refused_deletion_fails_and_says_the_role_is_the_problem( mcp_root_server, mock_api_client, mock_router ): - """`deleted: false` is the part the agent must not have to infer.""" + """A 403 is the same answer for every tool, so `lib.errors` writes it.""" mock_router.add( "DELETE", INTEGRATION, status=403, json=jsonapi_error(403, "Permission denied.") ) async with Client(mcp_root_server) as client: - result = await client.call_tool( - "prowler_delete_integration", {"integration_id": "i1"} - ) - - assert result.data["deleted"] is False - assert "Permission denied." in result.data["message"] + with pytest.raises(Exception, match="prowler_get_current_user"): + await client.call_tool( + "prowler_delete_integration", {"integration_id": "i1"} + ) async def test_checking_a_connection_surfaces_why_it_failed( @@ -1026,6 +1014,32 @@ async def test_an_accepted_dispatch_with_no_task_id_is_not_safe_to_retry( assert "task_id" not in result.data +async def test_a_dispatch_with_no_usable_credential_is_raised_not_called_unknown( + mcp_root_server, mock_api_client, mock_router, monkeypatch +): + """Authentication runs before the request, so nothing was ever queued. + + Reported as `unknown` it reads as "work items may exist in Jira, go and + look" -- for a call that never reached Prowler. It is not a dispatch outcome + at all: the credential has to be fixed, and no retry of this call does that. + """ + monkeypatch.setattr(mock_api_client.auth_manager, "mode", "http") + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="no usable credential"): + await client.call_tool( + "prowler_send_findings_to_jira", + { + "integration_id": "i1", + "project_key": "PROJ", + "issue_type": "Task", + "finding_ids": ["f1"], + }, + ) + + assert mock_router.paths() == [] + + async def test_a_dispatch_task_that_died_halfway_is_never_safe_to_retry( mcp_root_server, mock_api_client, mock_router ): diff --git a/mcp_server/tests/prowler_app/tools/test_muting.py b/mcp_server/tests/prowler_app/tools/test_muting.py new file mode 100644 index 0000000000..7c980600e3 --- /dev/null +++ b/mcp_server/tests/prowler_app/tools/test_muting.py @@ -0,0 +1,64 @@ +"""Tests for the muting tools. + +Muting is permanent and deleting a rule does not undo it, so the one thing +these assertions protect is that an agent is never told a deletion failed when +it did not: the old answer keyed off the *shape* of the API's reply rather than +off anything having gone wrong, and said nothing a caller could act on. +""" + +import pytest +from fastmcp import Client + +from tests.helpers.jsonapi import jsonapi_document, jsonapi_error, jsonapi_resource + +MUTE_RULE = "/api/v1/mute-rules/m1" + + +async def test_a_deleted_rule_is_reported_deleted_whatever_the_body( + mcp_root_server, mock_api_client, mock_router +): + """Prowler answers 204 with no body, but 200 with one is just as much a yes. + + The old check read ``success`` out of the parsed body, which only exists for + the empty-body case, so a deletion that worked could be reported as + "Failed to delete mute rule". + """ + mock_router.add( + "DELETE", + MUTE_RULE, + json=jsonapi_document(jsonapi_resource("mute-rules", "m1", {})), + ) + + async with Client(mcp_root_server) as client: + result = await client.call_tool("prowler_delete_mute_rule", {"rule_id": "m1"}) + + assert "m1 deleted successfully" in result.data["message"] + assert "stay muted" in result.data["message"] + # A flag with one reachable value is not a fact, it is an invitation to + # branch on a shape that does not exist. + assert "success" not in result.data + + +async def test_an_empty_body_deletion_is_reported_the_same_way( + mcp_root_server, mock_api_client, mock_router +): + """The 204 path, which is what Prowler actually sends today.""" + mock_router.add("DELETE", MUTE_RULE, status=204) + + async with Client(mcp_root_server) as client: + result = await client.call_tool("prowler_delete_mute_rule", {"rule_id": "m1"}) + + assert "m1 deleted successfully" in result.data["message"] + + +async def test_a_refused_deletion_is_an_error( + mcp_root_server, mock_api_client, mock_router +): + """A rule that was not deleted has to leave the tool as an error.""" + mock_router.add( + "DELETE", MUTE_RULE, status=404, json=jsonapi_error(404, "Not found.") + ) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="Not found"): + await client.call_tool("prowler_delete_mute_rule", {"rule_id": "m1"}) diff --git a/mcp_server/tests/prowler_app/tools/test_providers.py b/mcp_server/tests/prowler_app/tools/test_providers.py new file mode 100644 index 0000000000..b3438a79a7 --- /dev/null +++ b/mcp_server/tests/prowler_app/tools/test_providers.py @@ -0,0 +1,316 @@ +"""Tests for the provider tools. + +Two behaviours drive most of the assertions here, and both are about telling a +failure apart from an outcome that is merely not final yet: + +* Deleting a provider removes it together with its scans, findings and + resources, in a background task that routinely outlives the polling window. + A deletion still running is not a failure, and reporting it as one invites a + retry of a destructive call that is already in flight. +* ``connect_provider`` runs a connection check, and a check that could not be + run says nothing about the provider's credentials. Reporting it as ``failed`` + blames an AWS role for an expired Prowler credential and sends the user off to + fix something that works. + +Tools are driven through an in-memory MCP client so FastMCP resolves the +pydantic ``Field`` defaults and a raised failure arrives the way a client sees +it. +""" + +import pytest +from fastmcp import Client + +from tests.helpers.http import MockRouter +from tests.helpers.jsonapi import ( + jsonapi_collection, + jsonapi_document, + jsonapi_error, + jsonapi_resource, + task_document, +) + +PROVIDERS = "/api/v1/providers" +PROVIDER = f"{PROVIDERS}/p1" +CONNECTION = f"{PROVIDER}/connection" +SECRETS = f"{PROVIDERS}/secrets" +TASK = "/api/v1/tasks/t1" + +PROVIDER_ATTRIBUTES = { + "uid": "123456789012", + "provider": "aws", + "alias": "production", + "connection": {"connected": True}, +} + + +@pytest.fixture +def mock_fast_polling(monkeypatch, api_client): + """Run the real polling loop, with a timeout a test can afford to wait out. + + The timeout path is the one worth testing here -- it is what used to be + reported as a failed deletion -- so the loop, its exception and the fallback + that reads the task afterwards all stay real. Only the 60 seconds go. + """ + original = type(api_client).poll_task_until_complete + + async def _fast(self, task_id, **_overridden): + return await original(self, task_id, timeout=0.3, poll_interval=0.05) + + monkeypatch.setattr(type(api_client), "poll_task_until_complete", _fast) + + +@pytest.fixture +def mock_polling_timeout(monkeypatch, api_client): + """Make the polling window run out on the first call, without the wait. + + The clock is what ends the polling loop here, so a test about *what happens + afterwards* has no reason to spend it. The read the fallback then makes is + the real one. + """ + + async def _timeout(self, task_id, **_overridden): + raise TimeoutError(f"Task {task_id} polling timed out after 60 seconds.") + + monkeypatch.setattr(type(api_client), "poll_task_until_complete", _timeout) + + +def stub_deletion_start(mock_router: MockRouter) -> MockRouter: + """Serve the DELETE as Prowler does: a task to poll, not a finished deletion.""" + return mock_router.add( + "DELETE", PROVIDER, json=jsonapi_document(jsonapi_resource("tasks", "t1", {})) + ) + + +# --------------------------------------------------------------- deletion + + +async def test_a_refused_deletion_is_an_error_not_a_result( + mcp_root_server, mock_api_client, mock_router +): + """Nothing started, so the classifier owns the message. + + Returned as ``{"deleted": false}`` it arrives with ``isError: false`` and a + model has no reason to treat it as anything but a completed call. + """ + mock_router.add( + "DELETE", PROVIDER, status=403, json=jsonapi_error(403, "Not allowed.") + ) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="not allowed to do this"): + await client.call_tool("prowler_delete_provider", {"provider_id": "p1"}) + + +async def test_a_deletion_with_no_task_back_warns_before_a_blind_retry( + mcp_root_server, mock_api_client, mock_router +): + """Prowler accepted it, so the deletion is probably running. + + Without the task ID there is nothing to watch it with, which makes "check + whether it is gone" the only safe instruction. + """ + mock_router.add("DELETE", PROVIDER, json={"data": {}}) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="did not return the ID"): + await client.call_tool("prowler_delete_provider", {"provider_id": "p1"}) + + assert mock_router.paths() == [f"DELETE {PROVIDER}"] + + +async def test_a_finished_deletion_reports_it_plainly( + mcp_root_server, mock_api_client, mock_router +): + """The success path is unchanged: the provider is gone.""" + stub_deletion_start(mock_router) + mock_router.add("GET", TASK, json=task_document("t1", "completed")) + + async with Client(mcp_root_server) as client: + result = await client.call_tool( + "prowler_delete_provider", {"provider_id": "p1"} + ) + + assert result.data["status"] == "deleted" + + +async def test_a_deletion_still_running_is_not_reported_as_a_failure( + mcp_root_server, mock_api_client, mock_router, mock_fast_polling +): + """Outliving the polling window is normal for a provider with many findings. + + The task ID goes back so the deletion can be followed, and the message says + not to send it again -- which is the whole point of not calling this failed. + """ + stub_deletion_start(mock_router) + mock_router.add("GET", TASK, json=task_document("t1", "executing")) + + async with Client(mcp_root_server) as client: + result = await client.call_tool( + "prowler_delete_provider", {"provider_id": "p1"} + ) + + assert result.data["status"] == "in_progress" + assert result.data["task_id"] == "t1" + assert "Do not send the deletion again" in result.data["message"] + + +async def test_a_deletion_that_finished_just_after_the_wait_is_reported_as_deleted( + mcp_root_server, mock_api_client, mock_router, mock_polling_timeout +): + """Polling gives up on the clock, not on the task. + + A deletion that completed a moment after the last poll is a finished + deletion, and the read the fallback makes is what says so. Reporting it as + still running would send the caller off to watch a provider that is gone. + """ + stub_deletion_start(mock_router) + mock_router.add("GET", TASK, json=task_document("t1", "completed")) + + async with Client(mcp_root_server) as client: + result = await client.call_tool( + "prowler_delete_provider", {"provider_id": "p1"} + ) + + assert result.data["status"] == "deleted" + assert "task_id" not in result.data + + +async def test_a_deletion_task_that_stopped_is_an_error_naming_what_is_left( + mcp_root_server, mock_api_client, mock_router +): + """Here the provider really is still there, so this one is a failure. + + A provider is removed together with everything attached to it, so a task + that stopped halfway can leave part of that gone -- which is why the message + sends the caller to look rather than asserting the state. + """ + stub_deletion_start(mock_router) + mock_router.add("GET", TASK, json=task_document("t1", "failed", error="boom")) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="ended as 'failed'"): + await client.call_tool("prowler_delete_provider", {"provider_id": "p1"}) + + +async def test_a_deletion_task_failure_does_not_relay_the_upstream_text( + mcp_root_server, mock_api_client, mock_router +): + """The task's own error is a celery traceback; it stays in the log.""" + stub_deletion_start(mock_router) + mock_router.add( + "GET", + TASK, + json=task_document("t1", "failed", error="Traceback: secret-internal-detail"), + ) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception) as raised: + await client.call_tool("prowler_delete_provider", {"provider_id": "p1"}) + + assert "secret-internal-detail" not in str(raised.value) + + +async def test_a_deletion_whose_progress_cannot_be_read_still_says_do_not_retry( + mcp_root_server, mock_api_client, mock_router, mock_fast_polling +): + """The outcome is unknown, which for a destructive call means: do not repeat. + + Reading the task is what tells "still running" from "stopped", so when that + read fails too the message drops the claim rather than guessing at one. + """ + stub_deletion_start(mock_router) + mock_router.add("GET", TASK, status=503, json=jsonapi_error(503, "Unavailable.")) + + async with Client(mcp_root_server) as client: + result = await client.call_tool( + "prowler_delete_provider", {"provider_id": "p1"} + ) + + assert result.data["status"] == "in_progress" + assert "could not be read" in result.data["message"] + assert "Do not send the deletion again" in result.data["message"] + # The failure that got us here is the classifier's to phrase, so its raw + # text stays in the log rather than riding along in the message. + assert "API request failed" not in result.data["message"] + + +# ------------------------------------------------------- connection check + + +async def test_a_connection_check_that_cannot_run_is_not_reported_as_failed( + mcp_root_server, mock_api_client, mock_router +): + """`not_tested` says nothing about the credentials, and that is the point. + + A 401 here is this server's own credential, not the provider's. Calling it + `failed` tells the user their AWS role is broken when it is fine. + """ + # Registered in order and consumed in order: the lookup before the create + # finds nothing, the one after it finds the provider that was just made. + mock_router.add("GET", PROVIDERS, json=jsonapi_collection([])) + mock_router.add( + "GET", + PROVIDERS, + json=jsonapi_collection( + [jsonapi_resource("providers", "p1", PROVIDER_ATTRIBUTES)] + ), + ) + mock_router.add( + "POST", + PROVIDERS, + json=jsonapi_document(jsonapi_resource("providers", "p1", PROVIDER_ATTRIBUTES)), + ) + mock_router.add( + "POST", CONNECTION, status=401, json=jsonapi_error(401, "Token expired.") + ) + mock_router.add( + "GET", + PROVIDER, + json=jsonapi_document(jsonapi_resource("providers", "p1", PROVIDER_ATTRIBUTES)), + ) + + async with Client(mcp_root_server) as client: + result = await client.call_tool( + "prowler_connect_provider", + {"provider_uid": "123456789012", "provider_type": "aws"}, + ) + + assert result.data["connected"] == "not_tested" + assert "never tested" in result.data["error"] + + +async def test_a_secret_lookup_failure_does_not_pass_as_having_no_secret( + mcp_root_server, mock_api_client, mock_router +): + """Returning None here would send the write down the create branch. + + A provider holds at most one secret, so creating a second one is refused and + the caller would be told its credentials were rejected when all that + actually failed was this read. + """ + mock_router.add( + "GET", + PROVIDERS, + json=jsonapi_collection( + [jsonapi_resource("providers", "p1", PROVIDER_ATTRIBUTES)] + ), + ) + mock_router.add( + "GET", SECRETS, status=429, json=jsonapi_error(429, "Too many requests.") + ) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="rate limiting"): + await client.call_tool( + "prowler_connect_provider", + { + "provider_uid": "123456789012", + "provider_type": "aws", + "credentials": { + "aws_access_key_id": "AKIA", + "aws_secret_access_key": "s", + }, + }, + ) + + assert f"POST {SECRETS}" not in mock_router.paths() diff --git a/mcp_server/tests/prowler_app/tools/test_roles.py b/mcp_server/tests/prowler_app/tools/test_roles.py new file mode 100644 index 0000000000..f792cf0a04 --- /dev/null +++ b/mcp_server/tests/prowler_app/tools/test_roles.py @@ -0,0 +1,147 @@ +"""Tests for the role (RBAC) tools. + +``prowler_set_user_role`` replaces the single role a user holds, and the API +silently drops a role ID that does not exist in the tenant -- which would leave +the user with no role at all. So the tool reads the role first, and what that +read says has to be told apart carefully: only a not-found is about the role ID +the caller passed. A permission error, a rate limit or a server error is about +the request, and reporting either as "find a valid role ID" sends the user to +fix an ID that was fine. + +Tools are driven through an in-memory MCP client so FastMCP resolves the +pydantic ``Field`` defaults and a raised failure arrives the way a client sees +it. +""" + +import pytest +from fastmcp import Client + +from tests.helpers.http import MockRouter +from tests.helpers.jsonapi import ( + jsonapi_document, + jsonapi_error, + jsonapi_resource, +) + +USER = "/api/v1/users/u1" +USER_ROLES = f"{USER}/relationships/roles" +ROLE = "/api/v1/roles/r2" + +ROLE_ATTRIBUTES = {"name": "admin", "manage_account": True} + + +def stub_user_holding(mock_router: MockRouter, role_id: str) -> MockRouter: + """Serve ``GET /users/u1?include=roles`` with the user holding one role.""" + return mock_router.add( + "GET", + USER, + json=jsonapi_document( + jsonapi_resource("users", "u1", {"name": "Ada"}), + included=[jsonapi_resource("roles", role_id, ROLE_ATTRIBUTES)], + ), + ) + + +async def test_setting_a_role_the_user_already_holds_changes_nothing( + mcp_root_server, mock_api_client, mock_router +): + """Idempotent by design: no PATCH goes out, so nothing can be replaced.""" + stub_user_holding(mock_router, "r2") + + async with Client(mcp_root_server) as client: + result = await client.call_tool( + "prowler_set_user_role", {"user_id": "u1", "role_id": "r2"} + ) + + assert result.data["changed"] is False + assert mock_router.paths() == [f"GET {USER}"] + + +async def test_a_role_that_does_not_exist_is_named_as_the_bad_argument( + mcp_root_server, mock_api_client, mock_router +): + """404 is the one answer that really is about the role ID. + + The PATCH would accept the ID and drop it, leaving the user with no role, so + the read has to stop the call -- and say which ID to replace. + """ + stub_user_holding(mock_router, "r1") + mock_router.add("GET", ROLE, status=404, json=jsonapi_error(404, "Not found.")) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="does not exist in this tenant") as raised: + await client.call_tool( + "prowler_set_user_role", {"user_id": "u1", "role_id": "r2"} + ) + + assert "left unchanged" in str(raised.value) + assert f"PATCH {USER_ROLES}" not in mock_router.paths() + + +@pytest.mark.parametrize( + ("status", "detail", "expected"), + [ + (403, "Not allowed.", "not allowed to do this"), + (429, "Slow down.", "rate limiting"), + (500, "Boom.", "failed on Prowler's side"), + ], + ids=["forbidden", "rate-limited", "server-error"], +) +async def test_a_role_read_that_failed_for_another_reason_is_not_a_bad_role_id( + mcp_root_server, mock_api_client, mock_router, status, detail, expected +): + """These say nothing about the ID, so the classifier owns the message. + + Told "use prowler_list_roles to find a valid role ID", an agent goes looking + for a role that was never the problem -- and finds the same wall. + """ + stub_user_holding(mock_router, "r1") + mock_router.add("GET", ROLE, status=status, json=jsonapi_error(status, detail)) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match=expected) as raised: + await client.call_tool( + "prowler_set_user_role", {"user_id": "u1", "role_id": "r2"} + ) + + assert "valid role ID" not in str(raised.value) + assert f"PATCH {USER_ROLES}" not in mock_router.paths() + + +async def test_a_set_role_reports_the_role_the_user_holds_afterwards( + mcp_root_server, mock_api_client, mock_router +): + """The authoritative state is read back rather than assumed from the PATCH.""" + mock_router.add( + "GET", + USER, + json=jsonapi_document( + jsonapi_resource("users", "u1", {"name": "Ada"}), + included=[jsonapi_resource("roles", "r1", ROLE_ATTRIBUTES)], + ), + ) + mock_router.add( + "GET", + USER, + json=jsonapi_document( + jsonapi_resource("users", "u1", {"name": "Ada"}), + included=[jsonapi_resource("roles", "r2", ROLE_ATTRIBUTES)], + ), + ) + mock_router.add( + "GET", + ROLE, + json=jsonapi_document(jsonapi_resource("roles", "r2", ROLE_ATTRIBUTES)), + ) + mock_router.add("PATCH", USER_ROLES, status=204) + + async with Client(mcp_root_server) as client: + result = await client.call_tool( + "prowler_set_user_role", {"user_id": "u1", "role_id": "r2"} + ) + + assert result.data["changed"] is True + assert [role["id"] for role in result.data["roles"]] == ["r2"] + assert mock_router.json_body("PATCH", USER_ROLES) == { + "data": [{"type": "roles", "id": "r2"}] + } diff --git a/mcp_server/tests/prowler_app/tools/test_scans.py b/mcp_server/tests/prowler_app/tools/test_scans.py new file mode 100644 index 0000000000..93e18c9a69 --- /dev/null +++ b/mcp_server/tests/prowler_app/tools/test_scans.py @@ -0,0 +1,201 @@ +"""Tests for the scans tools. + +Both write tools here used to answer a failure with a result object that the +protocol, the client and the model all read as a success, and both are calls +whose outcome an agent acts on: + +* ``prowler_trigger_scan`` starts work. Anything that reads as "nothing + happened" invites a second scan of the same provider. +* ``prowler_schedule_daily_scan`` was deciding whether the schedule existed by + reading the state of a different thing entirely -- the first scan Prowler + starts alongside it -- so a schedule that had just been created could be + reported as a failure. Retrying that can only hit the 409 the API raises for a + provider that already has one. + +Tools are driven through an in-memory MCP client so FastMCP resolves the +pydantic ``Field`` defaults and a raised failure arrives the way a client sees +it. +""" + +import pytest +from fastmcp import Client + +from tests.helpers.http import MockRouter +from tests.helpers.jsonapi import ( + jsonapi_document, + jsonapi_error, + jsonapi_resource, +) + +SCANS = "/api/v1/scans" +SCAN = f"{SCANS}/s1" +DAILY = "/api/v1/schedules/daily" + +SCAN_ATTRIBUTES = { + "name": "Nightly", + "trigger": "manual", + "state": "executing", + "progress": 40, +} + + +def stub_scan_creation(mock_router: MockRouter, scan_id: str = "s1") -> MockRouter: + """Serve the creation as Prowler does: a task carrying the new scan's ID.""" + return mock_router.add( + "POST", + SCANS, + json=jsonapi_document( + jsonapi_resource("tasks", "t1", {"task_args": {"scan_id": scan_id}}) + ), + ) + + +# ------------------------------------------------------------- trigger_scan + + +async def test_a_refused_scan_is_an_error_not_a_failed_looking_result( + mcp_root_server, mock_api_client, mock_router +): + """A rejection has to leave the tool as an error. + + Returned as a result it arrives with ``isError: false``, and a model reading + a successful tool call has no reason to doubt that a scan is now running. + """ + mock_router.add( + "POST", SCANS, status=403, json=jsonapi_error(403, "Insufficient permissions.") + ) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="not allowed to do this"): + await client.call_tool("prowler_trigger_scan", {"provider_id": "p1"}) + + +async def test_a_scan_with_no_id_back_warns_before_a_blind_retry( + mcp_root_server, mock_api_client, mock_router +): + """Prowler accepted it, so a second call could start a duplicate scan. + + The message names the provider because that is what makes the suggested + check actionable without another lookup. + """ + mock_router.add("POST", SCANS, json={"data": {"attributes": {"task_args": {}}}}) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="did not return its ID"): + await client.call_tool("prowler_trigger_scan", {"provider_id": "p1"}) + + assert mock_router.paths() == [f"POST {SCANS}"] + + +async def test_a_scan_whose_read_back_fails_still_hands_over_the_id( + mcp_root_server, mock_api_client, mock_router +): + """The scan is running; only reading it back went wrong. + + Reporting the read failure alone would read as "the scan was not created" + and invite a duplicate, so the error carries the ID to monitor instead. + """ + stub_scan_creation(mock_router) + mock_router.add("GET", SCAN, status=400, json=jsonapi_error(400, "Server error.")) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="Scan s1 was created") as raised: + await client.call_tool("prowler_trigger_scan", {"provider_id": "p1"}) + + # Naming the scan is the whole reason this message exists, so it is written + # here rather than assembled from the failure. Splicing the failure text in + # would put whatever it happens to say in front of the model unclassified, + # which is the one thing the shared classifier exists to decide. + assert "API request failed" not in str(raised.value) + + +async def test_a_created_scan_comes_back_with_its_details( + mcp_root_server, mock_api_client, mock_router +): + """The success path still reports the scan, which is what gets monitored.""" + stub_scan_creation(mock_router) + mock_router.add( + "GET", + SCAN, + json=jsonapi_document(jsonapi_resource("scans", "s1", SCAN_ATTRIBUTES)), + ) + + async with Client(mcp_root_server) as client: + result = await client.call_tool("prowler_trigger_scan", {"provider_id": "p1"}) + + assert result.data["scan"]["id"] == "s1" + # No status flag: a scan that was not created is raised, so "success" could + # only ever be the one value and says nothing a reader can act on. + assert "status" not in result.data + + +# ------------------------------------------------------ schedule_daily_scan + + +@pytest.mark.parametrize("first_run_state", ["available", "scheduled", "executing"]) +async def test_a_schedule_is_reported_created_whatever_the_first_run_does( + mcp_root_server, mock_api_client, mock_router, first_run_state +): + """The task in the answer is the first scan run, not the schedule. + + Prowler commits the recurring schedule inside the request that serves this + call, so an answer at all means it exists. Reading that task's state as the + outcome of the scheduling reported a schedule that had just been created as + a failure. + """ + mock_router.add( + "POST", + DAILY, + json=jsonapi_document( + jsonapi_resource("tasks", "t1", {"state": first_run_state}) + ), + ) + + async with Client(mcp_root_server) as client: + result = await client.call_tool( + "prowler_schedule_daily_scan", {"provider_id": "p1"} + ) + + assert result.data["first_run_state"] == first_run_state + assert "every 24 hours" in result.data["message"] + assert "scheduled" not in result.data + + +async def test_a_failed_first_run_leaves_the_schedule_standing( + mcp_root_server, mock_api_client, mock_router +): + """Worth saying, but it is not the schedule that failed. + + The gap it leaves is real -- no findings until tomorrow -- so the message + points at the manual scan that fills it rather than at the scheduling. + """ + mock_router.add( + "POST", + DAILY, + json=jsonapi_document(jsonapi_resource("tasks", "t1", {"state": "failed"})), + ) + + async with Client(mcp_root_server) as client: + result = await client.call_tool( + "prowler_schedule_daily_scan", {"provider_id": "p1"} + ) + + assert result.data["first_run_state"] == "failed" + assert "schedule is unaffected" in result.data["message"] + assert "prowler_trigger_scan" in result.data["message"] + + +async def test_an_existing_schedule_is_relayed_as_the_api_explains_it( + mcp_root_server, mock_api_client, mock_router +): + """The 409 already says the useful thing, so the classifier relays it.""" + mock_router.add( + "POST", + DAILY, + status=409, + json=jsonapi_error(409, "There is already a scheduled scan for this provider."), + ) + + async with Client(mcp_root_server) as client: + with pytest.raises(Exception, match="already a scheduled scan"): + await client.call_tool("prowler_schedule_daily_scan", {"provider_id": "p1"}) diff --git a/mcp_server/tests/prowler_app/utils/test_auth.py b/mcp_server/tests/prowler_app/utils/test_auth.py index d39e5826d7..3c822cf2dc 100644 --- a/mcp_server/tests/prowler_app/utils/test_auth.py +++ b/mcp_server/tests/prowler_app/utils/test_auth.py @@ -6,8 +6,12 @@ Reference for later branches: ``ProwlerAppAuth`` resolves its ``mode`` and and ``base_url=`` explicitly, as these tests do. """ +import base64 +import json + import pytest +from prowler_mcp_server.lib.errors import CredentialError from prowler_mcp_server.prowler_app.utils.auth import ProwlerAppAuth from tests.helpers.tokens import FAKE_API_KEY, MALFORMED_API_KEY, fake_jwt @@ -42,13 +46,92 @@ async def test_http_mode_accepts_a_bearer_api_key(http_request_headers): assert await auth.get_valid_token() == FAKE_API_KEY +def _jwt_with_payload(payload: object) -> str: + """Mint an unsigned JWT carrying an arbitrary payload. + + ``fake_jwt`` always writes a well-formed object, so the malformed payloads + below are built here instead. + """ + encoded = ( + base64.urlsafe_b64encode(json.dumps(payload).encode()).decode().rstrip("=") + ) + return f"header.{encoded}.fake-signature-not-verified" + + +async def test_http_mode_accepts_a_lowercase_bearer_scheme(http_request_headers): + """Authentication scheme names are case-insensitive (RFC 7235).""" + http_request_headers(authorization=f"bearer {FAKE_API_KEY}") + + auth = ProwlerAppAuth(mode="http") + + assert await auth.get_valid_token() == FAKE_API_KEY + + +async def test_http_mode_strips_only_the_scheme_prefix(http_request_headers): + """A token that repeats the scheme keeps it: only the prefix is removed.""" + token = f"{FAKE_API_KEY}_Bearer_suffix" + http_request_headers(authorization=f"Bearer {token}") + + auth = ProwlerAppAuth(mode="http") + + assert await auth.get_valid_token() == token + + +async def test_http_mode_rejects_an_authorization_header_without_a_token( + http_request_headers, +): + """A bare scheme carries no credential to authenticate with.""" + http_request_headers(authorization="Bearer ") + + auth = ProwlerAppAuth(mode="http") + + with pytest.raises(CredentialError, match="'Bearer ' form"): + await auth.get_valid_token() + + +async def test_http_mode_rejects_a_jwt_whose_payload_is_not_an_object( + http_request_headers, +): + """A payload that decodes to a list has no claims, so it is a bad credential. + + Without the type check it would reach `payload.get` and fail as an + unclassified `AttributeError`, which the client only sees masked. + """ + http_request_headers(authorization=f"Bearer {_jwt_with_payload(['exp'])}") + + auth = ProwlerAppAuth(mode="http") + + with pytest.raises(CredentialError, match="not a readable JWT"): + await auth.get_valid_token() + + +@pytest.mark.parametrize( + ("payload", "case"), + [ + ({"sub": "user"}, "missing"), + ({"exp": "1700000000"}, "string"), + ({"exp": None}, "null"), + ], +) +async def test_http_mode_rejects_a_jwt_without_a_numeric_expiration( + http_request_headers, payload: dict, case: str +): + """`exp` is a numeric date: comparing anything else raises a `TypeError`.""" + http_request_headers(authorization=f"Bearer {_jwt_with_payload(payload)}") + + auth = ProwlerAppAuth(mode="http") + + with pytest.raises(CredentialError, match="no readable 'exp' expiration claim"): + await auth.get_valid_token() + + async def test_http_mode_rejects_an_expired_jwt(http_request_headers): """An expired JWT is refused locally instead of being forwarded to the API.""" http_request_headers(authorization=f"Bearer {fake_jwt(expires_in=-60)}") auth = ProwlerAppAuth(mode="http") - with pytest.raises(ValueError, match="Token has expired"): + with pytest.raises(CredentialError, match="The token has expired"): await auth.get_valid_token()