#!/usr/bin/env python3
#MISE description="Generate docs/matchers.md, the catalog of public matchers"
"""Generate docs/matchers.md from the public methods in src/.

Each entry takes its description from the first paragraph of the doc
comment, and its example from the first passing assertion in the tests
that calls it. Run with --check to fail when the catalog is out of date.
"""

import pathlib
import re
import sys

ROOT = pathlib.Path(__file__).resolve().parents[2]
SRC = ROOT / "src"
OUT = ROOT / "docs" / "matchers.md"

# Subject type of `self`, and the type parameter bound where it matters.
GROUPS = [
    ("Any value", lambda t, b: t == "T" and "Number" not in b and "Compare" not in b),
    ("Ordered values", lambda t, b: t == "T" and "Compare" in b and "Number" not in b),
    ("Numbers", lambda t, b: t == "T" and "Number" in b),
    ("Floating-point numbers", lambda t, b: t == "F"),
    ("Booleans", lambda t, b: t == "Bool"),
    ("Options", lambda t, b: t == "T?"),
    ("Results", lambda t, b: t.startswith("Result[")),
    ("Strings", lambda t, b: t == "String"),
    ("Chars", lambda t, b: t == "Char"),
    ("Collections with a length", lambda t, b: t == "C"),
    ("Arrays", lambda t, b: t.startswith("Array[")),
    ("Maps", lambda t, b: t.startswith("Map[")),
    ("Pairs", lambda t, b: t.startswith("(") and not t.startswith("() ->")),
    ("Json", lambda t, b: t == "Json"),
    ("Functions (`expect_call`)", lambda t, b: t.startswith("() ->")),
    ("Errors", lambda t, b: t == "Error"),
]

HEADER = re.compile(r"pub fn(\[[^\]]*\])? Expectation::(\w+)\(")

# Public methods that are not matchers. The README documents them.
SKIP = {"assert_that"}


def balanced(text, start):
    """The index of the bracket that closes the one at `start`."""
    depth = 0
    for i in range(start, len(text)):
        if text[i] in "([":
            depth += 1
        elif text[i] in ")]":
            depth -= 1
            if depth == 0:
                return i
    raise ValueError("unbalanced brackets")


def doc_summary(doc):
    lines = [line[3:].strip() for line in doc.strip().splitlines()]
    lines = [line for line in lines if line != "|"]
    paragraph = []
    for line in lines:
        if not line:
            break
        paragraph.append(line)
    return " ".join(paragraph)


def matchers():
    found = []
    for path in sorted(SRC.glob("*.mbt")):
        if path.name.endswith("_test.mbt") or path.name.endswith("_wbtest.mbt"):
            continue
        text = path.read_text()
        for m in HEADER.finditer(text):
            bounds, name = m.groups()
            if name in SKIP:
                continue
            # The doc comment and attributes are the lines just above.
            above = text[: m.start()].rstrip("\n").split("\n")
            doc, attributes = [], []
            while above and (above[-1].startswith("///") or above[-1].startswith("#")):
                line = above.pop()
                (doc if line.startswith("///") else attributes).insert(0, line)
            if any(a.startswith("#deprecated") for a in attributes):
                continue
            close = balanced(text, m.end() - 1)
            params = text[m.end() : close]
            subject_start = params.index("Expectation[") + len("Expectation")
            subject = params[subject_start + 1 : balanced(params, subject_start)]
            returns = text[close + 1 : text.index("{", close)]
            returns = returns.replace("->", "").replace("raise Error", "").strip()
            found.append(
                {
                    "name": name,
                    "subject": " ".join(subject.split()),
                    "bounds": bounds or "",
                    "summary": doc_summary("\n".join(doc)),
                    "returns": "" if returns == "Unit" else returns,
                }
            )
    return found


FREE = re.compile(r"pub fn(\[[^\]]*\])? (\w+)\(")


def matcher_functions():
    """Public free functions that return a `Matcher`."""
    found = []
    for path in sorted(SRC.glob("*.mbt")):
        if path.name.endswith("_test.mbt") or path.name.endswith("_wbtest.mbt"):
            continue
        text = path.read_text()
        for m in FREE.finditer(text):
            close = balanced(text, m.end() - 1)
            returns = text[close + 1 : text.index("{", close)]
            if "Matcher[" not in returns:
                continue
            above = text[: m.start()].rstrip("\n").split("\n")
            doc = []
            while above and (above[-1].startswith("///") or above[-1].startswith("#")):
                line = above.pop()
                if line.startswith("///"):
                    doc.insert(0, line)
            found.append({"name": m.group(2), "summary": doc_summary("\n".join(doc))})
    return found


def failure_regions(text):
    """Character ranges of `failure(...)` calls, whose assertions fail."""
    regions = []
    for m in re.finditer(r"failure(?:_message|_line|_lines)?\(", text):
        depth = 0
        for i in range(m.end() - 1, len(text)):
            if text[i] == "(":
                depth += 1
            elif text[i] == ")":
                depth -= 1
                if depth == 0:
                    regions.append((m.start(), i))
                    break
    return regions


LITERAL = re.compile(r"""^(?:[-0-9"'\[({]|b"|Some\(|None|true|false|Ok\(|Err\(|Set\()""")


