"""Tell whether a domain's vhost already serves its lineage exactly as `certbot install --redirect` leaves it.

Every `certbot install --nginx` saves a checkpoint copy of every nginx file on the host, even when it
changes nothing. The deploy skips that install only when this prints `current`: the managed file holds
Certbot's TLS block for exactly that lineage behind Certbot's redirect, the lineage's certificate and key
load as a pair, the running nginx answers the name with that lineage's certificate, and no other server
block in the whole configuration nginx loads names it. Every other state, including any doubt, keeps the
install. Only fixed state tokens are printed, never file content.

A skipped install still tests and reloads nginx (`ensure_certificate.sh`), as the install did.
"""
import os
from pathlib import Path
import re
import socket
import ssl
import subprocess
import sys

DOMAIN = re.compile(r"[a-z0-9](?:[a-z0-9.-]*[a-z0-9])?")
TLS_LISTENS = {("443", "ssl"), ("[::]:443", "ssl", "ipv6only=on")}
HTTP_LISTENS = {("80",), ("[::]:80",)}
REDIRECT = [("return", ("301", "https://$host$request_uri"), None)]
WORD_END = " \t\r\n;{"
# Wherever nginx could start a token spelled `server_name`, bare or quoted, and in comments too.
SERVER_NAME = re.compile(rb"""(?<![^\s;{}#"'])server_name(?=[\s"'])""")
# A word nginx reads as one, or whose content a pair of quotes holds whole: no `#` or other quote, and no
# backslash that could escape the blank or `;` after it. A word holding a backslash is never a domain name.
READABLE_NAME = re.compile(rb"""(?:[^#"'\\]|\\[^#"'])*|(["'])[^#"'\\]*\1""")
WORD = re.compile(rb"[^ \t\r\n]+")  # nginx separates words by these four only
PEM_CERTIFICATE = re.compile(r"-----BEGIN CERTIFICATE-----\n.+?\n-----END CERTIFICATE-----", re.S)
ORIGIN = ("127.0.0.1", 443)
LOADED = ("nginx", "-T")
DUMPED_FILE = re.compile(rb"^# configuration file (.*):$", re.M)


class Unparsable(ValueError):
    pass


def tokens(text):
    """Yield (is_word, value); refuse escapes, variables in braces and stray quotes or braces in words."""
    index, size = 0, len(text)
    while index < size:
        char = text[index]
        if char in " \t\r\n":
            index += 1
        elif char == "#":
            end = text.find("\n", index)
            index = size if end < 0 else end
        elif char in ";{}":
            yield False, char
            index += 1
        elif char in "\"'":
            end = text.find(char, index + 1)
            if end < 0 or "\\" in text[index + 1:end] or text[end + 1:end + 2] not in ("",) + tuple(WORD_END):
                raise Unparsable
            yield True, text[index + 1:end]
            index = end + 1
        else:
            end = index
            while end < size and text[end] not in WORD_END:
                end += 1
            word = text[index:end]
            if any(mark in word for mark in ("}", "\"", "'", "\\", "${")):
                raise Unparsable
            yield True, word
            index = end


def parse(text):
    """Return top-level (name, arguments, children) entries; children is None for a simple directive."""
    stack, words = [[]], []
    for is_word, value in tokens(text):
        if is_word:
            words.append(value)
        elif value == "}":
            if words or len(stack) == 1:
                raise Unparsable
            stack.pop()
        elif not words:
            raise Unparsable
        else:
            children = [] if value == "{" else None
            stack[-1].append((words[0], tuple(words[1:]), children))
            words = []
            if children is not None:
                stack.append(children)
    if words or len(stack) != 1:
        raise Unparsable
    return stack[0]


def values(block, name):
    found = [entry for entry in block if entry[0] == name]
    if any(children is not None for _, _, children in found):
        raise Unparsable
    return [arguments for _, arguments, _ in found]


def role(block, domain):
    """Classify one server block as Certbot leaves it: `tls` on 443, `http` on 80, else `ambiguous`."""
    listens = values(block, "listen")
    if values(block, "server_name") != [(domain,)] or values(block, "ssl") or len(set(listens)) != len(listens):
        return "ambiguous"
    if listens and set(listens) <= TLS_LISTENS and ("443", "ssl") in listens:
        return "tls"
    if listens and set(listens) <= HTTP_LISTENS and ("80",) in listens:
        return "http"
    return "ambiguous"


def redirected(block, domain):
    """Certbot's own redirect test: `if ($host = <domain>) { return 301 https://$host$request_uri; }`."""
    condition = ("($host", "=", domain + ")")
    return any(name == "if" and arguments == condition and children == REDIRECT
               for name, arguments, children in block)


def loaded_configuration(command=LOADED):
    """`nginx -T`: every file nginx loads from disk, each after a `# configuration file <path>:` line, and the
    test's warnings. The dump holds secrets such as ingress tokens, so it stays in this process's memory."""
    result = subprocess.run(command, capture_output=True, timeout=30, check=True)
    return result.stdout, result.stderr


