#!/usr/bin/env python3
#MISE description="Generate the arity-specific code for typed parameter types"
"""Generate the code that repeats for each arity from 1 to 8.

MoonBit has no variadic generics, so each arity needs its own type or
function. Writing eight near-identical copies by hand is error-prone, so
this script writes them:

- src/captures_arity.mbt: Captures2 to Captures8, zip, map and the fold
  after 8 values.
- src/define_arity.mbt: ParamTypeRegistry::define1 to define8.
- src/parameter_type_def.mbt: ParameterTypeDef, FromGroups1 to
  FromGroups8 and ParamTypeRegistry::define_type1 to define_type8.

Run it after a change to this file, then commit the generated files.
"""

from __future__ import annotations

import subprocess
from pathlib import Path

ROOT = Path(__file__).resolve().parents[2]
MAX_ARITY = 8
LETTERS = ["A", "B", "C", "D", "E", "F", "G", "H"]
HEADER = "// Code generated by `mise run codegen:captures`. DO NOT EDIT.\n\n"


def params(n: int) -> list[str]:
    return LETTERS[:n]


def string_chain(n: int) -> str:
    """Captures::string() zipped n times: a CapturesN of Strings."""
    return "Captures::string()" + ".zip(Captures::string())" * (n - 1)


def captures_type(n: int) -> str:
    ps = ", ".join(params(n))
    return f"Captures{n}[{ps}]"


def gen_captures_arity() -> str:
    out = [HEADER]
    # Captures1::zip
    out.append(
        """///|
/// Add the next part. The result decodes both values.
pub fn[A, B] Captures1::zip(
  self : Captures1[A],
  next : Captures1[B],
) -> Captures2[A, B] {
  { p1: self, p2: next }
}
"""
    )
    for n in range(2, MAX_ARITY + 1):
        ps = params(n)
        fields = "\n".join(f"  priv p{i + 1} : Captures1[{p}]" for i, p in enumerate(ps))
        out.append(
            f"""
///|
/// A decoder of {n} values. Add a value with `zip`, or make one value with
/// `map`.
pub struct {captures_type(n)} {{
{fields}
}}
"""
        )
        # map
        counts = "\n".join(f"  let c{i + 1} = self.p{i + 1}.count" for i in range(n))
        total = " + ".join(f"c{i + 1}" for i in range(n))
        decodes = []
        for i in range(n):
            decodes.append(
                f"      let v{i + 1} = (self.p{i + 1}.decode)(values[offset:offset + c{i + 1}], first + offset)"
            )
            if i < n - 1:
                decodes.append(f"      offset = offset + c{i + 1}")
        decodes_text = "\n".join(decodes)
        args = ", ".join(f"v{i + 1}" for i in range(n))
        out.append(
            f"""
///|
/// Make one value from the {n} values with `f`. The result decodes the same
/// capture groups, so it can be zipped into a bigger decoder.
pub fn[{", ".join(ps)}, T] Captures{n}::map(
  self : {captures_type(n)},
  f : ({", ".join(ps)}) -> T raise,
) -> Captures1[T] {{
{counts}
  {{
    count: {total},
    decode: (values, first) => {{
      let mut offset = 0
{decodes_text}
      f({args})
    }},
  }}
}}
"""
        )
        # zip
        if n < MAX_ARITY:
            nxt = LETTERS[n]
            copied = ", ".join(f"p{i + 1}: self.p{i + 1}" for i in range(n))
            out.append(
                f"""
///|
/// Add the next part. The result decodes {n + 1} values.
pub fn[{", ".join(ps)}, {nxt}] Captures{n}::zip(
  self : {captures_type(n)},
  next : Captures1[{nxt}],
) -> {captures_type(n + 1)} {{
  {{ {copied}, p{n + 1}: next }}
}}
"""
            )
        else:
            tuple_type = "(" + ", ".join(ps) + ")"
            names = ", ".join(p.lower() for p in ps)
            out.append(
                f"""
///|
/// Add the next part after {MAX_ARITY} values. The {MAX_ARITY} values fold
/// into one tuple, so the result is a `Captures2` of that tuple and the next
/// value. Zipping can continue with no limit.
pub fn[{", ".join(ps)}, X] Captures{n}::zip(
  self : {captures_type(n)},
  next : Captures1[X],
) -> Captures2[{tuple_type}, X] {{
  self.map(({names}) => ({names})).zip(next)
}}
"""
            )
    return "".join(out)


