"""Generate a compact Rust table from Unicode Character Database files.

Dependencies: Python standard library only.
Sources:
- https://www.unicode.org/Public/17.0.0/ucd/DerivedCoreProperties.txt
- https://www.unicode.org/license.txt
"""

import bisect
import sys
from collections.abc import Iterable
from dataclasses import dataclass
from pathlib import Path
from urllib.request import urlopen

PROPERTY_NAMES: tuple[str, ...] = (
    "XID_Continue",
    "XID_Start",
)
RELEVANT_DERIVED_PROPERTIES = set(PROPERTY_NAMES)
SCRIPT_DIR = Path(__file__).parent.resolve()
UCD_VERSION = "17.0.0"
DERIVED_CORE_PROPERTIES_URL = (
    f"https://www.unicode.org/Public/{UCD_VERSION}/ucd/DerivedCoreProperties.txt"
)
LICENSE_URL = "https://www.unicode.org/license.txt"
OUTPUT_PATH = SCRIPT_DIR / "table.rs"
TESTS = []


def main():
    run_tests()
    sources = UnicodeSources.download()
    properties = parse(sources.derived_core_properties.splitlines())
    properties = filter_properties(properties, RELEVANT_DERIVED_PROPERTIES)
    ranges = build_property_ranges(properties)

    with open(OUTPUT_PATH, "wt", encoding="utf-8") as fobj:
        format_table(ranges, sources.license, fobj)


@dataclass
class UnicodeSources:
    derived_core_properties: str
    license: str

    @classmethod
    def download(cls) -> "UnicodeSources":
        return cls(
            derived_core_properties=cls._download_file(DERIVED_CORE_PROPERTIES_URL),
            license=cls._download_file(LICENSE_URL),
        )

    @staticmethod
    def _download_file(url: str) -> str:
        print(f":: downloading {url}", flush=True)
        with urlopen(url, timeout=30) as response:
            return response.read().decode("utf-8")


def parse(lines: Iterable[str]) -> Iterable[tuple[int, int, str]]:
    for line in lines:
        if "#" in line:
            line, *_ = line.partition("#")
        line = line.strip()

        if not line:
            continue

        range, _, property = line.partition(";")
        if ".." in range:
            start, _, end = range.partition("..")
            start = int(start.strip(), 16)
            end = int(end.strip(), 16) + 1

        else:
            start = int(range.strip(), 16)
            end = start + 1

        property = property.strip()

        yield start, end, property


def filter_properties(
    it: Iterable[tuple[int, int, str]], subset: set[str]
) -> Iterable[tuple[int, int, str]]:
    return ((start, end, prop) for start, end, prop in it if prop in subset)


def build_property_ranges(
    it: Iterable[tuple[int, int, str]],
) -> list[tuple[int, int, set[str]]]:
    ranges = [(0, 0x10FFFF + 1, {})]

    for start, end, property in it:
        lower_idx = bisect.bisect_left(ranges, start, key=lambda rng: rng[0])
        upper_idx = bisect.bisect_left(ranges, end, key=lambda rng: rng[1]) + 1

        lower_idx = lower_idx - 1 if lower_idx != 0 else 0

        lower_start, _, _ = ranges[lower_idx]
        _, upper_end, _ = ranges[upper_idx - 1]

        assert lower_start <= start <= upper_end
        assert lower_start <= end <= upper_end

        offset = 0
        for idx in range(lower_idx, upper_idx):
            idx = idx + offset
            range_start, range_end, range_properties = ranges[idx]

            inter = range_intersection((start, end), (range_start, range_end))
            if inter is None:
                continue

            inter_start, inter_end = inter

            new_ranges = []
            if range_start < inter_start:
                new_ranges.append((range_start, inter_start, {*range_properties}))

            new_ranges.append((inter_start, inter_end, {*range_properties, property}))

            if inter_end < range_end:
                new_ranges.append((inter_end, range_end, {*range_properties}))

            ranges[idx : idx + 1] = new_ranges
            offset += len(new_ranges) - 1

    return ranges


def range_intersection(
    left: tuple[int, int], right: tuple[int, int]
) -> tuple[int, int] | None:
    s0, e0 = left
    s1, e1 = right

    s = max(s0, s1)
    e = min(e0, e1)

    return (s, e) if s < e else None


def format_table(ranges, license, fobj):
    for line in license.splitlines():
        print(f"// {line.rstrip()}".rstrip(), file=fobj)

    print(file=fobj)
    print("// Generated by gen_table.py.", file=fobj)
    print(f"// Unicode version: {UCD_VERSION}.", file=fobj)
    print(file=fobj)
    print(file=fobj)
    print(
        """\
use std::range::Range;

const fn to_range(start: u32, end: u32) -> Range<[u8; 3]> {
    let start = start.to_le_bytes();
    if start[3] != 0 {
        panic!("invalid start");
    }

    let end = end.to_le_bytes();
    if end[3] != 0 {
        panic!("invalid end");
    }

    Range {
        start: [start[0], start[1], start[2]],
        end: [end[0], end[1], end[2]],
    }
}
""",
        file=fobj,
    )
    print(file=fobj)
    for idx, property in enumerate(PROPERTY_NAMES):
        print(f"pub const {property.upper()}: u16 = 1 << {idx};", file=fobj)

    print("pub const PROPERTIES: &[(Range<[u8; 3]>, u16)] = &[", file=fobj)
    for start, end, properties in ranges:
        property = (
            " | ".join(p.upper() for p in PROPERTY_NAMES if p in properties)
            if properties
            else 0
        )
        print(f"    (to_range({hex(start)}, {hex(end)}), {property}),", file=fobj)

    print("];", file=fobj)


def test(func):
    TESTS.append(func)
    return func


def run_tests():
    for func in TESTS:
        print(f":: test {func.__name__}", flush=True)
        func()


@test
def test_range_intersection():
    assert range_intersection((0, 10), (0, 10)) == (0, 10)
    assert range_intersection((0, 10), (10, 20)) is None
    assert range_intersection((0, 10), (0, 5)) == (0, 5)
    assert range_intersection((0, 10), (5, 15)) == (5, 10)
    assert range_intersection((0, 10), (-2, 5)) == (0, 5)


@test
def test_parse():
    assert list(
        parse(
            [
                "# comment",
                "0041..005A ; XID_Start # Lu [26] LATIN CAPITAL LETTER A..Z",
                "005F ; XID_Continue # Pc LOW LINE",
                "# NOTE: See UAX #44",
            ]
        )
    ) == [
        (0x41, 0x5B, "XID_Start"),
        (0x5F, 0x60, "XID_Continue"),
    ]


@test
def test_build_property_ranges():
    ranges = build_property_ranges(
        [
            (10, 20, "XID_Start"),
            (15, 25, "XID_Continue"),
        ]
    )
    assert ranges[:5] == [
        (0, 10, set()),
        (10, 15, {"XID_Start"}),
        (15, 20, {"XID_Start", "XID_Continue"}),
        (20, 25, {"XID_Continue"}),
        (25, 0x10FFFF + 1, set()),
    ]


if __name__ == "__main__":
    if sys.argv[1:] == ["--check"]:
        run_tests()
    else:
        main()
