#!/usr/bin/env python3
"""
Tests for team_resolver.derive_team_id — the ASC-API-based Apple developer
team identifier derivation used when ``team_id`` is omitted from
``ci.config.yaml``.

Runs offline: every ASC API interaction is stubbed via unittest.mock.
"""

from __future__ import annotations

import base64
import datetime
import plistlib
import subprocess
import sys
import unittest
from pathlib import Path
from unittest import mock

from cryptography import x509
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from cryptography.x509.oid import NameOID

# Make sibling scripts importable regardless of the pytest invocation cwd.
sys.path.insert(0, str(Path(__file__).resolve().parent))

import team_resolver  # noqa: E402  (import-after-sys.path modification)


def _make_dist_cert_der(team_id: str) -> bytes:
    """Generate a throwaway DER-encoded Apple-Distribution-style cert with OU=<team_id>."""
    key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
    subject = issuer = x509.Name(
        [
            x509.NameAttribute(
                NameOID.COMMON_NAME,
                f"Apple Distribution: Test Team ({team_id})",
            ),
            x509.NameAttribute(NameOID.ORGANIZATIONAL_UNIT_NAME, team_id),
            x509.NameAttribute(NameOID.ORGANIZATION_NAME, "Test Team"),
            x509.NameAttribute(NameOID.COUNTRY_NAME, "US"),
        ]
    )
    now = datetime.datetime.now(datetime.timezone.utc)
    cert = (
        x509.CertificateBuilder()
        .subject_name(subject)
        .issuer_name(issuer)
        .public_key(key.public_key())
        .serial_number(x509.random_serial_number())
        .not_valid_before(now)
        .not_valid_after(now + datetime.timedelta(days=365))
        .sign(key, hashes.SHA256())
    )
    return cert.public_bytes(serialization.Encoding.DER)


def _make_profile_content(team_id: str) -> bytes:
    """Build a fake mobileprovision payload (plain plist, not CMS-signed).

    team_resolver falls back to ``plistlib.loads`` directly when the payload
    isn't a signed blob, so a raw plist is sufficient for the test path.
    """
    plist_dict = {
        "AppIDName": "Test",
        "TeamIdentifier": [team_id],
        "Name": "CI-com.example.test",
        "UUID": "00000000-0000-0000-0000-000000000000",
    }
    return plistlib.dumps(plist_dict)


class DeriveTeamIdFromCertificateTests(unittest.TestCase):
    """Happy path: ASC returns at least one Distribution cert."""

    def test_derives_team_from_cert_ou_when_cert_exists(self):
        team_id = "QYVMNH6654"
        cert_der = _make_dist_cert_der(team_id)
        cert_b64 = base64.b64encode(cert_der).decode()
        asc_response = {
            "data": [
                {
                    "id": "CERT123",
                    "attributes": {
                        "certificateType": "DISTRIBUTION",
                        "certificateContent": cert_b64,
                    },
                }
            ]
        }
        with mock.patch.object(
            team_resolver, "get_json", return_value=asc_response
        ) as m_get:
            derived = team_resolver.derive_team_id("jwt-token")
        self.assertEqual(derived, team_id)
        m_get.assert_any_call(
            "/certificates",
            "jwt-token",
            params={"limit": "1", "sort": "-id"},
        )

    def test_returns_team_even_when_cert_type_is_development(self):
        team_id = "VD3WNXP8BH"
        cert_der = _make_dist_cert_der(team_id)
        cert_b64 = base64.b64encode(cert_der).decode()
        asc_response = {
            "data": [
                {
                    "id": "DEVCERT",
                    "attributes": {
                        "certificateType": "DEVELOPMENT",
                        "certificateContent": cert_b64,
                    },
                }
            ]
        }
        with mock.patch.object(team_resolver, "get_json", return_value=asc_response):
            derived = team_resolver.derive_team_id("jwt-token")
        self.assertEqual(derived, team_id)


