#!/usr/bin/env python3
"""jwt_api_validation_check.py

A safe, vendor-neutral validation helper for teams securing a C# API with JWTs.

What it does:
- Validates a JWT bearer configuration file or a set of provided values.
- Checks for the minimum production-safe settings described in the source article.
- Produces a readable checklist-style report with pass/fail findings.
- Supports offline inspection of a JWT without contacting any identity provider.

What it does NOT do:
- It does not issue tokens.
- It does not call external authentication services.
- It does not modify application code or infrastructure.
- It does not handle secrets securely for you; keep signing keys out of source control.

Usage examples:

  # Validate a JSON config file
  python jwt_api_validation_check.py --config jwt-settings.json

  # Inspect a token locally and validate against expected values
  python jwt_api_validation_check.py \
      --token "<jwt>" \
      --expected-issuer "https://issuer.example" \
      --expected-audience "my-api" \
      --require-claim scope=reports.write

Example config file format:
{
  "issuer": "https://issuer.example",
  "audience": "my-api",
  "clock_skew_seconds": 60,
  "require_https": true,
  "allowed_algorithms": ["RS256"],
  "required_claims": {
    "scope": "reports.write"
  }
}
"""

from __future__ import annotations

import argparse
import base64
import json
import sys
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple


@dataclass
class CheckResult:
    name: str
    passed: bool
    message: str


@dataclass
class ValidationReport:
    results: List[CheckResult] = field(default_factory=list)

    def add(self, name: str, passed: bool, message: str) -> None:
        self.results.append(CheckResult(name=name, passed=passed, message=message))

    def exit_code(self) -> int:
        return 0 if all(r.passed for r in self.results) else 2

    def render(self) -> str:
        lines = []
        for r in self.results:
            status = "PASS" if r.passed else "FAIL"
            lines.append(f"[{status}] {r.name}: {r.message}")
        return "\n".join(lines)


def load_json_file(path: Path) -> Dict[str, Any]:
    if not path.exists():
        raise FileNotFoundError(f"Config file not found: {path}")
    if not path.is_file():
        raise ValueError(f"Config path is not a file: {path}")
    with path.open("r", encoding="utf-8") as f:
        data = json.load(f)
    if not isinstance(data, dict):
        raise ValueError("Config file must contain a JSON object at the top level")
    return data


def decode_b64url_segment(segment: str) -> Dict[str, Any]:
    padding = "=" * (-len(segment) % 4)
    raw = base64.urlsafe_b64decode((segment + padding).encode("ascii"))
    obj = json.loads(raw.decode("utf-8"))
    if not isinstance(obj, dict):
        raise ValueError("JWT segment is not a JSON object")
    return obj


def parse_jwt_unverified(token: str) -> Tuple[Dict[str, Any], Dict[str, Any], str]:
    parts = token.strip().split(".")
    if len(parts) != 3:
        raise ValueError("Token must have exactly three JWT segments")
    header = decode_b64url_segment(parts[0])
    payload = decode_b64url_segment(parts[1])
    signature = parts[2]
    return header, payload, signature


def get_claim(payload: Dict[str, Any], claim_name: str) -> Optional[Any]:
    return payload.get(claim_name)


def normalize_claim_value(value: Any) -> List[str]:
    if value is None:
        return []
    if isinstance(value, str):
        # Support space-delimited scope strings and single values.
        return [v for v in value.split() if v]
    if isinstance(value, list):
        return [str(v) for v in value]
    return [str(value)]


def check_required_field(report: ValidationReport, name: str, value: Any) -> None:
    report.add(name, value is not None, "present" if value is not None else "missing")


def check_string_match(report: ValidationReport, name: str, actual: Any, expected: Optional[str]) -> None:
    if expected is None:
        report.add(name, True, "not provided; skipped")
        return
    passed = str(actual) == expected
    report.add(name, passed, f"expected={expected!r}, actual={actual!r}")


def check_claim(report: ValidationReport, payload: Dict[str, Any], claim_name: str, expected_value: str) -> None:
    actual = get_claim(payload, claim_name)
    actual_values = normalize_claim_value(actual)
    expected_values = normalize_claim_value(expected_value)
    passed = all(ev in actual_values for ev in expected_values)
    report.add(
        f"claim:{claim_name}",
        passed,
        f"expected to include {expected_values!r}, actual={actual_values!r}",
    )


