#!/usr/bin/env python3
"""
Generate Java record declarations from GajumaruChainObjects.asn
(or any module produced by gmser_schema_export).

This is intentionally a thin, schema-shaped codegen: it does NOT emit a
BER/DER codec. The abstract types are for typed construction; RLP remains
the wire format (see swiss.qpq.gajumaru.core.encoding.RLP).

Usage (from gajumaru-core/):

  bin/asn1_to_java \\
    --asn ../../gajumaru/asn1_generated/GajumaruChainObjects.asn \\
    --types Id,SignedTxV1,SpendTxV1

  bin/asn1_to_java --asn ... --all
"""

from __future__ import annotations

import argparse
import re
import sys
from dataclasses import dataclass, field
from pathlib import Path
from typing import List, Optional, Tuple, Union


# ---------------------------------------------------------------------------
# ASN.1 model (minimal)
# ---------------------------------------------------------------------------

AsnType = Union[
    "PrimitiveType",
    "NamedType",
    "SequenceType",
    "SequenceOfType",
]


@dataclass
class PrimitiveType:
    kind: str  # integer, octet_string, boolean, fixed_integer
    constraint: Optional[str] = None  # e.g. "0..255" or "12"


@dataclass
class NamedType:
    name: str  # Id, BigInt, SpendTxV1, ...


@dataclass
class SequenceField:
    name: str
    type: AsnType


@dataclass
class SequenceType:
    fields: List[SequenceField]


@dataclass
class SequenceOfType:
    elem: AsnType


@dataclass
class TypeDef:
    name: str
    type: AsnType
    comment: str = ""


# ---------------------------------------------------------------------------
# Parser
# ---------------------------------------------------------------------------

COMMENT_RE = re.compile(r"--[^\n]*")
TYPE_DEF_RE = re.compile(
    r"([A-Z][A-Za-z0-9]*)\s*::=\s*",
)


def strip_comments(text: str) -> str:
    return COMMENT_RE.sub("", text)


def parse_module(text: str) -> dict[str, TypeDef]:
    text = strip_comments(text)
    # Drop module wrapper noise; keep type assignments only.
    begin = text.find("BEGIN")
    end = text.rfind("END")
    if begin >= 0 and end > begin:
        text = text[begin + 5 : end]

    defs: dict[str, TypeDef] = {}
    pos = 0
    while True:
        m = TYPE_DEF_RE.search(text, pos)
        if not m:
            break
        name = m.group(1)
        start = m.end()
        # Find next top-level type def or EOF
        nxt = TYPE_DEF_RE.search(text, start)
        body = text[start : nxt.start() if nxt else len(text)].strip()
        # Trim trailing junk
        body = body.rstrip().rstrip(";").strip()
        try:
            asn_type = parse_type(body)
        except Exception as e:
            raise SystemExit(f"Failed to parse type {name}: {e}\nBody:\n{body[:200]}") from e
        defs[name] = TypeDef(name=name, type=asn_type)
        pos = nxt.start() if nxt else len(text)
    return defs


def parse_type(s: str) -> AsnType:
    s = s.strip()
    if s.startswith("SEQUENCE OF"):
        rest = s[len("SEQUENCE OF") :].strip()
        return SequenceOfType(parse_type(rest))
    if s.startswith("SEQUENCE"):
        rest = s[len("SEQUENCE") :].strip()
        if not rest.startswith("{"):
            raise ValueError(f"expected SEQUENCE {{...}}, got: {s[:60]}")
        inner = extract_braces(rest)
        return SequenceType(parse_fields(inner))
    if s.startswith("INTEGER"):
        rest = s[len("INTEGER") :].strip()
        if rest.startswith("(") and rest.endswith(")"):
            c = rest[1:-1].strip()
            if ".." in c or c == "MAX" or c.endswith("MAX"):
                return PrimitiveType("integer", c)
            # single value constraint: fixed integer
            if re.fullmatch(r"\d+", c):
                return PrimitiveType("fixed_integer", c)
            return PrimitiveType("integer", c)
        return PrimitiveType("integer")
    if s.startswith("OCTET STRING"):
        rest = s[len("OCTET STRING") :].strip()
        c = None
        if rest.startswith("(") and ")" in rest:
            c = rest[rest.find("(") + 1 : rest.rfind(")")]
        return PrimitiveType("octet_string", c)
    if s == "BOOLEAN":
        return PrimitiveType("boolean")
    if re.fullmatch(r"[A-Z][A-Za-z0-9]*", s):
        return NamedType(s)
    raise ValueError(f"unsupported type expression: {s[:80]!r}")


def extract_braces(s: str) -> str:
    """s starts with '{'; return content inside matching braces."""
    assert s[0] == "{"
    depth = 0
    for i, ch in enumerate(s):
        if ch == "{":
            depth += 1
        elif ch == "}":
            depth -= 1
            if depth == 0:
                return s[1:i]
    raise ValueError("unbalanced braces")


