"""Parser for KiCad PCB layout files (.kicad_pcb).

Extracts footprints, nets, tracks, vias, zones, layers, and board
dimensions from KiCad PCB files using the S-expression parser.
"""

from __future__ import annotations

from pathlib import Path

from .sexpr import parse_file, find_node, find_nodes, node_value


def parse_pcb(filepath: str | Path) -> dict:
    """Parse a .kicad_pcb file and return structured data."""
    tree = parse_file(filepath)
    return _extract_pcb(tree, filepath)


def _extract_pcb(tree: list, filepath: str | Path) -> dict:
    """Extract all relevant data from a parsed PCB tree."""
    return {
        "file": str(filepath),
        "version": node_value(find_node(tree, "version")),
        "generator": node_value(find_node(tree, "generator")),
        "general": _extract_general(find_node(tree, "general")),
        "paper": node_value(find_node(tree, "paper")),
        "title_block": _extract_title_block(find_node(tree, "title_block")),
        "layers": _extract_layers(tree),
        "nets": _extract_nets(tree),
        "footprints": _extract_footprints(tree),
        "segments": _extract_segments(tree),
        "vias": _extract_vias(tree),
        "zones": _extract_zones(tree),
        "dimensions": _extract_board_dimensions(tree),
    }


def _extract_general(node: list | None) -> dict:
    """Extract general board settings."""
    if not node:
        return {}
    result = {}
    thickness = find_node(node, "thickness")
    if thickness:
        try:
            result["thickness"] = float(node_value(thickness))
        except (ValueError, TypeError):
            pass
    return result


def _extract_title_block(node: list | None) -> dict:
    """Extract title block metadata."""
    if not node:
        return {}
    result = {}
    for field in ("title", "date", "rev", "company"):
        val = node_value(find_node(node, field))
        if val:
            result[field] = val
    for child in node:
        if isinstance(child, list) and child and child[0] == "comment" and len(child) >= 3:
            result[f"comment_{child[1]}"] = child[2]
    return result


def _extract_layers(tree: list) -> list[dict]:
    """Extract layer definitions."""
    layers_node = find_node(tree, "layers")
    if not layers_node:
        return []
    layers = []
    for child in layers_node[1:]:
        if isinstance(child, list) and len(child) >= 3:
            layer = {
                "number": child[0] if isinstance(child[0], str) else str(child[0]),
                "name": child[1] if isinstance(child[1], str) else "",
                "type": child[2] if len(child) > 2 and isinstance(child[2], str) else "",
            }
            if len(child) > 3 and isinstance(child[3], str):
                layer["canonical_name"] = child[3]
            layers.append(layer)
    return layers


def _extract_nets(tree: list) -> list[dict]:
    """Extract net definitions."""
    nets = []
    for node in find_nodes(tree, "net"):
        if len(node) >= 3:
            try:
                nets.append({
                    "number": int(node[1]),
                    "name": node[2] if isinstance(node[2], str) else "",
                })
            except (ValueError, TypeError):
                pass
    return nets


def _extract_position(node: list) -> dict | None:
    """Extract position from an (at x y [rot]) node."""
    at_node = find_node(node, "at")
    if not at_node or len(at_node) < 3:
        return None
    try:
        result = {"x": float(at_node[1]), "y": float(at_node[2])}
        if len(at_node) >= 4:
            try:
                result["rotation"] = float(at_node[3])
            except (ValueError, TypeError):
                pass
        return result
    except (ValueError, TypeError):
        return None


def _extract_property(node: list) -> dict | None:
    """Extract a property node."""
    if not node or node[0] != "property" or len(node) < 3:
        return None
    return {"name": node[1], "value": node[2]}