def gen_define_arity() -> str:
    out = [HEADER]
    for n in range(1, MAX_ARITY + 1):
        arg_types = ", ".join(["String"] * n)
        out.append(
            f"""///|
/// Register a typed parameter type whose transformer takes {n} capture
/// group value{'s' if n > 1 else ''}. See `define_with` for the rules.
pub fn[T] ParamTypeRegistry::define{n}(
  self : ParamTypeRegistry,
  name : String,
  regexps : Array[RegexPattern],
  transformer : ({arg_types}) -> T raise,
  use_for_snippets? : Bool = true,
  prefer_for_regexp_match? : Bool = false,
) -> ParameterType[T] raise ParameterTypeError {{
  self.define_with(
    name,
    regexps,
    {string_chain(n)}.map(transformer),
    use_for_snippets~,
    prefer_for_regexp_match~,
  )
}}

"""
        )
    return "".join(out)


def gen_parameter_type_def() -> str:
    out = [HEADER]
    out.append(
        """///|
/// A typed parameter type defined by the type itself. Implement it and one
/// of `FromGroups1` to `FromGroups8`, then register the type with
/// `define_type1` to `define_type8`.
pub(open) trait ParameterTypeDef {
  /// The name used in expressions, for example "color" for `{color}`.
  name() -> String
  /// The regexps that match the parameter.
  regexps() -> Array[RegexPattern]
  /// Default: `true`.
  use_for_snippets() -> Bool = _
  /// Default: `false`.
  prefer_for_regexp_match() -> Bool = _
}

///|
impl ParameterTypeDef with use_for_snippets() {
  true
}

///|
impl ParameterTypeDef with prefer_for_regexp_match() {
  false
}

///|
priv struct TypeDef {
  name : String
  regexps : Array[RegexPattern]
  use_for_snippets : Bool
  prefer_for_regexp_match : Bool
}

///|
/// Read the definition of `T`. The argument only selects `T`.
fn[T : ParameterTypeDef] type_def(_proxy : T?) -> TypeDef {
  {
    name: T::name(),
    regexps: T::regexps(),
    use_for_snippets: T::use_for_snippets(),
    prefer_for_regexp_match: T::prefer_for_regexp_match(),
  }
}

"""
    )
    for n in range(1, MAX_ARITY + 1):
        arg_types = ", ".join(["String"] * n)
        names = ", ".join(f"g{i + 1}" for i in range(n))
        typed_args = ", ".join(f"g{i + 1} : String" for i in range(n))
        out.append(
            f"""///|
/// Make a value from {n} capture group value{'s' if n > 1 else ''}.
pub(open) trait FromGroups{n} {{
  from_groups({arg_types}) -> Self raise
}}

///|
fn[T : FromGroups{n}] from_groups{n}({typed_args}) -> T raise {{
  T::from_groups({names})
}}

///|
/// Register the typed parameter type `T`, which takes {n} capture group
/// value{'s' if n > 1 else ''}. Select `T` with a type annotation, for example
/// `let color : ParameterType[Color] = registry.define_type{n}()`.
pub fn[T : ParameterTypeDef + FromGroups{n}] ParamTypeRegistry::define_type{n}(
  self : ParamTypeRegistry,
) -> ParameterType[T] raise ParameterTypeError {{
  let def = type_def((None : T?))
  let captures : Captures1[T] = {string_chain(n)}.map(({names}) => from_groups{n}({names}))
  self.define_with(
    def.name,
    def.regexps,
    captures,
    use_for_snippets=def.use_for_snippets,
    prefer_for_regexp_match=def.prefer_for_regexp_match,
  )
}}

"""
        )
    return "".join(out)


def main() -> None:
    outputs = {
        "src/captures_arity.mbt": gen_captures_arity(),
        "src/define_arity.mbt": gen_define_arity(),
        "src/parameter_type_def.mbt": gen_parameter_type_def(),
    }
    for path, text in outputs.items():
        (ROOT / path).write_text(text, encoding="utf-8")
        print(f"Wrote {path}")
    subprocess.run(["moon", "fmt"], cwd=ROOT, check=True)


if __name__ == "__main__":
    main()
