fix(mcp): verify JWT signatures in HTTP transport mode

This commit is contained in:
pedrooot
2026-10-06 14:48:33 +02:00
parent d38e04dd3e
commit 9613f7787a
7 changed files with 284 additions and 63 deletions
+1
View File
@@ -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.
@@ -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
@@ -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
+3 -3
View File
@@ -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
+79 -14
View File
@@ -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)}"
+138 -19
View File
@@ -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 <token>' 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}"
+2
View File
@@ -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]