def _extract_footprints(tree: list) -> list[dict]:
    """Extract footprint instances from the PCB."""
    footprints = []
    for node in find_nodes(tree, "footprint"):
        if len(node) < 2:
            continue
        fp_name = node[1] if isinstance(node[1], str) else ""

        fp = {
            "footprint": fp_name,
            "layer": node_value(find_node(node, "layer")),
            "uuid": node_value(find_node(node, "uuid")),
            "position": _extract_position(node),
            "locked": find_node(node, "locked") is not None,
        }

        # Extract properties
        props = {}
        for child in node:
            p = _extract_property(child) if isinstance(child, list) else None
            if p:
                props[p["name"]] = p["value"]
        fp["properties"] = props
        fp["reference"] = props.get("Reference", "")
        fp["value"] = props.get("Value", "")

        # Path (links to schematic)
        path_node = find_node(node, "path")
        if path_node:
            fp["path"] = node_value(path_node)

        # Sheet info
        fp["sheetname"] = node_value(find_node(node, "sheetname")) or ""
        fp["sheetfile"] = node_value(find_node(node, "sheetfile")) or ""

        # Attributes
        attr_node = find_node(node, "attr")
        if attr_node and len(attr_node) > 1:
            fp["attr"] = node_value(attr_node)
        else:
            fp["attr"] = ""

        # Count pads
        pads = find_nodes(node, "pad")
        fp["pad_count"] = len(pads)

        # Extract pad details
        pad_list = []
        for pad in pads:
            if len(pad) >= 4:
                pad_info = {
                    "number": pad[1] if isinstance(pad[1], str) else "",
                    "type": pad[2] if isinstance(pad[2], str) else "",
                    "shape": pad[3] if isinstance(pad[3], str) else "",
                }
                pad_at = find_node(pad, "at")
                if pad_at and len(pad_at) >= 3:
                    try:
                        pad_info["x"] = float(pad_at[1])
                        pad_info["y"] = float(pad_at[2])
                    except (ValueError, TypeError):
                        pass
                size_node = find_node(pad, "size")
                if size_node and len(size_node) >= 3:
                    try:
                        pad_info["width"] = float(size_node[1])
                        pad_info["height"] = float(size_node[2])
                    except (ValueError, TypeError):
                        pass
                net_node = find_node(pad, "net")
                if net_node and len(net_node) >= 3:
                    try:
                        pad_info["net_number"] = int(net_node[1])
                        pad_info["net_name"] = net_node[2] if isinstance(net_node[2], str) else ""
                    except (ValueError, TypeError):
                        pass
                pad_list.append(pad_info)
        fp["pads"] = pad_list

        footprints.append(fp)
    return footprints


def _extract_segments(tree: list) -> list[dict]:
    """Extract track segments."""
    segments = []
    for node in find_nodes(tree, "segment"):
        seg = {"uuid": node_value(find_node(node, "uuid"))}
        start = find_node(node, "start")
        if start and len(start) >= 3:
            try:
                seg["start"] = {"x": float(start[1]), "y": float(start[2])}
            except (ValueError, TypeError):
                pass
        end = find_node(node, "end")
        if end and len(end) >= 3:
            try:
                seg["end"] = {"x": float(end[1]), "y": float(end[2])}
            except (ValueError, TypeError):
                pass
        width = find_node(node, "width")
        if width:
            try:
                seg["width"] = float(node_value(width))
            except (ValueError, TypeError):
                pass
        layer = find_node(node, "layer")
        if layer:
            seg["layer"] = node_value(layer)
        net = find_node(node, "net")
        if net:
            try:
                seg["net"] = int(node_value(net))
            except (ValueError, TypeError):
                pass
        segments.append(seg)
    return segments


def _extract_vias(tree: list) -> list[dict]:
    """Extract vias."""
    vias = []
    for node in find_nodes(tree, "via"):
        via = {"uuid": node_value(find_node(node, "uuid"))}
        # Via type is the second element if it's a string (e.g., "blind")
        if len(node) >= 2 and isinstance(node[1], str) and node[1] not in ("(", ")"):
            # Check if it looks like a type keyword (not a sub-node)
            at_node = find_node(node, "at")
            if at_node != node[1]:
                via["type"] = node[1]
        at = find_node(node, "at")
        if at and len(at) >= 3:
            try:
                via["position"] = {"x": float(at[1]), "y": float(at[2])}
            except (ValueError, TypeError):
                pass
        size = find_node(node, "size")
        if size:
            try:
                via["size"] = float(node_value(size))
            except (ValueError, TypeError):
                pass
        drill = find_node(node, "drill")
        if drill:
            try:
                via["drill"] = float(node_value(drill))
            except (ValueError, TypeError):
                pass
        layers = find_node(node, "layers")
        if layers:
            via["layers"] = [l for l in layers[1:] if isinstance(l, str)]
        net = find_node(node, "net")
        if net:
            try:
                via["net"] = int(node_value(net))
            except (ValueError, TypeError):
                pass
        vias.append(via)
    return vias


