feat(mcp): raise instead of returning error objects in the Prowler App tools (#12532)

This commit is contained in:
Rubén De la Torre Vico
2026-08-31 10:57:58 +02:00
committed by GitHub
parent 89988dada5
commit f05a490cd7
33 changed files with 2364 additions and 682 deletions
@@ -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
@@ -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
+28 -1
View File
@@ -19,10 +19,17 @@ class ProwlerAPIError(Exception):
Attributes: Attributes:
status_code: HTTP status the API answered with status_code: HTTP status the API answered with
detail: JSON:API `errors[0].detail`, None when there is none to trust 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__( 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: ) -> None:
super().__init__(message) super().__init__(message)
self.status_code: int = status_code 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 # that must never be repeated to a model -- and None for a 5xx, see
# `jsonapi_detail`. # `jsonapi_detail`.
self.detail: str | None = 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): class ProwlerAPIUnreachable(Exception):
@@ -71,6 +84,10 @@ class InvalidArgument(ValueError):
"""An argument this server rejected before any request went out.""" """An argument this server rejected before any request went out."""
class CredentialError(Exception):
"""The credential the caller sent is missing, malformed or expired."""
# ------------------------------------------------------------------- messages # ------------------------------------------------------------------- messages
@@ -154,6 +171,16 @@ def _describe_failure(exc: BaseException) -> str | None:
"current state before sending it again." "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 <token>' header holding a valid Prowler API "
"key or an unexpired JWT."
)
if isinstance(exc, ProwlerAPIUnreachable): if isinstance(exc, ProwlerAPIUnreachable):
# The only failure a model can turn into a duplicate write by repeating. # The only failure a model can turn into a duplicate write by repeating.
return ( return (
@@ -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)]
@@ -354,6 +354,14 @@ class AttackPathQueryResult(MinimalSerializerMixin, BaseModel):
relationships: list[AttackPathsGraphRelationship] = Field( relationships: list[AttackPathsGraphRelationship] = Field(
default_factory=list, description="Relationships connecting the nodes" 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 @classmethod
def from_api_response( def from_api_response(
@@ -368,7 +376,15 @@ class AttackPathQueryResult(MinimalSerializerMixin, BaseModel):
Returns: Returns:
AttackPathQueryResult with parsed data and summary 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", []) nodes_data = attributes.get("nodes", [])
relationships_data = attributes.get("relationships", []) relationships_data = attributes.get("relationships", [])
@@ -2,7 +2,7 @@
from typing import Any, Literal 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 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): class ProviderConnectionStatus(MinimalSerializerMixin, BaseModel):
"""Result of provider connection operation.""" """Result of provider connection operation."""
@@ -191,18 +191,18 @@ class ScansListResponse(BaseModel):
class ScanCreationResult(MinimalSerializerMixin, 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. Used by trigger_scan(). A scan that was not created leaves the tool as an
Status indicates whether scan was created successfully or failed. 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( scan: DetailedScan = Field(
default=None, description="Detailed information about the scan that was created"
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)"
) )
message: str = Field( message: str = Field(
description="Human-readable message describing the scan creation result" description="Human-readable message describing the scan creation result"
@@ -210,13 +210,26 @@ class ScanCreationResult(MinimalSerializerMixin, BaseModel):
class ScheduleCreationResult(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( first_run_state: (
description="Whether the daily scan schedule was created successfully" 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( message: str = Field(
description="Human-readable message describing the scheduling result" description="Human-readable message describing the scheduling result"
@@ -3,7 +3,7 @@ from fastmcp import FastMCP
from prowler_mcp_server.prowler_app.utils.tool_loader import load_all_tools from prowler_mcp_server.prowler_app.utils.tool_loader import load_all_tools
# Initialize MCP server # 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 # Auto-discover and load all tools from the tools package
load_all_tools(app_mcp_server) load_all_tools(app_mcp_server)
@@ -7,8 +7,11 @@ through cloud infrastructure relationships.
from typing import Any, Literal from typing import Any, Literal
from fastmcp.exceptions import ToolError
from pydantic import Field 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 ( from prowler_mcp_server.prowler_app.models.attack_paths import (
AttackPathCartographySchema, AttackPathCartographySchema,
AttackPathQuery, AttackPathQuery,
@@ -76,7 +79,6 @@ class AttackPathsTools(BaseTool):
2. Use prowler_list_attack_paths_queries to see available queries for a scan 2. Use prowler_list_attack_paths_queries to see available queries for a scan
3. Use prowler_run_attack_paths_query to execute analysis 3. Use prowler_run_attack_paths_query to execute analysis
""" """
try:
# Validate pagination # Validate pagination
self.api_client.validate_page_size(page_size) self.api_client.validate_page_size(page_size)
@@ -106,20 +108,18 @@ class AttackPathsTools(BaseTool):
) )
return simplified_response.model_dump() 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)}"}
async def list_attack_paths_queries( async def list_attack_paths_queries(
self, 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" 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]]: ) -> list[dict[str, Any]]:
"""Discover available Attack Paths queries for a completed scan. """Discover available Attack Paths queries for a completed scan.
IMPORTANT: The scan must be in 'completed' state to list queries. 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: Each query includes:
- id: Query identifier to use with run_attack_paths_query - id: Query identifier to use with run_attack_paths_query
@@ -141,23 +141,32 @@ class AttackPathsTools(BaseTool):
api_response = await self.api_client.get( api_response = await self.api_client.get(
f"/attack-paths-scans/{scan_id}/queries" 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 [ return [
AttackPathQuery.from_api_response(query).model_dump() AttackPathQuery.from_api_response(query).model_dump()
for query in api_response.get("data", []) 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)}"}]
async def run_attack_paths_query( async def run_attack_paths_query(
self, self,
scan_id: str = Field( scan_id: NonBlankStr = Field(
description="UUID of a COMPLETED attack paths scan. The scan must be in 'completed' state" 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" 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( parameters: dict[str, str] = Field(
@@ -198,7 +207,6 @@ class AttackPathsTools(BaseTool):
3. Execute this tool with appropriate parameters 3. Execute this tool with appropriate parameters
4. Analyze the returned graph for security insights 4. Analyze the returned graph for security insights
""" """
try:
# Build the request payload following JSON:API format # Build the request payload following JSON:API format
request_data: dict[str, Any] = { request_data: dict[str, Any] = {
"data": { "data": {
@@ -213,24 +221,47 @@ class AttackPathsTools(BaseTool):
if parameters: if parameters:
request_data["data"]["attributes"]["parameters"] = parameters request_data["data"]["attributes"]["parameters"] = parameters
try:
api_response = await self.api_client.post( api_response = await self.api_client.post(
f"/attack-paths-scans/{scan_id}/queries/run", f"/attack-paths-scans/{scan_id}/queries/run",
json_data=request_data, 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 # Parse the response
query_result = AttackPathQueryResult.from_api_response(api_response) query_result = AttackPathQueryResult.from_api_response(api_response)
return query_result.model_dump() if not query_result.nodes:
except Exception as e: query_result = query_result.model_copy(
self.logger.error( update={
f"Failed to run attack paths query '{query_id}' on scan {scan_id}: {e}" "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( async def get_attack_paths_cartography_schema(
self, 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" description="UUID of a COMPLETED attack paths scan. Use `prowler_list_attack_paths_scans` with state=['completed'] to find scan IDs"
), ),
) -> dict[str, Any]: ) -> dict[str, Any]:
@@ -262,18 +293,43 @@ class AttackPathsTools(BaseTool):
api_response = await self.api_client.get( api_response = await self.api_client.get(
f"/attack-paths-scans/{scan_id}/schema" 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_content = await self.api_client.fetch_external_url(schema.raw_schema_url)
schema.raw_schema_url
)
return schema.model_copy( return schema.model_copy(update={"schema_content": schema_content}).model_dump()
update={"schema_content": schema_content}
).model_dump() # Private helper methods
except Exception as e:
self.logger.error( @staticmethod
f"Failed to get cartography schema for scan {scan_id}: {e}" 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."
) )
return {"error": f"Failed to get cartography schema: {str(e)}"}
@@ -6,8 +6,11 @@ across all cloud providers.
from typing import Any from typing import Any
from fastmcp.exceptions import ToolError
from pydantic import Field 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 ( from prowler_mcp_server.prowler_app.models.compliance import (
ComplianceFrameworksListResponse, ComplianceFrameworksListResponse,
ComplianceRequirementAttributesListResponse, ComplianceRequirementAttributesListResponse,
@@ -34,7 +37,7 @@ class ComplianceTools(BaseTool):
The scan_id of the latest completed scan for the provider. The scan_id of the latest completed scan for the provider.
Raises: Raises:
ValueError: If no completed scans are found for the provider. ToolError: If no completed scans are found for the provider
""" """
scan_params = { scan_params = {
"filter[provider]": provider_id, "filter[provider]": provider_id,
@@ -48,7 +51,7 @@ class ComplianceTools(BaseTool):
scans_data = scans_response.get("data", []) scans_data = scans_response.get("data", [])
if not scans_data: if not scans_data:
raise ValueError( raise ToolError(
f"No completed scans found for provider {provider_id}. " f"No completed scans found for provider {provider_id}. "
"Run a scan first using prowler_trigger_scan." "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 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: if not scan_id and not provider_id:
return { raise InvalidArgument(
"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." "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: elif scan_id and provider_id:
return { raise InvalidArgument(
"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." "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: elif not scan_id and provider_id:
try:
scan_id = await self._get_latest_scan_id_for_provider(provider_id) scan_id = await self._get_latest_scan_id_for_provider(provider_id)
except ValueError as e:
return {"error": str(e)}
params: dict[str, Any] = {"filter[scan_id]": scan_id} params: dict[str, Any] = {"filter[scan_id]": scan_id}
@@ -253,16 +253,16 @@ class ComplianceTools(BaseTool):
async def get_compliance_framework_state_details( async def get_compliance_framework_state_details(
self, 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", 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( scan_id: str | None = Field(
default=None, 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( provider_id: str | None = Field(
default=None, 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]: ) -> dict[str, Any]:
"""Get detailed requirement-level breakdown for a specific compliance framework. """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 - Use prowler_get_finding_details with these finding IDs for more details and remediation guidance
Default behavior: Default behavior:
- Requires either scan_id OR provider_id - Requires exactly one of scan_id OR provider_id; providing both is rejected
- With provider_id (no scan_id): Automatically finds the latest completed scan for that provider - With provider_id: Automatically finds the latest completed scan for that provider
- With scan_id: Uses that specific scan's compliance data - With scan_id: Uses that specific scan's compliance data
- Only shows failed requirements with their associated failed finding IDs - 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 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 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: if not scan_id and not provider_id:
return { raise InvalidArgument(
"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." "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 # Resolve provider_id to latest scan_id if needed
resolved_scan_id = scan_id resolved_scan_id = scan_id
if not scan_id and provider_id: if not scan_id and provider_id:
try: resolved_scan_id = await self._get_latest_scan_id_for_provider(provider_id)
resolved_scan_id = await self._get_latest_scan_id_for_provider(
provider_id
)
except ValueError as e:
return {"error": str(e)}
# Build params for requirements endpoint # Build params for requirements endpoint
params: dict[str, Any] = { params: dict[str, Any] = {
@@ -6,8 +6,10 @@ This module provides read-only tools for finding group triage and drill-downs.
from typing import Any, Literal from typing import Any, Literal
from urllib.parse import quote from urllib.parse import quote
from fastmcp.exceptions import ToolError
from pydantic import Field from pydantic import Field
from prowler_mcp_server.lib.types import NonBlankStr
from prowler_mcp_server.prowler_app.models.finding_groups import ( from prowler_mcp_server.prowler_app.models.finding_groups import (
DetailedFindingGroup, DetailedFindingGroup,
FindingGroupResourcesListResponse, FindingGroupResourcesListResponse,
@@ -236,7 +238,6 @@ class FindingGroupsTools(BaseTool):
prowler_get_finding_group_details for complete counters or prowler_get_finding_group_details for complete counters or
prowler_list_finding_group_resources to drill into affected resources. prowler_list_finding_group_resources to drill into affected resources.
""" """
try:
self.api_client.validate_page_size(page_size) self.api_client.validate_page_size(page_size)
date_range, params = self._base_date_params(date_from, date_to) date_range, params = self._base_date_params(date_from, date_to)
endpoint = self._group_endpoint(date_range) endpoint = self._group_endpoint(date_range)
@@ -273,13 +274,10 @@ class FindingGroupsTools(BaseTool):
api_response = await self.api_client.get(endpoint, params=clean_params) api_response = await self.api_client.get(endpoint, params=clean_params)
response = FindingGroupsListResponse.from_api_response(api_response) response = FindingGroupsListResponse.from_api_response(api_response)
return response.model_dump() return response.model_dump()
except Exception as e:
self.logger.error(f"Error listing finding groups: {e}")
return {"error": str(e), "status": "failed"}
async def get_finding_group_details( async def get_finding_group_details(
self, self,
check_id: str = Field( check_id: NonBlankStr = Field(
description="Public check ID that identifies the finding group. This is not a UUID." description="Public check ID that identifies the finding group. This is not a UUID."
), ),
date_from: str | None = Field( date_from: str | None = Field(
@@ -297,7 +295,6 @@ class FindingGroupsTools(BaseTool):
or historical data when dates are provided. Fully muted groups are or historical data when dates are provided. Fully muted groups are
included by default so accepted risk does not look like a missing group. 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) date_range, params = self._base_date_params(date_from, date_to)
endpoint = self._group_endpoint(date_range) endpoint = self._group_endpoint(date_range)
@@ -316,20 +313,19 @@ class FindingGroupsTools(BaseTool):
data = api_response.get("data", []) data = api_response.get("data", [])
if not data: if not data:
return { # No `from`: this names the check and the tool that lists valid ones,
"error": f"Finding group '{check_id}' not found.", # neither of which the shared classifier can know.
"status": "not_found", 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."
)
group = DetailedFindingGroup.from_api_response(data[0]) group = DetailedFindingGroup.from_api_response(data[0])
return group.model_dump() return group.model_dump()
except Exception as e:
self.logger.error(f"Error getting finding group details: {e}")
return {"error": str(e), "status": "failed"}
async def list_finding_group_resources( async def list_finding_group_resources(
self, self,
check_id: str = Field( check_id: NonBlankStr = Field(
description="Public check ID that identifies the finding group. This is not a UUID." description="Public check ID that identifies the finding group. This is not a UUID."
), ),
provider: list[str] = Field( provider: list[str] = Field(
@@ -426,7 +422,6 @@ class FindingGroupsTools(BaseTool):
`finding_id`. Use `prowler_get_finding_details(finding_id)` to `finding_id`. Use `prowler_get_finding_details(finding_id)` to
retrieve complete remediation guidance for a specific resource finding. retrieve complete remediation guidance for a specific resource finding.
""" """
try:
self.api_client.validate_page_size(page_size) self.api_client.validate_page_size(page_size)
date_range, params = self._base_date_params(date_from, date_to) date_range, params = self._base_date_params(date_from, date_to)
endpoint = self._resource_endpoint(check_id, date_range) endpoint = self._resource_endpoint(check_id, date_range)
@@ -465,6 +460,3 @@ class FindingGroupsTools(BaseTool):
api_response = await self.api_client.get(endpoint, params=clean_params) api_response = await self.api_client.get(endpoint, params=clean_params)
response = FindingGroupResourcesListResponse.from_api_response(api_response) response = FindingGroupResourcesListResponse.from_api_response(api_response)
return response.model_dump() return response.model_dump()
except Exception as e:
self.logger.error(f"Error listing finding group resources: {e}")
return {"error": str(e), "status": "failed"}
@@ -8,6 +8,7 @@ from typing import Any, Literal
from pydantic import Field from pydantic import Field
from prowler_mcp_server.lib.types import NonBlankStr
from prowler_mcp_server.prowler_app.models.findings import ( from prowler_mcp_server.prowler_app.models.findings import (
DetailedFinding, DetailedFinding,
FindingsListResponse, FindingsListResponse,
@@ -180,7 +181,7 @@ class FindingsTools(BaseTool):
async def get_finding_details( async def get_finding_details(
self, 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." 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]: ) -> dict[str, Any]:
@@ -9,8 +9,11 @@ This module provides tools for managing where Prowler sends its results, includi
import json import json
from typing import Any from typing import Any
from fastmcp.exceptions import ToolError
from pydantic import Field 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 ( from prowler_mcp_server.prowler_app.models.integrations import (
DetailedIntegration, DetailedIntegration,
IntegrationConnectionStatus, IntegrationConnectionStatus,
@@ -126,7 +129,7 @@ class IntegrationsTools(BaseTool):
async def get_integration( async def get_integration(
self, 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." 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]: ) -> dict[str, Any]:
@@ -157,7 +160,7 @@ class IntegrationsTools(BaseTool):
async def create_amazon_s3_integration( async def create_amazon_s3_integration(
self, 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)." description="Name of the S3 bucket where Prowler will upload the scan outputs (CSV, HTML, OCSF JSON and compliance reports)."
), ),
output_directory: str = Field( output_directory: str = Field(
@@ -168,15 +171,15 @@ class IntegrationsTools(BaseTool):
default=[], 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.", 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, 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.", 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, default=None,
description="External ID required by the trust policy of the assumed role. In Prowler Cloud this is the tenant ID.", 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, default=None,
description="Identifier for the role session, useful to track it in AWS logs. Only letters, digits and the characters =,.@_- are allowed.", 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, default=3600,
description="Duration of the assumed role session in seconds. Must be between 900 and 43200. Defaults to 3600 when omitted.", 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, default=None,
description="AWS access key ID. Only needed when the Prowler deployment has no ambient AWS credentials.", 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, default=None,
description="AWS secret access key. Required when 'aws_access_key_id' is provided.", 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, default=None,
description="AWS session token, only for temporary credentials.", description="AWS session token, only for temporary credentials.",
), ),
@@ -244,7 +247,6 @@ class IntegrationsTools(BaseTool):
""" """
self.logger.info(f"Creating Amazon S3 integration for bucket {bucket_name}...") self.logger.info(f"Creating Amazon S3 integration for bucket {bucket_name}...")
try:
credentials = self._build_aws_credentials( credentials = self._build_aws_credentials(
role_arn=role_arn, role_arn=role_arn,
external_id=external_id, external_id=external_id,
@@ -265,13 +267,10 @@ class IntegrationsTools(BaseTool):
provider_ids=provider_ids, provider_ids=provider_ids,
enabled=enabled, enabled=enabled,
) )
except Exception as e:
self.logger.error(f"Amazon S3 integration creation failed: {e}")
return {"error": str(e), "status": "failed"}
async def create_aws_security_hub_integration( async def create_aws_security_hub_integration(
self, 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." 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( send_only_fails: bool = Field(
@@ -282,15 +281,15 @@ class IntegrationsTools(BaseTool):
default=False, default=False,
description="When true, findings that are no longer present in the latest scan are archived in Security Hub.", 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, 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.", 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, default=None,
description="External ID required by the trust policy of the assumed role.", 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, default=None,
description="Identifier for the role session, useful to track it in AWS logs. Only letters, digits and the characters =,.@_- are allowed.", 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, default=None,
description="Duration of the assumed role session in seconds. Must be between 900 and 43200. Defaults to 3600 when omitted.", 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." 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, default=None,
description="AWS secret access key. Required when 'aws_access_key_id' is provided.", 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, default=None,
description="AWS session token, only for temporary credentials.", description="AWS session token, only for temporary credentials.",
), ),
@@ -344,7 +343,6 @@ class IntegrationsTools(BaseTool):
f"Creating AWS Security Hub integration for provider {provider_id}..." f"Creating AWS Security Hub integration for provider {provider_id}..."
) )
try:
credentials = self._build_aws_credentials( credentials = self._build_aws_credentials(
role_arn=role_arn, role_arn=role_arn,
external_id=external_id, external_id=external_id,
@@ -365,19 +363,16 @@ class IntegrationsTools(BaseTool):
provider_ids=[provider_id], provider_ids=[provider_id],
enabled=enabled, enabled=enabled,
) )
except Exception as e:
self.logger.error(f"AWS Security Hub integration creation failed: {e}")
return {"error": str(e), "status": "failed"}
async def create_jira_integration( async def create_jira_integration(
self, 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." 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." 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." 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( enabled: bool = Field(
@@ -416,11 +411,8 @@ class IntegrationsTools(BaseTool):
3. Use prowler_get_jira_issue_types with that project key to pick an issue type 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 4. Use prowler_send_findings_to_jira to create the work items
""" """
try:
normalized_domain = self._normalize_atlassian_domain(domain) normalized_domain = self._normalize_atlassian_domain(domain)
self.logger.info( self.logger.info(f"Creating Jira integration for domain {normalized_domain}...")
f"Creating Jira integration for domain {normalized_domain}..."
)
return await self._create_integration( return await self._create_integration(
integration_type="jira", integration_type="jira",
@@ -434,13 +426,10 @@ class IntegrationsTools(BaseTool):
provider_ids=[], provider_ids=[],
enabled=enabled, enabled=enabled,
) )
except Exception as e:
self.logger.error(f"Jira integration creation failed: {e}")
return {"error": str(e), "status": "failed"}
async def update_integration( async def update_integration(
self, self,
integration_id: str = Field( integration_id: NonBlankStr = Field(
description="UUID of the integration to update. Use prowler_list_integrations to find it." description="UUID of the integration to update. Use prowler_list_integrations to find it."
), ),
enabled: bool | None = Field( enabled: bool | None = Field(
@@ -494,7 +483,6 @@ class IntegrationsTools(BaseTool):
""" """
self.logger.info(f"Updating integration {integration_id}...") self.logger.info(f"Updating integration {integration_id}...")
try:
current = DetailedIntegration.from_api_response( current = DetailedIntegration.from_api_response(
await self._get_integration_raw(integration_id) await self._get_integration_raw(integration_id)
) )
@@ -502,11 +490,11 @@ class IntegrationsTools(BaseTool):
if provider_ids is not None: if provider_ids is not None:
if integration_type == "jira": if integration_type == "jira":
raise ValueError( raise InvalidArgument(
"Jira integrations are tenant-wide and cannot be attached to providers." "Jira integrations are tenant-wide and cannot be attached to providers."
) )
if integration_type == "aws_security_hub" and len(provider_ids) != 1: if integration_type == "aws_security_hub" and len(provider_ids) != 1:
raise ValueError( raise InvalidArgument(
"AWS Security Hub integrations must stay attached to exactly one AWS " "AWS Security Hub integrations must stay attached to exactly one AWS "
f"provider, got {len(provider_ids)}. Pass a single provider ID, or use " f"provider, got {len(provider_ids)}. Pass a single provider ID, or use "
"prowler_delete_integration to stop sending findings to Security Hub." "prowler_delete_integration to stop sending findings to Security Hub."
@@ -523,7 +511,7 @@ class IntegrationsTools(BaseTool):
if configuration is not None: if configuration is not None:
if integration_type == "jira": if integration_type == "jira":
raise ValueError( raise InvalidArgument(
"Jira integrations do not accept a configuration: it is generated by Prowler. " "Jira integrations do not accept a configuration: it is generated by Prowler. "
"Update the credentials instead, or run prowler_test_integration_connection to " "Update the credentials instead, or run prowler_test_integration_connection to "
"refresh the available projects and issue types." "refresh the available projects and issue types."
@@ -547,9 +535,7 @@ class IntegrationsTools(BaseTool):
} }
} }
if provider_ids is not None: if provider_ids is not None:
update_body["data"]["relationships"] = _providers_relationship( update_body["data"]["relationships"] = _providers_relationship(provider_ids)
provider_ids
)
await self.api_client.patch( await self.api_client.patch(
f"/integrations/{integration_id}", json_data=update_body f"/integrations/{integration_id}", json_data=update_body
@@ -561,14 +547,10 @@ class IntegrationsTools(BaseTool):
current.provider_ids current.provider_ids
) )
recheck_connection = ( recheck_connection = (
credentials is not None credentials is not None or configuration is not None or providers_changed
or configuration is not None
or providers_changed
) )
connection_status = ( connection_status = (
await self._test_connection(integration_id) await self._test_connection(integration_id) if recheck_connection else None
if recheck_connection
else None
) )
updated = await self._get_integration_raw(integration_id) updated = await self._get_integration_raw(integration_id)
@@ -577,13 +559,10 @@ class IntegrationsTools(BaseTool):
updated, connection_status updated, connection_status
).model_dump() ).model_dump()
return DetailedIntegration.from_api_response(updated).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"}
async def delete_integration( async def delete_integration(
self, self,
integration_id: str = Field( integration_id: NonBlankStr = Field(
description="UUID of the integration to permanently remove. Use prowler_list_integrations to find it." description="UUID of the integration to permanently remove. Use prowler_list_integrations to find it."
), ),
) -> dict[str, Any]: ) -> dict[str, Any]:
@@ -606,22 +585,15 @@ class IntegrationsTools(BaseTool):
""" """
self.logger.info(f"Deleting integration {integration_id}...") self.logger.info(f"Deleting integration {integration_id}...")
try:
await self.api_client.delete(f"/integrations/{integration_id}") await self.api_client.delete(f"/integrations/{integration_id}")
return { # No `deleted` flag: an integration that was not deleted leaves this tool
"deleted": True, # as an error, so the flag could only ever be True and a reader branching
"message": f"Integration {integration_id} deleted successfully", # on it would be looking for a shape that does not exist.
} return {"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)}",
}
async def test_integration_connection( async def test_integration_connection(
self, self,
integration_id: str = Field( integration_id: NonBlankStr = Field(
description="UUID of the integration to check. Use prowler_list_integrations to find it." description="UUID of the integration to check. Use prowler_list_integrations to find it."
), ),
) -> dict[str, Any]: ) -> dict[str, Any]:
@@ -654,10 +626,10 @@ class IntegrationsTools(BaseTool):
async def get_jira_issue_types( async def get_jira_issue_types(
self, 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." 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." 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]: ) -> dict[str, Any]:
@@ -692,13 +664,13 @@ class IntegrationsTools(BaseTool):
async def send_findings_to_jira( async def send_findings_to_jira(
self, self,
integration_id: str = Field( integration_id: NonBlankStr = Field(
description="UUID of the Jira integration to send the findings through. It must be enabled." 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." 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." 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( finding_ids: list[str] = Field(
@@ -783,20 +755,26 @@ class IntegrationsTools(BaseTool):
return self._jira_dispatch_unknown( return self._jira_dispatch_unknown(
task_id=None, task_id=None,
error=( 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." "It may have been queued anyway."
), ),
) )
self.logger.error(f"Jira dispatch was rejected by Prowler: {e}") self.logger.error(f"Jira dispatch was rejected by Prowler: {e}")
return self._jira_dispatch_rejected(str(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: except Exception as e:
# No answer came back, so the request may still have been accepted # No answer came back, so the request may still have been accepted
self.logger.error(f"Jira dispatch could not be started: {e}") self.logger.error(f"Jira dispatch could not be started: {e}")
return self._jira_dispatch_unknown( return self._jira_dispatch_unknown(
task_id=None, task_id=None,
error=( 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." "It may have been accepted anyway."
), ),
) )
@@ -866,7 +844,7 @@ class IntegrationsTools(BaseTool):
normalized = normalized.removesuffix(".atlassian.net") normalized = normalized.removesuffix(".atlassian.net")
if not normalized: if not normalized:
raise ValueError( raise InvalidArgument(
f"Invalid Jira domain: {domain}. Provide the Atlassian site name, for example " f"Invalid Jira domain: {domain}. Provide the Atlassian site name, for example "
"'acme' for the site 'https://acme.atlassian.net'." "'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 not isinstance(credentials.get(key), str) or not credentials[key].strip()
] ]
if missing: if missing:
raise ValueError( raise InvalidArgument(
"Jira credentials are replaced as a whole, so 'domain', 'user_mail' and " "Jira credentials are replaced as a whole, so 'domain', 'user_mail' and "
f"'api_token' are all required. Missing or empty: {', '.join(missing)}. " f"'api_token' are all required. Missing or empty: {', '.join(missing)}. "
"Sending an incomplete object would destroy the stored credentials and break " "Sending an incomplete object would destroy the stored credentials and break "
@@ -908,29 +886,33 @@ class IntegrationsTools(BaseTool):
try: try:
value = json.loads(value) value = json.loads(value)
except json.JSONDecodeError as e: 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): 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 return value
async def _get_integration_raw(self, integration_id: str) -> dict[str, Any]: async def _get_integration_raw(self, integration_id: str) -> dict[str, Any]:
"""Fetch the raw JSON:API resource of an integration. """Fetch the raw JSON:API resource of an integration.
Raises: 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}") response = await self.api_client.get(f"/integrations/{integration_id}")
integration = response.get("data") integration = response.get("data")
if not isinstance(integration, dict) or not integration.get("id"): 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 " f"Integration {integration_id} was not found. Use prowler_list_integrations "
"to get a valid integration ID." "to get a valid integration ID."
) )
if not isinstance(integration.get("attributes"), dict): if not isinstance(integration.get("attributes"), dict):
raise ValueError( raise ToolError(
f"Prowler returned integration {integration_id} without its attributes, so " f"Prowler returned integration {integration_id} without its attributes, so "
"its state cannot be read." "its state cannot be read."
) )
@@ -970,7 +952,9 @@ class IntegrationsTools(BaseTool):
integration_id = api_response.get("data", {}).get("id") integration_id = api_response.get("data", {}).get("id")
if not integration_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 " "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 " "connection could not be checked. Use prowler_list_integrations to see whether "
"the integration exists before creating it again." "the integration exists before creating it again."
@@ -981,11 +965,17 @@ class IntegrationsTools(BaseTool):
try: try:
integration = await self._get_integration_raw(integration_id) integration = await self._get_integration_raw(integration_id)
except Exception as e: except Exception as e:
# The integration exists, so surface its ID instead of a plain read failure # The integration exists, so surface its ID instead of a plain read
raise ValueError( # failure. No `from` clause: a cause would let the shared classifier
f"Integration {integration_id} was created, but reading its state failed: {e} " # 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." "Use prowler_get_integration with that ID to check it."
) from e )
return IntegrationConnectionStatus.create( return IntegrationConnectionStatus.create(
integration, connection_status integration, connection_status
@@ -1030,7 +1020,7 @@ class IntegrationsTools(BaseTool):
return { return {
"connected": None, "connected": None,
"error": ( "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 " "about the stored credentials, run prowler_test_integration_connection "
"to check them again." "to check them again."
), ),
@@ -8,8 +8,11 @@ This module provides tools for managing finding muting in Prowler, including:
import json import json
from typing import Any from typing import Any
from fastmcp.exceptions import ToolError
from pydantic import Field 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 ( from prowler_mcp_server.prowler_app.models.muting import (
DetailedMuteRule, DetailedMuteRule,
MutelistResponse, MutelistResponse,
@@ -28,10 +31,31 @@ class MutingTools(BaseTool):
# ===== MUTELIST TOOLS ===== # ===== 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]: async def get_mutelist(self) -> dict[str, Any]:
"""Retrieve the current mutelist configuration for the tenant. """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 For detailed information about mutelist structure and configuration, search Prowler documentation
using prowler_docs_search tool available in this MCP Server. using prowler_docs_search tool available in this MCP Server.
@@ -47,26 +71,15 @@ class MutingTools(BaseTool):
""" """
self.logger.info("Retrieving mutelist configuration...") self.logger.info("Retrieving mutelist configuration...")
# Query processors filtered by type=mutelist mutelist = await self._get_mutelist_raw()
params = { if mutelist is None:
"filter[processor_type]": "mutelist", # No `from`: this names the tool that creates one, which the shared
"fields[processors]": "processor_type,configuration,inserted_at,updated_at", # classifier cannot know.
} raise ToolError(
"No mutelist configuration exists for this tenant. Use "
clean_params = self.api_client.build_filter_params(params) "prowler_set_mutelist to create one."
api_response = await self.api_client.get("/processors", params=clean_params) )
return mutelist
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()
async def set_mutelist( async def set_mutelist(
self, self,
@@ -128,9 +141,9 @@ Structure:
configuration = json.loads(configuration) configuration = json.loads(configuration)
# Check if mutelist already exists # 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 # Create new mutelist
self.logger.info("Creating new mutelist...") self.logger.info("Creating new mutelist...")
create_body = { create_body = {
@@ -183,21 +196,22 @@ Structure:
self.logger.info("Deleting mutelist configuration...") self.logger.info("Deleting mutelist configuration...")
# Get existing mutelist # Get existing mutelist
existing_mutelist = await self.get_mutelist() existing_mutelist = await self._get_mutelist_raw()
if "error" in existing_mutelist: if existing_mutelist is None:
return { raise ToolError(
"success": False, "There is no mutelist configuration to delete. Use "
"message": "No mutelist found to delete", "prowler_get_mutelist to confirm the current state."
} )
# Delete the mutelist # Delete the mutelist
mutelist_id = existing_mutelist["id"] mutelist_id = existing_mutelist["id"]
await self.api_client.delete(f"/processors/{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 { return {
"success": True, "message": "Mutelist deleted successfully. Findings it had muted stay muted."
"message": "Mutelist deleted successfully",
} }
# ===== MUTE RULES TOOLS ===== # ===== MUTE RULES TOOLS =====
@@ -268,7 +282,7 @@ Structure:
elif enabled.lower() == "false": elif enabled.lower() == "false":
params["filter[enabled]"] = False params["filter[enabled]"] = False
else: else:
raise ValueError( raise InvalidArgument(
f"Invalid enabled value: {enabled}. Valid values are True, False, 'true', 'false' or None." f"Invalid enabled value: {enabled}. Valid values are True, False, 'true', 'false' or None."
) )
if search: if search:
@@ -282,7 +296,7 @@ Structure:
async def get_mute_rule( async def get_mute_rule(
self, 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')." description="UUID of the mute rule to retrieve. Must be a valid UUID format (e.g., '019ac0d6-90d5-73e9-9acf-c22e256f1bac')."
), ),
) -> dict[str, Any]: ) -> dict[str, Any]:
@@ -316,10 +330,10 @@ Structure:
async def create_mute_rule( async def create_mute_rule(
self, 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')." 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')." 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( finding_ids: list[str] = Field(
@@ -367,14 +381,14 @@ Structure:
async def update_mute_rule( async def update_mute_rule(
self, self,
rule_id: str = Field( rule_id: NonBlankStr = Field(
description="UUID of the mute rule to update. Must be a valid UUID format." description="UUID of the mute rule to update. Must be a valid UUID format."
), ),
name: str | None = Field( name: NonBlankStr | None = Field(
default=None, default=None,
description="New name for the rule. If not specified, name remains unchanged.", description="New name for the rule. If not specified, name remains unchanged.",
), ),
reason: str | None = Field( reason: NonBlankStr | None = Field(
default=None, default=None,
description="New reason for the rule. If not specified, reason remains unchanged.", description="New reason for the rule. If not specified, reason remains unchanged.",
), ),
@@ -435,7 +449,7 @@ Structure:
async def delete_mute_rule( async def delete_mute_rule(
self, self,
rule_id: str = Field( rule_id: NonBlankStr = Field(
description="UUID of the mute rule to delete. Must be a valid UUID format." description="UUID of the mute rule to delete. Must be a valid UUID format."
), ),
) -> dict[str, Any]: ) -> dict[str, Any]:
@@ -457,15 +471,18 @@ Structure:
""" """
self.logger.info(f"Deleting mute rule {rule_id}...") 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 { return {
"success": True, "message": (
"message": "Mute rule deleted successfully", f"Mute rule {rule_id} deleted successfully. The findings it muted stay "
} "muted."
else: )
return {
"success": False,
"message": "Failed to delete mute rule",
} }
@@ -6,10 +6,14 @@ including searching, connecting, and deleting providers.
from typing import Any from typing import Any
from fastmcp.exceptions import ToolError
from pydantic import Field 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 ( from prowler_mcp_server.prowler_app.models.providers import (
ProviderConnectionStatus, ProviderConnectionStatus,
ProviderDeletionResult,
ProvidersListResponse, ProvidersListResponse,
) )
from prowler_mcp_server.prowler_app.tools.base import BaseTool from prowler_mcp_server.prowler_app.tools.base import BaseTool
@@ -95,7 +99,7 @@ class ProvidersTools(BaseTool):
elif connected.lower() == "false": elif connected.lower() == "false":
params["filter[connected]"] = False params["filter[connected]"] = False
else: else:
raise ValueError( raise InvalidArgument(
f"Invalid connected value: {connected}. Valid values are True, False, 'true', 'false' or None." 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( async def connect_provider(
self, 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" 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." 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, 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.", 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( async def delete_provider(
self, 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)" 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]: ) -> 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 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. 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}...") self.logger.info(f"Deleting provider {provider_id}...")
try:
# Initiate the deletion task # 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_response = await self.api_client.delete(f"/providers/{provider_id}")
task_id = task_response.get("data", {}).get("id") task_id = task_response.get("data", {}).get("id")
# Poll until task completes (with 60 second timeout) 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( await self.api_client.poll_task_until_complete(
task_id=task_id, timeout=60, poll_interval=1.0 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: except Exception as e:
self.logger.error(f"Provider deletion failed: {e}") self.logger.error(f"Provider deletion did not complete cleanly: {e}")
return { return await self._provider_deletion_fallback(provider_id, task_id)
"deleted": False,
"message": f"Provider {provider_id} deletion failed: {str(e)}", return ProviderDeletionResult(
} status="deleted",
message=f"Provider {provider_id} deleted successfully",
).model_dump()
# Private helper methods # 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: async def _check_provider_exists(self, provider_uid: str) -> str | None:
"""Check if a provider already exists by its UID. """Check if a provider already exists by its UID.
@@ -357,7 +448,7 @@ class ProvidersTools(BaseTool):
return prowler_provider_id return prowler_provider_id
else: else:
# Multiple providers with the same UID is a data integrity issue # 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"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." 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) provider_id = await self._check_provider_exists(provider_uid)
if provider_id is None: 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 return provider_id
async def _update_provider_alias( async def _update_provider_alias(
@@ -418,7 +513,10 @@ class ProvidersTools(BaseTool):
f"/providers/{prowler_provider_id}", json_data=update_body f"/providers/{prowler_provider_id}", json_data=update_body
) )
if result.get("data", {}).get("attributes", {}).get("alias") != alias: 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: def _determine_secret_type(self, credentials: dict[str, Any]) -> str:
"""Determine the secret type from credentials structure. """Determine the secret type from credentials structure.
@@ -443,9 +541,17 @@ class ProvidersTools(BaseTool):
prowler_provider_id: The Prowler-generated provider ID prowler_provider_id: The Prowler-generated provider ID
Returns: Returns:
The secret ID if exists, None otherwise The secret ID if the provider has one, None if it has none
Raises:
Exception: If the lookup itself failed, so that "no secret" is never
reported for a provider whose secret could not be read
""" """
try: # 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( response = await self.api_client.get(
"/providers/secrets", "/providers/secrets",
params={"filter[provider]": prowler_provider_id}, params={"filter[provider]": prowler_provider_id},
@@ -458,13 +564,8 @@ class ProvidersTools(BaseTool):
f"Found existing secret {secret_id} for provider {prowler_provider_id}" f"Found existing secret {secret_id} for provider {prowler_provider_id}"
) )
return secret_id return secret_id
else:
self.logger.info( self.logger.info(f"No existing secret found for provider {prowler_provider_id}")
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 return None
async def _get_secret_type(self, secret_id: str) -> str | None: async def _get_secret_type(self, secret_id: str) -> str | None:
@@ -573,13 +674,24 @@ class ProvidersTools(BaseTool):
raise raise
async def _test_connection(self, prowler_provider_id: str) -> dict[str, Any]: 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: Args:
prowler_provider_id: The Prowler-generated provider ID prowler_provider_id: The Prowler-generated provider ID
Returns: 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}...") self.logger.info(f"Testing connection for provider {prowler_provider_id}...")
try: try:
@@ -589,6 +701,11 @@ class ProvidersTools(BaseTool):
) )
task_id = task_response.get("data", {}).get("id") 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) # Poll until task completes (with 60 second timeout)
completed_task = await self.api_client.poll_task_until_complete( completed_task = await self.api_client.poll_task_until_complete(
task_id=task_id, timeout=60, poll_interval=1.0 task_id=task_id, timeout=60, poll_interval=1.0
@@ -596,13 +713,26 @@ class ProvidersTools(BaseTool):
# Extract the result from the completed task # Extract the result from the completed task
task_result = ( 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 return task_result
except Exception as e: except Exception as e:
self.logger.error(f"Connection test failed: {e}") self.logger.error(f"Connection test could not be completed: {e}")
return {"connected": False, "error": str(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( async def _get_final_provider_state(
self, prowler_provider_id: str self, prowler_provider_id: str
@@ -8,6 +8,7 @@ from typing import Any
from pydantic import Field from pydantic import Field
from prowler_mcp_server.lib.types import NonBlankStr
from prowler_mcp_server.prowler_app.models.resources import ( from prowler_mcp_server.prowler_app.models.resources import (
DetailedResource, DetailedResource,
ResourceEventsResponse, ResourceEventsResponse,
@@ -176,7 +177,7 @@ class ResourcesTools(BaseTool):
async def get_resource( async def get_resource(
self, 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" 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]: ) -> dict[str, Any]:
@@ -347,7 +348,7 @@ class ResourcesTools(BaseTool):
async def get_resource_events( async def get_resource_events(
self, 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`." 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( lookback_days: int = Field(
@@ -11,8 +11,11 @@ adding to it.
from typing import Any from typing import Any
from fastmcp.exceptions import ToolError
from pydantic import Field 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 ( from prowler_mcp_server.prowler_app.models.roles import (
DetailedRole, DetailedRole,
RolesListResponse, RolesListResponse,
@@ -70,7 +73,7 @@ class RolesTools(BaseTool):
async def get_role( async def get_role(
self, 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." 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]: ) -> dict[str, Any]:
@@ -98,7 +101,7 @@ class RolesTools(BaseTool):
async def get_user_roles( async def get_user_roles(
self, 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." 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]: ) -> dict[str, Any]:
@@ -124,10 +127,10 @@ class RolesTools(BaseTool):
async def set_user_role( async def set_user_role(
self, 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." 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." description="Prowler's internal UUID (v4) for the role the user should hold. Use `prowler_list_roles` to find role IDs."
), ),
) -> dict[str, Any]: ) -> dict[str, Any]:
@@ -166,11 +169,20 @@ class RolesTools(BaseTool):
# user with no role at all. Confirm the role exists before replacing. # user with no role at all. Confirm the role exists before replacing.
try: try:
await self.api_client.get(f"/roles/{role_id}") await self.api_client.get(f"/roles/{role_id}")
except Exception as e: except ProwlerAPIError as e:
raise ValueError( if e.status_code != 404:
f"Role {role_id} could not be read ({e}), so user {user_id} was left " # Only a not-found says anything about the role ID. A permission
f"unchanged. Use `prowler_list_roles` to find a valid role ID." # error, a rate limit or a server error is about the request, so
) from e # 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 # 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. # same call the Prowler UI makes when changing a user's role.
@@ -5,8 +5,10 @@ This module provides tools for managing and monitoring Prowler security scans.
from typing import Any, Literal from typing import Any, Literal
from fastmcp.exceptions import ToolError
from pydantic import Field from pydantic import Field
from prowler_mcp_server.lib.types import NonBlankStr
from prowler_mcp_server.prowler_app.models.scans import ( from prowler_mcp_server.prowler_app.models.scans import (
DetailedScan, DetailedScan,
ScanCreationResult, ScanCreationResult,
@@ -127,7 +129,7 @@ class ScansTools(BaseTool):
async def get_scan( async def get_scan(
self, 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" 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]: ) -> dict[str, Any]:
@@ -171,10 +173,10 @@ class ScansTools(BaseTool):
async def trigger_scan( async def trigger_scan(
self, 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" 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, 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'", 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,7 +193,6 @@ class ScansTools(BaseTool):
3. Use `prowler_get_scan` with the returned scan 'id' to monitor progress 3. Use `prowler_get_scan` with the returned scan 'id' to monitor progress
4. Once completed, use `prowler_search_security_findings` to analyze results 4. Once completed, use `prowler_search_security_findings` to analyze results
""" """
try:
# Build request data # Build request data
request_data: dict[str, Any] = { request_data: dict[str, Any] = {
"data": { "data": {
@@ -222,29 +223,40 @@ class ScansTools(BaseTool):
) )
if not scan_id: if not scan_id:
raise Exception("No scan_id returned from scan creation") # 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."
)
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_response = await self.api_client.get(f"/scans/{scan_id}")
scan_info = DetailedScan.from_api_response(scan_response["data"]) scan_info = DetailedScan.from_api_response(scan_response["data"])
except Exception as e:
# 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( return ScanCreationResult(
scan=scan_info, 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.", 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() ).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()
async def schedule_daily_scan( async def schedule_daily_scan(
self, 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" 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]: ) -> 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) task_response.get("data", {}).get("attributes", {}).get("state", None)
) )
if task_state == "available": message = (
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." f"Daily schedule created for provider {provider_id}. Prowler will scan it "
else: "every 24 hours until the provider is deleted. Use prowler_list_scans with "
return_message = "Daily schedule creation failed. Please try again later." "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( return ScheduleCreationResult(
scheduled=(task_state == "available"), first_run_state=first_run_state,
message=return_message, message=message,
).model_dump() ).model_dump()
async def update_scan( async def update_scan(
self, 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." 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." 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]: ) -> dict[str, Any]:
@@ -9,6 +9,7 @@ from typing import Any
from pydantic import Field from pydantic import Field
from prowler_mcp_server.lib.types import NonBlankStr
from prowler_mcp_server.prowler_app.models.users import ( from prowler_mcp_server.prowler_app.models.users import (
DetailedUser, DetailedUser,
UsersListResponse, UsersListResponse,
@@ -79,7 +80,7 @@ class UsersTools(BaseTool):
async def get_user( async def get_user(
self, 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." 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]: ) -> dict[str, Any]:
@@ -118,7 +118,21 @@ class ProwlerAPIClient(metaclass=SingletonMeta):
if detail: if detail:
message = f"{message} - {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: except httpx.RequestError as e:
# No answer came back, so whether the request was applied is unknown. # No answer came back, so whether the request was applied is unknown.
logger.error(f"Error during {method.value} {path}: {e}") logger.error(f"Error during {method.value} {path}: {e}")
@@ -6,6 +6,7 @@ from datetime import datetime
from fastmcp.server.dependencies import get_http_headers from fastmcp.server.dependencies import get_http_headers
from prowler_mcp_server import __version__ from prowler_mcp_server import __version__
from prowler_mcp_server.lib.errors import CredentialError
from prowler_mcp_server.lib.logger import logger from prowler_mcp_server.lib.logger import logger
@@ -64,7 +65,12 @@ class ProwlerAppAuth:
# Decode and parse JSON # Decode and parse JSON
decoded = base64.b64decode(base64_payload).decode("utf-8") 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: except Exception as e:
logger.warning(f"Failed to parse JWT token: {e}") logger.warning(f"Failed to parse JWT token: {e}")
return None return None
@@ -76,14 +82,16 @@ class ProwlerAppAuth:
authorization_header = headers.get("authorization", None) authorization_header = headers.get("authorization", None)
if not authorization_header: if not authorization_header:
raise ValueError("No authorization header provided") raise CredentialError("No Authorization header was sent")
# Extract token from Bearer header # Extract token from Bearer header. Authentication scheme names are
if authorization_header.startswith("Bearer "): # case-insensitive (RFC 7235), and only the scheme prefix is removed:
token = authorization_header.replace("Bearer ", "") # a token that happens to contain the word again keeps it.
else: scheme, _, credential = authorization_header.partition(" ")
raise ValueError( token = credential.strip()
"Invalid authorization header format. Expected 'Bearer <token>'" if scheme.lower() != "bearer" or not token:
raise CredentialError(
"The Authorization header is not in 'Bearer <token>' form"
) )
# Check if it's an API key or JWT token # Check if it's an API key or JWT token
@@ -94,17 +102,29 @@ class ProwlerAppAuth:
# JWT token - validate and check expiration # JWT token - validate and check expiration
payload = self._parse_jwt(token) payload = self._parse_jwt(token)
if not payload: 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()) now = int(datetime.now().timestamp())
exp = payload.get("exp", 0)
if exp <= now: if exp <= now:
raise ValueError("Token has expired") raise CredentialError("The token has expired")
return token return token
else: 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: async def get_valid_token(self) -> str:
"""Get a valid token (API key or JWT token).""" """Get a valid token (API key or JWT token)."""
@@ -3,6 +3,7 @@ from typing import Any
from fastmcp import FastMCP from fastmcp import FastMCP
from pydantic import Field from pydantic import Field
from prowler_mcp_server.lib.types import NonBlankStr
from prowler_mcp_server.prowler_documentation.search_engine import ( from prowler_mcp_server.prowler_documentation.search_engine import (
ProwlerDocsSearchEngine, ProwlerDocsSearchEngine,
) )
@@ -14,7 +15,9 @@ prowler_docs_search_engine = ProwlerDocsSearchEngine()
@docs_mcp_server.tool() @docs_mcp_server.tool()
def search( 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( page_size: int = Field(
5, 5,
description="Number of top results to return. It must be between 1 and 20.", description="Number of top results to return. It must be between 1 and 20.",
@@ -39,7 +42,7 @@ def search(
@docs_mcp_server.tool() @docs_mcp_server.tool()
def get_document( 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." 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]: ) -> dict[str, str]:
@@ -9,6 +9,7 @@ from fastmcp import FastMCP
from pydantic import Field from pydantic import Field
from prowler_mcp_server import __version__ from prowler_mcp_server import __version__
from prowler_mcp_server.lib.types import NonBlankStr
# Initialize FastMCP for Prowler Hub # Initialize FastMCP for Prowler Hub
hub_mcp_server = FastMCP("prowler-hub") hub_mcp_server = FastMCP("prowler-hub")
@@ -149,7 +150,7 @@ async def list_checks(
@hub_mcp_server.tool() @hub_mcp_server.tool()
async def semantic_search_checks( async def semantic_search_checks(
term: str = Field( term: NonBlankStr = Field(
description="Search term. Examples: 'public access', 'encryption', 'MFA', 'logging'.", description="Search term. Examples: 'public access', 'encryption', 'MFA', 'logging'.",
), ),
) -> dict: ) -> dict:
@@ -208,7 +209,7 @@ async def semantic_search_checks(
@hub_mcp_server.tool() @hub_mcp_server.tool()
async def get_check_details( 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'" description="The check ID to retrieve details for. Example: 's3_bucket_level_public_access_block'"
), ),
) -> dict: ) -> dict:
@@ -346,10 +347,10 @@ async def get_check_details(
@hub_mcp_server.tool() @hub_mcp_server.tool()
async def get_check_code( 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.", 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`.", description="The check ID. Example: 's3_bucket_public_access'. Get IDs from `prowler_hub_list_checks` or `prowler_hub_search_checks`.",
), ),
) -> dict: ) -> dict:
@@ -392,10 +393,10 @@ async def get_check_code(
@hub_mcp_server.tool() @hub_mcp_server.tool()
async def get_check_fixer( 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.", 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`.", description="The check ID. Example: 's3_bucket_public_access'. Get IDs from `prowler_hub_list_checks` or `prowler_hub_search_checks`.",
), ),
) -> dict: ) -> dict:
@@ -517,7 +518,7 @@ async def list_compliances(
@hub_mcp_server.tool() @hub_mcp_server.tool()
async def semantic_search_compliances( async def semantic_search_compliances(
term: str = Field( term: NonBlankStr = Field(
description="Search term. Examples: 'CIS', 'HIPAA', 'PCI', 'GDPR', 'SOC2', 'NIST'.", description="Search term. Examples: 'CIS', 'HIPAA', 'PCI', 'GDPR', 'SOC2', 'NIST'.",
), ),
) -> dict: ) -> dict:
@@ -568,7 +569,7 @@ async def semantic_search_compliances(
@hub_mcp_server.tool() @hub_mcp_server.tool()
async def get_compliance_details( 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.", 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: ) -> dict:
@@ -708,7 +709,7 @@ async def list_providers() -> dict:
@hub_mcp_server.tool() @hub_mcp_server.tool()
async def get_provider_services( 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.", description="The provider ID to get services for. Example: 'aws', 'azure', 'gcp', 'kubernetes'. Use `prowler_hub_list_providers` to get available provider IDs.",
), ),
) -> dict: ) -> dict:
+33 -1
View File
@@ -11,7 +11,11 @@ import pytest
from fastmcp import Client from fastmcp import Client
from pydantic import BaseModel, ValidationError 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 ( from prowler_mcp_server.prowler_app.utils.api_client import (
ProwlerAPIError, ProwlerAPIError,
ProwlerAPIInvalidResponse, 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." 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(): def test_a_pydantic_rejection_names_the_field_without_echoing_the_value():
"""Pydantic quotes `input_value` back, and these tools take credentials.""" """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 result.isError is True
assert "gateway timeout" not in result.content[0].text assert "gateway timeout" not in result.content[0].text
assert "argument" 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
+125
View File
@@ -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 == []
@@ -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"}
)
@@ -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"))
@@ -312,7 +312,8 @@ async def test_creating_a_jira_integration_rejects_an_empty_domain(
): ):
"""A domain that normalizes to nothing is caught before the round trip.""" """A domain that normalizes to nothing is caught before the round trip."""
async with Client(mcp_root_server) as client: async with Client(mcp_root_server) as client:
result = await client.call_tool( with pytest.raises(Exception, match="Invalid Jira domain"):
await client.call_tool(
"prowler_create_jira_integration", "prowler_create_jira_integration",
{ {
"domain": "https://", "domain": "https://",
@@ -321,8 +322,6 @@ async def test_creating_a_jira_integration_rejects_an_empty_domain(
}, },
) )
assert result.data["status"] == "failed"
assert "Invalid Jira domain" in result.data["error"]
assert mock_router.requests == [] assert mock_router.requests == []
@@ -334,14 +333,10 @@ async def test_creating_a_jira_integration_rejects_an_empty_domain(
], ],
ids=["security-hub", "amazon-s3"], 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 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 refused creation is a tool error, and it still carries the API's reason."""
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.
"""
mock_router.add( mock_router.add(
"POST", "POST",
INTEGRATIONS, INTEGRATIONS,
@@ -350,10 +345,8 @@ async def test_a_rejected_creation_is_reported_rather_than_raised(
) )
async with Client(mcp_root_server) as client: async with Client(mcp_root_server) as client:
result = await client.call_tool(tool, arguments) with pytest.raises(Exception, match="already has this integration"):
await client.call_tool(tool, arguments)
assert result.data["status"] == "failed"
assert "already has this integration" in result.data["error"]
async def test_a_creation_with_no_id_back_warns_before_a_blind_retry( 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": {}}) mock_router.add("POST", INTEGRATIONS, json={"data": {}})
async with Client(mcp_root_server) as client: async with Client(mcp_root_server) as client:
result = await client.call_tool( with pytest.raises(Exception, match="did not return its ID"):
await client.call_tool(
"prowler_create_amazon_s3_integration", {"bucket_name": "my-reports"} "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}"] assert mock_router.paths() == [f"POST {INTEGRATIONS}"]
@@ -395,13 +387,11 @@ async def test_a_creation_whose_read_back_fails_still_hands_over_the_id(
) )
async with Client(mcp_root_server) as client: async with Client(mcp_root_server) as client:
result = await client.call_tool( with pytest.raises(Exception, match="Integration i1 was created"):
await client.call_tool(
"prowler_create_amazon_s3_integration", {"bucket_name": "my-reports"} "prowler_create_amazon_s3_integration", {"bucket_name": "my-reports"}
) )
assert result.data["status"] == "failed"
assert "Integration i1 was created" in result.data["error"]
async def test_a_connection_check_that_cannot_run_is_not_reported_as_a_failure( async def test_a_connection_check_that_cannot_run_is_not_reported_as_a_failure(
mcp_root_server, mock_api_client, mock_router mcp_root_server, mock_api_client, mock_router
@@ -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) stub_integration(mock_router, S3_ATTRIBUTES)
async with Client(mcp_root_server) as client: async with Client(mcp_root_server) as client:
result = await client.call_tool( with pytest.raises(Exception, match=message):
await client.call_tool(
"prowler_update_integration", "prowler_update_integration",
{"integration_id": "i1", "configuration": configuration}, {"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() 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) stub_integration(mock_router, JIRA_ATTRIBUTES)
async with Client(mcp_root_server) as client: async with Client(mcp_root_server) as client:
result = await client.call_tool( with pytest.raises(Exception, match="do not accept a configuration"):
await client.call_tool(
"prowler_update_integration", "prowler_update_integration",
{"integration_id": "i1", "configuration": {"domain": "other"}}, {"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() 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) stub_integration(mock_router, JIRA_ATTRIBUTES)
async with Client(mcp_root_server) as client: async with Client(mcp_root_server) as client:
result = await client.call_tool( with pytest.raises(Exception, match="tenant-wide"):
await client.call_tool(
"prowler_update_integration", "prowler_update_integration",
{"integration_id": "i1", "provider_ids": ["p1"]}, {"integration_id": "i1", "provider_ids": ["p1"]},
) )
assert "tenant-wide" in result.data["error"]
assert f"PATCH {INTEGRATION}" not in mock_router.paths() 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",)) stub_integration(mock_router, SECURITY_HUB_ATTRIBUTES, provider_ids=("p1",))
async with Client(mcp_root_server) as client: async with Client(mcp_root_server) as client:
result = await client.call_tool( with pytest.raises(Exception, match="exactly one AWS provider"):
await client.call_tool(
"prowler_update_integration", "prowler_update_integration",
{"integration_id": "i1", "provider_ids": provider_ids}, {"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() 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) stub_integration(mock_router, JIRA_ATTRIBUTES)
async with Client(mcp_root_server) as client: async with Client(mcp_root_server) as client:
result = await client.call_tool( with pytest.raises(Exception, match="replaced as a whole"):
await client.call_tool(
"prowler_update_integration", "prowler_update_integration",
{"integration_id": "i1", "credentials": credentials}, {"integration_id": "i1", "credentials": credentials},
) )
assert "replaced as a whole" in result.data["error"]
assert f"PATCH {INTEGRATION}" not in mock_router.paths() 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 # ------------------------------------------------- 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 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 It says so in the message and nowhere else: a `deleted: true` flag could only
gone, and a retry of a delete that actually succeeded reads as a new failure. ever be true, because an integration that was not deleted leaves the tool as
an error.
""" """
mock_router.add("DELETE", INTEGRATION, status=204) mock_router.add("DELETE", INTEGRATION, status=204)
@@ -766,25 +755,24 @@ async def test_deleting_an_integration_reports_the_outcome_either_way(
"prowler_delete_integration", {"integration_id": "i1"} "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 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( mock_router.add(
"DELETE", INTEGRATION, status=403, json=jsonapi_error(403, "Permission denied.") "DELETE", INTEGRATION, status=403, json=jsonapi_error(403, "Permission denied.")
) )
async with Client(mcp_root_server) as client: async with Client(mcp_root_server) as client:
result = await client.call_tool( with pytest.raises(Exception, match="prowler_get_current_user"):
await client.call_tool(
"prowler_delete_integration", {"integration_id": "i1"} "prowler_delete_integration", {"integration_id": "i1"}
) )
assert result.data["deleted"] is False
assert "Permission denied." in result.data["message"]
async def test_checking_a_connection_surfaces_why_it_failed( async def test_checking_a_connection_surfaces_why_it_failed(
mcp_root_server, mock_api_client, mock_router mcp_root_server, mock_api_client, mock_router
@@ -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 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( async def test_a_dispatch_task_that_died_halfway_is_never_safe_to_retry(
mcp_root_server, mock_api_client, mock_router mcp_root_server, mock_api_client, mock_router
): ):
@@ -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"})
@@ -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()
@@ -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"}]
}
@@ -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"})
@@ -6,8 +6,12 @@ Reference for later branches: ``ProwlerAppAuth`` resolves its ``mode`` and
and ``base_url=`` explicitly, as these tests do. and ``base_url=`` explicitly, as these tests do.
""" """
import base64
import json
import pytest import pytest
from prowler_mcp_server.lib.errors import CredentialError
from prowler_mcp_server.prowler_app.utils.auth import ProwlerAppAuth from prowler_mcp_server.prowler_app.utils.auth import ProwlerAppAuth
from tests.helpers.tokens import FAKE_API_KEY, MALFORMED_API_KEY, fake_jwt 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 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 <token>' 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): 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.""" """An expired JWT is refused locally instead of being forwarded to the API."""
http_request_headers(authorization=f"Bearer {fake_jwt(expires_in=-60)}") http_request_headers(authorization=f"Bearer {fake_jwt(expires_in=-60)}")
auth = ProwlerAppAuth(mode="http") 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() await auth.get_valid_token()