"""Standalone security checks.

This runner intentionally avoids pytest because the project package named
``platform`` currently collides with Python's standard-library module.
"""

from __future__ import annotations

import os
import sys
import tempfile
import time
import unittest
from pathlib import Path


PROJECT_ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(PROJECT_ROOT))
loaded_platform = sys.modules.get("platform")
if loaded_platform is not None and not hasattr(loaded_platform, "__path__"):
    del sys.modules["platform"]

import jwt
from cryptography.hazmat.primitives.asymmetric import rsa
from cryptography.hazmat.primitives.serialization import (
    Encoding,
    NoEncryption,
    PrivateFormat,
    PublicFormat,
)
from flask import Flask

from platform.gateway.platform_api import platform_api
from platform.auth import jwt_auth


class JWTConnectionChecks(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        cls.temp_dir = tempfile.TemporaryDirectory()
        cls.keys_dir = Path(cls.temp_dir.name)
        private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
        cls.private_pem = private_key.private_bytes(
            Encoding.PEM,
            PrivateFormat.PKCS8,
            NoEncryption(),
        )
        public_pem = private_key.public_key().public_bytes(
            Encoding.PEM,
            PublicFormat.SubjectPublicKeyInfo,
        )
        cls.public_pem = public_pem.decode("utf-8")
        (cls.keys_dir / "test-key.pem").write_bytes(public_pem)
        os.environ["JWT_PUBLIC_KEYS_DIR"] = str(cls.keys_dir)
        os.environ["JWT_AUDIENCE"] = "absolutems-api"
        os.environ["JWT_ALLOWED_ISSUERS"] = "https://absolutems.com.au"

        app = Flask(__name__)
        app.config["REGISTER_BOOTSTRAP_SITES"] = False
        app.register_blueprint(platform_api)
        cls.client = app.test_client()

    @classmethod
    def tearDownClass(cls):
        cls.temp_dir.cleanup()

    @classmethod
    def make_token(cls, **overrides) -> str:
        now = int(time.time())
        claims = {
            "iss": "https://absolutems.com.au",
            "aud": "absolutems-api",
            "sub": "wp-user-1",
            "tenant_id": "absolutems",
            "site_id": "test-site",
            "roles": ["administrator"],
            "scopes": ["platform:connect"],
            "jti": "test-request",
            "iat": now,
            "nbf": now - 1,
            "exp": now + 300,
        }
        claims.update(overrides)
        return jwt.encode(
            claims,
            cls.private_pem,
            algorithm="RS256",
            headers={"kid": "test-key"},
        )

    def test_missing_token_is_rejected(self):
        response = self.client.get("/api/v1/platform/connection")
        self.assertEqual(401, response.status_code)
        self.assertEqual("Missing Bearer token", response.get_json()["message"])

    def test_valid_token_is_accepted(self):
        response = self.client.get(
            "/api/v1/platform/connection",
            headers={"Authorization": f"Bearer {self.make_token()}"},
        )
        payload = response.get_json()
        self.assertEqual(200, response.status_code)
        self.assertTrue(payload["authenticated"])
        self.assertEqual("absolutems", payload["identity"]["tenant_id"])
        self.assertEqual("wp-user-1", payload["identity"]["subject"])

    def test_wrong_scope_is_rejected(self):
        response = self.client.get(
            "/api/v1/platform/connection",
            headers={
                "Authorization": f"Bearer {self.make_token(scopes=['wp_invoices:read'])}",
            },
        )
        self.assertEqual(403, response.status_code)

    def test_wrong_issuer_is_rejected(self):
        response = self.client.get(
            "/api/v1/platform/connection",
            headers={
                "Authorization": f"Bearer {self.make_token(iss='https://attacker.example')}",
            },
        )
        self.assertEqual(401, response.status_code)

    def test_registered_revoked_key_never_falls_back_to_file(self):
        original = jwt_auth._registered_key
        jwt_auth._registered_key = lambda kid: {
            "key_id": kid,
            "site_id": "test-site",
            "tenant_id": "absolutems",
            "issuer": "https://absolutems.com.au",
            "audience": "absolutems-api",
            "public_key_pem": self.public_pem,
            "key_status": "revoked",
            "site_status": "active",
        }
        try:
            response = self.client.get(
                "/api/v1/platform/connection",
                headers={"Authorization": f"Bearer {self.make_token()}"},
            )
        finally:
            jwt_auth._registered_key = original
        self.assertEqual(401, response.status_code)
        self.assertIn("revoked", response.get_json()["message"])


if __name__ == "__main__":
    unittest.main(verbosity=2)