def _extract_zones(tree: list) -> list[dict]:
    """Extract zone definitions (simplified — no polygon data)."""
    zones = []
    for node in find_nodes(tree, "zone"):
        zone = {"uuid": node_value(find_node(node, "uuid"))}
        net = find_node(node, "net")
        if net:
            try:
                zone["net"] = int(node_value(net))
            except (ValueError, TypeError):
                pass
        net_name = find_node(node, "net_name")
        if net_name:
            zone["net_name"] = node_value(net_name)
        layer = find_node(node, "layer")
        if layer:
            zone["layer"] = node_value(layer)
        layers = find_node(node, "layers")
        if layers:
            zone["layers"] = [l for l in layers[1:] if isinstance(l, str)]
        name = find_node(node, "name")
        if name:
            zone["name"] = node_value(name)
        priority = find_node(node, "priority")
        if priority:
            try:
                zone["priority"] = int(node_value(priority))
            except (ValueError, TypeError):
                pass
        zones.append(zone)
    return zones


def _extract_board_dimensions(tree: list) -> dict | None:
    """Calculate board dimensions from Edge.Cuts layer graphics.

    Scans for gr_line, gr_rect, gr_arc, and gr_circle on Edge.Cuts.
    Returns bounding box in mm.
    """
    min_x = float("inf")
    min_y = float("inf")
    max_x = float("-inf")
    max_y = float("-inf")
    found = False

    def update_bounds(x, y):
        nonlocal min_x, min_y, max_x, max_y, found
        min_x = min(min_x, x)
        min_y = min(min_y, y)
        max_x = max(max_x, x)
        max_y = max(max_y, y)
        found = True

    for tag in ("gr_line", "gr_rect", "gr_arc", "gr_circle", "gr_poly"):
        for node in find_nodes(tree, tag):
            layer = find_node(node, "layer")
            if not layer or node_value(layer) != "Edge.Cuts":
                continue
            # Extract coordinates from start/end/center nodes
            for coord_tag in ("start", "end", "center", "mid"):
                coord = find_node(node, coord_tag)
                if coord and len(coord) >= 3:
                    try:
                        update_bounds(float(coord[1]), float(coord[2]))
                    except (ValueError, TypeError):
                        pass
            # For polygons, scan pts
            pts = find_node(node, "pts")
            if pts:
                for xy in find_nodes(pts, "xy"):
                    if len(xy) >= 3:
                        try:
                            update_bounds(float(xy[1]), float(xy[2]))
                        except (ValueError, TypeError):
                            pass

    # Also check fp_line/fp_rect on Edge.Cuts inside footprints
    for fp in find_nodes(tree, "footprint"):
        fp_at = find_node(fp, "at")
        fp_x, fp_y, fp_rot = 0.0, 0.0, 0.0
        if fp_at and len(fp_at) >= 3:
            try:
                fp_x = float(fp_at[1])
                fp_y = float(fp_at[2])
                if len(fp_at) >= 4:
                    fp_rot = float(fp_at[3])
            except (ValueError, TypeError):
                pass
        for tag in ("fp_line", "fp_rect", "fp_arc"):
            for node in find_nodes(fp, tag):
                layer = find_node(node, "layer")
                if not layer or node_value(layer) != "Edge.Cuts":
                    continue
                for coord_tag in ("start", "end", "center", "mid"):
                    coord = find_node(node, coord_tag)
                    if coord and len(coord) >= 3:
                        try:
                            # fp_line coordinates are relative to footprint
                            x = float(coord[1]) + fp_x
                            y = float(coord[2]) + fp_y
                            update_bounds(x, y)
                        except (ValueError, TypeError):
                            pass

    if not found:
        return None

    return {
        "min_x": round(min_x, 3),
        "min_y": round(min_y, 3),
        "max_x": round(max_x, 3),
        "max_y": round(max_y, 3),
        "width": round(max_x - min_x, 3),
        "height": round(max_y - min_y, 3),
    }


