"""Atlassian Connect (Confluence) trust surface: qsh, JWT verify, secret store. Private Connect app, symmetric (HS256) shared-secret flow. Atlassian mints a ``sharedSecret`` per install and signs each request with it; we verify the signature, the ``exp``, and the ``qsh`` (query-string-hash) claim that binds a token to one exact method+path+query (defeating cross-endpoint replay). The shared secret is stored encrypted-at-rest in the LangGraph store, keyed by the tenant ``clientKey``. ``qsh`` is Connect-specific (no JWT library provides it) so it is hand-rolled here from the Atlassian spec and pinned to the official test vector in ``tests/test_atlassian_connect.py``. A subtle bug here is a silent auth bypass. """ from __future__ import annotations import hashlib import hmac import logging import os import re import time from typing import Any from urllib.parse import parse_qs, quote import httpx import jwt from langgraph_sdk import get_client from ..encryption import decrypt_token, encrypt_token from .http import DEFAULT_HTTP_TIMEOUT logger = logging.getLogger(__name__) LANGGRAPH_URL = os.environ.get("LANGGRAPH_URL") or os.environ.get( "LANGGRAPH_URL_PROD", "http://localhost:2024" ) # The app's public origin — must equal the descriptor baseUrl and the `aud` in # Atlassian's signed-install lifecycle JWTs. CONNECT_BASE_URL = os.environ.get("CONNECT_BASE_URL", "").rstrip("/") # Atlassian's CDN of public keys for signed-install (asymmetric) lifecycle JWTs. _CONNECT_INSTALL_KEYS_BASE = "https://connect-install-keys.atlassian.com" # Tenant binding (REQUIRED): the signature-verified clientKey(s) — i.e. the JWT # `iss` — we accept installs from. signed-install proves the caller is *an* # Atlassian tenant, not *ours*, and the descriptor is served publicly, so # without this any tenant could install the app and drive runs. Empty => reject # ALL installs (fail closed). The install `baseUrl` is an untrusted body field # and is NOT a valid binding; only the signed `iss` is. Bootstrap: attempt an # install, read the rejected clientKey from the logs, add it here, re-install. CONNECT_EXPECTED_CLIENT_KEYS: frozenset[str] = frozenset( key.strip() for key in os.environ.get("CONNECT_EXPECTED_CLIENT_KEYS", "").replace(",", " ").split() if key.strip() ) # Optional defense-in-depth on the (untrusted) install baseUrl host. Enforced # only when set — the clientKey allowlist above is the real tenant gate. CONNECT_EXPECTED_BASE_URL_HOSTS: frozenset[str] = frozenset( host.strip().lower() for host in os.environ.get("CONNECT_EXPECTED_BASE_URL", "").replace(",", " ").split() if host.strip() ) _JWT_LEEWAY_SECONDS = 10 _INSTALL_NS = ("atlassian_connect", "installations") # --- qsh (query-string hash) ----------------------------------------------- def _encode(value: Any) -> str: """RFC-3986 encode a single component (Atlassian encodeRfc3986: space->%20).""" return quote(str(value), safe="") def _canonical_uri(path: str) -> str: if not path: return "/" if len(path) > 1 and path.endswith("/"): path = path[:-1] return path.replace("&", "%26") def _canonical_query(query: str) -> str: if not query: return "" parsed = parse_qs(query, keep_blank_values=True) parsed.pop("jwt", None) pairs = [ f"{_encode(key)}={','.join(sorted(_encode(v) for v in values))}" for key, values in parsed.items() ] pairs.sort() return "&".join(pairs) def canonical_request(method: str, path: str, query: str = "") -> str: return "&".join([method.upper(), _canonical_uri(path), _canonical_query(query)]) def compute_qsh(method: str, path: str, query: str = "") -> str: """SHA-256 hex of the canonical `METHOD&path&query` request string.""" return hashlib.sha256(canonical_request(method, path, query).encode()).hexdigest() # --- token extraction + JWT verification ----------------------------------- def extract_connect_token(request: Any) -> str | None: """Pull the Connect JWT from `Authorization: JWT ` or the `?jwt=` param.""" header = request.headers.get("Authorization") or request.headers.get("authorization") or "" if header[:4].upper() == "JWT ": token = header[4:].strip() if token: return token query_token = request.query_params.get("jwt") if hasattr(request, "query_params") else None return query_token or None def verify_connect_jwt( request: Any, *, shared_secret: str | None = None, expected_client_key: str | None = None, qsh_required: bool = True, ) -> dict[str, Any] | None: """Verify a Connect JWT. Returns the claims on success, None on any failure. ``shared_secret`` is passed explicitly by the lifecycle routes (the stored secret); the webhook resolves it from the token's ``iss`` via the store. ``qsh_required`` is True for the webhook and False (verify-if-present) for lifecycle callbacks, whose auth strength comes from the signature + issuer binding rather than qsh. Every failure path returns None with no side effect. """ token = extract_connect_token(request) if not token: logger.warning("Connect JWT missing — rejecting") return None # Fail-fast alg pin before any lookup: blocks alg=none and RS256/ES256 # algorithm-confusion against a symmetric secret. try: alg = jwt.get_unverified_header(token).get("alg") except jwt.PyJWTError: logger.warning("Connect JWT header undecodable — rejecting") return None if alg != "HS256": logger.warning("Connect JWT alg %r is not HS256 — rejecting", alg) return None # Read the (untrusted) issuer to resolve the secret; trust nothing yet. try: unverified = jwt.decode(token, options={"verify_signature": False}) except jwt.PyJWTError: logger.warning("Connect JWT undecodable — rejecting") return None issuer = unverified.get("iss") if not issuer: logger.warning("Connect JWT missing iss — rejecting") return None # The caller always supplies the secret: lifecycle passes the stored secret; # the webhook resolves it from `iss` via get_shared_secret and passes it. # A missing secret fails closed (never verify against nothing / a default). if not shared_secret: logger.warning("No shared secret provided for Connect JWT — rejecting") return None try: claims = jwt.decode( token, shared_secret, algorithms=["HS256"], options={ "require": ["exp", "iss"], "verify_signature": True, "verify_exp": True, "verify_nbf": True, }, leeway=_JWT_LEEWAY_SECONDS, ) except jwt.PyJWTError as exc: logger.warning("Connect JWT signature/claims invalid: %s", exc.__class__.__name__) return None if expected_client_key is not None and claims.get("iss") != expected_client_key: logger.warning("Connect JWT iss does not match expected client key — rejecting") return None # qsh verified LAST — it is a signed claim, only trustworthy post-signature. qsh_claim = claims.get("qsh") if qsh_claim == "context-qsh": logger.warning("Connect JWT carries context-qsh (iframe token) — rejecting") return None if qsh_claim is None: if qsh_required: logger.warning("Connect JWT missing required qsh — rejecting") return None else: expected_qsh = compute_qsh(request.method, request.url.path, request.url.query) if not hmac.compare_digest(expected_qsh, qsh_claim): logger.warning("Connect JWT qsh mismatch — rejecting") return None return claims _KID_RE = re.compile(r"^[A-Za-z0-9._-]+$") async def _fetch_atlassian_public_key(kid: str) -> str | None: """Fetch Atlassian's PEM public key for a signed-install ``kid``. ``kid`` is validated to a strict charset before use (defense-in-depth on top of the fixed host + percent-encoding) so a malformed kid fails fast without a network call and can never influence the request path. """ if not _KID_RE.match(kid): logger.warning("Rejecting Connect install: malformed kid") return None async with httpx.AsyncClient(timeout=DEFAULT_HTTP_TIMEOUT) as client: try: response = await client.get(f"{_CONNECT_INSTALL_KEYS_BASE}/{quote(kid, safe='')}") response.raise_for_status() return response.text except Exception as exc: # noqa: BLE001 logger.warning("Failed to fetch Atlassian install public key: %s", exc) return None async def verify_asymmetric_install_jwt( request: Any, *, expected_client_key: str | None = None ) -> dict[str, Any] | None: """Verify a signed-install (RS256) lifecycle JWT against Atlassian's keys. signed-install cryptographically authenticates the install/uninstall callbacks — including the FIRST install — so first install is not trust-on-first-use. The token is RS256-signed by Atlassian; its ``kid`` selects a published public key, and its ``aud`` must equal this app's baseUrl (blocking tokens minted for another Connect app). qsh is verify-if-present on lifecycle. Returns claims on success, None otherwise. """ if not CONNECT_BASE_URL: logger.warning("CONNECT_BASE_URL unset — cannot verify signed-install aud; rejecting") return None token = extract_connect_token(request) if not token: return None try: header = jwt.get_unverified_header(token) except jwt.PyJWTError: return None if header.get("alg") != "RS256": logger.warning("Signed-install JWT alg %r is not RS256 — rejecting", header.get("alg")) return None kid = header.get("kid") if not kid: return None public_key = await _fetch_atlassian_public_key(kid) if not public_key: return None try: claims = jwt.decode( token, public_key, algorithms=["RS256"], audience=CONNECT_BASE_URL, options={ "require": ["exp", "iss", "aud"], "verify_signature": True, "verify_exp": True, "verify_nbf": True, "verify_aud": True, }, leeway=_JWT_LEEWAY_SECONDS, ) except jwt.PyJWTError as exc: logger.warning("Signed-install JWT invalid: %s", exc.__class__.__name__) return None if expected_client_key is not None and claims.get("iss") != expected_client_key: logger.warning("Signed-install JWT iss does not match clientKey — rejecting") return None qsh_claim = claims.get("qsh") if qsh_claim == "context-qsh": return None if qsh_claim is not None: expected_qsh = compute_qsh(request.method, request.url.path, request.url.query) if not hmac.compare_digest(expected_qsh, qsh_claim): logger.warning("Signed-install JWT qsh mismatch — rejecting") return None return claims async def verify_connect_webhook(request: Any) -> dict[str, Any] | None: """Resolve the shared secret from the token's issuer, then verify (qsh required). The async counterpart to verify_connect_jwt for the webhook route, where the secret must be looked up from the store by the (untrusted-until-verified) issuer. Returns claims on success, None on any failure. """ token = extract_connect_token(request) if not token: return None try: issuer = jwt.decode(token, options={"verify_signature": False}).get("iss") except jwt.PyJWTError: return None if not issuer: return None secret = await get_shared_secret(issuer) if not secret: logger.warning("No installation for Connect issuer — rejecting webhook") return None return verify_connect_jwt( request, shared_secret=secret, expected_client_key=issuer, qsh_required=True ) # --- encrypted installation store ------------------------------------------ def _client() -> Any: return get_client(url=LANGGRAPH_URL) async def get_installation(client_key: str) -> dict[str, Any] | None: """Return the stored installation record (secret decrypted), or None.""" item = await _client().store.get_item(_INSTALL_NS, client_key) if not item: return None value = item.get("value") if not isinstance(value, dict): return None record = dict(value) try: record["shared_secret"] = decrypt_token(record["shared_secret_enc"]) except Exception: # noqa: BLE001 — fail closed on decrypt/missing-key failure logger.warning("Failed to decrypt stored Connect shared secret for %s", client_key) return None return record async def get_shared_secret(client_key: str) -> str | None: record = await get_installation(client_key) return (record or {}).get("shared_secret") or None async def put_installation( client_key: str, shared_secret: str, base_url: str, product_type: str, *, first_install: bool, ) -> None: now_ms = int(time.time() * 1000) installed_at = now_ms if not first_install: existing = await get_installation(client_key) installed_at = (existing or {}).get("installed_at_ms", now_ms) record = { "client_key": client_key, "shared_secret_enc": encrypt_token(shared_secret), "base_url": base_url, "product_type": product_type, "installed_at_ms": installed_at, "updated_at_ms": now_ms, } await _client().store.put_item(_INSTALL_NS, client_key, record) async def delete_installation(client_key: str) -> None: await _client().store.delete_item(_INSTALL_NS, client_key) def client_key_allowed(client_key: str) -> bool: """Whether a signature-verified clientKey (JWT iss) is an accepted tenant. Fail closed: an empty allowlist accepts no installs (the mandatory tenant binding — see CONNECT_EXPECTED_CLIENT_KEYS). """ return bool(client_key) and client_key in CONNECT_EXPECTED_CLIENT_KEYS def base_url_host_allowed(base_url: str) -> bool: """Whether an install baseUrl's host is in the optional CONNECT_EXPECTED_BASE_URL allowlist. Defense-in-depth only (the baseUrl is an untrusted body field); the real tenant gate is client_key_allowed. Returns False when the allowlist is empty. """ if not CONNECT_EXPECTED_BASE_URL_HOSTS: return False from urllib.parse import urlparse host = (urlparse(base_url).hostname or "").lower() return bool(host) and host in CONNECT_EXPECTED_BASE_URL_HOSTS