def score(code, name, public):
    """Lower is better: a literal subject, `name` as the first method, and
    no helpers that exist only in the tests."""
    call = re.match(r"expect(?:_call)?\(", code)
    close = balanced(code, call.end() - 1)
    subject = code[call.end() : close]
    first = re.match(r"\.(\w+)\(", code[close + 1 :])
    methods = re.findall(r"\.(\w+)\(", code[close + 1 :])
    result = 0
    if not LITERAL.match(subject) or re.search(r"\bfails?\(|parse_number", subject):
        result += 4
    if not first or first.group(1) != name:
        result += 2
    if name != "not" and "not" in methods:
        result += 1
    if any(m not in public for m in methods):
        result += 8
    if re.search(r"\.\w+\([a-z_]\w*\)", code):
        result += 3
    calls = re.findall(r"(?<![.\w@])([a-z_]\w*)\(", code)
    if any(c not in public for c in calls):
        result += 8
    return result


def examples(public):
    """The best passing one-line assertion from the tests, by method name."""
    candidates = {}
    for path in sorted(SRC.glob("*_test.mbt")) + sorted(SRC.glob("*_wbtest.mbt")):
        text = path.read_text()
        regions = failure_regions(text)
        offset = 0
        for line in text.splitlines(keepends=True):
            start = offset
            offset += len(line)
            stripped = line.strip()
            if any(a <= start <= b for a, b in regions):
                continue
            match = re.match(r"(?:@expect\.)?(expect(?:_call)?\(.*)$", stripped)
            if not match or stripped.count("(") != stripped.count(")"):
                continue
            code = match.group(1)
            for name in set(re.findall(r"(?<![\w@])\.?(\w+)\(", code)):
                candidates.setdefault(name, []).append(code)
    result = {}
    for name, codes in candidates.items():
        best = min(codes, key=lambda code: (score(code, name, public), len(code)))
        # Users call the free functions of the package with `@expect.`.
        qualified = re.sub(
            r"(?<![.\w@])(\w+)\(",
            lambda m: ("@expect." if m.group(1) in public else "") + m.group(0),
            best,
        )
        result[name] = qualified
    return result


def render():
    found = matchers()
    functions = matcher_functions()
    public = (
        {m["name"] for m in found}
        | {f["name"] for f in functions}
        | {"not", "because", "assert_that", "expect", "expect_call"}
    )
    samples = examples(public)
    out = [
        "# Matchers",
        "",
        "<!-- Generated by mise-tasks/docs/catalog. Do not edit. -->",
        "",
        "Every public method on `Expectation`, grouped by the type of the value",
        "under test. Methods with a return type navigate: they return an",
        "expectation on a part of the value.",
        "",
        "The examples come from the tests, so they compile and pass. `fails` and",
        "`parse_number` are small helpers in the tests: `fails(message)` raises a",
        "`Failure`, and `parse_number(text)` raises a `ParseError` for text that",
        "is not a number.",
        "",
    ]
    missing = []
    for title, test in GROUPS:
        group = [m for m in found if test(m["subject"], m["bounds"])]
        if not group:
            continue
        out += [f"## {title}", "", "| Method | Description | Example |", "|---|---|---|"]
        for m in group:
            name = f"`{m['name']}`"
            if m["returns"]:
                name += f" → `{m['returns']}`"
            example = samples.get(m["name"])
            if example is None:
                missing.append(m["name"])
            example_text = f"`{example}`" if example else ""
            summary = m["summary"].replace("|", "\\|")
            out.append(f"| {name} | {summary} | {example_text.replace('|', chr(92) + '|')} |")
        out.append("")
    if functions:
        out += [
            "## Matcher values",
            "",
            "Functions that return a `Matcher[T]`. Run one with `to`, or pass it to",
            "another matcher such as `to_contain_element_matching`.",
            "",
            "| Function | Description | Example |",
            "|---|---|---|",
        ]
        for f in functions:
            example = samples.get(f["name"])
            if example is None:
                missing.append(f["name"])
            example_text = f"`{example}`" if example else ""
            summary = f["summary"].replace("|", "\\|")
            out.append(f"| `{f['name']}` | {summary} | {example_text.replace('|', chr(92) + '|')} |")
        out.append("")
    grouped = {m["name"] for title, test in GROUPS for m in found if test(m["subject"], m["bounds"])}
    ungrouped = [m["name"] for m in found if m["name"] not in grouped]
    return "\n".join(out), missing, ungrouped


def main():
    text, missing, ungrouped = render()
    problems = []
    if missing:
        problems.append("no example in the tests for: " + ", ".join(missing))
    if ungrouped:
        problems.append("no group for: " + ", ".join(ungrouped))
    if problems:
        for problem in problems:
            print(f"error: {problem}", file=sys.stderr)
        sys.exit(1)
    if "--check" in sys.argv:
        if not OUT.exists() or OUT.read_text() != text:
            print("error: docs/matchers.md is out of date. Run mise run docs:catalog.", file=sys.stderr)
            sys.exit(1)
        return
    OUT.write_text(text)
    print(f"wrote {OUT.relative_to(ROOT)}")


if __name__ == "__main__":
    main()