# ── High-level query functions ──────────────────────────────────────


def list_pcb_footprints(filepath: str | Path) -> list[dict]:
    """List all footprint instances on the PCB.

    Returns a list of dicts with: reference, value, footprint, layer,
    position, rotation, pad_count, attr.
    """
    data = parse_pcb(filepath)
    result = []
    for fp in data["footprints"]:
        pos = fp.get("position") or {}
        result.append({
            "reference": fp["reference"],
            "value": fp["value"],
            "footprint": fp["footprint"],
            "layer": fp["layer"],
            "x": pos.get("x"),
            "y": pos.get("y"),
            "rotation": pos.get("rotation", 0),
            "pad_count": fp["pad_count"],
            "attr": fp["attr"],
        })
    return result


def get_pcb_statistics(filepath: str | Path) -> dict:
    """Get summary statistics for a PCB layout."""
    data = parse_pcb(filepath)

    # Count copper layers
    copper_layers = [l for l in data["layers"] if l.get("type") in ("signal", "power", "mixed", "copper")]

    # Count unique nets (exclude net 0 which is unconnected)
    active_nets = [n for n in data["nets"] if n["number"] != 0]

    # Count total pads
    total_pads = sum(fp["pad_count"] for fp in data["footprints"])

    # SMD vs through-hole
    smd_count = sum(1 for fp in data["footprints"] if fp.get("attr") == "smd")
    th_count = sum(1 for fp in data["footprints"] if fp.get("attr") == "through_hole")

    return {
        "file": str(filepath),
        "title": data["title_block"].get("title", ""),
        "board_thickness": data["general"].get("thickness"),
        "dimensions": data["dimensions"],
        "copper_layer_count": len(copper_layers),
        "total_layer_count": len(data["layers"]),
        "footprint_count": len(data["footprints"]),
        "smd_footprints": smd_count,
        "through_hole_footprints": th_count,
        "total_pads": total_pads,
        "net_count": len(active_nets),
        "track_count": len(data["segments"]),
        "via_count": len(data["vias"]),
        "zone_count": len(data["zones"]),
    }


def list_pcb_nets(filepath: str | Path) -> list[dict]:
    """List all nets in the PCB with connected pad count.

    Returns nets sorted by number, excluding net 0 (unconnected).
    """
    data = parse_pcb(filepath)

    # Build pad count per net from footprint pads
    net_pad_count = {}
    for fp in data["footprints"]:
        for pad in fp.get("pads", []):
            net_num = pad.get("net_number")
            if net_num is not None and net_num != 0:
                net_pad_count[net_num] = net_pad_count.get(net_num, 0) + 1

    result = []
    for net in data["nets"]:
        if net["number"] == 0:
            continue
        result.append({
            "number": net["number"],
            "name": net["name"],
            "pad_count": net_pad_count.get(net["number"], 0),
        })

    return sorted(result, key=lambda n: n["number"])


def find_tracks_by_net(filepath: str | Path, net_name: str) -> dict:
    """Find all track segments and vias for a specific net.

    Args:
        filepath: Path to .kicad_pcb file
        net_name: Net name to search for

    Returns dict with segments, vias, and zone info for the net.
    """
    data = parse_pcb(filepath)

    # Find net number from name
    net_number = None
    for net in data["nets"]:
        if net["name"] == net_name:
            net_number = net["number"]
            break

    if net_number is None:
        return {"error": f"Net '{net_name}' not found", "segments": [], "vias": [], "zones": []}

    segments = [s for s in data["segments"] if s.get("net") == net_number]
    vias = [v for v in data["vias"] if v.get("net") == net_number]
    zones = [z for z in data["zones"] if z.get("net") == net_number]

    return {
        "net_name": net_name,
        "net_number": net_number,
        "segment_count": len(segments),
        "via_count": len(vias),
        "zone_count": len(zones),
        "segments": segments,
        "vias": vias,
        "zones": [{"layer": z.get("layer"), "name": z.get("name"), "priority": z.get("priority")} for z in zones],
    }