"""Offline RFC 3161 TimeStampToken overenie (vendored do verifikatora).

Toto je samostatna kopia OVEROVACEJ casti api/app/tsa.py — verifikator nepotrebuje
request-building (nikdy si nepyta novy token), len validaciu existujuceho tokenu
na hlave retazca. Vendorujeme ju (namiesto `from app.tsa import verify_timestamp`),
aby forenzny verifikator NEZAVISEL od beziacej appky — to je cely pointa
nezavisleho dokazu (DESIGN-proof-core.md §0, §1: spor musi vediet overit aj ten,
kto appku nema).

Overuje (zhodne s app.tsa.verify_timestamp):
  1) token existuje,
  2) hashAlgorithm v TSTInfo je SHA256,
  3) hashedMessage == SHA256(RAW record_hash) — token kryje PRESNE tuto hlavu,
  4) CMS podpis tokenu plati (embedded TSA cert; cryptography),
  5) signedAttrs vazba (RFC 5652 §5.4): message-digest == hash(eContent) a
     content-type == eContentType — bez toho by sa dal TSTInfo vymenit
     (content-splice) pri platnom podpise.

Pozn.: messageImprint hashuje RAW 32 bajtov record_hashu (bytes.fromhex), zhodne
s tym, co KMS podpisuje (MessageType=DIGEST) a co app.tsa posiela TSA.
"""
from __future__ import annotations

from dataclasses import dataclass
from datetime import UTC, datetime, timedelta

from cryptography import x509
from cryptography.exceptions import InvalidSignature
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.asymmetric import ec, padding
from cryptography.hazmat.primitives.asymmetric import rsa as _rsa
from cryptography.hazmat.primitives.serialization import pkcs7

import _asn1 as asn1  # vendored TLV kodek (verifier/ je na sys.path; standalone)

# OID-y (RFC 3161 / PKCS#7).
_OID_SHA256 = "2.16.840.1.101.3.4.2.1"
_OID_CONTENT_TYPE = "1.2.840.113549.1.9.3"     # signedAttr: content-type
_OID_MESSAGE_DIGEST = "1.2.840.113549.1.9.4"   # signedAttr: message-digest (RFC 5652 §5.4)

# DER tag pre [0] IMPLICIT signedAttrs v SignerInfo (context-specific, constructed).
_TAG_SIGNED_ATTRS = 0xA0
# DER tag pre GeneralizedTime (TSTInfo.genTime).
_GENERALIZED_TIME = 0x18


class TsaError(RuntimeError):
    """TSA token ma neocakavanu strukturu (parse zlyhal)."""


@dataclass(frozen=True)
class TimestampResult:
    """Vysledok overenia TSA tokenu."""

    ok: bool
    reason: str | None = None


def _digest_for_imprint(record_hash: str) -> bytes:
    """SHA256 nad RAW bajtami record_hashu (messageImprint hashedMessage)."""
    raw = bytes.fromhex(record_hash)
    digest = hashes.Hash(hashes.SHA256())
    digest.update(raw)
    return digest.finalize()


def _tstinfo_message_imprint(token_der: bytes) -> tuple[str, bytes]:
    """Z TimeStampToken vytiahni (hashAlgorithm OID, hashedMessage) z TSTInfo.

    TimeStampToken = ContentInfo { contentType id-signedData, content SignedData }
    SignedData.encapContentInfo.eContent = OCTET STRING obsahujuci DER TSTInfo.
    TSTInfo ::= SEQUENCE { version, policy, messageImprint MessageImprint, ... }
    """
    (content_info,) = asn1.parse_all(token_der)
    ci_children = content_info.children()
    # ci_children[0] = OID signedData; ci_children[1] = [0] EXPLICIT SignedData
    signed_data = ci_children[1].children()[0]
    sd_children = signed_data.children()
    # SignedData: version, digestAlgorithms SET, encapContentInfo SEQUENCE, ...
    encap = sd_children[2]
    encap_children = encap.children()
    # encapContentInfo: eContentType OID, [0] EXPLICIT eContent OCTET STRING
    econtent_explicit = encap_children[1]
    econtent_octets = econtent_explicit.children()[0]
    tstinfo, _ = asn1.parse_one(econtent_octets.content)
    tst_children = tstinfo.children()
    # TSTInfo: version, policy OID, messageImprint SEQUENCE, serial, genTime, ...
    message_imprint = tst_children[2]
    mi_children = message_imprint.children()
    algo_oid = asn1.decode_oid(mi_children[0].children()[0].content)
    hashed_message = mi_children[1].content
    return algo_oid, hashed_message