class DeriveTeamIdFromProfileFallbackTests(unittest.TestCase):
    """Fallback when no certs exist: use an existing provisioning profile's plist."""

    def test_falls_back_to_profile_when_no_certs(self):
        team_id = "ABC1234567"
        profile_b64 = base64.b64encode(_make_profile_content(team_id)).decode()
        cert_response = {"data": []}
        profile_response = {
            "data": [
                {
                    "id": "PROF123",
                    "attributes": {
                        "name": "CI-com.example.test",
                        "profileContent": profile_b64,
                    },
                }
            ]
        }

        def fake_get_json(path, token, *, params=None):
            if path.startswith("/certificates"):
                return cert_response
            if path.startswith("/profiles"):
                return profile_response
            raise AssertionError(f"unexpected path {path}")

        with mock.patch.object(team_resolver, "get_json", side_effect=fake_get_json):
            derived = team_resolver.derive_team_id("jwt-token")
        self.assertEqual(derived, team_id)

    def test_handles_cms_signed_profile_via_security_tool(self):
        """Realistic profiles are CMS-signed; plistlib.loads raises.

        team_resolver must shell out to ``security cms -D`` to decode them.
        """
        team_id = "TEAM987654"
        # Produce a blob that plistlib.loads can't parse directly.
        fake_cms_blob = b"\x30\x82\x01\x00" + b"NOT-A-PLIST" * 20
        profile_response = {
            "data": [
                {
                    "id": "P1",
                    "attributes": {
                        "name": "CI-com.example.test",
                        "profileContent": base64.b64encode(fake_cms_blob).decode(),
                    },
                }
            ]
        }

        def fake_get_json(path, token, *, params=None):
            if path.startswith("/certificates"):
                return {"data": []}
            return profile_response

        decoded_plist = plistlib.dumps({"TeamIdentifier": [team_id]})

        def fake_check_output(args, *a, **kw):
            # args[0] == "security", args[1:3] == ("cms", "-D")
            self.assertEqual(args[0], "security")
            self.assertIn("cms", args)
            return decoded_plist

        with mock.patch.object(team_resolver, "get_json", side_effect=fake_get_json):
            with mock.patch.object(
                subprocess, "check_output", side_effect=fake_check_output
            ):
                derived = team_resolver.derive_team_id("jwt-token")
        self.assertEqual(derived, team_id)


class DeriveTeamIdFailureTests(unittest.TestCase):
    """Absolute failure path — no certs AND no profiles."""

    def test_returns_empty_when_neither_cert_nor_profile_found(self):
        def fake_get_json(path, token, *, params=None):
            return {"data": []}

        with mock.patch.object(team_resolver, "get_json", side_effect=fake_get_json):
            derived = team_resolver.derive_team_id("jwt-token")
        self.assertEqual(derived, "")

    def test_returns_empty_on_certificate_without_ou(self):
        """Malformed cert with no OU attribute — must not crash, must return ''."""
        key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
        subject = x509.Name(
            [x509.NameAttribute(NameOID.COMMON_NAME, "Malformed")]
        )
        now = datetime.datetime.now(datetime.timezone.utc)
        cert = (
            x509.CertificateBuilder()
            .subject_name(subject)
            .issuer_name(subject)
            .public_key(key.public_key())
            .serial_number(1)
            .not_valid_before(now)
            .not_valid_after(now + datetime.timedelta(days=30))
            .sign(key, hashes.SHA256())
        )
        cert_b64 = base64.b64encode(
            cert.public_bytes(serialization.Encoding.DER)
        ).decode()
        cert_response = {
            "data": [
                {
                    "id": "BAD",
                    "attributes": {
                        "certificateType": "DISTRIBUTION",
                        "certificateContent": cert_b64,
                    },
                }
            ]
        }

        def fake_get_json(path, token, *, params=None):
            if path.startswith("/certificates"):
                return cert_response
            return {"data": []}

        with mock.patch.object(team_resolver, "get_json", side_effect=fake_get_json):
            derived = team_resolver.derive_team_id("jwt-token")
        self.assertEqual(derived, "")


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