def parse_fields(body: str) -> List[SequenceField]:
    """Parse 'name Type, name Type' with nested braces."""
    fields: List[SequenceField] = []
    parts = split_top_level(body, ",")
    for part in parts:
        part = part.strip()
        if not part:
            continue
        # field name is first token (lowercase start in our export)
        m = re.match(r"([A-Za-z][A-Za-z0-9]*)\s+(.+)$", part, re.DOTALL)
        if not m:
            raise ValueError(f"bad field: {part[:60]!r}")
        fname, ftype = m.group(1), m.group(2).strip()
        fields.append(SequenceField(fname, parse_type(ftype)))
    return fields


def split_top_level(s: str, sep: str) -> List[str]:
    parts: List[str] = []
    depth = 0
    start = 0
    for i, ch in enumerate(s):
        if ch == "{":
            depth += 1
        elif ch == "}":
            depth -= 1
        elif ch == sep and depth == 0:
            parts.append(s[start:i])
            start = i + 1
    parts.append(s[start:])
    return parts


# ---------------------------------------------------------------------------
# Java emission
# ---------------------------------------------------------------------------

HEADER = """\
/*
 * Generated by bin/asn1_to_java from GajumaruChainObjects.asn.
 * Do not edit by hand — regenerate from the ASN.1 schema.
 *
 * These types are abstract syntax (headers). Wire encoding is RLP
 * via swiss.qpq.gajumaru.core.encoding.RLP / Asn1Rlp, not BER/DER.
 */

"""


def java_type_name(name: str) -> str:
    return name


def field_to_java(
    f: SequenceField,
    owner: str,
    nested: List[Tuple[str, SequenceType]],
) -> Tuple[str, str]:
    """Return (javaType, fieldName). May append nested type defs."""
    jtype = asn_to_java(f.type, owner, f.name, nested)
    return jtype, f.name


def asn_to_java(
    t: AsnType,
    owner: str,
    field_name: str,
    nested: List[Tuple[str, SequenceType]],
) -> str:
    if isinstance(t, NamedType):
        # Built-in aliases
        if t.name == "BigInt":
            return "java.math.BigInteger"
        if t.name in ("Uint8", "Uint16"):
            return "int"
        if t.name == "Uint32":
            return "long"
        if t.name in ("Uint64", "Uint128"):
            return "java.math.BigInteger"
        return t.name
    if isinstance(t, PrimitiveType):
        if t.kind == "boolean":
            return "boolean"
        if t.kind == "octet_string":
            return "byte[]"
        if t.kind == "fixed_integer":
            return "int"
        if t.kind == "integer":
            # constrained small ranges map to int/long when obvious
            if t.constraint in ("0..255", "0..65535"):
                return "int"
            if t.constraint == "0..4294967295":
                return "long"
            return "java.math.BigInteger"
        raise ValueError(t)
    if isinstance(t, SequenceOfType):
        elem = asn_to_java(t.elem, owner, field_name + "Elem", nested)
        return f"java.util.List<{box(elem)}>"
    if isinstance(t, SequenceType):
        nested_name = f"{owner}_{capitalize(field_name)}"
        nested.append((nested_name, t))
        # also need to resolve nested field types recursively for emission
        return nested_name
    raise ValueError(f"unknown type {t}")


def box(jtype: str) -> str:
    return {
        "int": "Integer",
        "long": "Long",
        "boolean": "Boolean",
        "byte[]": "byte[]",  # List<byte[]> is awkward but ok for now
    }.get(jtype, jtype)


def capitalize(s: str) -> str:
    return s[:1].upper() + s[1:] if s else s


def emit_sequence_record(
    name: str,
    seq: SequenceType,
    package: str,
    all_nested: List[Tuple[str, SequenceType]],
) -> str:
    nested: List[Tuple[str, SequenceType]] = []
    components = []
    constants = []
    for f in seq.fields:
        jtype, jname = field_to_java(f, name, nested)
        components.append(f"    {jtype} {jname}")
        if isinstance(f.type, PrimitiveType) and f.type.kind == "fixed_integer":
            if f.name == "tag":
                constants.append(f"    public static final int TAG = {f.type.constraint};")
            elif f.name == "vsn":
                constants.append(f"    public static final int VSN = {f.type.constraint};")

    all_nested.extend(nested)

    body = ",\n".join(components)
    const_block = ("\n" + "\n".join(constants) + "\n") if constants else ""
    return (
        f"package {package};\n\n"
        f"{HEADER}"
        f"public record {name}(\n{body}\n) {{\n"
        f"{const_block}"
        f"}}\n"
    )