def _tstinfo_gentime(token_der: bytes) -> datetime:
    """Vytiahni genTime z TSTInfo (GeneralizedTime) ako tz-aware UTC datetime."""
    (content_info,) = asn1.parse_all(token_der)
    signed_data = content_info.children()[1].children()[0]
    encap = signed_data.children()[2]
    econtent_octets = encap.children()[1].children()[0]
    tstinfo, _ = asn1.parse_one(econtent_octets.content)
    tst_children = tstinfo.children()
    gentime_node = tst_children[4]
    if gentime_node.tag != _GENERALIZED_TIME:
        msg = f"genTime is not a GeneralizedTime (tag={gentime_node.tag:#x})"
        raise TsaError(msg)
    return _parse_generalized_time(gentime_node.content)


def _parse_generalized_time(raw: bytes) -> datetime:
    """Parsuj DER GeneralizedTime (YYYYMMDDHHMMSS[.fff]Z) na tz-aware UTC datetime."""
    text_value = raw.decode("ascii")
    if not text_value.endswith("Z"):
        msg = f"genTime is not in UTC (missing 'Z'): {text_value!r}"
        raise TsaError(msg)
    core = text_value[:-1]
    fmt = "%Y%m%d%H%M%S.%f" if "." in core else "%Y%m%d%H%M%S"
    return datetime.strptime(core, fmt).replace(tzinfo=UTC)


# Podporovane digest algoritmy TSA podpisu (zrkadlo api/app/tsa.py — DRZAT V SYNCU).
# SHA-1 zamerne chyba; nikdy nedefaultovat na SHA-256 (freetsa podpisuje SHA-512).
_DIGEST_OID_TO_HASH: dict[str, type[hashes.HashAlgorithm]] = {
    "2.16.840.1.101.3.4.2.1": hashes.SHA256,
    "2.16.840.1.101.3.4.2.2": hashes.SHA384,
    "2.16.840.1.101.3.4.2.3": hashes.SHA512,
}
_SIGALG_OID_TO_HASH: dict[str, type[hashes.HashAlgorithm]] = {
    "1.2.840.10045.4.3.2": hashes.SHA256,  # ecdsa-with-SHA256
    "1.2.840.10045.4.3.3": hashes.SHA384,  # ecdsa-with-SHA384
    "1.2.840.10045.4.3.4": hashes.SHA512,  # ecdsa-with-SHA512
    "1.2.840.113549.1.1.11": hashes.SHA256,  # sha256WithRSAEncryption
    "1.2.840.113549.1.1.12": hashes.SHA384,  # sha384WithRSAEncryption
    "1.2.840.113549.1.1.13": hashes.SHA512,  # sha512WithRSAEncryption
}


