From 9613f7787ab55560abcff6e13b80865cc7d10d6e Mon Sep 17 00:00:00 2001 From: pedrooot Date: Tue, 6 Oct 2026 14:48:33 +0200 Subject: [PATCH] fix(mcp): verify JWT signatures in HTTP transport mode --- mcp_server/README.md | 1 + ...mcp-jwt-signature-verification.security.md | 1 + .../prowler_app/utils/auth.py | 87 +++++++--- mcp_server/pyproject.toml | 6 +- mcp_server/tests/helpers/tokens.py | 93 +++++++++-- .../tests/prowler_app/utils/test_auth.py | 157 +++++++++++++++--- mcp_server/uv.lock | 2 + 7 files changed, 284 insertions(+), 63 deletions(-) create mode 100644 mcp_server/changelog.d/mcp-jwt-signature-verification.security.md diff --git a/mcp_server/README.md b/mcp_server/README.md index a4c5a736b0..11f5901dfe 100644 --- a/mcp_server/README.md +++ b/mcp_server/README.md @@ -105,6 +105,7 @@ Deploy your own remote MCP server: - Full control over deployment - Requires Python 3.12+ or Docker +- Set `DJANGO_TOKEN_VERIFYING_KEY` to the Prowler API's JWT public key (PEM, `\n`-escaped newlines allowed) so the server verifies the signature of user tokens before forwarding them; without it only their expiration is checked See the [Installation Guide](https://docs.prowler.com/getting-started/installation/prowler-mcp) for complete instructions. diff --git a/mcp_server/changelog.d/mcp-jwt-signature-verification.security.md b/mcp_server/changelog.d/mcp-jwt-signature-verification.security.md new file mode 100644 index 0000000000..21cb5f9ac3 --- /dev/null +++ b/mcp_server/changelog.d/mcp-jwt-signature-verification.security.md @@ -0,0 +1 @@ +JWT signatures in HTTP transport mode are verified against the Prowler API RS256 public key when `DJANGO_TOKEN_VERIFYING_KEY` is set, refusing forged, `alg: none` and HMAC-keyed tokens diff --git a/mcp_server/prowler_mcp_server/prowler_app/utils/auth.py b/mcp_server/prowler_mcp_server/prowler_app/utils/auth.py index 1af63920ec..e72a933af1 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/utils/auth.py +++ b/mcp_server/prowler_mcp_server/prowler_app/utils/auth.py @@ -3,12 +3,17 @@ import json import os from datetime import datetime +import jwt from fastmcp.server.dependencies import get_http_headers from prowler_mcp_server import __version__ from prowler_mcp_server.lib.errors import CredentialError from prowler_mcp_server.lib.logger import logger +# The Prowler API signs its JWTs with RS256. Pinning the list keeps a token that +# declares `alg: none`, or an HMAC algorithm keyed with the public key, out. +JWT_ALGORITHMS = ["RS256"] + class ProwlerAppAuth: """Handles authentication for Prowler API using API keys or JWT tokens.""" @@ -17,12 +22,20 @@ class ProwlerAppAuth: self, mode: str = os.getenv("PROWLER_MCP_TRANSPORT_MODE", "stdio"), base_url: str = os.getenv("API_BASE_URL", "https://api.prowler.com/api/v1"), + jwt_verifying_key: str | None = os.getenv("DJANGO_TOKEN_VERIFYING_KEY"), ): self.base_url = base_url.rstrip("/") logger.info(f"Using Prowler API base URL: {self.base_url}") self.mode = mode self.access_token: str | None = None self.api_key: str | None = None + # Env files cannot hold a multi-line PEM, so the same escaped-newline + # form the API accepts for this variable is accepted here. + self.jwt_verifying_key = ( + jwt_verifying_key.replace("\\n", "\n").strip() or None + if jwt_verifying_key + else None + ) if mode == "stdio": # STDIO mode # PROWLER_API_KEY is the current variable; PROWLER_APP_API_KEY is kept @@ -36,16 +49,14 @@ class ProwlerAppAuth: if not self.api_key.startswith("pk_"): raise ValueError("Prowler API key format is incorrect") + elif mode == "http" and not self.jwt_verifying_key: + logger.warning( + "DJANGO_TOKEN_VERIFYING_KEY is not set: JWT signatures will not be " + "verified by the MCP server, only their expiration" + ) def _parse_jwt(self, token: str) -> dict | None: - """Parse JWT token and return payload - - Args: - token: JWT token to parse - - Returns: - Parsed JWT payload, or None if parsing fails - """ + """Decode a JWT payload without verifying it; None if it is unreadable.""" if not token: return None @@ -75,6 +86,42 @@ class ProwlerAppAuth: logger.warning(f"Failed to parse JWT token: {e}") return None + def _verify_jwt(self, token: str) -> dict: + """Verify the signature and standard time claims; raise CredentialError otherwise.""" + try: + return jwt.decode( + token, + self.jwt_verifying_key, + algorithms=JWT_ALGORITHMS, + # The API's audience is deployment-specific and unknown here; the + # API checks it on every forwarded request. + options={"require": ["exp"], "verify_aud": False}, + ) + except jwt.ExpiredSignatureError: + raise CredentialError("The token has expired") + except jwt.PyJWTError as e: + logger.warning(f"Rejected JWT: {type(e).__name__}: {e}") + raise CredentialError("The token could not be verified") + + def _check_jwt_expiration(self, token: str) -> None: + """Fallback when no verifying key is configured: readable and not expired.""" + payload = self._parse_jwt(token) + if not payload: + raise CredentialError("The token is not a readable JWT") + + # `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" + ) + + now = int(datetime.now().timestamp()) + if exp <= now: + raise CredentialError("The token has expired") + async def authenticate(self) -> str: """Authenticate and return token (API key for STDIO, API key or JWT for HTTP).""" if self.mode == "http": @@ -98,27 +145,13 @@ class ProwlerAppAuth: if token.startswith("pk_"): # API key - no expiration check needed return token + + if self.jwt_verifying_key: + self._verify_jwt(token) else: - # JWT token - validate and check expiration - payload = self._parse_jwt(token) - if not payload: - raise CredentialError("The token is not a readable JWT") + self._check_jwt_expiration(token) - # 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" - ) - - now = int(datetime.now().timestamp()) - if exp <= now: - raise CredentialError("The token has expired") - - return token + return token else: # PROWLER_MCP_TRANSPORT_MODE holds something this server does not # support. Nothing about a call caused it and nothing about a call diff --git a/mcp_server/pyproject.toml b/mcp_server/pyproject.toml index e6ae26fe32..aba8837957 100644 --- a/mcp_server/pyproject.toml +++ b/mcp_server/pyproject.toml @@ -17,7 +17,8 @@ dev = [ [project] dependencies = [ "fastmcp==3.4.5", - "httpx==0.28.1" + "httpx==0.28.1", + "pyjwt[crypto]==2.14.0" ] description = "MCP server for Prowler ecosystem" name = "prowler-mcp" @@ -74,8 +75,6 @@ extend-select = [ ] [tool.uv] -package = true - # Transitive pins fastmcp does not raise on its own; each carries a known HIGH. constraint-dependencies = [ "cryptography==50.0.0", @@ -84,3 +83,4 @@ constraint-dependencies = [ "pyjwt==2.14.0", "python-multipart==0.0.30" ] +package = true diff --git a/mcp_server/tests/helpers/tokens.py b/mcp_server/tests/helpers/tokens.py index 93af70ad91..d7a9ba03e6 100644 --- a/mcp_server/tests/helpers/tokens.py +++ b/mcp_server/tests/helpers/tokens.py @@ -1,34 +1,99 @@ """Obviously-fake credentials for tests. Deliberately unrealistic so repository secret scanning does not flag them. Never -put a value here that could be mistaken for a real key. +put a value here that could be mistaken for a real key. The RSA key pairs below +are generated fresh on every import and never leave the test process. """ import base64 +import hashlib +import hmac import json import time +import jwt +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa + # Prowler API keys are recognised by their `pk_` prefix; anything else is rejected. FAKE_API_KEY = "pk_fake_api_key_for_unit_testing_only" FAKE_LEGACY_API_KEY = "pk_fake_legacy_api_key_for_unit_testing_only" MALFORMED_API_KEY = "not_a_prowler_api_key" -def fake_jwt(expires_in: int = 3600, **claims: object) -> str: - """Mint an unsigned JWT whose ``exp`` is ``expires_in`` seconds from now. +def _generate_rsa_pem_pair() -> tuple[str, str]: + """Return a (private, public) PEM pair shaped like the API's generated keys.""" + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + private_pem = private_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ).decode() + public_pem = ( + private_key.public_key() + .public_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PublicFormat.SubjectPublicKeyInfo, + ) + .decode() + ) + return private_pem, public_pem - Pass a negative ``expires_in`` for an already-expired token. - ``ProwlerAppAuth._parse_jwt`` only base64url-decodes the payload and reads - ``exp`` -- it never verifies the signature, because the Prowler API is what - validates the token. A placeholder signature is therefore enough, and avoids - adding a JWT library just for tests. +# The pair the MCP server under test trusts, and an unrelated pair an attacker +# could hold. Module-level so the 2048-bit generation runs once per session. +JWT_SIGNING_KEY, JWT_VERIFYING_KEY = _generate_rsa_pem_pair() +ROGUE_JWT_SIGNING_KEY, _ = _generate_rsa_pem_pair() + + +def _jwt_payload(expires_in: int, claims: dict[str, object]) -> dict[str, object]: + """The claim set the Prowler API issues, so the verifier sees a realistic token.""" + now = int(time.time()) + return { + "typ": "access", + "iss": "https://api.testing.invalid", + "aud": "https://api.testing.invalid", + "iat": now, + "exp": now + expires_in, + "jti": "0f0f0f0f0f0f4f0f8f0f0f0f0f0f0f0f", + "sub": "00000000-0000-4000-8000-000000000000", + "tenant_id": "00000000-0000-4000-8000-000000000001", + **claims, + } + + +def fake_jwt( + expires_in: int = 3600, + *, + signing_key: str = JWT_SIGNING_KEY, + **claims: object, +) -> str: + """Mint an RS256 JWT whose ``exp`` is ``expires_in`` seconds from now. + + Pass a negative ``expires_in`` for an already-expired token, or + ``signing_key=ROGUE_JWT_SIGNING_KEY`` for one the server must not trust. + """ + return jwt.encode(_jwt_payload(expires_in, claims), signing_key, algorithm="RS256") + + +def unsigned_jwt(expires_in: int = 3600, **claims: object) -> str: + """Mint a JWT declaring ``alg: none`` with an empty signature segment.""" + return jwt.encode(_jwt_payload(expires_in, claims), key=None, algorithm="none") + + +def hmac_jwt_keyed_with(secret: str, expires_in: int = 3600, **claims: object) -> str: + """Mint an HS256 JWT keyed with ``secret``, e.g. the server's public key. + + Built by hand because ``jwt.encode`` refuses a PEM as an HMAC secret, which + is exactly the algorithm-confusion token the server must refuse too. """ - def _segment(payload: dict[str, object]) -> str: - raw = json.dumps(payload, separators=(",", ":")).encode() - return base64.urlsafe_b64encode(raw).decode().rstrip("=") + def _segment(data: bytes) -> str: + return base64.urlsafe_b64encode(data).decode().rstrip("=") - header = _segment({"alg": "HS256", "typ": "JWT"}) - body = _segment({"exp": int(time.time()) + expires_in, **claims}) - return f"{header}.{body}.fake-signature-not-verified" + header = _segment(json.dumps({"alg": "HS256", "typ": "JWT"}).encode()) + body = _segment(json.dumps(_jwt_payload(expires_in, claims)).encode()) + signature = hmac.new( + secret.encode(), f"{header}.{body}".encode(), hashlib.sha256 + ).digest() + return f"{header}.{body}.{_segment(signature)}" diff --git a/mcp_server/tests/prowler_app/utils/test_auth.py b/mcp_server/tests/prowler_app/utils/test_auth.py index 3c822cf2dc..01a107a135 100644 --- a/mcp_server/tests/prowler_app/utils/test_auth.py +++ b/mcp_server/tests/prowler_app/utils/test_auth.py @@ -1,19 +1,29 @@ """Tests for Prowler API authentication. -Reference for later branches: ``ProwlerAppAuth`` resolves its ``mode`` and -``base_url`` in default arguments, which Python evaluates once at module import. -``monkeypatch.setenv`` therefore has no effect on them -- always pass ``mode=`` -and ``base_url=`` explicitly, as these tests do. +Reference for later branches: ``ProwlerAppAuth`` resolves its ``mode``, +``base_url`` and ``jwt_verifying_key`` in default arguments, which Python +evaluates once at module import. ``monkeypatch.setenv`` therefore has no effect +on them -- always pass them explicitly, as these tests do. """ import base64 import json +import jwt import pytest from prowler_mcp_server.lib.errors import CredentialError from prowler_mcp_server.prowler_app.utils.auth import ProwlerAppAuth -from tests.helpers.tokens import FAKE_API_KEY, MALFORMED_API_KEY, fake_jwt +from tests.helpers.tokens import ( + FAKE_API_KEY, + JWT_SIGNING_KEY, + JWT_VERIFYING_KEY, + MALFORMED_API_KEY, + ROGUE_JWT_SIGNING_KEY, + fake_jwt, + hmac_jwt_keyed_with, + unsigned_jwt, +) async def test_stdio_mode_reads_the_api_key_from_the_environment(): @@ -41,7 +51,7 @@ async def test_http_mode_accepts_a_bearer_api_key(http_request_headers): """In HTTP transport the token comes from the request's Authorization header.""" http_request_headers(authorization=f"Bearer {FAKE_API_KEY}") - auth = ProwlerAppAuth(mode="http") + auth = ProwlerAppAuth(mode="http", jwt_verifying_key=JWT_VERIFYING_KEY) assert await auth.get_valid_token() == FAKE_API_KEY @@ -62,7 +72,7 @@ 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") + auth = ProwlerAppAuth(mode="http", jwt_verifying_key=JWT_VERIFYING_KEY) assert await auth.get_valid_token() == FAKE_API_KEY @@ -72,7 +82,7 @@ async def test_http_mode_strips_only_the_scheme_prefix(http_request_headers): token = f"{FAKE_API_KEY}_Bearer_suffix" http_request_headers(authorization=f"Bearer {token}") - auth = ProwlerAppAuth(mode="http") + auth = ProwlerAppAuth(mode="http", jwt_verifying_key=JWT_VERIFYING_KEY) assert await auth.get_valid_token() == token @@ -83,12 +93,121 @@ async def test_http_mode_rejects_an_authorization_header_without_a_token( """A bare scheme carries no credential to authenticate with.""" http_request_headers(authorization="Bearer ") - auth = ProwlerAppAuth(mode="http") + auth = ProwlerAppAuth(mode="http", jwt_verifying_key=JWT_VERIFYING_KEY) with pytest.raises(CredentialError, match="'Bearer ' form"): await auth.get_valid_token() +# ------------------------------------------------- JWT with a verifying key + + +async def test_http_mode_accepts_a_jwt_signed_by_the_api_key_pair( + http_request_headers, +): + """A token signed with the private half of the configured key pair passes.""" + token = fake_jwt() + http_request_headers(authorization=f"Bearer {token}") + + auth = ProwlerAppAuth(mode="http", jwt_verifying_key=JWT_VERIFYING_KEY) + + assert await auth.get_valid_token() == token + + +async def test_http_mode_rejects_a_jwt_with_a_forged_signature(http_request_headers): + """A well-formed, unexpired token signed by another key pair is refused.""" + token = fake_jwt(signing_key=ROGUE_JWT_SIGNING_KEY) + http_request_headers(authorization=f"Bearer {token}") + + auth = ProwlerAppAuth(mode="http", jwt_verifying_key=JWT_VERIFYING_KEY) + + with pytest.raises(CredentialError, match="could not be verified"): + await auth.get_valid_token() + + +async def test_http_mode_rejects_a_jwt_declaring_the_none_algorithm( + http_request_headers, +): + """`alg: none` is not in the pinned algorithm list, so the token is refused.""" + http_request_headers(authorization=f"Bearer {unsigned_jwt()}") + + auth = ProwlerAppAuth(mode="http", jwt_verifying_key=JWT_VERIFYING_KEY) + + with pytest.raises(CredentialError, match="could not be verified"): + await auth.get_valid_token() + + +async def test_http_mode_rejects_an_hmac_jwt_keyed_with_the_public_key( + http_request_headers, +): + """An HS256 token signed with the public key as the secret is refused. + + This is the algorithm-confusion attack: the public key is not secret, so a + verifier that honoured the token's own `alg` would accept it. + """ + token = hmac_jwt_keyed_with(JWT_VERIFYING_KEY) + http_request_headers(authorization=f"Bearer {token}") + + auth = ProwlerAppAuth(mode="http", jwt_verifying_key=JWT_VERIFYING_KEY) + + with pytest.raises(CredentialError, match="could not be verified"): + await auth.get_valid_token() + + +async def test_http_mode_rejects_an_expired_jwt_with_a_valid_signature( + http_request_headers, +): + """A correctly signed but expired token is refused locally.""" + http_request_headers(authorization=f"Bearer {fake_jwt(expires_in=-60)}") + + auth = ProwlerAppAuth(mode="http", jwt_verifying_key=JWT_VERIFYING_KEY) + + with pytest.raises(CredentialError, match="The token has expired"): + await auth.get_valid_token() + + +async def test_http_mode_rejects_a_signed_jwt_without_an_expiration( + http_request_headers, +): + """`exp` is required: a token that never expires is refused even if signed.""" + token = jwt.encode({"sub": "user"}, JWT_SIGNING_KEY, algorithm="RS256") + http_request_headers(authorization=f"Bearer {token}") + + auth = ProwlerAppAuth(mode="http", jwt_verifying_key=JWT_VERIFYING_KEY) + + with pytest.raises(CredentialError, match="could not be verified"): + await auth.get_valid_token() + + +async def test_http_mode_accepts_the_verifying_key_with_escaped_newlines( + http_request_headers, +): + """The PEM arrives through an env file, where newlines are written as `\\n`.""" + token = fake_jwt() + http_request_headers(authorization=f"Bearer {token}") + + auth = ProwlerAppAuth( + mode="http", jwt_verifying_key=JWT_VERIFYING_KEY.replace("\n", "\\n") + ) + + assert await auth.get_valid_token() == token + + +# ----------------------------------------------- JWT without a verifying key + + +async def test_http_mode_without_a_verifying_key_only_checks_expiration( + http_request_headers, +): + """With no key configured the token is forwarded for the API to verify.""" + token = fake_jwt(signing_key=ROGUE_JWT_SIGNING_KEY) + http_request_headers(authorization=f"Bearer {token}") + + auth = ProwlerAppAuth(mode="http", jwt_verifying_key=None) + + assert await auth.get_valid_token() == token + + async def test_http_mode_rejects_a_jwt_whose_payload_is_not_an_object( http_request_headers, ): @@ -99,27 +218,27 @@ async def test_http_mode_rejects_a_jwt_whose_payload_is_not_an_object( """ http_request_headers(authorization=f"Bearer {_jwt_with_payload(['exp'])}") - auth = ProwlerAppAuth(mode="http") + auth = ProwlerAppAuth(mode="http", jwt_verifying_key=None) with pytest.raises(CredentialError, match="not a readable JWT"): await auth.get_valid_token() @pytest.mark.parametrize( - ("payload", "case"), + "payload", [ - ({"sub": "user"}, "missing"), - ({"exp": "1700000000"}, "string"), - ({"exp": None}, "null"), + pytest.param({"sub": "user"}, id="missing"), + pytest.param({"exp": "1700000000"}, id="string"), + pytest.param({"exp": None}, id="null"), ], ) async def test_http_mode_rejects_a_jwt_without_a_numeric_expiration( - http_request_headers, payload: dict, case: str + http_request_headers, payload: dict ): """`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") + auth = ProwlerAppAuth(mode="http", jwt_verifying_key=None) with pytest.raises(CredentialError, match="no readable 'exp' expiration claim"): await auth.get_valid_token() @@ -129,7 +248,7 @@ async def test_http_mode_rejects_an_expired_jwt(http_request_headers): """An expired JWT is refused locally instead of being forwarded to the API.""" http_request_headers(authorization=f"Bearer {fake_jwt(expires_in=-60)}") - auth = ProwlerAppAuth(mode="http") + auth = ProwlerAppAuth(mode="http", jwt_verifying_key=None) with pytest.raises(CredentialError, match="The token has expired"): await auth.get_valid_token() @@ -141,5 +260,5 @@ def test_api_keys_and_jwts_use_different_authorization_schemes(): assert auth.get_headers(FAKE_API_KEY)["Authorization"] == f"Api-Key {FAKE_API_KEY}" - jwt = fake_jwt() - assert auth.get_headers(jwt)["Authorization"] == f"Bearer {jwt}" + token = fake_jwt() + assert auth.get_headers(token)["Authorization"] == f"Bearer {token}" diff --git a/mcp_server/uv.lock b/mcp_server/uv.lock index e529c9c3c4..da7623a164 100644 --- a/mcp_server/uv.lock +++ b/mcp_server/uv.lock @@ -792,6 +792,7 @@ source = { editable = "." } dependencies = [ { name = "fastmcp" }, { name = "httpx" }, + { name = "pyjwt", extra = ["crypto"] }, ] [package.dev-dependencies] @@ -810,6 +811,7 @@ dev = [ requires-dist = [ { name = "fastmcp", specifier = "==3.4.5" }, { name = "httpx", specifier = "==0.28.1" }, + { name = "pyjwt", extras = ["crypto"], specifier = "==2.14.0" }, ] [package.metadata.requires-dev]