"""Verify Apple's PEM request boundary with real CSR, certificate and PKCS12 bytes."""
from __future__ import annotations

import base64
import datetime as dt
import json
import unittest
from contextlib import redirect_stdout
from io import StringIO
from unittest import mock

import requests
from cryptography import x509
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.serialization import pkcs12

import cert_factory


class CertificateRequestTests(unittest.TestCase):
    def test_generated_pem_has_a_valid_signature_and_matching_public_key(self):
        key, content = cert_factory.generate_key_and_csr()
        csr = x509.load_pem_x509_csr(content)
        self.assertTrue(csr.is_signature_valid)
        self.assertEqual(csr.public_key().public_numbers(), key.public_key().public_numbers())

    def test_posted_pem_and_returned_certificate_preserve_the_signing_identity(self):
        key, content = cert_factory.generate_key_and_csr()
        received = []

        def apple_post(method, path, token, *, json_body, allow_status):
            self.assertEqual((method, path), ("POST", "/certificates"))
            attributes = json_body["data"]["attributes"]
            self.assertEqual(attributes["certificateType"], "DISTRIBUTION")
            csr = x509.load_pem_x509_csr(attributes["csrContent"].encode("ascii"))
            self.assertTrue(csr.is_signature_valid)
            now = dt.datetime.now(dt.timezone.utc)
            cert = (x509.CertificateBuilder().subject_name(csr.subject).issuer_name(csr.subject)
                    .public_key(csr.public_key()).serial_number(x509.random_serial_number())
                    .not_valid_before(now - dt.timedelta(minutes=1))
                    .not_valid_after(now + dt.timedelta(days=1)).sign(key, hashes.SHA256()))
            received.append(cert)
            response = requests.Response()
            response.status_code = 201
            response._content = json.dumps({"data": {"id": "TESTCERT", "attributes": {
                "certificateContent": base64.b64encode(cert.public_bytes(serialization.Encoding.DER)).decode()
            }}}).encode()
            return response

        with mock.patch.object(cert_factory, "request", side_effect=apple_post) as request:
            with redirect_stdout(StringIO()):
                cert_id, der = cert_factory.create_distribution_cert("test-token", content.decode("ascii"))
        request.assert_called_once()
        self.assertEqual(cert_id, "TESTCERT")
        p12 = cert_factory.serialize_p12(key, der, "test-only-password")
        restored_key, restored_cert, _ = pkcs12.load_key_and_certificates(p12, b"test-only-password")
        self.assertEqual(restored_key.public_key().public_numbers(), key.public_key().public_numbers())
        self.assertEqual(restored_cert.fingerprint(hashes.SHA256()), received[0].fingerprint(hashes.SHA256()))


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