from __future__ import annotations
import datetime
import socket
import ssl
from typing import Iterable, List, Tuple

import requests

from ..localization.helpers import create_translated_finding
from ..model import Finding, ScanRequest

LEGACY_TLS_VERSIONS: Tuple[Tuple[str, ssl.TLSVersion], ...] = (
    ("TLS 1.0", ssl.TLSVersion.TLSv1),
    ("TLS 1.1", ssl.TLSVersion.TLSv1_1),
)


def evaluate_tls(request: ScanRequest) -> Iterable[Finding]:
    domain = request.domain
    context = ssl.create_default_context()
    findings: List[Finding] = []

    try:
        with socket.create_connection((domain, request.https_port), timeout=request.timeout) as sock:
            with context.wrap_socket(sock, server_hostname=domain) as tls_socket:
                cert = tls_socket.getpeercert()
    except ssl.SSLError as exc:
        findings.append(
            create_translated_finding(
                check="tls",
                i18n_key="tls.handshake_failed",
                severity="high",
                i18n_params={
                    "domain": domain,
                    "port": request.https_port,
                    "error": str(exc),
                },
            )
        )
        return findings
    except OSError as exc:
        findings.append(
            create_translated_finding(
                check="tls",
                i18n_key="tls.https_port_unreachable",
                severity="medium",
                i18n_params={
                    "domain": domain,
                    "port": request.https_port,
                    "error": str(exc),
                },
            )
        )
        return findings

    if not cert:
        findings.append(
            create_translated_finding(
                check="tls",
                i18n_key="tls.certificate_missing",
                severity="high",
            )
        )
        return findings

    _check_cert_dates(cert, findings)
    _check_cert_hostname(cert, domain, findings)
    findings.extend(_check_legacy_tls(domain, request))
    findings.extend(_check_http_downgrade(domain, request))

    return findings


def _check_cert_dates(cert: dict, findings: List[Finding]) -> None:
    not_after = cert.get("notAfter")
    not_before = cert.get("notBefore")

    def _parse(value: str) -> datetime.datetime:
        parsed = datetime.datetime.strptime(value, "%b %d %H:%M:%S %Y %Z")
        if parsed.tzinfo is None:
            parsed = parsed.replace(tzinfo=datetime.UTC)
        return parsed

    now = datetime.datetime.now(datetime.UTC)

    if not_before:
        nb = _parse(not_before)
        if nb > now:
            findings.append(
                create_translated_finding(
                    check="tls",
                    i18n_key="tls.certificate_not_yet_valid",
                    severity="medium",
                    i18n_params={"date": f"{nb.isoformat()}Z"},
                )
            )

    if not_after:
        na = _parse(not_after)
        delta = na - now
        if delta.total_seconds() < 0:
            findings.append(
                create_translated_finding(
                    check="tls",
                    i18n_key="tls.certificate_expired",
                    severity="high",
                    i18n_params={"date": f"{na.isoformat()}Z"},
                )
            )
        elif delta.days < 14:
            findings.append(
                create_translated_finding(
                    check="tls",
                    i18n_key="tls.certificate_near_expiry",
                    severity="medium",
                    i18n_params={
                        "date": f"{na.isoformat()}Z",
                        "days": delta.days,
                    },
                )
            )


def _check_cert_hostname(cert: dict, domain: str, findings: List[Finding]) -> None:
    alt_names = []
    for entry in cert.get("subjectAltName", []):
        if entry[0] == "DNS":
            alt_names.append(entry[1])

    if not alt_names:
        findings.append(
            create_translated_finding(
                check="tls",
                i18n_key="tls.no_san",
                severity="medium",
            )
        )
        return

    if not any(_matches_hostname(domain, alt) for alt in alt_names):
        findings.append(
            create_translated_finding(
                check="tls",
                i18n_key="tls.hostname_not_covered",
                severity="high",
                i18n_params={
                    "domain": domain,
                    "san_list": ", ".join(alt_names[:5]),
                },
            )
        )


def _matches_hostname(domain: str, pattern: str) -> bool:
    if pattern.startswith("*."):
        suffix = pattern[1:]
        return domain.endswith(suffix) and domain.count(".") >= pattern.count(".")
    return domain == pattern


def _check_legacy_tls(domain: str, request: ScanRequest) -> List[Finding]:
    findings: List[Finding] = []
    for name, version in LEGACY_TLS_VERSIONS:
        try:
            if _supports_version(domain, request.https_port, version, request.timeout):
                findings.append(
                    create_translated_finding(
                        check="tls",
                        i18n_key="tls.legacy_protocol",
                        severity="medium",
                        i18n_params={"protocol": name},
                    )
                )
        except Exception:
            # If the handshake fails the version is effectively disabled.
            continue
    return findings


def _supports_version(domain: str, port: int, version: ssl.TLSVersion, timeout: float) -> bool:
    ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
    ctx.minimum_version = version
    ctx.maximum_version = version
    ctx.check_hostname = False
    ctx.verify_mode = ssl.CERT_NONE

    with socket.create_connection((domain, port), timeout=timeout) as sock:
        with ctx.wrap_socket(sock, server_hostname=domain) as tls_socket:
            tls_socket.do_handshake()
            return True


def _check_http_downgrade(domain: str, request: ScanRequest) -> List[Finding]:
    findings: List[Finding] = []
    url = f"http://{domain}:{request.http_port}"
    try:
        response = requests.get(
            url,
            headers={"User-Agent": request.user_agent},
            timeout=request.timeout,
            allow_redirects=True,
        )
    except requests.RequestException as exc:
        findings.append(
            create_translated_finding(
                check="tls",
                i18n_key="tls.http_redirect_check_failed",
                severity="low",
                i18n_params={"error": str(exc)},
            )
        )
        return findings

    if response.url.startswith("http://"):
        findings.append(
            create_translated_finding(
                check="tls",
                i18n_key="tls.https_redirect_missing",
                severity="high",
                i18n_params={"url": response.url},
            )
        )
    elif response.status_code not in (301, 302, 308) and response.history:
        first = response.history[0]
        if first.status_code not in (301, 302, 308):
            findings.append(
                create_translated_finding(
                    check="tls",
                    i18n_key="tls.http_redirect_not_permanent",
                    severity="low",
                    i18n_params={"status": first.status_code},
                )
            )
    return findings