def _signed_attrs_and_signature(token_der: bytes) -> tuple[bytes, bytes, hashes.HashAlgorithm]:
    """Vrati (signature, podpisane_bajty, digest_algo) z prveho SignerInfo.

    Ak su pritomne signedAttrs ([0] IMPLICIT), podpisuju sa ony (re-enkodovane
    ako SET, RFC 5652 §5.4); inak sa podpisuje eContent. TSA tokeny signedAttrs
    vzdy maju, takze overujeme nad nimi.

    Hash podpisu sa cita zo SignerInfo (signatureAlgorithm, fallback
    digestAlgorithm) — NIE natvrdo SHA-256 (freetsa podpisuje ecdsa-with-SHA512).
    """
    (content_info,) = asn1.parse_all(token_der)
    signed_data = content_info.children()[1].children()[0]
    sd_children = signed_data.children()
    # posledny prvok SignedData je SignerInfos SET
    signer_infos = sd_children[-1]
    signer_info = signer_infos.children()[0]
    si_children = signer_info.children()
    # SignerInfo: version, sid, digestAlgorithm, [0] signedAttrs OPTIONAL,
    #             signatureAlgorithm, signature OCTET STRING, ...
    signed_attrs = None
    signed_attrs_idx = None
    sig_octet = None
    for idx, node in enumerate(si_children):
        if node.tag == _TAG_SIGNED_ATTRS:  # [0] IMPLICIT signedAttrs
            signed_attrs = node
            signed_attrs_idx = idx
        if node.tag == asn1._OCTET_STRING and sig_octet is None:  # noqa: SLF001
            sig_octet = node
    if signed_attrs is None or sig_octet is None or signed_attrs_idx is None:
        msg = "SignerInfo without signedAttrs/signature"
        raise TsaError(msg)

    # Hash: primarne zo signatureAlgorithm (SEQUENCE hned za signedAttrs),
    # fallback digestAlgorithm (index 2).
    hash_cls: type[hashes.HashAlgorithm] | None = None
    if signed_attrs_idx + 1 < len(si_children):
        sig_alg_children = si_children[signed_attrs_idx + 1].children()
        if sig_alg_children:
            hash_cls = _SIGALG_OID_TO_HASH.get(asn1.decode_oid(sig_alg_children[0].content))
    if hash_cls is None:
        digest_oid = asn1.decode_oid(si_children[2].children()[0].content)
        hash_cls = _DIGEST_OID_TO_HASH.get(digest_oid)
    if hash_cls is None:
        msg = "unsupported digest/signature algorithm in SignerInfo"
        raise TsaError(msg)

    # signedAttrs sa podpisuju ako SET (tag 0x31), nie ako [0] IMPLICIT (0xA0).
    signed_bytes = bytes([0x31]) + asn1._der_len(len(signed_attrs.content)) + signed_attrs.content  # noqa: SLF001
    return sig_octet.content, signed_bytes, hash_cls()


def _sha256(data: bytes) -> bytes:
    """SHA256 nad danymi bajtami."""
    digest = hashes.Hash(hashes.SHA256())
    digest.update(data)
    return digest.finalize()


def _encap_content(token_der: bytes) -> tuple[str, bytes]:
    """Vrati (eContentType OID, eContent bajty = DER TSTInfo) z encapContentInfo."""
    (content_info,) = asn1.parse_all(token_der)
    signed_data = content_info.children()[1].children()[0]
    encap = signed_data.children()[2]
    encap_children = encap.children()
    econtent_type = asn1.decode_oid(encap_children[0].content)
    econtent_octets = encap_children[1].children()[0]
    return econtent_type, econtent_octets.content


def _signer_digest_algo_oid(token_der: bytes) -> str:
    """OID digestAlgorithm z prveho SignerInfo (algoritmus messageDigestu)."""
    (content_info,) = asn1.parse_all(token_der)
    signed_data = content_info.children()[1].children()[0]
    signer_info = signed_data.children()[-1].children()[0]
    digest_algid = signer_info.children()[2]
    return asn1.decode_oid(digest_algid.children()[0].content)


def _signed_attrs_node(token_der: bytes) -> asn1.DerNode:
    """Vrati [0] IMPLICIT signedAttrs uzol z prveho SignerInfo."""
    (content_info,) = asn1.parse_all(token_der)
    signed_data = content_info.children()[1].children()[0]
    signer_info = signed_data.children()[-1].children()[0]
    for node in signer_info.children():
        if node.tag == _TAG_SIGNED_ATTRS:
            return node
    msg = "SignerInfo without signedAttrs"
    raise TsaError(msg)


def _signed_attr_values(signed_attrs: asn1.DerNode, oid: str) -> list[asn1.DerNode] | None:
    """Hodnoty (SET OF) signedAttr podla OID; None ak atribut chyba.

    Attribute ::= SEQUENCE { attrType OID, attrValues SET OF AttributeValue }
    """
    for attr in signed_attrs.children():
        attr_children = attr.children()
        if asn1.decode_oid(attr_children[0].content) == oid:
            return attr_children[1].children()
    return None


