From 1679094f22bbb9029a56ea524189d487b3d8be3c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Rub=C3=A9n=20De=20la=20Torre=20Vico?= Date: Mon, 31 Aug 2026 13:09:34 +0200 Subject: [PATCH] feat(mcp): raise instead of returning error objects in the Prowler Hub tools (#12533) --- .../mcp-hub-check-provider-mismatch.fixed.md | 1 + mcp_server/prowler_mcp_server/lib/errors.py | 64 ++ mcp_server/prowler_mcp_server/lib/urls.py | 36 + .../prowler_mcp_server/prowler_hub/server.py | 679 ++++++++++-------- mcp_server/tests/lib/test_errors.py | 78 ++ mcp_server/tests/lib/test_urls.py | 59 ++ mcp_server/tests/prowler_hub/test_server.py | 335 +++++++++ 7 files changed, 949 insertions(+), 303 deletions(-) create mode 100644 mcp_server/changelog.d/mcp-hub-check-provider-mismatch.fixed.md create mode 100644 mcp_server/prowler_mcp_server/lib/urls.py create mode 100644 mcp_server/tests/lib/test_urls.py create mode 100644 mcp_server/tests/prowler_hub/test_server.py diff --git a/mcp_server/changelog.d/mcp-hub-check-provider-mismatch.fixed.md b/mcp_server/changelog.d/mcp-hub-check-provider-mismatch.fixed.md new file mode 100644 index 0000000000..da39e914b1 --- /dev/null +++ b/mcp_server/changelog.d/mcp-hub-check-provider-mismatch.fixed.md @@ -0,0 +1 @@ +`prowler_hub_get_check_code` and `prowler_hub_get_check_fixer` now report a check ID that belongs to another provider as such, naming that provider, instead of reporting the ID as one that does not exist diff --git a/mcp_server/prowler_mcp_server/lib/errors.py b/mcp_server/prowler_mcp_server/lib/errors.py index 9491bf87eb..1aa3a25607 100644 --- a/mcp_server/prowler_mcp_server/lib/errors.py +++ b/mcp_server/prowler_mcp_server/lib/errors.py @@ -55,6 +55,59 @@ class ProwlerAPIInvalidResponse(Exception): """The API answered, but with a body this server could not read as JSON.""" +class UpstreamInvalidResponse(Exception): + """An upstream this server reads directly answered with a body that is not JSON. + + Raised in place of the `json.JSONDecodeError` httpx would otherwise let out. + That one is a ValueError this module reads as a malformed argument, which is + the opposite story: it sends a model off to fix a call that was fine. + + Attributes: + host: Host that answered, so the message can name what has to be fixed + """ + + def __init__(self, message: str, *, host: str) -> None: + super().__init__(message) + self.host: str = host + + +def parse_json_response(response: httpx.Response) -> Any: + """Parse an upstream answer as JSON, telling an unreadable body from a bad + argument. + + For every upstream a sub-server reads with an httpx client of its own -- + Prowler Hub, the documentation site. `httpx` lets a body it cannot decode + out as a `json.JSONDecodeError`, which is a ValueError this module reads as + a malformed argument. Coming from an upstream -- an HTML error page from an + edge, a truncated body -- that is the wrong story, and the caller has no + argument to fix. + + The Prowler API client parses its own answers and raises + `ProwlerAPIInvalidResponse` instead: it also carries writes, where an + unreadable answer leaves the outcome unknown rather than merely absent. + + Args: + response: The answer to parse. + + Returns: + The parsed body. + + Raises: + UpstreamInvalidResponse: The body is not JSON. + """ + try: + return response.json() + except ValueError as e: + # `.request` raises rather than returning None when it was never set. + request = getattr(response, "_request", None) + host = request.url.host if request is not None else "The upstream service" + # Status only: the decoder's own message quotes the body it choked on, + # and that body is the upstream text this server never relays. + raise UpstreamInvalidResponse( + f"{response.status_code} body is not JSON", host=host + ) from e + + def jsonapi_detail(response: httpx.Response) -> str | None: """Return the API's own JSON:API error detail, when there is one to trust. @@ -171,6 +224,17 @@ def _describe_failure(exc: BaseException) -> str | None: "current state before sending it again." ) + if isinstance(exc, UpstreamInvalidResponse): + # The counterpart of the `json.JSONDecodeError` branch below: the same + # decode failure is a malformed argument on one side of this server and + # an upstream fault on the other, and only the type tells them apart. + return ( + f"{exc.host} answered with a body this server could not read as JSON, " + "so the call has no result to return. Nothing in the arguments caused " + f"this and changing them will not help -- {exc.host} is answering with " + "something other than the JSON it documents. Retry later." + ) + 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. diff --git a/mcp_server/prowler_mcp_server/lib/urls.py b/mcp_server/prowler_mcp_server/lib/urls.py new file mode 100644 index 0000000000..0c24a04417 --- /dev/null +++ b/mcp_server/prowler_mcp_server/lib/urls.py @@ -0,0 +1,36 @@ +"""URL construction shared by every sub-server. + +An identifier joined into a path unencoded is not sent as itself: httpx resolves +the URL per RFC 3986, so "../" walks the request onto another endpoint. +""" + +from urllib.parse import quote + +_DOT_SEGMENTS = frozenset({".", ".."}) + + +def path_segment(value: str) -> str: + """Encode one path segment, so an identifier names a resource and nothing else. + + Args: + value: The segment to encode, taken as a name in full. + + Returns: + The segment percent-encoded, with the dots escaped when it is only dots. + """ + encoded = quote(value, safe="") + # A dot is legal in a name, so `quote` keeps it: a segment of nothing but + # dots would still resolve away rather than name anything. + return encoded.replace(".", "%2E") if encoded in _DOT_SEGMENTS else encoded + + +def url_path(*segments: str) -> str: + """Build a URL path from one argument per segment, each of them encoded. + + Args: + *segments: The path segments, in order. + + Returns: + The joined path, with a leading slash. + """ + return "/" + "/".join(path_segment(segment) for segment in segments) diff --git a/mcp_server/prowler_mcp_server/prowler_hub/server.py b/mcp_server/prowler_mcp_server/prowler_hub/server.py index 8d18a2c240..2c78ac7c36 100644 --- a/mcp_server/prowler_mcp_server/prowler_hub/server.py +++ b/mcp_server/prowler_mcp_server/prowler_hub/server.py @@ -6,13 +6,19 @@ Provides access to Prowler Hub API for security checks and compliance frameworks import httpx from fastmcp import FastMCP +from fastmcp.exceptions import ToolError from pydantic import Field from prowler_mcp_server import __version__ +from prowler_mcp_server.lib.errors import ( + UpstreamInvalidResponse, + parse_json_response, +) from prowler_mcp_server.lib.types import NonBlankStr +from prowler_mcp_server.lib.urls import url_path # Initialize FastMCP for Prowler Hub -hub_mcp_server = FastMCP("prowler-hub") +hub_mcp_server = FastMCP("prowler-hub", mask_error_details=True) # API base URL BASE_URL = "https://hub.prowler.com/api" @@ -27,6 +33,19 @@ prowler_hub_client = httpx.Client( }, ) +# Sentences for the not-found cases. They are authored here, and raised as a +# ToolError without a `from` clause, because they name the resource the caller +# asked for and the next tool to reach for -- neither of which the shared +# classifier in lib/errors.py can know. +_CHECK_NOT_FOUND = ( + "No check with the ID '{check_id}' exists in Prowler Hub. Use " + "prowler_hub_semantic_search_checks to find the right ID." +) +_COMPLIANCE_NOT_FOUND = ( + "No compliance framework with the ID '{compliance_id}' exists in Prowler Hub. " + "Use prowler_hub_semantic_search_compliances to find the right ID." +) + # GitHub raw content base URL for Prowler checks GITHUB_RAW_BASE = ( "https://raw.githubusercontent.com/prowler-cloud/prowler/refs/heads/master/" @@ -43,6 +62,21 @@ github_raw_client = httpx.Client( ) +def _get_hub_endpoint( + *path_segments: str, params: dict[str, str] | None = None +) -> httpx.Response: + """GET a Prowler Hub endpoint, named as one argument per path segment. + + Args: + *path_segments: The endpoint path segments, in order. + params: Query parameters for the request. + + Returns: + The response unread, so a caller can tell a 404 from a failed request. + """ + return prowler_hub_client.get(url_path(*path_segments), params=params) + + def github_check_path(provider_id: str, check_id: str, suffix: str) -> str: """Build the GitHub raw URL for a given check artifact suffix using provider and check_id. @@ -53,7 +87,83 @@ def github_check_path(provider_id: str, check_id: str, suffix: str) -> str: service_id = check_id.split("_", 1)[0] except IndexError: service_id = check_id - return f"{GITHUB_RAW_BASE}/{provider_id}/services/{service_id}/{check_id}/{check_id}{suffix}" + path = url_path(provider_id, "services", service_id, check_id, check_id) + return f"{GITHUB_RAW_BASE}{path}{suffix}" + + +def _hub_provider_for_check(check_id: str) -> str | None: + """Ask Prowler Hub which provider it lists a check under. + + Args: + check_id: Check ID the caller asked for + + Returns: + The provider the Hub lists the check under, or None when the Hub knows + no such check. + + Raises: + httpx.HTTPError: The Hub could not be reached. + UpstreamInvalidResponse: The Hub answered with a body that is not JSON. + ValueError: The Hub answered with something that names no provider. + """ + response = _get_hub_endpoint("check", check_id) + if response.status_code == 404: + return None + response.raise_for_status() + check = parse_json_response(response) + + # An empty body is how the Hub reports an unknown ID on some routes, so it + # is read the same way get_check_details reads it: no such check. + if not isinstance(check, dict) or not check: + return None + + provider = check.get("provider") + if isinstance(provider, str) and provider.strip(): + return provider + # A check the Hub returned without a provider tells us nothing about the + # provider the caller asked for, so it counts as unanswered rather than as + # a check that does not exist. + raise ValueError(f"Prowler Hub listed check '{check_id}' without a provider") + + +def _explain_missing_check_file( + provider_id: str, + check_id: str, + *, + when_check_belongs_here: str, + when_unverified: str, +) -> str: + """Explain a 404 from GitHub for one of a check's source files. + + GitHub answers 404 to three different mistakes, an ID that exists nowhere, + an ID that exists under a different provider, and an ID that exists right + here whose file is simply absent, and cannot tell them apart. Prowler Hub + can, so it is asked before anything is claimed about the ID. + + Args: + provider_id: Provider the caller asked for + check_id: Check the caller asked for + when_check_belongs_here: Message for the case where the Hub confirms the + check does belong to this provider + when_unverified: Message for the case where the Hub could not be asked + + Returns: + The sentence to fail the tool with + """ + try: + hub_provider = _hub_provider_for_check(check_id) + except (httpx.HTTPError, UpstreamInvalidResponse, ValueError): + return when_unverified + + if hub_provider is None: + return _CHECK_NOT_FOUND.format(check_id=check_id) + if hub_provider != provider_id: + return ( + f"Provider '{provider_id}' has no check '{check_id}'. Prowler Hub lists " + f"that check under provider '{hub_provider}', so retry with " + f"provider_id='{hub_provider}'." + ) + return when_check_belongs_here # Security Check Tools @@ -123,29 +233,22 @@ async def list_checks( if compliances: params["compliances"] = ",".join(compliances) - try: - response = prowler_hub_client.get("/check", params=params) - response.raise_for_status() - checks = response.json() + response = _get_hub_endpoint("check", params=params) + response.raise_for_status() + checks = parse_json_response(response) - # Return checks as a lightweight list - checks_list = [] - for check in checks: - check_data = { - "id": check["id"], - "provider": check["provider"], - "title": check["title"], - "severity": check["severity"], - } - checks_list.append(check_data) - - return {"count": len(checks), "checks": checks_list} - except httpx.HTTPStatusError as e: - return { - "error": f"HTTP error {e.response.status_code}: {e.response.text}", + # Return checks as a lightweight list + checks_list = [] + for check in checks: + check_data = { + "id": check["id"], + "provider": check["provider"], + "title": check["title"], + "severity": check["severity"], } - except Exception as e: - return {"error": str(e)} + checks_list.append(check_data) + + return {"count": len(checks), "checks": checks_list} @hub_mcp_server.tool() @@ -182,29 +285,22 @@ async def semantic_search_checks( 2. Use `prowler_hub_list_checks` with filters for more targeted browsing 3. Use `prowler_hub_get_check_details` to get complete information for a specific check """ - try: - response = prowler_hub_client.get("/check/search", params={"term": term}) - response.raise_for_status() - checks = response.json() + response = _get_hub_endpoint("check", "search", params={"term": term}) + response.raise_for_status() + checks = parse_json_response(response) - # Return checks as a lightweight list - checks_list = [] - for check in checks: - check_data = { - "id": check["id"], - "provider": check["provider"], - "title": check["title"], - "severity": check["severity"], - } - checks_list.append(check_data) - - return {"count": len(checks), "checks": checks_list} - except httpx.HTTPStatusError as e: - return { - "error": f"HTTP error {e.response.status_code}: {e.response.text}", + # Return checks as a lightweight list + checks_list = [] + for check in checks: + check_data = { + "id": check["id"], + "provider": check["provider"], + "title": check["title"], + "severity": check["severity"], } - except Exception as e: - return {"error": str(e)} + checks_list.append(check_data) + + return {"count": len(checks), "checks": checks_list} @hub_mcp_server.tool() @@ -274,75 +370,75 @@ async def get_check_details( 2. Use this tool with the check 'id' to get complete information including remediation guidance """ try: - response = prowler_hub_client.get(f"/check/{check_id}") + response = _get_hub_endpoint("check", check_id) response.raise_for_status() - check = response.json() - - if not check: - return {"error": f"Check '{check_id}' not found"} - - # Build response with only non-empty fields to save tokens - result = {} - - # Core fields - result["id"] = check["id"] - if check.get("title"): - result["title"] = check["title"] - if check.get("description"): - result["description"] = check["description"] - if check.get("provider"): - result["provider"] = check["provider"] - if check.get("service"): - result["service"] = check["service"] - if check.get("severity"): - result["severity"] = check["severity"] - if check.get("risk"): - result["risk"] = check["risk"] - if check.get("resource_type"): - result["resource_type"] = check["resource_type"] - - # List fields - if check.get("reference"): - result["reference"] = check["reference"] - if check.get("additional_urls"): - result["additional_urls"] = check["additional_urls"] - if check.get("services_required"): - result["services_required"] = check["services_required"] - if check.get("categories"): - result["categories"] = check["categories"] - if check.get("compliances"): - result["compliances"] = check["compliances"] - - # Other fields - if check.get("notes"): - result["notes"] = check["notes"] - if check.get("related_url"): - result["related_url"] = check["related_url"] - if check.get("fixer") is not None: - result["fixer"] = check["fixer"] - - # Remediation - filter out empty nested values - remediation = check.get("remediation", {}) - if remediation: - filtered_remediation = {} - for key, value in remediation.items(): - if value and isinstance(value, dict): - # Filter out empty values within nested dict - filtered_value = {k: v for k, v in value.items() if v} - if filtered_value: - filtered_remediation[key] = filtered_value - elif value: - filtered_remediation[key] = value - if filtered_remediation: - result["remediation"] = filtered_remediation - - return result except httpx.HTTPStatusError as e: - return { - "error": f"HTTP error {e.response.status_code}: {e.response.text}", - } - except Exception as e: - return {"error": str(e)} + if e.response.status_code == 404: + # No `from`: this names the check, which the shared classifier cannot. + raise ToolError(_CHECK_NOT_FOUND.format(check_id=check_id)) + raise + + check = parse_json_response(response) + + if not check: + raise ToolError(_CHECK_NOT_FOUND.format(check_id=check_id)) + + # Build response with only non-empty fields to save tokens + result = {} + + # Core fields + result["id"] = check["id"] + if check.get("title"): + result["title"] = check["title"] + if check.get("description"): + result["description"] = check["description"] + if check.get("provider"): + result["provider"] = check["provider"] + if check.get("service"): + result["service"] = check["service"] + if check.get("severity"): + result["severity"] = check["severity"] + if check.get("risk"): + result["risk"] = check["risk"] + if check.get("resource_type"): + result["resource_type"] = check["resource_type"] + + # List fields + if check.get("reference"): + result["reference"] = check["reference"] + if check.get("additional_urls"): + result["additional_urls"] = check["additional_urls"] + if check.get("services_required"): + result["services_required"] = check["services_required"] + if check.get("categories"): + result["categories"] = check["categories"] + if check.get("compliances"): + result["compliances"] = check["compliances"] + + # Other fields + if check.get("notes"): + result["notes"] = check["notes"] + if check.get("related_url"): + result["related_url"] = check["related_url"] + if check.get("fixer") is not None: + result["fixer"] = check["fixer"] + + # Remediation - filter out empty nested values + remediation = check.get("remediation", {}) + if remediation: + filtered_remediation = {} + for key, value in remediation.items(): + if value and isinstance(value, dict): + # Filter out empty values within nested dict + filtered_value = {k: v for k, v in value.items() if v} + if filtered_value: + filtered_remediation[key] = filtered_value + elif value: + filtered_remediation[key] = value + if filtered_remediation: + result["remediation"] = filtered_remediation + + return result @hub_mcp_server.tool() @@ -364,31 +460,38 @@ async def get_check_code( "content": "Python source code of the check implementation" } """ - if provider_id and check_id: - url = github_check_path(provider_id, check_id, ".py") - try: - resp = github_raw_client.get(url) - resp.raise_for_status() - return { - "content": resp.text, - } - except httpx.HTTPStatusError as e: - if e.response.status_code == 404: - return { - "error": f"Check {check_id} not found in Prowler", - } - else: - return { - "error": f"HTTP error {e.response.status_code}: {e.response.text}", - } - except Exception as e: - return { - "error": str(e), - } - else: - return { - "error": "Provider ID and check ID are required", - } + url = github_check_path(provider_id, check_id, ".py") + try: + resp = github_raw_client.get(url) + resp.raise_for_status() + except httpx.HTTPStatusError as e: + if e.response.status_code == 404: + # No `from`: this names the check and the provider that does have + # it, neither of which the shared classifier in lib/errors.py knows. + raise ToolError( + _explain_missing_check_file( + provider_id, + check_id, + when_check_belongs_here=( + f"Prowler Hub lists check '{check_id}' under provider " + f"'{provider_id}', but prowler-cloud/prowler has no source file " + "for it on the master branch. The check may have been renamed or " + "moved since the Hub last indexed it." + ), + when_unverified=( + f"Provider '{provider_id}' has no check '{check_id}' in " + "prowler-cloud/prowler, and Prowler Hub could not be asked which " + "provider does. Either the ID is wrong or the check belongs to " + "another provider, prowler_hub_get_check_details reports the " + "provider a check belongs to." + ), + ) + ) + raise + + return { + "content": resp.text, + } @hub_mcp_server.tool() @@ -402,8 +505,9 @@ async def get_check_fixer( ) -> dict: """Fetch the auto-remediation (fixer) code for a Prowler security check. - IMPORTANT: Not all checks have fixers. A "fixer not found" response means the check - doesn't have auto-remediation code - this is normal for many checks. + IMPORTANT: Not all checks have fixers. A check with no auto-remediation code fails + with a message saying so - this is normal for many checks and not a problem to + report or retry. Fixer code provides automated remediation that can fix security issues detected by checks. Use this to understand how to programmatically remediate findings. @@ -412,40 +516,37 @@ async def get_check_fixer( { "content": "Python source code of the auto-remediation implementation" } - Or if no fixer exists: - { - "error": "Fixer not found for check {check_id}" - } """ - if provider_id and check_id: - url = github_check_path(provider_id, check_id, "_fixer.py") - try: - resp = github_raw_client.get(url) - if resp.status_code == 404: - return { - "error": f"Fixer not found for check {check_id}", - } - resp.raise_for_status() - return { - "content": resp.text, - } - except httpx.HTTPStatusError as e: - if e.response.status_code == 404: - return { - "error": f"Check {check_id} not found in Prowler", - } - else: - return { - "error": f"HTTP error {e.response.status_code}: {e.response.text}", - } - except Exception as e: - return { - "error": str(e), - } - else: - return { - "error": "Provider ID and check ID are required", - } + url = github_check_path(provider_id, check_id, "_fixer.py") + try: + resp = github_raw_client.get(url) + resp.raise_for_status() + except httpx.HTTPStatusError as e: + if e.response.status_code == 404: + # "No fixer" is only one of the reasons the file is missing, and the + # others are the caller's to fix, so they are told apart first. + raise ToolError( + _explain_missing_check_file( + provider_id, + check_id, + when_check_belongs_here=( + f"Check {check_id} has no auto-remediation code. Many checks do " + "not, and that is normal." + ), + when_unverified=( + f"Provider '{provider_id}' has no auto-remediation code for " + f"check '{check_id}'. Many checks have none, and that is normal, " + f"but Prowler Hub could not be asked whether the check belongs " + f"to '{provider_id}' at all. Confirm it with " + "prowler_hub_get_check_details if you expected a fixer." + ), + ) + ) + raise + + return { + "content": resp.text, + } # Compliance Framework Tools @@ -492,28 +593,21 @@ async def list_compliances( if provider: params["provider"] = ",".join(provider) - try: - response = prowler_hub_client.get("/compliance", params=params) - response.raise_for_status() - compliances = response.json() + response = _get_hub_endpoint("compliance", params=params) + response.raise_for_status() + compliances = parse_json_response(response) - # Return compliances as a lightweight list - compliances_list = [] - for compliance in compliances: - compliance_data = { - "id": compliance["id"], - "name": compliance["name"], - "provider": compliance["provider"], - } - compliances_list.append(compliance_data) - - return {"count": len(compliances), "compliances": compliances_list} - except httpx.HTTPStatusError as e: - return { - "error": f"HTTP error {e.response.status_code}: {e.response.text}", + # Return compliances as a lightweight list + compliances_list = [] + for compliance in compliances: + compliance_data = { + "id": compliance["id"], + "name": compliance["name"], + "provider": compliance["provider"], } - except Exception as e: - return {"error": str(e)} + compliances_list.append(compliance_data) + + return {"count": len(compliances), "compliances": compliances_list} @hub_mcp_server.tool() @@ -543,28 +637,21 @@ async def semantic_search_compliances( ] } """ - try: - response = prowler_hub_client.get("/compliance/search", params={"term": term}) - response.raise_for_status() - compliances = response.json() + response = _get_hub_endpoint("compliance", "search", params={"term": term}) + response.raise_for_status() + compliances = parse_json_response(response) - # Return compliances as a lightweight list - compliances_list = [] - for compliance in compliances: - compliance_data = { - "id": compliance["id"], - "name": compliance["name"], - "provider": compliance["provider"], - } - compliances_list.append(compliance_data) - - return {"count": len(compliances), "compliances": compliances_list} - except httpx.HTTPStatusError as e: - return { - "error": f"HTTP error {e.response.status_code}: {e.response.text}", + # Return compliances as a lightweight list + compliances_list = [] + for compliance in compliances: + compliance_data = { + "id": compliance["id"], + "name": compliance["name"], + "provider": compliance["provider"], } - except Exception as e: - return {"error": str(e)} + compliances_list.append(compliance_data) + + return {"count": len(compliances), "compliances": compliances_list} @hub_mcp_server.tool() @@ -599,63 +686,60 @@ async def get_compliance_details( } """ try: - response = prowler_hub_client.get(f"/compliance/{compliance_id}") + response = _get_hub_endpoint("compliance", compliance_id) response.raise_for_status() - compliance = response.json() - - if not compliance: - return {"error": f"Compliance '{compliance_id}' not found"} - - # Build response with only non-empty fields to save tokens - result = {} - - # Core fields - result["id"] = compliance["id"] - if compliance.get("name"): - result["name"] = compliance["name"] - if compliance.get("framework"): - result["framework"] = compliance["framework"] - if compliance.get("provider"): - result["provider"] = compliance["provider"] - if compliance.get("version"): - result["version"] = compliance["version"] - if compliance.get("description"): - result["description"] = compliance["description"] - - # Numeric fields - if compliance.get("total_checks"): - result["total_checks"] = compliance["total_checks"] - if compliance.get("total_requirements"): - result["total_requirements"] = compliance["total_requirements"] - - # Requirements - filter out empty nested values - requirements = compliance.get("requirements", []) - if requirements: - filtered_requirements = [] - for req in requirements: - filtered_req = {} - if req.get("id"): - filtered_req["id"] = req["id"] - if req.get("name"): - filtered_req["name"] = req["name"] - if req.get("description"): - filtered_req["description"] = req["description"] - if req.get("checks"): - filtered_req["checks"] = req["checks"] - if filtered_req: - filtered_requirements.append(filtered_req) - if filtered_requirements: - result["requirements"] = filtered_requirements - - return result except httpx.HTTPStatusError as e: if e.response.status_code == 404: - return {"error": f"Compliance '{compliance_id}' not found"} - return { - "error": f"HTTP error {e.response.status_code}: {e.response.text}", - } - except Exception as e: - return {"error": str(e)} + raise ToolError(_COMPLIANCE_NOT_FOUND.format(compliance_id=compliance_id)) + raise + + compliance = parse_json_response(response) + + if not compliance: + raise ToolError(_COMPLIANCE_NOT_FOUND.format(compliance_id=compliance_id)) + + # Build response with only non-empty fields to save tokens + result = {} + + # Core fields + result["id"] = compliance["id"] + if compliance.get("name"): + result["name"] = compliance["name"] + if compliance.get("framework"): + result["framework"] = compliance["framework"] + if compliance.get("provider"): + result["provider"] = compliance["provider"] + if compliance.get("version"): + result["version"] = compliance["version"] + if compliance.get("description"): + result["description"] = compliance["description"] + + # Numeric fields + if compliance.get("total_checks"): + result["total_checks"] = compliance["total_checks"] + if compliance.get("total_requirements"): + result["total_requirements"] = compliance["total_requirements"] + + # Requirements - filter out empty nested values + requirements = compliance.get("requirements", []) + if requirements: + filtered_requirements = [] + for req in requirements: + filtered_req = {} + if req.get("id"): + filtered_req["id"] = req["id"] + if req.get("name"): + filtered_req["name"] = req["name"] + if req.get("description"): + filtered_req["description"] = req["description"] + if req.get("checks"): + filtered_req["checks"] = req["checks"] + if filtered_req: + filtered_requirements.append(filtered_req) + if filtered_requirements: + result["requirements"] = filtered_requirements + + return result # Provider Tools @@ -684,27 +768,20 @@ async def list_providers() -> dict: ] } """ - try: - response = prowler_hub_client.get("/providers") - response.raise_for_status() - providers = response.json() + response = _get_hub_endpoint("providers") + response.raise_for_status() + providers = parse_json_response(response) - providers_list = [] - for provider in providers: - providers_list.append( - { - "id": provider["id"], - "name": provider.get("name", ""), - } - ) + providers_list = [] + for provider in providers: + providers_list.append( + { + "id": provider["id"], + "name": provider.get("name", ""), + } + ) - return {"count": len(providers), "providers": providers_list} - except httpx.HTTPStatusError as e: - return { - "error": f"HTTP error {e.response.status_code}: {e.response.text}", - } - except Exception as e: - return {"error": str(e)} + return {"count": len(providers), "providers": providers_list} @hub_mcp_server.tool() @@ -728,24 +805,20 @@ async def get_provider_services( "services": ["s3", "ec2", "iam", "rds", "lambda", ...] } """ - try: - response = prowler_hub_client.get("/providers") - response.raise_for_status() - providers = response.json() + response = _get_hub_endpoint("providers") + response.raise_for_status() + providers = parse_json_response(response) - for provider in providers: - if provider["id"] == provider_id: - return { - "provider_id": provider["id"], - "provider_name": provider.get("name", ""), - "count": len(provider.get("services", [])), - "services": provider.get("services", []), - } + for provider in providers: + if provider["id"] == provider_id: + return { + "provider_id": provider["id"], + "provider_name": provider.get("name", ""), + "count": len(provider.get("services", [])), + "services": provider.get("services", []), + } - return {"error": f"Provider '{provider_id}' not found"} - except httpx.HTTPStatusError as e: - return { - "error": f"HTTP error {e.response.status_code}: {e.response.text}", - } - except Exception as e: - return {"error": str(e)} + known = ", ".join(sorted(str(provider["id"]) for provider in providers)) + raise ToolError( + f"Prowler has no provider with the ID '{provider_id}'. Available: {known}." + ) diff --git a/mcp_server/tests/lib/test_errors.py b/mcp_server/tests/lib/test_errors.py index cd32752b42..a682034bc8 100644 --- a/mcp_server/tests/lib/test_errors.py +++ b/mcp_server/tests/lib/test_errors.py @@ -7,6 +7,7 @@ reaches a model is text this server produced. import json +import httpx import pytest from fastmcp import Client from pydantic import BaseModel, ValidationError @@ -14,7 +15,9 @@ from pydantic import BaseModel, ValidationError from prowler_mcp_server.lib.errors import ( CredentialError, InvalidArgument, + UpstreamInvalidResponse, _describe_failure, + parse_json_response, ) from prowler_mcp_server.prowler_app.utils.api_client import ( ProwlerAPIError, @@ -26,6 +29,42 @@ from tests.helpers.jsonapi import jsonapi_error LATEST = "/api/v1/findings/latest" +# --------------------------------------------------------------- json bodies + + +def _answer( + body: str, *, url: str = "https://hub.prowler.com/api/check" +) -> httpx.Response: + """An answer as a client would hand it back, request attached.""" + return httpx.Response(200, text=body, request=httpx.Request("GET", url)) + + +def test_a_json_body_is_returned_as_it_is(): + """The helper only classifies the failure; the success path is untouched.""" + assert parse_json_response(_answer('{"id": "s3_bucket_public_access"}')) == { + "id": "s3_bucket_public_access" + } + + +def test_a_body_that_is_not_json_names_the_host_that_answered(): + """Which upstream is misbehaving is the one useful fact here, and the shared + helper is reached from every sub-server that reads an upstream directly.""" + with pytest.raises(UpstreamInvalidResponse) as raised: + parse_json_response(_answer("502 Bad Gateway")) + + assert raised.value.host == "hub.prowler.com" + assert "Bad Gateway" not in str(raised.value) + + +def test_a_body_that_is_not_json_is_not_a_valueerror(): + """`JSONDecodeError` is a ValueError, and callers tell an upstream fault from + a bad argument by type alone.""" + with pytest.raises(UpstreamInvalidResponse) as raised: + parse_json_response(_answer("not json")) + + assert not isinstance(raised.value, ValueError) + + # ------------------------------------------------------------ classification @@ -94,6 +133,29 @@ def test_an_unreadable_api_answer_is_never_called_safe_to_repeat(): assert "check the current state" in message +def test_an_unreadable_upstream_answer_is_not_blamed_on_the_arguments(): + """A `JSONDecodeError` from an upstream and one from an argument are the same + exception and opposite instructions.""" + message = _describe_failure( + UpstreamInvalidResponse("200 body is not JSON", host="hub.prowler.com") + ) + + assert "hub.prowler.com" in message + assert "could not read as JSON" in message + assert "changing them will not help" in message + + +def test_an_unreadable_upstream_answer_never_quotes_the_body(): + """The body is someone else's text, so only the host and the status leave here.""" + message = _describe_failure( + UpstreamInvalidResponse( + "502 body is not JSON", host="raw.githubusercontent.com" + ) + ) + + assert "body is not JSON" not in message + + def test_an_argument_this_server_rejected_is_repeated_verbatim(): """`InvalidArgument` exists to mark a message as one we wrote.""" message = _describe_failure( @@ -227,3 +289,19 @@ async def test_a_tool_specific_message_survives_masking( assert result.isError is True assert "prowler_list_integrations" in result.content[0].text + + +async def test_a_hub_tool_failure_says_which_host_refused_it( + mcp_root_server, hub_router +): + """Hub failures arrive as raw httpx errors: host and status relayed, body not.""" + hub_router.add( + "GET", "/api/check", status=503, text="upstream nginx 10.1.2.3" + ) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp("prowler_hub_list_checks", {}) + + assert result.isError is True + assert "hub.prowler.com" in result.content[0].text + assert "10.1.2.3" not in result.content[0].text diff --git a/mcp_server/tests/lib/test_urls.py b/mcp_server/tests/lib/test_urls.py new file mode 100644 index 0000000000..64798d5445 --- /dev/null +++ b/mcp_server/tests/lib/test_urls.py @@ -0,0 +1,59 @@ +"""Tests for the shared URL path builder. + +The bug these pin: an identifier interpolated into a path was resolved away by +httpx per RFC 3986, so "../" reached an endpoint no tool meant to call. +""" + +import pytest + +from prowler_mcp_server.lib.urls import path_segment, url_path + + +@pytest.mark.parametrize( + ("value", "expected"), + [ + ("s3_bucket_public_access", "s3_bucket_public_access"), + ("cis_4.0_aws", "cis_4.0_aws"), + ("../../evil", "..%2F..%2Fevil"), + ("....//evil", "....%2F%2Fevil"), + ("%2e%2e%2f", "%252e%252e%252f"), + ("..;/", "..%3B%2F"), + ("s3/../evil", "s3%2F..%2Fevil"), + ("evil?fields=all", "evil%3Ffields%3Dall"), + ("evil#frag", "evil%23frag"), + ("evil\\wrong", "evil%5Cwrong"), + ("two words", "two%20words"), + ], + ids=[ + "plain", + "dots-in-a-name", + "traversal", + "stripped-filter-bypass", + "already-encoded", + "path-parameter", + "mid-path", + "query", + "fragment", + "backslash", + "space", + ], +) +def test_a_segment_survives_as_a_name_and_never_as_syntax(value, expected): + """A real ID passes through untouched; URL syntax comes back as characters.""" + assert path_segment(value) == expected + + +@pytest.mark.parametrize("value", [".", ".."], ids=["here", "up-one"]) +def test_a_segment_of_nothing_but_dots_is_escaped_rather_than_left_to_resolve(value): + """`quote` keeps a dot, so a segment of only dots would still resolve away.""" + assert path_segment(value) == value.replace(".", "%2E") + + +def test_a_path_is_the_segments_it_was_given_and_no_others(): + """One argument per segment, so no call site has to encode anything.""" + assert url_path("users", "../../evil", "roles") == "/users/..%2F..%2Fevil/roles" + + +def test_a_single_segment_path_keeps_its_leading_slash(): + """Every caller joins this onto a base URL that ends without a slash.""" + assert url_path("providers") == "/providers" diff --git a/mcp_server/tests/prowler_hub/test_server.py b/mcp_server/tests/prowler_hub/test_server.py new file mode 100644 index 0000000000..50711aa1fa --- /dev/null +++ b/mcp_server/tests/prowler_hub/test_server.py @@ -0,0 +1,335 @@ +"""Tests for the Prowler Hub tools. + +The Hub sub-server uses its own httpx clients, so its failures never pass through +the Prowler API client. They still have to arrive as tool errors rather than as a +result object, which the protocol, the client and the model all read as a success. +""" + +import pytest +from fastmcp import Client + +CHECKS = "/api/check" +PROVIDERS = "/api/providers" +COMPLIANCE = "/api/compliance" +CHECK_ID = "s3_bucket_public_access" +HUB_CHECK = f"{CHECKS}/{CHECK_ID}" + + +def github_check(provider: str, suffix: str = ".py") -> str: + """The raw.githubusercontent path a check artifact is fetched from.""" + return ( + f"/prowler-cloud/prowler/refs/heads/master/prowler/providers/{provider}" + f"/services/s3/{CHECK_ID}/{CHECK_ID}{suffix}" + ) + + +GITHUB_CHECK = github_check("aws") +GITHUB_FIXER = github_check("aws", "_fixer.py") + + +async def test_listing_checks_returns_the_lightweight_shape( + mcp_root_server, hub_router +): + """The happy path, so the failure tests below are not the only coverage.""" + hub_router.add( + "GET", + CHECKS, + json=[ + { + "id": "s3_bucket_public_access", + "provider": "aws", + "title": "S3 buckets should block public access", + "severity": "high", + } + ], + ) + + async with Client(mcp_root_server) as client: + result = await client.call_tool("prowler_hub_list_checks", {}) + + assert result.data["count"] == 1 + assert result.data["checks"][0]["id"] == "s3_bucket_public_access" + + +async def test_an_unknown_check_fails_and_names_the_tool_that_finds_one( + mcp_root_server, hub_router +): + """A 404 is the caller's mistake, and the fix is a different tool.""" + hub_router.add("GET", f"{CHECKS}/nope", status=404) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_hub_get_check_details", {"check_id": "nope"} + ) + + assert result.isError is True + assert "prowler_hub_semantic_search_checks" in result.content[0].text + + +async def test_an_unknown_provider_fails_and_lists_the_real_ones( + mcp_root_server, hub_router +): + """The valid values are already in hand, so withholding them wastes a call.""" + hub_router.add( + "GET", + PROVIDERS, + json=[{"id": "aws", "name": "Amazon Web Services", "services": ["s3"]}], + ) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_hub_get_provider_services", {"provider_id": "alicloud"} + ) + + assert result.isError is True + assert "aws" in result.content[0].text + + +async def test_a_check_without_a_fixer_says_that_is_normal(mcp_root_server, hub_router): + """Most checks have no auto-remediation, so this must not read as a defect.""" + hub_router.add("GET", GITHUB_FIXER, status=404) + hub_router.add("GET", HUB_CHECK, json={"id": CHECK_ID, "provider": "aws"}) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_hub_get_check_fixer", + {"provider_id": "aws", "check_id": CHECK_ID}, + ) + + assert result.isError is True + message = result.content[0].text + assert "normal" in message + # The Hub confirmed the check is an aws check, so nothing is left to verify. + assert "prowler_hub_get_check_details" not in message + + +async def test_a_check_from_another_provider_names_the_provider_that_has_it( + mcp_root_server, hub_router +): + """The ID exists; only the provider is wrong. Saying otherwise sends the + caller off to search for an ID they already hold.""" + hub_router.add("GET", github_check("azure"), status=404) + hub_router.add("GET", HUB_CHECK, json={"id": CHECK_ID, "provider": "aws"}) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_hub_get_check_code", + {"provider_id": "azure", "check_id": CHECK_ID}, + ) + + assert result.isError is True + message = result.content[0].text + assert "provider_id='aws'" in message + assert "No check with the ID" not in message + + +async def test_a_fixer_from_another_provider_is_not_reported_as_a_missing_fixer( + mcp_root_server, hub_router +): + """'That check has no fixer' about a check the provider never had is a lie + the caller cannot act on.""" + hub_router.add("GET", github_check("azure", "_fixer.py"), status=404) + hub_router.add("GET", HUB_CHECK, json={"id": CHECK_ID, "provider": "aws"}) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_hub_get_check_fixer", + {"provider_id": "azure", "check_id": CHECK_ID}, + ) + + assert result.isError is True + message = result.content[0].text + assert "provider_id='aws'" in message + assert "auto-remediation" not in message + + +async def test_a_check_id_that_exists_nowhere_is_still_reported_as_unknown( + mcp_root_server, hub_router +): + """The Hub not having the ID either is the one case that does justify the + original message.""" + hub_router.add("GET", github_check("azure"), status=404) + hub_router.add("GET", HUB_CHECK, status=404) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_hub_get_check_code", + {"provider_id": "azure", "check_id": CHECK_ID}, + ) + + assert result.isError is True + assert "No check with the ID" in result.content[0].text + + +async def test_an_unreachable_hub_leaves_the_cause_open_rather_than_guessing( + mcp_root_server, hub_router +): + """With nothing to distinguish the causes, naming one of them is a guess.""" + hub_router.add("GET", github_check("azure"), status=404) + hub_router.add("GET", HUB_CHECK, status=503) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_hub_get_check_code", + {"provider_id": "azure", "check_id": CHECK_ID}, + ) + + assert result.isError is True + message = result.content[0].text + assert "No check with the ID" not in message + assert "prowler_hub_get_check_details" in message + + +async def test_a_check_code_hit_never_asks_the_hub(mcp_root_server, hub_router): + """The Hub lookup exists to explain a 404. On the happy path it is dead + weight -- a second round trip for every call that already succeeded.""" + hub_router.add("GET", GITHUB_CHECK, text="class s3_bucket_public_access: ...") + + async with Client(mcp_root_server) as client: + result = await client.call_tool( + "prowler_hub_get_check_code", + {"provider_id": "aws", "check_id": CHECK_ID}, + ) + + assert "class s3_bucket_public_access" in result.data["content"] + assert hub_router.paths() == [f"GET {GITHUB_CHECK}"] + + +async def test_a_hub_outage_is_reported_rather_than_returned_as_an_empty_list( + mcp_root_server, hub_router +): + """An empty result set and a failed request are different answers.""" + hub_router.add("GET", CHECKS, status=500, json={"detail": "boom"}) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp("prowler_hub_list_checks", {}) + + assert result.isError is True + assert result.structuredContent is None + + +@pytest.mark.parametrize( + ("check_id", "routed_as", "sent_as"), + [ + ("../../evil", f"{CHECKS}/../../evil", b"/api/check/..%2F..%2Fevil"), + ("s3/../evil", f"{CHECKS}/s3/../evil", b"/api/check/s3%2F..%2Fevil"), + ("..", f"{CHECKS}/..", b"/api/check/%2E%2E"), + ( + "s3_x?fields=all", + f"{CHECKS}/s3_x?fields=all", + b"/api/check/s3_x%3Ffields%3Dall", + ), + ("s3_x#frag", f"{CHECKS}/s3_x#frag", b"/api/check/s3_x%23frag"), + ], + ids=["traversal", "mid-path", "dot-segment", "query", "fragment"], +) +async def test_an_id_names_a_check_and_cannot_name_an_endpoint( + mcp_root_server, hub_router, check_id, routed_as, sent_as +): + """The bug this pins: httpx resolved "../.." away and the request left + /api/check for another endpoint of the Hub.""" + hub_router.add("GET", routed_as, status=404) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_hub_get_check_details", {"check_id": check_id} + ) + + assert hub_router.requests[0].url.raw_path == sent_as + assert result.isError is True + assert "No check with the ID" in result.content[0].text + + +async def test_a_compliance_id_cannot_name_an_endpoint_either( + mcp_root_server, hub_router +): + """Every Hub path is built by the same helper, so this holds without its own + guard.""" + hub_router.add("GET", f"{COMPLIANCE}/../../evil", status=404) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_hub_get_compliance_details", {"compliance_id": "../../evil"} + ) + + assert hub_router.requests[0].url.raw_path == b"/api/compliance/..%2F..%2Fevil" + assert result.isError is True + assert "No compliance framework with the ID" in result.content[0].text + + +async def test_a_check_source_url_confines_the_provider_and_the_check_alike( + mcp_root_server, hub_router +): + """Both halves of the GitHub raw URL come from the caller, so both are + confined.""" + hub_router.add("GET", github_check("../../../../evil"), status=404) + hub_router.add("GET", HUB_CHECK, json={"id": CHECK_ID, "provider": "aws"}) + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_hub_get_check_code", + {"provider_id": "../../../../evil", "check_id": CHECK_ID}, + ) + + assert ( + hub_router.requests[0].url.raw_path + == github_check("..%2F..%2F..%2F..%2Fevil").encode() + ) + assert result.isError is True + + +async def test_a_hub_body_that_is_not_json_is_not_blamed_on_the_arguments( + mcp_root_server, hub_router +): + """An edge answering 200 with an HTML page decodes to the same + `JSONDecodeError` a malformed argument does, and the two mean opposite + things: nothing in this call can be corrected.""" + hub_router.add("GET", CHECKS, text="502 Bad Gateway") + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp("prowler_hub_list_checks", {}) + + assert result.isError is True + message = result.content[0].text + assert "hub.prowler.com" in message + assert "could not read as JSON" in message + assert "Bad Gateway" not in message + assert "Send it as a real object" not in message + + +async def test_an_unreadable_hub_answer_does_not_become_an_unknown_check( + mcp_root_server, hub_router +): + """The 404 branch is the only one that may claim the ID does not exist. A + body that could not be read says nothing about the ID.""" + hub_router.add("GET", HUB_CHECK, text="not json") + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_hub_get_check_details", {"check_id": CHECK_ID} + ) + + assert result.isError is True + message = result.content[0].text + assert "could not read as JSON" in message + assert "No check with the ID" not in message + + +async def test_an_unreadable_hub_answer_leaves_a_missing_check_file_unexplained( + mcp_root_server, hub_router +): + """The Hub is asked which provider owns the check. A body it could not read + is no more of an answer than an outage, so it hedges the same way.""" + hub_router.add("GET", github_check("azure"), status=404) + hub_router.add("GET", HUB_CHECK, text="not json") + + async with Client(mcp_root_server) as client: + result = await client.call_tool_mcp( + "prowler_hub_get_check_code", + {"provider_id": "azure", "check_id": CHECK_ID}, + ) + + assert result.isError is True + message = result.content[0].text + assert "No check with the ID" not in message + assert "prowler_hub_get_check_details" in message