#!/usr/bin/env -S uv run --script
#MISE description="Generate src/conformance_wbtest.mbt from the reference Cucumber Expressions testdata"
# /// script
# requires-python = ">=3.9"
# dependencies = ["pyyaml>=6"]
# ///
"""Generate the conformance suite from the reference testdata.

The reference implementation (https://github.com/cucumber/cucumber-expressions)
keeps language-neutral YAML test cases in `testdata/`. This script turns them
into MoonBit whitebox tests in `src/conformance_wbtest.mbt`.

To update the reference version, change REF_SHA, run the task, and review the
new failures. Add each known gap to KNOWN_GAPS with its issue number.
"""

from __future__ import annotations

import argparse
import io
import json
import os
import subprocess
import sys
import tarfile
import tempfile
import urllib.request
from pathlib import Path

import yaml

REF_REPO = "cucumber/cucumber-expressions"
REF_SHA = "14929870f701318a2a881873f9c7cff4bd0af201"

ROOT = Path(__file__).resolve().parents[2]
OUTPUT = ROOT / "src" / "conformance_wbtest.mbt"

# Reference cases that fail today. Each one is generated with #skip and the
# issue that tracks the gap. Remove an entry when its issue is fixed.
KNOWN_GAPS: dict[str, int] = {
    "matching/matches-bigdecimal": 26,
}



def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--testdata", type=Path)
    parser.add_argument("--ref", default=REF_SHA)
    args = parser.parse_args()

    if args.testdata:
        write_suite(args.testdata, source=str(args.testdata))
    else:
        with tempfile.TemporaryDirectory() as tmp:
            testdata = download_testdata(args.ref, Path(tmp))
            write_suite(testdata, source=f"{REF_REPO}@{args.ref}")

    subprocess.run(["moon", "fmt"], cwd=ROOT, check=True)
    print(f"Wrote {OUTPUT.relative_to(ROOT)}")


def download_testdata(ref: str, dest: Path) -> Path:
    url = f"https://codeload.github.com/{REF_REPO}/tar.gz/{ref}"
    print(f"Downloading {url}", file=sys.stderr)
    with urllib.request.urlopen(url) as response:
        data = response.read()
    with tarfile.open(fileobj=io.BytesIO(data), mode="r:gz") as tar:
        members = [m for m in tar.getmembers() if "/testdata/" in m.name]
        tar.extractall(dest, members=members)
    (top,) = [p for p in dest.iterdir() if p.is_dir()]
    return top / "testdata"


# ---------------------------------------------------------------------------
# MoonBit rendering
# ---------------------------------------------------------------------------


def lit(value: str) -> str:
    """Render a MoonBit string literal."""
    return json.dumps(value, ensure_ascii=False)


def skip_attr(case_id: str) -> str:
    issue = KNOWN_GAPS.get(case_id)
    return f'#skip("Known gap: #{issue}")\n' if issue else ""


def test_block(case_id: str, body: list[str]) -> str:
    lines = "\n".join(f"  {line}" for line in body)
    return f'///|\n{skip_attr(case_id)}test "conformance/{case_id}" {{\n{lines}\n}}\n'


def expect_error(call: str, expected: str) -> list[str]:
    return [
        f"let message = try {call} catch {{",
        "  err => err.message()",
        "} noraise {",
        '  _ => fail("expected an error")',
        "}",
        f"assert_eq(message, {lit(expected)})",
    ]


def matching_case(case_id: str, case: dict) -> str:
    expression = lit(case["expression"])
    if "exception" in case:
        return test_block(
            case_id, expect_error(f"Expression::parse({expression})", case["exception"])
        )

    body = [
        f"let expr = Expression::parse({expression})",
        f"let result = expr.match_({lit(case['text'])})",
    ]
    args = case["expected_args"]
    if args is None:
        body.append('guard result is None else { fail("expected no match") }')
        return test_block(case_id, body)

    body += [
        'guard result is Some(m) else { fail("expected a match") }',
        f"assert_eq(m.params.length(), {len(args)})",
    ]
    for i, arg in enumerate(args):
        value = f"m.params[{i}].value"
        if isinstance(arg, bool) or arg is None:
            raise ValueError(f"{case_id}: unsupported expected arg {arg!r}")
        if isinstance(arg, int):
            body.append(f"assert_eq(conformance_integer({value}), Some({lit(str(arg))}))")
        elif isinstance(arg, float):
            body.append(f"assert_true(conformance_close({value}, {arg!r}))")
        else:
            body.append(f"assert_eq(conformance_text({value}), Some({lit(arg)}))")
    return test_block(case_id, body)


def regular_expression_case(case_id: str, case: dict) -> str:
    body = [
        f"let expr = RegularExpression::new({lit(case['expression'])})",
        f"let result = expr.match_({lit(case['text'])})",
    ]
    args = case["expected_args"]
    if args is None:
        body.append('guard result is None else { fail("expected no match") }')
        return test_block(case_id, body)
    body += [
        'guard result is Some(m) else { fail("expected a match") }',
        f"assert_eq(m.params.length(), {len(args)})",
    ]
    for i, arg in enumerate(args):
        value = f"m.params[{i}].value"
        if arg is None:
            body.append(f'guard {value} is NullVal else {{ fail("arg {i}: expected NullVal") }}')
        else:
            body.append(f"assert_eq(conformance_text({value}), Some({lit(str(arg))}))")
    return test_block(case_id, body)


