#!/usr/bin/env python3
"""ast_to_facts.py -- Python AST fact extractor for Souffle-LintQ.

Visits a Qiskit Python script with the standard-library ``ast`` module and
emits tab-separated ``.facts`` files matching the EDB schema from the
Souffle-LintQ proposal:

    Stmt(id, line, func)
    CFGEdge(from_stmt, to_stmt)
    Assign(stmt, dest_var, src_var)
    CircuitAlloc(stmt, var_name, num_qubits, num_clbits)
    GateOp(stmt, circuit_var, gate_name, qubit_idx)
    MeasureOp(stmt, circuit_var, qubit_idx, clbit_idx)
    CircuitCall(stmt, method, target_var, arg_var)

Usage:
    python3 ast_to_facts.py <target.py> [--out <facts_dir>]
"""

from __future__ import annotations

import argparse
import ast
import os
import sys

# Gate method names emitted as GateOp rows.
GATE_NAMES = {
    "h", "x", "y", "z", "s", "t", "sdg", "tdg",
    "rx", "ry", "rz", "p", "u", "u1", "u2", "u3",
    "cx", "cz", "cy", "ch", "swap", "cswap",
    "cp", "cphase", "crz", "crx", "cry", "ccx",
    "rxx", "ryy", "rzz", "rzx", "iswap", "dcx", "ecr",
    "reset",
}

# Circuit method calls emitted as CircuitCall rows (GhostCompose watches compose).
CIRCUIT_METHODS = {
    "compose", "append", "extend", "add_register",
    "to_gate", "control", "reverse_ops", "measure_active",
    "measure_all", "barrier", "copy", "inverse",
}