def emit_alias(name: str, t: AsnType, package: str) -> Optional[str]:
    """Emit a tiny holder for INTEGER aliases we don't map away."""
    # We map BigInt/Uint* at use sites; still emit Id as a record.
    if name in ("BigInt", "Uint8", "Uint16", "Uint32", "Uint64", "Uint128"):
        return None
    if isinstance(t, SequenceType):
        return None  # handled elsewhere
    if isinstance(t, PrimitiveType) and t.kind == "integer":
        # skip pure aliases
        return None
    return None


def collect_dependencies(name: str, defs: dict[str, TypeDef], acc: set[str]) -> None:
    if name in acc or name not in defs:
        return
    acc.add(name)
    t = defs[name].type
    walk_deps(t, defs, acc)


def walk_deps(t: AsnType, defs: dict[str, TypeDef], acc: set[str]) -> None:
    if isinstance(t, NamedType):
        if t.name in defs:
            collect_dependencies(t.name, defs, acc)
    elif isinstance(t, SequenceOfType):
        walk_deps(t.elem, defs, acc)
    elif isinstance(t, SequenceType):
        for f in t.fields:
            walk_deps(f.type, defs, acc)


def generate(
    defs: dict[str, TypeDef],
    wanted: List[str],
    package: str,
    out_dir: Path,
) -> List[Path]:
    # Always include named deps of wanted types
    selected: set[str] = set()
    for w in wanted:
        if w not in defs:
            raise SystemExit(f"Unknown type {w}. Available: {', '.join(sorted(defs))}")
        collect_dependencies(w, defs, selected)

    out_dir.mkdir(parents=True, exist_ok=True)
    written: List[Path] = []

    # Emit in dependency-friendly order: common first
    order = sorted(selected, key=lambda n: (0 if n == "Id" else 1, n))

    for name in order:
        tdef = defs[name]
        t = tdef.type
        if isinstance(t, SequenceType):
            nested_acc: List[Tuple[str, SequenceType]] = []
            src = emit_sequence_record(name, t, package, nested_acc)
            path = out_dir / f"{name}.java"
            path.write_text(src)
            written.append(path)
            # nested anonymous sequences as separate top-level records
            for nname, nseq in nested_acc:
                more: List[Tuple[str, SequenceType]] = []
                nsrc = emit_sequence_record(nname, nseq, package, more)
                npath = out_dir / f"{nname}.java"
                npath.write_text(nsrc)
                written.append(npath)
                # flatten one level of nesting iteratively
                queue = list(more)
                while queue:
                    qn, qs = queue.pop(0)
                    more2: List[Tuple[str, SequenceType]] = []
                    qsrc = emit_sequence_record(qn, qs, package, more2)
                    qpath = out_dir / f"{qn}.java"
                    qpath.write_text(qsrc)
                    written.append(qpath)
                    queue.extend(more2)
        else:
            alias = emit_alias(name, t, package)
            if alias:
                path = out_dir / f"{name}.java"
                path.write_text(alias)
                written.append(path)

    return written


def main(argv: List[str]) -> int:
    ap = argparse.ArgumentParser(description=__doc__)
    ap.add_argument(
        "--asn",
        type=Path,
        required=True,
        help="Path to GajumaruChainObjects.asn",
    )
    ap.add_argument(
        "--package",
        default="swiss.qpq.gajumaru.core.asn1",
        help="Java package for generated sources",
    )
    ap.add_argument(
        "--out",
        type=Path,
        default=None,
        help="Output directory for .java files "
        "(default: src/main/java/<package path> under gajumaru-core)",
    )
    ap.add_argument(
        "--types",
        default="Id,SignedTxV1,SpendTxV1",
        help="Comma-separated type names to generate (plus dependencies)",
    )
    ap.add_argument(
        "--all",
        action="store_true",
        help="Generate all SEQUENCE types in the module",
    )
    args = ap.parse_args(argv)

    asn_path = args.asn.resolve()
    if not asn_path.is_file():
        print(f"ASN.1 file not found: {asn_path}", file=sys.stderr)
        return 1

    text = asn_path.read_text()
    defs = parse_module(text)
    if not defs:
        print("No type definitions parsed", file=sys.stderr)
        return 1

    if args.all:
        wanted = [n for n, d in defs.items() if isinstance(d.type, SequenceType)]
    else:
        wanted = [t.strip() for t in args.types.split(",") if t.strip()]

    if args.out is None:
        script_dir = Path(__file__).resolve().parent
        project_dir = script_dir.parent
        pkg_path = args.package.replace(".", "/")
        out_dir = project_dir / "src" / "main" / "java" / pkg_path
    else:
        out_dir = args.out

    written = generate(defs, wanted, args.package, out_dir)
    print(f"Parsed {len(defs)} ASN.1 types from {asn_path}")
    print(f"Wrote {len(written)} Java file(s) to {out_dir}:")
    for p in written:
        print(f"  {p.name}")
    return 0


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