def transformation_case(case_id: str, case: dict) -> str:
    expression = lit(case["expression"])
    return test_block(
        case_id,
        [f"assert_eq(compile_expression({expression}), {lit(case['expected_regex'])})"],
    )


def parser_case(case_id: str, case: dict) -> str:
    expression = lit(case["expression"])
    if "exception" in case:
        return test_block(
            case_id, expect_error(f"parse_expression({expression})", case["exception"])
        )
    return test_block(
        case_id,
        [
            f"let ast = parse_expression({expression})",
            f"let expected = {render_node(case['expected_ast'])}",
            'if ast != expected { fail("got " + @debug.to_string(ast)) }',
        ],
    )


def tokenizer_case(case_id: str, case: dict) -> str:
    expression = lit(case["expression"])
    if "exception" in case:
        return test_block(
            case_id, expect_error(f"tokenize({expression})", case["exception"])
        )
    tokens = ", ".join(render_token(token) for token in case["expected_tokens"])
    return test_block(
        case_id,
        [
            f"let tokens = tokenize({expression})",
            f"let expected : Array[Token] = [{tokens}]",
            'if tokens != expected { fail("got " + @debug.to_string(tokens)) }',
        ],
    )


NODE_TYPES = {
    "EXPRESSION_NODE": "ExpressionNode",
    "OPTIONAL_NODE": "OptionalNode",
    "ALTERNATION_NODE": "AlternationNode",
    "ALTERNATIVE_NODE": "AlternativeNode",
    "PARAMETER_NODE": "ParameterNode",
}


def render_node(node: dict) -> str:
    start, end = node["start"], node["end"]
    if node["type"] == "TEXT_NODE":
        return f"conformance_text_node({lit(node['token'])}, {start}, {end})"
    nodes = ", ".join(render_node(child) for child in node.get("nodes", []))
    return f"conformance_node({NODE_TYPES[node['type']]}, {start}, {end}, [{nodes}])"


TOKEN_TYPES = {
    "START_OF_LINE": "StartOfLine",
    "END_OF_LINE": "EndOfLine",
    "WHITE_SPACE": "WhiteSpace",
    "BEGIN_OPTIONAL": "BeginOptional",
    "END_OPTIONAL": "EndOptional",
    "BEGIN_PARAMETER": "BeginParameter",
    "END_PARAMETER": "EndParameter",
    "ALTERNATION": "Alternation",
    "TEXT": "Text",
}


def render_token(token: dict) -> str:
    return (
        f"{{ type_: {TOKEN_TYPES[token['type']]}, text: {lit(token['text'])}, "
        f"start: {token['start']}, end: {token['end']} }}"
    )


# Each suite directory, the prefix of its test names, and its renderer.
SUITES = {
    "cucumber-expression/matching": ("matching", matching_case),
    "cucumber-expression/transformation": ("transformation", transformation_case),
    "cucumber-expression/parser": ("parser", parser_case),
    "cucumber-expression/tokenizer": ("tokenizer", tokenizer_case),
    "regular-expression/matching": ("regular-expression", regular_expression_case),
}

HELPERS = """\
///|
/// Integer value of an argument, as a decimal string.
fn conformance_integer(value : ParamValue) -> String? {
  match value {
    IntVal(n) | ShortVal(n) => Some(n.to_string())
    LongVal(n) => Some(n.to_string())
    ByteVal(b) => Some(b.to_int().to_string())
    BigIntegerVal(n) => Some(n.to_string())
    _ => None
  }
}

///|
/// True when a floating point argument is close to the expected value.
fn conformance_close(value : ParamValue, expected : Double) -> Bool {
  match value {
    FloatVal(d) | DoubleVal(d) => (d - expected).abs() <= 1.0e-9
    _ => false
  }
}

///|
/// Text of an argument that the reference gives as a string.
fn conformance_text(value : ParamValue) -> String? {
  match value {
    StringVal(s) | WordVal(s) | AnonymousVal(s) => Some(s)
    BigDecimalVal(d) => Some(d.to_string())
    BigIntegerVal(n) => Some(n.to_string())
    _ => None
  }
}

///|
fn conformance_node(
  type_ : NodeType,
  start : Int,
  end : Int,
  nodes : Array[Node],
) -> Node {
  { type_, nodes, token: "", start, end }
}

///|
fn conformance_text_node(token : String, start : Int, end : Int) -> Node {
  { type_: TextNode, nodes: [], token, start, end }
}
"""


def write_suite(testdata: Path, source: str) -> None:
    unused = set(KNOWN_GAPS)
    blocks: list[str] = []
    for suite, (section, render) in SUITES.items():
        suite_dir = testdata / suite
        if not suite_dir.is_dir():
            raise SystemExit(f"missing testdata directory: {suite_dir}")
        for path in sorted(suite_dir.glob("*.yaml")):
            case_id = f"{section}/{path.stem}"
            unused.discard(case_id)
            case = yaml.safe_load(path.read_text(encoding="utf-8"))
            blocks.append(render(case_id, case))
    if unused:
        raise SystemExit(f"KNOWN_GAPS has cases that do not exist: {sorted(unused)}")

    header = f"""\
// Code generated by `mise run conformance:generate`. DO NOT EDIT.
//
// Conformance suite from the reference Cucumber Expressions testdata:
//   {source}
// The testdata is Copyright (c) Cucumber Ltd and contributors, MIT License.
//
// Known gaps are generated with #skip and the issue that tracks them.

"""
    OUTPUT.write_text(header + HELPERS + "\n" + "\n".join(blocks), encoding="utf-8")


if __name__ == "__main__":
    main()