def named(text, domain):
    """How many `server_name` words name the domain, or None when that count could be wrong.

    Directives are only scanned, never parsed. nginx loads a `server_name` directive only where a token can start,
    spelled bare or quoted, and ends it at a `;`; the scan takes every such place, in comments and strings too, so
    a commented-out name counts, while a reference such as an upstream host is not a server_name. Up to the first
    `;`, nginx reads exactly the words between blanks, unquoting a word a pair of quotes holds whole, unless a
    `#`, another quote or a backslash before a blank or that `;` could end the statement later or change a name
    (`server_name a # b;` goes on to the next line). Any of those is doubt, whatever the statement names, and so is
    a quoted directive name, whose closing quote stands as a word of its own; a regular expression's own escapes
    are not. nginx and Certbot take `.<domain>` for the domain itself. Only exact names count: a wildcard or a
    regular expression that also matches the domain is not a duplicate, because Certbot (like nginx) ranks the
    exact-name block above it for both the certificate and the port-80 redirect, so the install would edit the
    managed block anyway and skipping it changes neither.
    """
    wanted, count = {domain.encode(), b"." + domain.encode()}, 0
    for match in SERVER_NAME.finditer(text):
        end = match.end()
        stop = text.find(b";", end)
        if stop < 0:
            continue  # no directive nginx loads lacks its `;`
        for word in WORD.findall(text, end, stop):
            if not READABLE_NAME.fullmatch(word):
                return None
            count += word.strip(b"\"'").lower() in wanted
    return count


def named_elsewhere(path, domain, loaded):
    """Doubt unless the loaded configuration names the domain only in this file, loaded once.

    Certbot weighs every server block nginx loads, whichever `include` reaches it, so this reads the whole
    configuration as `nginx -T` dumps it, not one directory, and a statement `named` cannot count anywhere in it,
    this file included, is doubt. The whole dump is counted against this file, so a section line written inside
    another file cannot hide a name; nginx dumps a file once however often it is included, and warns when a
    second copy conflicts.
    """
    dump, warnings = loaded
    here = named(path.read_bytes(), domain)
    return (DUMPED_FILE.findall(dump).count(os.fsencode(path)) != 1 or here is None or here != named(dump, domain)
            or b'conflicting server name "' + domain.encode() + b'"' in warnings.lower())


def configured(servers, domain, fullchain, privkey):
    roles = [role(block, domain) for block in servers]
    if not servers or "ambiguous" in roles or roles.count("tls") > 1 or roles.count("http") > 1:
        return "ambiguous"
    if "tls" not in roles:
        return "http_only"
    tls = servers[roles.index("tls")]
    certificates, keys = values(tls, "ssl_certificate"), values(tls, "ssl_certificate_key")
    if len(certificates) != 1 or len(keys) != 1:
        return "ambiguous"
    if certificates != [(fullchain,)] or keys != [(privkey,)]:
        return "other_certificate"
    if "http" not in roles or not redirected(servers[roles.index("http")], domain):
        return "no_redirect"
    return "current"


def leaf_certificate(fullchain):
    """DER of the lineage's first certificate, which nginx presents as the leaf."""
    match = PEM_CERTIFICATE.search(Path(fullchain).read_text(encoding="ascii"))
    return ssl.PEM_cert_to_DER_cert(match.group(0)) if match else b""


def paired(fullchain, privkey):
    """Load the lineage as nginx does: a key that does not match the leaf fails the next reload host-wide.

    The empty password refuses an encrypted key at once instead of prompting for one.
    """
    try:
        ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER).load_cert_chain(fullchain, privkey, password=b"")
    except ssl.SSLError:
        return False
    return True


def served_certificate(domain, origin=ORIGIN):
    """DER of the certificate the running nginx presents for this name; the file alone may not be loaded."""
    context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
    context.check_hostname = False
    context.verify_mode = ssl.CERT_NONE
    with socket.create_connection(origin, timeout=5) as connection:
        with context.wrap_socket(connection, server_hostname=domain) as tls:
            return tls.getpeercert(binary_form=True)


def state(vhost, domain, fullchain, privkey, probe=served_certificate, dump=loaded_configuration):
    path = Path(vhost)
    if not DOMAIN.fullmatch(domain) or not all(value.startswith("/") for value in (vhost, fullchain, privkey)):
        return "ambiguous"
    if path.is_symlink() or (path.exists() and not path.is_file()):
        return "ambiguous"
    if not path.exists():
        return "absent"
    entries = parse(path.read_text(encoding="utf-8"))
    if any(name not in ("map", "server") or children is None for name, _, children in entries):
        return "ambiguous"
    servers = [children for name, _, children in entries if name == "server"]
    result = configured(servers, domain, fullchain, privkey)
    if result != "current":
        return result
    leaf = leaf_certificate(fullchain) if os.path.isfile(fullchain) and os.path.isfile(privkey) else b""
    if not leaf or not paired(fullchain, privkey):
        return "lineage_unreadable"
    try:
        live = probe(domain)
    except OSError:
        return "unreachable"
    if live != leaf:
        return "not_served"
    # Last: the handshake is cheap, while the dump loads and tests the host's whole configuration.
    try:
        loaded = dump()
    except (OSError, subprocess.SubprocessError):
        return "config_unknown"
    return "duplicate" if named_elsewhere(path, domain, loaded) else "current"


def main(arguments, probe=served_certificate, dump=loaded_configuration):
    if len(arguments) != 4:
        return "ambiguous"
    try:
        return state(*arguments, probe=probe, dump=dump)
    except ValueError:
        return "unparsable"
    except OSError:
        return "unreadable"


if __name__ == "__main__":
    print(main(sys.argv[1:]))