def _check_signed_attrs_binding(token_der: bytes) -> str | None:  # noqa: PLR0911 — guard clause na kazdu vazbu
    """Over vazbu signedAttrs -> eContent (RFC 5652 §5.4). None = OK, inak reason.

    CMS podpis kryje LEN signedAttrs (SET), NIE priamo eContent. Bez tejto vazby by
    sa dal eContent (TSTInfo, z ktoreho `_tstinfo_message_imprint` cita messageImprint)
    VYMENIT pri zachovani povodnych signedAttrs+podpisu+certov (content-splice) — verify
    by prijal jeden pravy token pre LUBOVOLNY record_hash. Preto viazeme:
      - message-digest signedAttr (1.2.840.113549.1.9.4) MUSI == digest(eContent),
      - content-type signedAttr (1.2.840.113549.1.9.3) MUSI == eContentType.
    """
    try:
        signed_attrs = _signed_attrs_node(token_der)
        econtent_type, econtent = _encap_content(token_der)
        digest_algo_oid = _signer_digest_algo_oid(token_der)
    except Exception:  # noqa: BLE001 — chybna struktura => neoverene
        return "tsa_signed_attrs_parse_error"

    # sha256/384/512 OK (freetsa pouziva sha512); nezavisle od messageImprint
    # algoritmu record_hash (ten sa vynucuje na SHA-256 zvlast).
    digest_cls = _DIGEST_OID_TO_HASH.get(digest_algo_oid)
    if digest_cls is None:
        return f"tsa_unsupported_digest_algo:{digest_algo_oid}"

    ct_values = _signed_attr_values(signed_attrs, _OID_CONTENT_TYPE)
    if not ct_values:
        return "tsa_missing_content_type_attr"
    if asn1.decode_oid(ct_values[0].content) != econtent_type:
        return "tsa_content_type_mismatch"

    md_values = _signed_attr_values(signed_attrs, _OID_MESSAGE_DIGEST)
    if not md_values:
        return "tsa_missing_message_digest_attr"
    econtent_digest = hashes.Hash(digest_cls())
    econtent_digest.update(econtent)
    if md_values[0].content != econtent_digest.finalize():
        return "tsa_message_digest_mismatch"
    return None


def _load_trust_anchors(trust_pem: str) -> list[x509.Certificate]:
    """Naparsuj PEM bundle pinnutych TSA certov (CA alebo leaf) na trust anchory."""
    data = trust_pem.encode("utf-8")
    try:
        return list(x509.load_pem_x509_certificates(data))
    except Exception:  # noqa: BLE001 — neplatny PEM => ziadne anchory
        return []


def _cert_matches_anchor(cert: x509.Certificate, anchor: x509.Certificate) -> bool:
    """Zisti, ci `cert` JE pinnuty anchor (zhoda fingerprintu) alebo je nim PODPISANY."""
    if cert.fingerprint(hashes.SHA256()) == anchor.fingerprint(hashes.SHA256()):
        return True
    if cert.issuer != anchor.subject:
        return False
    pub = anchor.public_key()
    try:
        if isinstance(pub, ec.EllipticCurvePublicKey):
            pub.verify(
                cert.signature,
                cert.tbs_certificate_bytes,
                ec.ECDSA(cert.signature_hash_algorithm),  # type: ignore[arg-type]
            )
        elif isinstance(pub, _rsa.RSAPublicKey):
            pub.verify(
                cert.signature,
                cert.tbs_certificate_bytes,
                padding.PKCS1v15(),
                cert.signature_hash_algorithm,  # type: ignore[arg-type]
            )
        else:
            return False
    except Exception:  # noqa: BLE001
        return False
    return True