class QiskitFactExtractor:
    """Walks statement lists in source order, assigning statement ids,
    recording domain facts, and building a (best-effort) control-flow graph."""

    def __init__(self) -> None:
        self.stmt_id = 0
        self.node_ids: dict[ast.AST, int] = {}

        self.stmts: list[tuple[int, int, str]] = []
        self.cfg_edges: list[tuple[int, int]] = []
        self.assigns: list[tuple[int, str, str]] = []
        self.circuit_allocs: list[tuple[int, str, int, int]] = []
        self.gate_ops: list[tuple[int, str, str, int]] = []
        self.measure_ops: list[tuple[int, str, int, int]] = []
        self.circuit_calls: list[tuple[int, str, str, str]] = []

    # ------------------------------------------------------------------ ids
    def id_of(self, node: ast.AST) -> int:
        if node not in self.node_ids:
            self.stmt_id += 1
            self.node_ids[node] = self.stmt_id
        return self.node_ids[node]

    def _name_of(self, node: ast.AST) -> str:
        """Best-effort dotted name for a Name/Attribute/Call receiver."""
        if isinstance(node, ast.Name):
            return node.id
        if isinstance(node, ast.Attribute):
            base = self._name_of(node.value)
            return f"{base}.{node.attr}" if base else node.attr
        if isinstance(node, ast.Call):
            return self._name_of(node.func)
        return "unknown"

    # -------------------------------------------------------------- visitors
    def visit_stmts(self, stmts, func: str, succ: list[int]) -> list[int]:
        """Visit a block of statements.

        Returns the entry statement ids of this block (the ids a predecessor
        should connect its CFG edges to). ``succ`` are the statement ids that
        control flows to once the block completes.
        """
        if not stmts:
            return list(succ)

        first = stmts[0]
        fid = self.id_of(first)
        self._record_stmt(first, func)
        self._extract_facts(first, func)

        rest = stmts[1:]
        nxt = self.visit_stmts(rest, func, succ)
        self._link_stmt(first, func, nxt)
        return [fid]

    def _record_stmt(self, node: ast.AST, func: str) -> None:
        sid = self.id_of(node)
        line = getattr(node, "lineno", 0)
        self.stmts.append((sid, line, func))

    def _link_stmt(self, s: ast.AST, func: str, nxt: list[int]) -> None:
        fid = self.id_of(s)

        if isinstance(s, ast.If):
            outs: list[int] = []
            if s.body:
                self.cfg_edges.append((fid, self.id_of(s.body[0])))
                outs += self.visit_stmts(s.body, func, nxt)
            if s.orelse:
                self.cfg_edges.append((fid, self.id_of(s.orelse[0])))
                outs += self.visit_stmts(s.orelse, func, nxt)
            if not s.orelse:
                outs += nxt
            return

        if isinstance(s, (ast.For, ast.While, ast.AsyncFor)):
            outs = []
            if s.body:
                self.cfg_edges.append((fid, self.id_of(s.body[0])))
                # back-edge: loop body can return to the header
                outs += self.visit_stmts(s.body, func, [fid])
            else:
                outs.append(fid)
            if s.orelse:
                outs += self.visit_stmts(s.orelse, func, nxt)
            else:
                outs += nxt
            return

        if isinstance(s, (ast.With, ast.AsyncWith)):
            if s.body:
                self.cfg_edges.append((fid, self.id_of(s.body[0])))
            self.visit_stmts(s.body, func, nxt)
            return

        if isinstance(s, (ast.FunctionDef, ast.AsyncFunctionDef)):
            # body executes only when called; no inline CFG linkage
            self.visit_stmts(s.body, s.name, [])
            for t in nxt:
                self.cfg_edges.append((fid, t))
            return

        if isinstance(s, ast.ClassDef):
            for item in s.body:
                self.visit_stmts([item], func, [])
            for t in nxt:
                self.cfg_edges.append((fid, t))
            return

        # simple statement: fall through to successor block
        for t in nxt:
            self.cfg_edges.append((fid, t))

    # -------------------------------------------------------- domain facts
    def _extract_facts(self, node: ast.AST, func: str) -> None:
        sid = self.id_of(node)

        if isinstance(node, ast.Assign):
            self._fact_assign(sid, node)
        elif isinstance(node, ast.Expr) and isinstance(node.value, ast.Call):
            self._fact_expr_call(sid, node.value)

    def _fact_assign(self, sid: int, node: ast.Assign) -> None:
        target = node.targets[0]
        dest_var = target.id if isinstance(target, ast.Name) else "unknown"
        value = node.value

        if isinstance(value, ast.Call) and isinstance(value.func, ast.Name):
            if value.func.id == "QuantumCircuit":
                num_q, num_c = self._count_qubits(value)
                self.circuit_allocs.append((sid, dest_var, num_q, num_c))
                self.assigns.append((sid, dest_var, "call_result"))
                return

        if isinstance(value, ast.Name):
            self.assigns.append((sid, dest_var, value.id))
        elif isinstance(value, ast.Call):
            self.assigns.append((sid, dest_var, "call_result"))
        else:
            self.assigns.append((sid, dest_var, "expr"))

    def _count_qubits(self, call: ast.Call) -> tuple[int, int]:
        num_q, num_c = 0, 0
        for arg in call.args:
            if isinstance(arg, ast.Constant) and isinstance(arg.value, int):
                if num_q == 0:
                    num_q = arg.value
                else:
                    num_c = arg.value
            elif isinstance(arg, ast.Call) and isinstance(arg.func, ast.Name):
                fname = arg.func.id
                n = 0
                if arg.args and isinstance(arg.args[0], ast.Constant) \
                        and isinstance(arg.args[0].value, int):
                    n = arg.args[0].value
                if fname == "QuantumRegister":
                    num_q += n
                elif fname == "ClassicalRegister":
                    num_c += n
        return num_q, num_c

    def _fact_expr_call(self, sid: int, call: ast.Call) -> None:
        if not isinstance(call.func, ast.Attribute):
            return
        method = call.func.attr
        circuit_var = self._name_of(call.func.value)

        if method in GATE_NAMES:
            for arg in call.args:
                if isinstance(arg, ast.Constant) and isinstance(arg.value, int):
                    self.gate_ops.append((sid, circuit_var, method, arg.value))
            return

        if method == "measure":
            ints = [a.value for a in call.args
                    if isinstance(a, ast.Constant) and isinstance(a.value, int)]
            q_idx = ints[0] if len(ints) > 0 else 0
            c_idx = ints[1] if len(ints) > 1 else 0
            self.measure_ops.append((sid, circuit_var, q_idx, c_idx))
            return

        if method in CIRCUIT_METHODS:
            arg_var = "unknown"
            if call.args:
                a0 = call.args[0]
                if isinstance(a0, ast.Name):
                    arg_var = a0.id
                else:
                    arg_var = self._name_of(a0)
            self.circuit_calls.append((sid, method, circuit_var, arg_var))

    # -------------------------------------------------------------- output
    def write_facts(self, output_dir: str) -> None:
        os.makedirs(output_dir, exist_ok=True)

        def dump(fname, rows):
            path = os.path.join(output_dir, fname)
            with open(path, "w") as fh:
                for row in rows:
                    fh.write("\t".join(str(c) for c in row) + "\n")

        dump("Stmt.facts", self.stmts)
        dump("CFGEdge.facts", self.cfg_edges)
        dump("Assign.facts", self.assigns)
        dump("CircuitAlloc.facts", self.circuit_allocs)
        dump("GateOp.facts", self.gate_ops)
        dump("MeasureOp.facts", self.measure_ops)
        dump("CircuitCall.facts", self.circuit_calls)

    # ----------------------------------------------------------------- api
    def run(self, source: str, output_dir: str) -> None:
        tree = ast.parse(source)
        self.visit_stmts(tree.body, "global", [])
        self.write_facts(output_dir)


def main(argv=None) -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("target", help="path to the Qiskit .py script")
    parser.add_argument("--out", default="facts",
                        help="output directory for .facts (default: facts)")
    args = parser.parse_args(argv)

    with open(args.target, "r") as fh:
        source = fh.read()

    ext = QiskitFactExtractor()
    ext.run(source, args.out)

    counts = {
        "Stmt": len(ext.stmts),
        "CFGEdge": len(ext.cfg_edges),
        "Assign": len(ext.assigns),
        "CircuitAlloc": len(ext.circuit_allocs),
        "GateOp": len(ext.gate_ops),
        "MeasureOp": len(ext.measure_ops),
        "CircuitCall": len(ext.circuit_calls),
    }
    print(f"EDB facts written to ./{args.out}")
    for name, n in counts.items():
        print(f"  {name:12s} {n}")
    return 0


if __name__ == "__main__":
    sys.exit(main())