def validate_config(config: Dict[str, Any]) -> ValidationReport:
    report = ValidationReport()

    issuer = config.get("issuer")
    audience = config.get("audience")
    clock_skew = config.get("clock_skew_seconds", None)
    require_https = config.get("require_https", True)
    allowed_algorithms = config.get("allowed_algorithms", [])
    required_claims = config.get("required_claims", {})

    check_required_field(report, "issuer", issuer)
    check_required_field(report, "audience", audience)

    report.add(
        "require_https",
        bool(require_https) is True,
        "HTTPS is required" if bool(require_https) is True else "HTTPS requirement is disabled",
    )

    if clock_skew is None:
        report.add("clock_skew_seconds", True, "not provided; relying on framework default")
    else:
        try:
            clock_skew_int = int(clock_skew)
            passed = 0 <= clock_skew_int <= 300
            report.add(
                "clock_skew_seconds",
                passed,
                f"value={clock_skew_int}; recommended range is 0..300 seconds",
            )
        except (TypeError, ValueError):
            report.add("clock_skew_seconds", False, f"invalid integer value: {clock_skew!r}")

    if not isinstance(allowed_algorithms, list):
        report.add("allowed_algorithms", False, "must be a list of algorithms")
    else:
        passed = len(allowed_algorithms) > 0
        report.add(
            "allowed_algorithms",
            passed,
            f"configured={allowed_algorithms!r}" if passed else "no algorithms configured",
        )

    if isinstance(required_claims, dict) and required_claims:
        for claim_name, expected_value in required_claims.items():
            check_claim(report, {claim_name: expected_value}, claim_name, str(expected_value))
    else:
        report.add("required_claims", True, "no claim requirements configured")

    return report


def validate_token_against_expectations(
    token: str,
    expected_issuer: Optional[str],
    expected_audience: Optional[str],
    require_claims: List[str],
) -> ValidationReport:
    report = ValidationReport()

    header, payload, _signature = parse_jwt_unverified(token)

    report.add("jwt_structure", True, "token has three segments")

    alg = header.get("alg")
    typ = header.get("typ")
    check_required_field(report, "header.alg", alg)
    report.add("header.typ", typ in (None, "JWT"), f"value={typ!r}")

    check_string_match(report, "issuer", payload.get("iss"), expected_issuer)
    check_string_match(report, "audience", payload.get("aud"), expected_audience)

    if "exp" in payload:
        report.add("exp_present", True, "expiration claim found")
    else:
        report.add("exp_present", False, "missing expiration claim")

    if "nbf" in payload:
        report.add("nbf_present", True, "not-before claim found")
    else:
        report.add("nbf_present", True, "not-before claim not required")

    for requirement in require_claims:
        if "=" not in requirement:
            report.add("required_claim_format", False, f"invalid claim requirement: {requirement!r}")
            continue
        claim_name, expected_value = requirement.split("=", 1)
        check_claim(report, payload, claim_name.strip(), expected_value.strip())

    return report


def parse_required_claims(values: List[str]) -> List[str]:
    return values


def main() -> int:
    parser = argparse.ArgumentParser(
        description="Validate JWT API security expectations for a C# ASP.NET Core API.")
    parser.add_argument("--config", type=Path, help="Path to a JSON config file")
    parser.add_argument("--token", help="JWT to inspect locally (unverified decoding only)")
    parser.add_argument("--expected-issuer", help="Expected issuer value for the token")
    parser.add_argument("--expected-audience", help="Expected audience value for the token")
    parser.add_argument(
        "--require-claim",
        action="append",
        default=[],
        help="Required claim in the form name=value; may be repeated",
    )
    parser.add_argument(
        "--print-json",
        action="store_true",
        help="Print the results as JSON instead of text",
    )

    args = parser.parse_args()

    overall = ValidationReport()

    if args.config:
        try:
            config = load_json_file(args.config)
            config_report = validate_config(config)
            overall.results.extend(config_report.results)
        except Exception as exc:
            overall.add("config_load", False, str(exc))

    if args.token:
        try:
            token_report = validate_token_against_expectations(
                args.token,
                args.expected_issuer,
                args.expected_audience,
                parse_required_claims(args.require_claim),
            )
            overall.results.extend(token_report.results)
        except Exception as exc:
            overall.add("token_parse", False, str(exc))
    else:
        overall.add("token_parse", True, "no token provided; skipped local inspection")

    if args.print_json:
        print(json.dumps(
            {
                "passed": all(r.passed for r in overall.results),
                "results": [r.__dict__ for r in overall.results],
            },
            indent=2,
        ))
    else:
        print(overall.render())

    return overall.exit_code()


if __name__ == "__main__":
    raise SystemExit(main())