#!/usr/bin/env python3
"""ASP.NET Core JWT Validation Checklist

A non-destructive helper script for reviewing JWT authentication settings and
endpoint authorization requirements before production use.

This script does NOT validate real tokens. Instead, it helps teams review a
JWT configuration and compare it against a practical security checklist drawn
from common ASP.NET Core API pitfalls.

Usage examples:
  python jwt_validation_checklist.py --issuer https://issuer.example
  python jwt_validation_checklist.py --issuer https://issuer.example --audience api://orders --require-tenant --required-claims scope roles

The script prints a readable assessment and exits with code 0 when all selected
checks pass, or 1 when one or more checks fail.
"""

from __future__ import annotations

import argparse
import json
import sys
from dataclasses import dataclass, field
from typing import List, Optional
from urllib.parse import urlparse


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


@dataclass
class ChecklistConfig:
    issuer: str
    audience: str
    required_claims: List[str] = field(default_factory=list)
    require_tenant: bool = False
    tenant_claim: str = "tenant_id"
    tenant_source: str = "request context"
    require_https: bool = True
    allowed_clock_skew_seconds: int = 300
    token_lifetime_minutes: Optional[int] = None


def validate_https_url(value: str, label: str) -> Optional[str]:
    if not value:
        return f"{label} is required."

    parsed = urlparse(value)
    if parsed.scheme != "https":
        return f"{label} should use https:// to protect tokens in transit."
    if not parsed.netloc:
        return f"{label} must include a valid host."
    return None


def validate_audience(value: str) -> Optional[str]:
    if not value:
        return "Audience is required so tokens intended for other APIs are rejected."
    if value.strip() != value:
        return "Audience should not contain leading or trailing whitespace."
    if " " in value:
        return "Audience should not contain spaces."
    return None


def validate_required_claims(claims: List[str]) -> Optional[str]:
    if not claims:
        return "At least one required claim should be specified for authorization policy checks."
    normalized = [c.strip() for c in claims if c.strip()]
    if len(normalized) != len(claims):
        return "Required claims should not contain empty values."
    return None


def validate_tenant_settings(require_tenant: bool, tenant_claim: str, tenant_source: str) -> Optional[str]:
    if require_tenant:
        if not tenant_claim.strip():
            return "Tenant claim name is required when tenant validation is enabled."
        if not tenant_source.strip():
            return "Tenant source is required when tenant validation is enabled."
    return None


def validate_clock_skew(seconds: int) -> Optional[str]:
    if seconds < 0:
        return "Allowed clock skew cannot be negative."
    if seconds > 300:
        return "Allowed clock skew is high; consider reducing it to avoid widening the acceptance window."
    return None


def validate_token_lifetime(minutes: Optional[int]) -> Optional[str]:
    if minutes is None:
        return None
    if minutes <= 0:
        return "Token lifetime must be greater than zero."
    if minutes > 60 * 24:
        return "Token lifetime is long; shorter lifetimes are usually safer for bearer tokens."
    return None