def _select_trusted_signer(
    token_der: bytes, trust_pem: str, *, trust_required: bool
) -> x509.Certificate | None:
    """Vyber signer cert a over jeho DOVERU proti pinnutym anchorom (FIX #2).

    NEberieme slepo certs[0]. Ak `trust_required`, signer cert MUSI matchovat /
    chainovat na anchor z `trust_pem`; inak None (token odmietnuty).
    """
    try:
        certs = pkcs7.load_der_pkcs7_certificates(token_der)
    except Exception:  # noqa: BLE001
        return None
    if not certs:
        return None
    signer_cert = certs[0]

    if not trust_required:
        return signer_cert

    anchors = _load_trust_anchors(trust_pem)
    if not anchors:
        return None
    for anchor in anchors:
        if _cert_matches_anchor(signer_cert, anchor):
            return signer_cert
    return None


def _verify_cms_signature(token_der: bytes, signer_cert: x509.Certificate) -> bool:
    """Over CMS podpis TimeStampTokenu cez DOVERYHODNY signer cert (cryptography)."""
    try:
        sig, signed_bytes, digest_algo = _signed_attrs_and_signature(token_der)
    except Exception:  # noqa: BLE001
        return False

    public_key = signer_cert.public_key()
    try:
        if isinstance(public_key, ec.EllipticCurvePublicKey):
            public_key.verify(sig, signed_bytes, ec.ECDSA(digest_algo))
        elif isinstance(public_key, _rsa.RSAPublicKey):
            public_key.verify(sig, signed_bytes, padding.PKCS1v15(), digest_algo)
        else:
            return False
    except InvalidSignature:
        return False
    except Exception:  # noqa: BLE001
        return False
    return True


def verify_timestamp(  # noqa: PLR0911 — guard clauses na kazdy dovod zlyhania
    record_hash: str,
    token: bytes | None,
    *,
    trust_pem: str = "",
    clock_skew_seconds: int = 300,
    trust_required: bool | None = None,
) -> TimestampResult:
    """Over RFC 3161 token na hlave retazca (zhodne s app.tsa.verify_timestamp).

    `trust_pem` = PEM pinnutych TSA anchorov; ak `trust_required` (default =
    bool(trust_pem)), signer cert MUSI chainovat na anchor (FIX #2). genTime musi
    byt pritomny a v platnosti certu (FIX #2/#7). Nevyhadzuje na tamper.
    """
    enforce_trust = bool(trust_pem) if trust_required is None else trust_required

    if not token:
        return TimestampResult(ok=False, reason="missing_tsa_token")
    try:
        algo_oid, hashed_message = _tstinfo_message_imprint(token)
    except Exception as exc:  # noqa: BLE001
        return TimestampResult(ok=False, reason=f"tsa_token_parse_error:{exc}")
    if algo_oid != _OID_SHA256:
        return TimestampResult(ok=False, reason=f"tsa_hash_algo_mismatch:{algo_oid}")
    expected = _digest_for_imprint(record_hash)
    if hashed_message != expected:
        return TimestampResult(ok=False, reason="tsa_message_imprint_mismatch")

    signer_cert = _select_trusted_signer(token, trust_pem, trust_required=enforce_trust)
    if signer_cert is None:
        return TimestampResult(ok=False, reason="tsa_untrusted_cert")

    try:
        gentime = _tstinfo_gentime(token)
    except Exception as exc:  # noqa: BLE001
        return TimestampResult(ok=False, reason=f"tsa_gentime_parse_error:{exc}")
    skew = timedelta(seconds=clock_skew_seconds)
    if not (
        signer_cert.not_valid_before_utc - skew
        <= gentime
        <= signer_cert.not_valid_after_utc + skew
    ):
        return TimestampResult(ok=False, reason="tsa_gentime_outside_cert_validity")

    if not _verify_cms_signature(token, signer_cert):
        return TimestampResult(ok=False, reason="tsa_signature_invalid")

    # signedAttrs -> eContent vazba (RFC 5652 §5.4). Podpis kryje len signedAttrs;
    # bez tejto vazby by sa dal eContent (TSTInfo) vymenit -> content-splice.
    binding_reason = _check_signed_attrs_binding(token)
    if binding_reason is not None:
        return TimestampResult(ok=False, reason=binding_reason)
    return TimestampResult(ok=True)