def run_checks(cfg: ChecklistConfig) -> List[ValidationResult]:
    results: List[ValidationResult] = []

    issuer_err = validate_https_url(cfg.issuer, "Issuer")
    results.append(
        ValidationResult(
            name="Issuer validation",
            passed=issuer_err is None,
            message=issuer_err or "Issuer looks valid and uses HTTPS.",
        )
    )

    audience_err = validate_audience(cfg.audience)
    results.append(
        ValidationResult(
            name="Audience validation",
            passed=audience_err is None,
            message=audience_err or "Audience is specified for API-bound tokens.",
        )
    )

    claims_err = validate_required_claims(cfg.required_claims)
    results.append(
        ValidationResult(
            name="Required claims",
            passed=claims_err is None,
            message=claims_err or f"Required claims specified: {', '.join(cfg.required_claims)}",
        )
    )

    tenant_err = validate_tenant_settings(cfg.require_tenant, cfg.tenant_claim, cfg.tenant_source)
    results.append(
        ValidationResult(
            name="Tenant binding",
            passed=tenant_err is None,
            message=tenant_err or (
                f"Tenant validation enabled using claim '{cfg.tenant_claim}' against {cfg.tenant_source}."
                if cfg.require_tenant
                else "Tenant validation not enabled."
            ),
        )
    )

    skew_err = validate_clock_skew(cfg.allowed_clock_skew_seconds)
    results.append(
        ValidationResult(
            name="Clock skew",
            passed=skew_err is None,
            message=skew_err or f"Allowed clock skew: {cfg.allowed_clock_skew_seconds} seconds.",
        )
    )

    lifetime_err = validate_token_lifetime(cfg.token_lifetime_minutes)
    results.append(
        ValidationResult(
            name="Token lifetime",
            passed=lifetime_err is None,
            message=lifetime_err or (
                f"Token lifetime: {cfg.token_lifetime_minutes} minutes." if cfg.token_lifetime_minutes else "Token lifetime not specified."
            ),
        )
    )

    return results


def print_report(results: List[ValidationResult]) -> int:
    passed = 0
    failed = 0
    print("JWT Validation Checklist Report")
    print("=" * 32)
    for result in results:
        status = "PASS" if result.passed else "FAIL"
        if result.passed:
            passed += 1
        else:
            failed += 1
        print(f"[{status}] {result.name}: {result.message}")

    print("\nSummary")
    print("-" * 7)
    print(f"Passed: {passed}")
    print(f"Failed: {failed}")
    return 0 if failed == 0 else 1


def parse_args(argv: List[str]) -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Review JWT validation settings for an ASP.NET Core API before production use.",
        formatter_class=argparse.ArgumentDefaultsHelpFormatter,
    )
    parser.add_argument("--issuer", required=True, help="Expected token issuer (use an https:// URL).")
    parser.add_argument("--audience", required=True, help="Expected token audience for this API.")
    parser.add_argument(
        "--required-claims",
        nargs="*",
        default=[],
        help="Claims required by endpoint authorization policies, such as scope, role, or client_id.",
    )
    parser.add_argument(
        "--require-tenant",
        action="store_true",
        help="Enable tenant binding checks for multi-tenant APIs.",
    )
    parser.add_argument(
        "--tenant-claim",
        default="tenant_id",
        help="Claim used to identify tenant context.",
    )
    parser.add_argument(
        "--tenant-source",
        default="request context",
        help="Source used to determine the expected tenant, such as host, path, or access policy.",
    )
    parser.add_argument(
        "--allowed-clock-skew-seconds",
        type=int,
        default=300,
        help="Maximum acceptable clock skew for token lifetime checks.",
    )
    parser.add_argument(
        "--token-lifetime-minutes",
        type=int,
        default=None,
        help="Optional expected token lifetime in minutes for review purposes.",
    )
    parser.add_argument(
        "--json",
        action="store_true",
        help="Output results as JSON instead of a text report.",
    )
    return parser.parse_args(argv)


def main(argv: List[str]) -> int:
    args = parse_args(argv)

    cfg = ChecklistConfig(
        issuer=args.issuer,
        audience=args.audience,
        required_claims=args.required_claims,
        require_tenant=args.require_tenant,
        tenant_claim=args.tenant_claim,
        tenant_source=args.tenant_source,
        allowed_clock_skew_seconds=args.allowed_clock_skew_seconds,
        token_lifetime_minutes=args.token_lifetime_minutes,
    )

    results = run_checks(cfg)

    if args.json:
        payload = {
            "passed": sum(1 for r in results if r.passed),
            "failed": sum(1 for r in results if not r.passed),
            "results": [r.__dict__ for r in results],
        }
        print(json.dumps(payload, indent=2))
        return 0 if payload["failed"] == 0 else 1

    return print_report(results)


if __name__ == "__main__":
    raise SystemExit(main(sys.argv[1:]))