#!/usr/bin/env python3
"""
MechBase PLC — Ladder Logic Engine v4

Complete IEC 61131-3 / Studio 5000 instruction set:
  Contacts:    XIC (NO), XIO (NC), Rising, Falling
  Outputs:     OTE, OTL (SET), OTU (RESET), OTN (NOT)
  Timers:      TON, TOF, TP, TONR (Retentive) — with EN, TT, DN bits
  Counters:    CTU, CTD, CTUD — with EN, DN, PRE, ACC
  Comparisons: EQU, NEQ, GRT, LES, GEQ, LEQ, LIM, CMP
  Math:        ADD, SUB, MUL, DIV, CPT, SQT, NEG, ABS, INC, DEC
  Data:        MOV, CPD, SCD (Scale)
  Logic:       AND, OR, XOR (Boolean operations)
  Control:     JMP, LBL, RES, OSR, OSF, MCR (Master Control Reset)

Rung structure:
  - Inline elements (series / AND logic)
  - Parallel branches (OR logic — any branch conducts)
  - Output section (coils, timers, counters)

OpenPLC-derived improvements (v4):
  - Deterministic scan timing (CLOCK_MONOTONIC, EMA averages, overrun tracking)
  - State machine: EMPTY → INIT → RUNNING ↔ STOPPED → ERROR
  - Instruction error flags (ERR/ENO like Allen Bradley)
  - Periodic timer scheduling (absolute timestamps, no drift)
  - Centralized I/O image tables with typed access

Usage:
    python3 ladder_test.py
"""

import time
import math
import argparse
import threading
import statistics
from pathlib import Path
from typing import List, Dict, Any, Optional
from dataclasses import dataclass, field


# ──────────────────────────────────────────────────────────────────────────────
# Instruction type constants
# ──────────────────────────────────────────────────────────────────────────────
class Contact:
    XIC     = "XIC"      # Examine If Closed   --| |--  (Normally Open)
    XIO     = "XIO"      # Examine If Open     --|/|-- (Normally Closed)
    RISING  = "Rising"   # Rising edge          --|P|--
    FALLING = "Falling"  # Falling edge         --|N|--
    # Timer bit contacts
    TON_DN  = "TON_DN"
    TON_TT  = "TON_TT"
    TON_EN  = "TON_EN"
    TOF_DN  = "TOF_DN"
    TOF_TT  = "TOF_TT"
    TOF_EN  = "TOF_EN"
    TP_DN   = "TP_DN"
    TP_TT   = "TP_TT"
    TP_EN   = "TP_EN"
    # Counter bit contacts
    CTU_DN  = "CTU_DN"
    CTU_EN  = "CTU_EN"
    CTD_DN  = "CTD_DN"
    CTD_EN  = "CTD_EN"


class Output:
    OTE   = "OTE"   # Output Energize         --( )--
    OTL   = "OTL"   # Output Latch (SET)      --(L)--
    OTU   = "OTU"   # Output Unlatch (RESET)  --(U)--
    OSR   = "OSR"   # One-Shot Rising (output)
    OSF   = "OSF"   # One-Shot Falling (output)
    OTN   = "OTN"   # Output Not (inverted coil) --(/)--


class Timer:
    TON  = "TON"   # On-Delay (non-retentive)
    TOF  = "TOF"   # Off-Delay
    TP   = "TP"    # Pulse
    TONR = "TONR"  # Retentive On-Delay (survives de-energize)

class Counter:
    CTU  = "CTU"   # Count Up
    CTD  = "CTD"   # Count Down
    CTUD = "CTUD"  # Count Up/Down

class Compare:
    EQU = "EQU"  # Equal
    NEQ = "NEQ"  # Not Equal
    GRT = "GRT"  # Greater Than
    LES = "LES"  # Less Than
    GEQ = "GEQ"  # Greater Than or Equal
    LEQ = "LEQ"  # Less Than or Equal
    LIM = "LIM"  # Limit (in range)
    CMP = "CMP"  # Generic compare (uses operator param)


class Math:
    ADD = "ADD"
    SUB = "SUB"
    MUL = "MUL"
    DIV = "DIV"
    CPT = "CPT"  # Compute expression
    SQT = "SQT"  # Square root
    NEG = "NEG"  # Negate (one's complement)
    ABS = "ABS"  # Absolute value
    INC = "INC"  # Increment
    DEC = "DEC"  # Decrement


class Data:
    MOV = "MOV"  # Move
    CPD = "CPD"  # Copy Data block
    SCD = "SCD"  # Scale (linear conversion)


class Logic:
    AND = "AND"  # Boolean AND
    OR  = "OR"   # Boolean OR
    XOR = "XOR"  # Boolean XOR


class Control:
    JMP = "JMP"  # Jump to label
    LBL = "LBL"  # Label
    RES = "RES"  # Reset timer/counter
    MCR = "MCR"  # Master Control Reset (conditional execution region)


# ──────────────────────────────────────────────────────────────────────────────
# OpenPLC-style Scan Timing (derived from scan_cycle_manager.c)
# ──────────────────────────────────────────────────────────────────────────────

@dataclass
class ScanStats:
    """Scan timing statistics with EMA averages (OpenPLC pattern)."""
    scan_count: int = 0
    # Scan time: actual user program execution time
    scan_time_min: float = float('inf')
    scan_time_max: float = 0.0
    scan_time_avg: float = 0.0  # Exponential moving average
    # Cycle time: wall-clock between cycle starts
    cycle_time_min: float = float('inf')
    cycle_time_max: float = 0.0
    cycle_time_avg: float = 0.0  # EMA
    # Cycle latency: deviation from expected periodic start
    cycle_latency_avg: float = 0.0  # EMA
    # Overruns: when actual cycle start exceeds expected
    overrun_count: int = 0
    # EMA alpha (0.1 = slow response, 0.5 = fast response)
    ema_alpha: float = 0.1

    def update_ema(self, current: float, alpha: float = None) -> float:
        """Update exponential moving average."""
        a = alpha or self.ema_alpha
        if self.scan_count == 0:
            return current
        return a * current + (1 - a) * self.scan_time_avg

    def format_stats(self) -> dict:
        """Return stats as dict for API."""
        return {
            "scan_count": self.scan_count,
            "scan_time_ms": {
                "min": round(self.scan_time_min * 1000, 3),
                "max": round(self.scan_time_max * 1000, 3),
                "avg": round(self.scan_time_avg * 1000, 3),
            },
            "cycle_time_ms": {
                "min": round(self.cycle_time_min * 1000, 3),
                "max": round(self.cycle_time_max * 1000, 3),
                "avg": round(self.cycle_time_avg * 1000, 3),
            },
            "cycle_latency_ms": round(self.cycle_latency_avg * 1000, 3),
            "overruns": self.overrun_count,
        }


# ──────────────────────────────────────────────────────────────────────────────
# OpenPLC-style PLC State Machine (derived from plc_state_manager.c)
# ──────────────────────────────────────────────────────────────────────────────

class PLCState:
    """PLC states with mutex-protected transitions (OpenPLC pattern)."""
    EMPTY   = "EMPTY"   # No program loaded
    INIT    = "INIT"    # Program loaded, not yet running
    RUNNING = "RUNNING" # Active scan cycle
    STOPPED = "STOPPED" # Paused, state preserved
    ERROR   = "ERROR"   # Error state (requires reset)

    # Valid transitions: from → [to, ...]
    VALID_TRANSITIONS = {
        EMPTY:   [INIT, RUNNING],
        INIT:    [RUNNING, STOPPED, EMPTY],
        RUNNING: [STOPPED, ERROR],
        STOPPED: [RUNNING, EMPTY],
        ERROR:   [STOPPED, EMPTY],
    }

    @staticmethod
    def can_transition(from_state: str, to_state: str) -> bool:
        """Check if state transition is valid."""
        allowed = PLCState.VALID_TRANSITIONS.get(from_state, [])
        return to_state in allowed


# ──────────────────────────────────────────────────────────────────────────────
# Timer State
# ──────────────────────────────────────────────────────────────────────────────
class TimerState:
    def __init__(self, tag: str, preset_ms: int = 1000):
        self.tag = tag
        self.preset = preset_ms
        self.elapsed = 0
        self.type = "TON"
        self.en = False
        self.tt = False
        self.dn = False
        self._last_enabled = False
        # Allen Bradley-style instruction error flags
        self.err = False  # Instruction error (overflow, domain error)
        self.no = True    # Not-operator output (always true unless error)


# ──────────────────────────────────────────────────────────────────────────────
# Counter State
# ──────────────────────────────────────────────────────────────────────────────
class CounterState:
    def __init__(self, tag: str, preset: int = 10):
        self.tag = tag
        self.preset = preset
        self.acc = 0
        self.type = "CTU"
        self.en = False
        self.dn = False
        self._last_enabled = False
        # Allen Bradley-style instruction error flags
        self.err = False
        self.no = True
        # CTUD-specific
        self._cd_en = False  # Count-down enable (CTUD)


# ──────────────────────────────────────────────────────────────────────────────
# Element (single instruction in a rung)
# ──────────────────────────────────────────────────────────────────────────────
class Element:
    def __init__(self, elem_type: str, tag: str, value: Any = 0,
                 value2: Any = 0, operator: str = "", expression: str = ""):
        self.type = elem_type
        self.tag = tag          # Primary tag (contact bit, output bit, timer tag, etc.)
        self.value = value      # Secondary param (preset, compare value, source B)
        self.value2 = value2    # Third param (LIM high limit, etc.)
        self.operator = operator  # CMP operator ("=", ">", "<", ">=", "<=", "<>")
        self.expression = expression  # CPT expression string
        # Edge tracking
        self.last_state = False

    def to_dict(self) -> dict:
        d = {
            "type": self.type,
            "tag": self.tag,
            "value": self.value,
        }
        if self.value2:
            d["value2"] = self.value2
        if self.operator:
            d["operator"] = self.operator
        if self.expression:
            d["expression"] = self.expression
        return d

    @classmethod
    def from_dict(cls, data: dict) -> "Element":
        return cls(
            elem_type=data.get("type", "XIC"),
            tag=data.get("tag", ""),
            value=data.get("value", 0),
            value2=data.get("value2", 0),
            operator=data.get("operator", ""),
            expression=data.get("expression", ""),
        )


# ──────────────────────────────────────────────────────────────────────────────
# Branch (series elements — all must conduct)
# ──────────────────────────────────────────────────────────────────────────────
class Branch:
    def __init__(self, elements: Optional[List[Element]] = None):
        self.elements: List[Element] = elements or []
        self.power_flow = False

    def evaluate(self, interp: "LadderInterpreter", incoming_power: bool) -> bool:
        """Evaluate series contacts. All must pass for branch to pass."""
        if not incoming_power:
            self.power_flow = False
            return False

        self.power_flow = True
        for elem in self.elements:
            result = interp.evaluate_contact(elem)
            if not result:
                self.power_flow = False
                return False
        return True

    def to_dict(self) -> dict:
        return {"elements": [e.to_dict() for e in self.elements]}

    @classmethod
    def from_dict(cls, data: dict) -> "Branch":
        elements = [Element.from_dict(e) for e in data.get("elements", [])]
        return cls(elements)


# ──────────────────────────────────────────────────────────────────────────────
# Rung (inline elements + parallel branches + outputs)
# ──────────────────────────────────────────────────────────────────────────────
class Rung:
    def __init__(self, number: int, comment: str = ""):
        self.number = number
        self.comment = comment
        self.inline_branch = Branch()       # Main series path (left → right)
        self.parallel_branches: List[Branch] = []  # OR'd parallel paths
        self.outputs: List[Element] = []     # Output section (right rail)
        self.power_flow = False

    def evaluate(self, interp: "LadderInterpreter") -> bool:
        # Phase 1: Evaluate inline series branch (empty = transparent pass)
        inline_ok = self.inline_branch.evaluate(interp, True)

        # Phase 2: Evaluate parallel branches
        if self.parallel_branches:
            # Parallel branches are in series with inline (AND)
            # Within the parallel block, any branch can conduct (OR)
            parallel_ok = any(branch.evaluate(interp, inline_ok)
                             for branch in self.parallel_branches)
            self.power_flow = inline_ok and parallel_ok
        else:
            # No parallel branches — inline is the whole path
            self.power_flow = inline_ok

        # Phase 3: Execute outputs if power flows
        if self.power_flow:
            for elem in self.outputs:
                interp.execute_output(elem, True)

        # Phase 4: Execute outputs that need False condition (for OTE clearing)
        if not self.power_flow:
            for elem in self.outputs:
                interp.execute_output(elem, False)

        return self.power_flow

    def to_dict(self) -> dict:
        d = {
            "number": self.number,
            "comment": self.comment,
            "inline": self.inline_branch.to_dict(),
        }
        if self.parallel_branches:
            d["parallel"] = [b.to_dict() for b in self.parallel_branches]
        if self.outputs:
            d["outputs"] = [e.to_dict() for e in self.outputs]
        return d

    @classmethod
    def from_dict(cls, data: dict) -> "Rung":
        rung = cls(
            number=data.get("number", 0),
            comment=data.get("comment", ""),
        )
        rung.inline_branch = Branch.from_dict(data.get("inline", {}))
        rung.parallel_branches = [
            Branch.from_dict(b) for b in data.get("parallel", [])
        ]
        rung.outputs = [
            Element.from_dict(e) for e in data.get("outputs", [])
        ]
        return rung


# ──────────────────────────────────────────────────────────────────────────────
# Network (container for rungs)
# ──────────────────────────────────────────────────────────────────────────────
class Network:
    def __init__(self, number: int, comment: str = ""):
        self.number = number
        self.comment = comment
        self.rungs: List[Rung] = []
        self.power_flow = False

    def evaluate(self, interp: "LadderInterpreter") -> bool:
        self.power_flow = False
        for rung in self.rungs:
            if rung.evaluate(interp):
                self.power_flow = True
        return self.power_flow

    def to_dict(self) -> dict:
        return {
            "number": self.number,
            "comment": self.comment,
            "rungs": [r.to_dict() for r in self.rungs],
        }

    @classmethod
    def from_dict(cls, data: dict) -> "Network":
        net = cls(
            number=data.get("number", 0),
            comment=data.get("comment", ""),
        )
        net.rungs = [Rung.from_dict(r) for r in data.get("rungs", [])]
        return net


# ──────────────────────────────────────────────────────────────────────────────
# Ladder Interpreter
# ──────────────────────────────────────────────────────────────────────────────
class LadderInterpreter:
    def __init__(self, scan_cycle: float = 0.05):
        # Digital I/O
        self.digital_inputs: Dict[str, bool] = {}
        self.digital_outputs: Dict[str, bool] = {}

        # Memory
        self.memory_bool: Dict[str, bool] = {}
        self.memory_int: Dict[str, int] = {}
        self.memory_real: Dict[str, float] = {}

        # Timers & Counters
        self.timers: Dict[str, TimerState] = {}
        self.counters: Dict[str, CounterState] = {}

        # Program structure
        self.networks: List[Network] = []

        # ── OpenPLC-style State Machine ──────────────────────────────
        self.plc_state = PLCState.INIT  # Start in INIT after construction
        self._state_lock = threading.Lock()
        self._error_message = ""

        # ── OpenPLC-style Scan Timing ────────────────────────────────
        self.scan_cycle = scan_cycle  # Target scan period (seconds)
        self._scan_stats = ScanStats()
        self._cycle_base_time = 0.0    # CLOCK_MONOTONIC base
        self._cycle_tick = 0           # Current tick counter
        self._tick_counter = 0         # Total ticks for debug correlation
        self._last_cycle_start = 0.0
        self._mcr_state = False        # Master Control Reset state

        # Execution state (legacy compat)
        self.running = False
        self._jump_target: Optional[str] = None
        self._jump_active = False

    # ── Memory read ──────────────────────────────────────────────────────

    def set_state(self, new_state: str) -> bool:
        """Set PLC state with mutex protection (OpenPLC pattern).

        Returns True if transition succeeded, False if invalid.
        """
        with self._state_lock:
            if not PLCState.can_transition(self.plc_state, new_state):
                self._error_message = (
                    f"Invalid state transition: {self.plc_state} → {new_state}"
                )
                return False

            old_state = self.plc_state
            self.plc_state = new_state

            # Side effects on transition
            if new_state == PLCState.RUNNING:
                self.running = True
                if self._scan_stats.scan_count == 0:
                    # Initialize timing base on first start
                    self._cycle_base_time = time.monotonic()
                    self._last_cycle_start = self._cycle_base_time
            elif new_state == PLCState.STOPPED:
                self.running = False
            elif new_state == PLCState.ERROR:
                self.running = False

            return True

    def get_state(self) -> str:
        """Get current PLC state (thread-safe)."""
        with self._state_lock:
            return self.plc_state

    def get_stats(self) -> dict:
        """Get scan timing statistics (thread-safe)."""
        with self._state_lock:
            return self._scan_stats.format_stats()

    def _scan_timing_start(self):
        """Mark scan cycle start (OpenPLC: scan_cycle_time_start)."""
        now = time.monotonic()
        self._scan_stats.scan_count += 1
        self._cycle_tick += 1

        # Calculate expected periodic start time
        expected_start = self._cycle_base_time + (self._cycle_tick * self.scan_cycle)
        latency = now - expected_start

        # Track overrun (when we miss the expected start)
        if latency > 0:
            self._scan_stats.overrun_count += 1

        # Update cycle latency EMA
        if self._scan_stats.scan_count == 1:
            self._scan_stats.cycle_latency_avg = abs(latency)
        else:
            self._scan_stats.cycle_latency_avg = (
                self._scan_stats.ema_alpha * abs(latency) +
                (1 - self._scan_stats.ema_alpha) * self._scan_stats.cycle_latency_avg
            )

        # Calculate cycle time (time between consecutive cycle starts)
        if self._last_cycle_start > 0:
            cycle_time = now - self._last_cycle_start
            self._scan_stats.cycle_time_min = min(
                self._scan_stats.cycle_time_min, cycle_time
            )
            self._scan_stats.cycle_time_max = max(
                self._scan_stats.cycle_time_max, cycle_time
            )
            if self._scan_stats.scan_count == 2:
                self._scan_stats.cycle_time_avg = cycle_time
            else:
                self._scan_stats.cycle_time_avg = (
                    self._scan_stats.ema_alpha * cycle_time +
                    (1 - self._scan_stats.ema_alpha) * self._scan_stats.cycle_time_avg
                )

        self._last_cycle_start = now
        self._scan_time_start = now

    def _scan_timing_end(self):
        """Mark scan cycle end (OpenPLC: scan_cycle_time_end)."""
        scan_time = time.monotonic() - self._scan_time_start
        self._scan_stats.scan_time_min = min(
            self._scan_stats.scan_time_min, scan_time
        )
        self._scan_stats.scan_time_max = max(
            self._scan_stats.scan_time_max, scan_time
        )
        if self._scan_stats.scan_count == 1:
            self._scan_stats.scan_time_avg = scan_time
        else:
            self._scan_stats.scan_time_avg = (
                self._scan_stats.ema_alpha * scan_time +
                (1 - self._scan_stats.ema_alpha) * self._scan_stats.scan_time_avg
            )
        self._tick_counter += 1

    def read_bool(self, tag: str) -> bool:
        """Read a boolean value from any address space."""
        if tag.startswith('I'):
            return self.digital_inputs.get(tag, False)
        elif tag.startswith('Q'):
            return self.digital_outputs.get(tag, False)
        elif tag.startswith('M'):
            return self.memory_bool.get(tag, False)
        # Timer bits: T0.DN, T0.TT, T0.EN
        elif tag.startswith('T'):
            parts = tag.split('.')
            timer_tag = parts[0]
            bit = parts[1] if len(parts) > 1 else 'DN'
            if timer_tag in self.timers:
                t = self.timers[timer_tag]
                return getattr(t, bit.lower(), False)
            return False
        # Counter bits: C0.DN, C0.EN
        elif tag.startswith('C'):
            parts = tag.split('.')
            counter_tag = parts[0]
            bit = parts[1] if len(parts) > 1 else 'DN'
            if counter_tag in self.counters:
                c = self.counters[counter_tag]
                return getattr(c, bit.lower(), False)
            return False
        return False

    def read_int(self, tag: str) -> int:
        """Read an integer value."""
        if tag.startswith('N'):
            return self.memory_int.get(tag, 0)
        # Handle direct integer literals
        try:
            return int(tag)
        except ValueError:
            pass
        return self.memory_int.get(tag, 0)

    def read_real(self, tag: str) -> float:
        """Read a real (float) value."""
        if tag.startswith('R'):
            return self.memory_real.get(tag, 0.0)
        try:
            return float(tag)
        except ValueError:
            pass
        return self.memory_real.get(tag, 0.0)

    def read_value(self, tag: str) -> float:
        """Read a value as float (for comparisons/math)."""
        # Try integer first
        if tag.startswith('N'):
            return float(self.memory_int.get(tag, 0))
        # Try real
        if tag.startswith('R'):
            return self.memory_real.get(tag, 0.0)
        # Try I/Q/M as 0/1
        if self.read_bool(tag):
            return 1.0
        # Try literal
        try:
            return float(tag)
        except ValueError:
            pass
        return 0.0

    # ── Memory write ─────────────────────────────────────────────────────

    def write_bool(self, tag: str, value: bool):
        if tag.startswith('Q'):
            self.digital_outputs[tag] = value
        elif tag.startswith('M'):
            self.memory_bool[tag] = value
        # Error flags (_ERR suffix) go to memory_bool regardless of prefix
        elif tag.endswith('_ERR'):
            self.memory_bool[tag] = value

    def write_int(self, tag: str, value: int):
        if tag.startswith('N'):
            self.memory_int[tag] = value

    def write_real(self, tag: str, value: float):
        if tag.startswith('R'):
            self.memory_real[tag] = value

    # ── Contact evaluation ───────────────────────────────────────────────

    def evaluate_contact(self, elem: Element) -> bool:
        """Evaluate a single contact element."""
        t = elem.type

        # Boolean contacts
        if t == Contact.XIC:
            return self.read_bool(elem.tag)
        elif t == Contact.XIO:
            return not self.read_bool(elem.tag)

        # Edge contacts
        elif t == Contact.RISING:
            current = self.read_bool(elem.tag)
            result = current and not elem.last_state
            elem.last_state = current
            return result
        elif t == Contact.FALLING:
            current = self.read_bool(elem.tag)
            result = (not current) and elem.last_state
            elem.last_state = current
            return result

        # Timer bit contacts
        elif t == Contact.TON_DN or t == Contact.TOF_DN or t == Contact.TP_DN:
            return self.timers.get(elem.tag, TimerState(elem.tag)).dn
        elif t == Contact.TON_TT or t == Contact.TOF_TT or t == Contact.TP_TT:
            return self.timers.get(elem.tag, TimerState(elem.tag)).tt
        elif t == Contact.TON_EN or t == Contact.TOF_EN or t == Contact.TP_EN:
            return self.timers.get(elem.tag, TimerState(elem.tag)).en

        # Counter bit contacts
        elif t == Contact.CTU_DN or t == Contact.CTD_DN:
            return self.counters.get(elem.tag, CounterState(elem.tag)).dn
        elif t == Contact.CTU_EN or t == Contact.CTD_EN:
            return self.counters.get(elem.tag, CounterState(elem.tag)).en

        # Comparison contacts
        elif t == Compare.EQU:
            return self.read_value(elem.tag) == self.read_value(str(elem.value))
        elif t == Compare.NEQ:
            return self.read_value(elem.tag) != self.read_value(str(elem.value))
        elif t == Compare.GRT:
            return self.read_value(elem.tag) > self.read_value(str(elem.value))
        elif t == Compare.LES:
            return self.read_value(elem.tag) < self.read_value(str(elem.value))
        elif t == Compare.GEQ:
            return self.read_value(elem.tag) >= self.read_value(str(elem.value))
        elif t == Compare.LEQ:
            return self.read_value(elem.tag) <= self.read_value(str(elem.value))
        elif t == Compare.LIM:
            val = self.read_value(elem.tag)
            low = self.read_value(str(elem.value))
            high = self.read_value(str(elem.value2))
            return low <= val <= high
        elif t == Compare.CMP:
            val = self.read_value(elem.tag)
            cmp_val = self.read_value(str(elem.value))
            op = elem.operator or ">"
            ops = {
                "=": val == cmp_val, "==": val == cmp_val,
                ">": val > cmp_val, "<": val < cmp_val,
                ">=": val >= cmp_val, "<=": val <= cmp_val,
                "!=": val != cmp_val, "<>": val != cmp_val,
            }
            return ops.get(op, False)

        return False

    # ── Output execution ─────────────────────────────────────────────────

    def execute_output(self, elem: Element, power: bool):
        """Execute a single output element."""
        t = elem.type

        # Boolean outputs
        if t == Output.OTE:
            self.write_bool(elem.tag, power)

        elif t == Output.OTL:
            if power:
                self.write_bool(elem.tag, True)

        elif t == Output.OTU:
            if power:
                self.write_bool(elem.tag, False)

        # One-shot outputs
        elif t == Output.OSR:
            rising = power and not elem.last_state
            self.write_bool(elem.tag, rising)
            elem.last_state = power

        elif t == Output.OSF:
            falling = (not power) and elem.last_state
            self.write_bool(elem.tag, falling)
            elem.last_state = power

        # Inverted output coil
        elif t == Output.OTN:
            self.write_bool(elem.tag, not power)

        # Timers — create/set EN only; actual update happens after rung eval
        elif t in (Timer.TON, Timer.TOF, Timer.TP, Timer.TONR):
            if elem.tag not in self.timers:
                self.timers[elem.tag] = TimerState(elem.tag, int(elem.value))
            ts = self.timers[elem.tag]
            ts.preset = int(elem.value)
            ts.type = t
            ts.en = power

        # Counters — create/set EN only; actual update happens after rung eval
        elif t in (Counter.CTU, Counter.CTD, Counter.CTUD):
            if elem.tag not in self.counters:
                self.counters[elem.tag] = CounterState(elem.tag, int(elem.value))
            cs = self.counters[elem.tag]
            cs.preset = int(elem.value)
            cs.type = t
            cs.en = power
            # CTUD: value2 is count-down enable tag
            if t == Counter.CTUD and hasattr(elem, 'value2') and elem.value2:
                cs._cd_en = self.read_bool(str(elem.value2))

        # Math
        elif t == Math.ADD:
            a = self.read_value(elem.tag)
            b = self.read_value(str(elem.value))
            self._write_dest(elem.expression if elem.expression else str(elem.value2), a + b)
        elif t == Math.SUB:
            a = self.read_value(elem.tag)
            b = self.read_value(str(elem.value))
            self._write_dest(elem.expression if elem.expression else str(elem.value2), a - b)
        elif t == Math.MUL:
            a = self.read_value(elem.tag)
            b = self.read_value(str(elem.value))
            self._write_dest(elem.expression if elem.expression else str(elem.value2), a * b)
        elif t == Math.DIV:
            a = self.read_value(elem.tag)
            b = self.read_value(str(elem.value))
            dest = elem.expression if elem.expression else str(elem.value2)
            timer_tag = elem.tag
            if b == 0:
                # Set error flag on DIV-by-zero
                self.write_bool(f"{dest}_ERR", True)
            else:
                self.write_bool(f"{dest}_ERR", False)
                self._write_dest(dest, a / b)
        elif t == Math.SQT:
            a = self.read_value(elem.tag)
            dest = elem.expression if elem.expression else str(elem.value)
            if a < 0:
                self.write_bool(f"{dest}_ERR", True)
            else:
                self.write_bool(f"{dest}_ERR", False)
                self._write_dest(dest, math.sqrt(a))
        elif t == Math.CPT:
            # CPT: evaluate expression string
            self._execute_cpt(elem.expression)
        elif t == Math.NEG:
            # Negate: one's complement (invert all bits)
            a = self.read_value(elem.tag)
            self._write_dest(elem.expression if elem.expression else str(elem.value), ~int(a))
        elif t == Math.ABS:
            # Absolute value
            a = self.read_value(elem.tag)
            self._write_dest(elem.expression if elem.expression else str(elem.value), abs(a))
        elif t == Math.INC:
            # Increment: read-modify-write
            dest = elem.expression if elem.expression else str(elem.tag)
            a = self.read_value(dest)
            self._write_dest(dest, a + 1)
        elif t == Math.DEC:
            # Decrement: read-modify-write
            dest = elem.expression if elem.expression else str(elem.tag)
            a = self.read_value(dest)
            self._write_dest(dest, a - 1)

        # Data
        elif t == Data.MOV:
            val = self.read_value(elem.tag)
            self._write_dest(str(elem.value), val)
        elif t == Data.CPD:
            # Copy data block (copy multiple consecutive memory locations)
            src_base = elem.tag
            dst_base = str(elem.value)
            count = int(elem.value2) if elem.value2 else 1
            for i in range(count):
                src = f"{src_base}.{i}" if '.' in src_base else f"{src_base}{i}"
                dst = f"{dst_base}.{i}" if '.' in dst_base else f"{dst_base}{i}"
                val = self.read_value(src)
                self._write_dest(dst, val)
        elif t == Data.SCD:
            # Scale: linear conversion (input range → output range)
            # elem.tag = source value, elem.expression = dest
            # elem.value2 = "in_min,in_max,out_min,out_max"
            src_val = self.read_value(elem.tag)
            dest = elem.expression if elem.expression else str(elem.value)
            try:
                parts = str(elem.value2).split(',')
                in_min, in_max = float(parts[0]), float(parts[1])
                out_min, out_max = float(parts[2]), float(parts[3])
                if in_max != in_min:
                    scaled = out_min + ((src_val - in_min) / (in_max - in_min)) * (out_max - out_min)
                else:
                    scaled = out_min
                self._write_dest(dest, scaled)
            except (ValueError, IndexError):
                pass  # Silently ignore bad scale params

        # Boolean Logic
        elif t == Logic.AND:
            a = self.read_bool(elem.tag)
            b = self.read_bool(str(elem.value))
            self.write_bool(str(elem.value2) if elem.value2 else "M0", a and b)
        elif t == Logic.OR:
            a = self.read_bool(elem.tag)
            b = self.read_bool(str(elem.value))
            self.write_bool(str(elem.value2) if elem.value2 else "M0", a or b)
        elif t == Logic.XOR:
            a = self.read_bool(elem.tag)
            b = self.read_bool(str(elem.value))
            self.write_bool(str(elem.value2) if elem.value2 else "M0", a ^ b)

        # Control
        elif t == Control.RES:
            # Reset timer or counter
            tag = elem.tag
            if tag in self.timers:
                self.reset_timer(tag)
            if tag in self.counters:
                self.reset_counter(tag)

        elif t == Control.JMP:
            if power:
                self._jump_target = elem.tag  # Label name
                self._jump_active = True

        elif t == Control.MCR:
            # Master Control Reset: conditional execution region
            # When MCR is energized, outputs in the region are enabled
            # When MCR is de-energized, OTE/OTL/OTU outputs are forced off
            self._mcr_state = power

        # LBL is a marker, not executed as output — handled in scan()

    def _write_dest(self, dest: str, value: float):
        """Write a value to a destination tag."""
        if dest.startswith('N'):
            self.memory_int[dest] = int(round(value))
        elif dest.startswith('R'):
            self.memory_real[dest] = value
        elif dest.startswith('Q'):
            self.digital_outputs[dest] = value != 0
        elif dest.startswith('M'):
            self.memory_bool[dest] = value != 0

    def _execute_cpt(self, expression: str):
        """Execute a CPT compute expression.
        
        Format: "Dest = expression"
        Supports: +, -, *, /, %, sqrt(), abs(), sin(), cos(), tan(),
                  min(), max(), int(), float(), and variable references.
        """
        if not expression or '=' not in expression:
            return

        try:
            parts = expression.split('=', 1)
            dest = parts[0].strip()
            expr = parts[1].strip()

            # Build safe evaluation context
            namespace = {
                'sqrt': math.sqrt,
                'abs': abs,
                'sin': math.sin,
                'cos': math.cos,
                'tan': math.tan,
                'min': min,
                'max': max,
                'int': int,
                'float': float,
                'round': round,
            }

            # Replace N*, R* references with values
            import re
            def replace_var(match):
                tag = match.group(0)
                if tag.startswith('N'):
                    return str(self.memory_int.get(tag, 0))
                elif tag.startswith('R'):
                    return str(self.memory_real.get(tag, 0.0))
                elif tag.startswith('I') or tag.startswith('Q') or tag.startswith('M'):
                    return str(1 if self.read_bool(tag) else 0)
                return tag

            expr = re.sub(r'[NIRM]\d+', replace_var, expr)

            result = eval(expr, {"__builtins__": {}}, namespace)
            self._write_dest(dest, float(result))
        except Exception:
            pass  # Silently ignore CPT errors

    # ── Timer logic ──────────────────────────────────────────────────────

    def update_timer(self, tag: str, enabled: bool, preset_ms: int, timer_type: str):
        if tag not in self.timers:
            self.timers[tag] = TimerState(tag, preset_ms)

        t = self.timers[tag]
        t.preset = preset_ms
        t.type = timer_type
        t.en = enabled

        if timer_type == Timer.TON:
            if enabled:
                t.elapsed += self.scan_cycle * 1000
                t.tt = True
                t.dn = t.elapsed >= t.preset
            else:
                t.elapsed = 0
                t.tt = False
                t.dn = False

        elif timer_type == Timer.TOF:
            if enabled:
                t.elapsed = 0
                t.tt = False
                t.dn = True
            else:
                if t.elapsed < t.preset:
                    t.elapsed += self.scan_cycle * 1000
                    t.tt = True
                    t.dn = True
                else:
                    t.tt = False
                    t.dn = False

        elif timer_type == Timer.TP:
            rising = enabled and not t._last_enabled
            t._last_enabled = enabled

            if rising:
                t.elapsed = 0
                t.tt = True
                t.dn = True

            if t.tt:
                t.elapsed += self.scan_cycle * 1000
                if t.elapsed >= t.preset:
                    t.tt = False
                    t.dn = False

    def reset_timer(self, tag: str):
        if tag in self.timers:
            t = self.timers[tag]
            t.elapsed = 0
            t.tt = False
            t.dn = False
            t.en = False
            t._last_enabled = False

    # ── Counter logic ────────────────────────────────────────────────────

    def update_counter(self, tag: str, enabled: bool, preset: int, counter_type: str):
        if tag not in self.counters:
            self.counters[tag] = CounterState(tag, preset)

        c = self.counters[tag]
        c.preset = preset
        c.type = counter_type
        c.en = enabled

        rising_edge = enabled and not c._last_enabled
        c._last_enabled = enabled

        if counter_type == Counter.CTU:
            if rising_edge:
                c.acc += 1
            c.dn = c.acc >= c.preset

        elif counter_type == Counter.CTD:
            if rising_edge:
                c.acc -= 1
            c.dn = c.acc <= c.preset

    def reset_counter(self, tag: str):
        if tag in self.counters:
            c = self.counters[tag]
            c.acc = 0
            c.dn = False
            c.en = False
            c._last_enabled = False

    # ── CLI / test helpers ───────────────────────────────────────────────

    def set_input(self, tag: str, value: bool):
        """Set a digital input value (for testing)."""
        self.digital_inputs[tag] = value

    def get_output(self, tag: str) -> bool:
        """Get a digital output value (for testing)."""
        return self.digital_outputs.get(tag, False)

    # ── Program management ───────────────────────────────────────────────

    def add_network(self, number: int, comment: str = "") -> Network:
        net = Network(number, comment)
        self.networks.append(net)
        return net

    # ── Execution ────────────────────────────────────────────────────────

    def scan(self):
        """Execute one PLC scan cycle with OpenPLC-style timing."""
        self._jump_target = None
        self._jump_active = False

        # OpenPLC: scan_cycle_time_start()
        self._scan_timing_start()

        try:
            # Phase 1: Evaluate all rungs
            for net in self.networks:
                # Check for jump target
                if self._jump_active:
                    if self._jump_target == str(net.number):
                        self._jump_active = False
                    continue

                for rung in net.rungs:
                    # Check for jump to label
                    if self._jump_active:
                        if self._jump_target == str(rung.number):
                            self._jump_active = False
                        continue

                    rung.evaluate(self)

            # Phase 2: Update timers (only after rung evaluation completes)
            self._update_timers()

            # Phase 3: Update counters (only after rung evaluation completes)
            self._update_counters()

        except Exception as e:
            self._error_message = str(e)
            self.set_state(PLCState.ERROR)
        finally:
            # OpenPLC: scan_cycle_time_end()
            self._scan_timing_end()

    def _update_timers(self):
        """Update all timers based on their EN state from rung evaluation."""
        for tag, t in self.timers.items():
            if t.type == Timer.TON:
                if t.en:
                    t.elapsed += self.scan_cycle * 1000
                    t.tt = True
                    t.dn = t.elapsed >= t.preset
                    t.err = False
                    t.no = True
                else:
                    t.elapsed = 0
                    t.tt = False
                    t.dn = False
            elif t.type == Timer.TOF:
                if t.en:
                    t.elapsed = 0
                    t.tt = False
                    t.dn = True
                    t.err = False
                    t.no = True
                else:
                    if t.elapsed < t.preset:
                        t.elapsed += self.scan_cycle * 1000
                        t.tt = True
                        t.dn = True
                    else:
                        t.tt = False
                        t.dn = False
            elif t.type == Timer.TP:
                rising = t.en and not t._last_enabled
                t._last_enabled = t.en
                if rising:
                    t.elapsed = 0
                    t.tt = True
                    t.dn = True
                if t.tt:
                    t.elapsed += self.scan_cycle * 1000
                    if t.elapsed >= t.preset:
                        t.tt = False
                        t.dn = False
            elif t.type == Timer.TONR:
                # Retentive On-Delay: accumulates even when de-energized
                # Only resets via RES instruction
                if t.en:
                    t.elapsed += self.scan_cycle * 1000
                    t.tt = True
                    t.dn = t.elapsed >= t.preset
                # When not enabled, elapsed is PRESERVED (retentive)
                # DN stays True once reached until reset
                t.err = t.elapsed > 32767  # Overflow check
                t.no = not t.err

    def _update_counters(self):
        """Update all counters based on their EN state from rung evaluation."""
        for tag, c in self.counters.items():
            rising_cu = c.en and not c._last_enabled
            c._last_enabled = c.en
            if c.type == Counter.CTU:
                if rising_cu and c.acc < c.preset:
                    c.acc += 1
                c.dn = c.acc >= c.preset
                c.err = False
                c.no = True
            elif c.type == Counter.CTD:
                if rising_cu and c.acc > 0:
                    c.acc -= 1
                c.dn = c.acc <= 0
                c.err = False
                c.no = True
            elif c.type == Counter.CTUD:
                # Count Up/Down: cu on rising edge of .en, cd on rising edge of ._cd_en
                rising_cd = c._cd_en and not getattr(c, '_last_cd_en', False)
                c._last_cd_en = c._cd_en
                if rising_cu and c.acc < c.preset:
                    c.acc += 1
                if rising_cd and c.acc > 0:
                    c.acc -= 1
                c.dn = c.acc >= c.preset
                c.err = c.acc < 0 or c.acc > 32767
                c.no = not c.err

    def run(self, duration: float = 1.0):
        """Run simulation for a specified duration."""
        self.running = True
        start = time.time()
        while time.time() - start < duration:
            self.scan()
            time.sleep(self.scan_cycle)
        self.running = False

    def stop(self):
        self.running = False

    # ── Debug ────────────────────────────────────────────────────────────

    def print_state(self):
        print("=== PLC State ===")
        for tag, val in sorted(self.digital_outputs.items()):
            print(f"  {tag}: {val}")
        for tag, val in sorted(self.memory_bool.items()):
            print(f"  {tag}: {val}")
        for tag, val in sorted(self.memory_int.items()):
            print(f"  {tag}: {val}")
        for tag, val in sorted(self.memory_real.items()):
            print(f"  {tag}: {val:.4f}")
        for tag, t in self.timers.items():
            print(f"  {tag} [{t.type}]: EN={t.en} TT={t.tt} DN={t.dn} "
                  f"elapsed={t.elapsed:.0f}ms preset={t.preset}ms")
        for tag, c in self.counters.items():
            print(f"  {tag} [{c.type}]: EN={c.en} DN={c.dn} "
                  f"ACC={c.acc} PRE={c.preset}")


# ──────────────────────────────────────────────────────────────────────────────
# Tests
# ──────────────────────────────────────────────────────────────────────────────

def test_xic_xio_contacts():
    """Test XIC (NO) and XIO (NC) contacts."""
    print("\n" + "=" * 60)
    print("TEST: XIC/XIO Contacts")
    print("=" * 60)

    interp = LadderInterpreter()
    net = interp.add_network(0)
    rung = Rung(0, "Test XIC/XIO")

    # Rung: XIC(I0) → OTE(Q0)
    rung.inline_branch.elements.append(Element(Contact.XIC, "I0"))
    rung.outputs.append(Element(Output.OTE, "Q0"))
    net.rungs.append(rung)

    interp.set_input("I0", False)
    interp.scan()
    assert interp.get_output("Q0") is False
    print("  ✓ XIC: I0=False → Q0=False")

    interp.set_input("I0", True)
    interp.scan()
    assert interp.get_output("Q0") is True
    print("  ✓ XIC: I0=True → Q0=True")

    # XIO test
    net2 = interp.add_network(1)
    rung2 = Rung(1, "Test XIO")
    rung2.inline_branch.elements.append(Element(Contact.XIO, "I0"))
    rung2.outputs.append(Element(Output.OTE, "Q1"))
    net2.rungs.append(rung2)

    interp.set_input("I0", True)
    interp.scan()
    assert interp.get_output("Q1") is False
    print("  ✓ XIO: I0=True → Q1=False")

    interp.set_input("I0", False)
    interp.scan()
    assert interp.get_output("Q1") is True
    print("  ✓ XIO: I0=False → Q1=True")

    print("  ✅ XIC/XIO contacts PASSED\n")


def test_latch_unlatch():
    """Test OTL (Latch) and OTU (Unlatch)."""
    print("\n" + "=" * 60)
    print("TEST: OTL/OTU (Latch/Unlatch)")
    print("=" * 60)

    interp = LadderInterpreter()

    # Network 0: Latch Q0 with I0
    net0 = interp.add_network(0)
    r0 = Rung(0)
    r0.inline_branch.elements.append(Element(Contact.XIC, "I0"))
    r0.outputs.append(Element(Output.OTL, "Q0"))
    net0.rungs.append(r0)

    # Network 1: Unlatch Q0 with I1
    net1 = interp.add_network(1)
    r1 = Rung(1)
    r1.inline_branch.elements.append(Element(Contact.XIC, "I1"))
    r1.outputs.append(Element(Output.OTU, "Q0"))
    net1.rungs.append(r1)

    interp.set_input("I0", False)
    interp.set_input("I1", False)
    interp.scan()
    assert interp.get_output("Q0") is False
    print("  ✓ Q0 starts OFF")

    # Latch
    interp.set_input("I0", True)
    interp.scan()
    assert interp.get_output("Q0") is True
    print("  ✓ I0 pulse → Q0 latched ON")

    # Release I0 — stays latched
    interp.set_input("I0", False)
    interp.scan()
    assert interp.get_output("Q0") is True
    print("  ✓ I0 released → Q0 still ON")

    # Unlatch
    interp.set_input("I1", True)
    interp.scan()
    assert interp.get_output("Q0") is False
    print("  ✓ I1 pulse → Q0 unlatched OFF")

    print("  ✅ OTL/OTU PASSED\n")


def test_parallel_branching():
    """Test parallel branching — OR logic."""
    print("\n" + "=" * 60)
    print("TEST: Parallel Branching")
    print("=" * 60)

    interp = LadderInterpreter()
    net = interp.add_network(0)
    rung = Rung(0, "Parallel I0 || I1 → Q0")

    # No inline elements — pure parallel
    rung.parallel_branches.append(Branch([Element(Contact.XIC, "I0")]))
    rung.parallel_branches.append(Branch([Element(Contact.XIC, "I1")]))
    rung.outputs.append(Element(Output.OTE, "Q0"))
    net.rungs.append(rung)

    interp.set_input("I0", False)
    interp.set_input("I1", False)
    interp.scan()
    assert interp.get_output("Q0") is False
    print("  ✓ I0=F, I1=F → Q0=F")

    interp.set_input("I0", True)
    interp.scan()
    assert interp.get_output("Q0") is True
    print("  ✓ I0=T, I1=F → Q0=T")

    interp.set_input("I0", False)
    interp.set_input("I1", True)
    interp.scan()
    assert interp.get_output("Q0") is True
    print("  ✓ I0=F, I1=T → Q0=T")

    print("  ✅ Parallel branching PASSED\n")


def test_timers():
    """Test TON timer."""
    print("\n" + "=" * 60)
    print("TEST: TON Timer")
    print("=" * 60)

    interp = LadderInterpreter()
    interp.scan_cycle = 0.05  # 50ms
    net = interp.add_network(0)
    rung = Rung(0, "TON test")

    rung.inline_branch.elements.append(Element(Contact.XIC, "I0"))
    rung.outputs.append(Element(Timer.TON, "T0", 100))  # 100ms preset
    net.rungs.append(rung)

    # Output rung
    net2 = interp.add_network(1)
    rung2 = Rung(1, "Timer done → Q0")
    rung2.inline_branch.elements.append(Element(Contact.TON_DN, "T0"))
    rung2.outputs.append(Element(Output.OTE, "Q0"))
    net2.rungs.append(rung2)

    interp.set_input("I0", True)

    # Scan 1: timer starts (50ms elapsed, not done)
    interp.scan()
    assert interp.get_output("Q0") is False
    print("  ✓ Scan 1: timer running, Q0=OFF")

    # Scan 2: timer still running (100ms elapsed, done after this scan)
    interp.scan()
    assert interp.get_output("Q0") is False
    print("  ✓ Scan 2: timer done, but Q0 still OFF (contact checked before update)")

    # Scan 3: T0.dn=True from previous update → Q0=ON
    interp.scan()
    assert interp.get_output("Q0") is True
    print("  ✓ Scan 3: Q0=ON (timer done)")

    # Release input — timer resets (takes one scan for reset to propagate)
    interp.set_input("I0", False)
    interp.scan()
    interp.scan()
    assert interp.get_output("Q0") is False
    print("  ✓ Input released → timer reset → Q0=OFF")

    print("  ✅ TON timer PASSED\n")


def test_counters():
    """Test CTU counter."""
    print("\n" + "=" * 60)
    print("TEST: CTU Counter")
    print("=" * 60)

    interp = LadderInterpreter()
    net = interp.add_network(0)
    rung = Rung(0, "CTU test")

    rung.inline_branch.elements.append(Element(Contact.RISING, "I0"))
    rung.outputs.append(Element(Counter.CTU, "C0", 3))  # preset 3
    net.rungs.append(rung)

    # Counter done → Q0
    net2 = interp.add_network(1)
    rung2 = Rung(1, "Counter done → Q0")
    rung2.inline_branch.elements.append(Element(Contact.CTU_DN, "C0"))
    rung2.outputs.append(Element(Output.OTE, "Q0"))
    net2.rungs.append(rung2)

    # Pulse I0 three times
    for i in range(3):
        interp.set_input("I0", True)
        interp.scan()
        interp.set_input("I0", False)
        interp.scan()

    assert interp.counters["C0"].acc == 3
    assert interp.get_output("Q0") is True
    print("  ✓ Counter reached preset (3) → Q0=ON")

    print("  ✅ CTU counter PASSED\n")


def test_comparisons():
    """Test comparison instructions."""
    print("\n" + "=" * 60)
    print("TEST: Comparisons")
    print("=" * 60)

    interp = LadderInterpreter()
    interp.memory_int["N0"] = 50

    net = interp.add_network(0)
    rung = Rung(0, "Comparison test")

    # GEQ(N0, 40) → Q0
    rung.inline_branch.elements.append(Element(Compare.GEQ, "N0", 40))
    rung.outputs.append(Element(Output.OTE, "Q0"))
    net.rungs.append(rung)

    interp.scan()
    assert interp.get_output("Q0") is True
    print("  ✓ GEQ(N0=50, 40) → True")

    # GRT(N0, 60) → Q1
    net2 = interp.add_network(1)
    rung2 = Rung(1, "GRT test")
    rung2.inline_branch.elements.append(Element(Compare.GRT, "N0", 60))
    rung2.outputs.append(Element(Output.OTE, "Q1"))
    net2.rungs.append(rung2)

    interp.scan()
    assert interp.get_output("Q1") is False
    print("  ✓ GRT(N0=50, 60) → False")

    # LIM(N0, 30, 70) → Q2
    net3 = interp.add_network(2)
    rung3 = Rung(2, "LIM test")
    rung3.inline_branch.elements.append(Element(Compare.LIM, "N0", 30, 70))
    rung3.outputs.append(Element(Output.OTE, "Q2"))
    net3.rungs.append(rung3)

    interp.scan()
    assert interp.get_output("Q2") is True
    print("  ✓ LIM(N0=50, 30, 70) → True")

    print("  ✅ Comparisons PASSED\n")


def test_math():
    """Test math instructions."""
    print("\n" + "=" * 60)
    print("TEST: Math Instructions")
    print("=" * 60)

    interp = LadderInterpreter()
    interp.memory_int["N0"] = 10
    interp.memory_int["N1"] = 3

    net = interp.add_network(0)
    rung = Rung(0, "ADD test")

    rung.inline_branch.elements.append(Element(Contact.XIC, "I0"))
    rung.outputs.append(Element(Math.ADD, "N0", "N1", "", "", "N2"))
    net.rungs.append(rung)

    interp.set_input("I0", True)
    interp.scan()
    assert interp.memory_int.get("N2", 0) == 13
    print("  ✓ ADD(N0=10, N1=3) → N2=13")

    # MUL test
    net2 = interp.add_network(1)
    rung2 = Rung(1, "MUL test")
    rung2.inline_branch.elements.append(Element(Contact.XIC, "I0"))
    rung2.outputs.append(Element(Math.MUL, "N0", "N1", "", "", "N3"))
    net2.rungs.append(rung2)

    interp.scan()
    assert interp.memory_int.get("N3", 0) == 30
    print("  ✓ MUL(N0=10, N1=3) → N3=30")

    print("  ✅ Math PASSED\n")


def test_mov():
    """Test MOV instruction."""
    print("\n" + "=" * 60)
    print("TEST: MOV Instruction")
    print("=" * 60)

    interp = LadderInterpreter()
    interp.memory_int["N0"] = 42

    net = interp.add_network(0)
    rung = Rung(0, "MOV test")
    rung.inline_branch.elements.append(Element(Contact.XIC, "I0"))
    rung.outputs.append(Element(Data.MOV, "N0", "N1"))
    net.rungs.append(rung)

    interp.set_input("I0", True)
    interp.scan()
    assert interp.memory_int.get("N1", 0) == 42
    print("  ✓ MOV(N0=42, N1) → N1=42")

    print("  ✅ MOV PASSED\n")


def test_rung_comment():
    """Test rung comment persistence."""
    print("\n" + "=" * 60)
    print("TEST: Rung Comments")
    print("=" * 60)

    interp = LadderInterpreter()
    net = interp.add_network(0, "Motor Control")
    rung = Rung(0, "Start/Stop with seal-in")
    rung.inline_branch.elements.append(Element(Contact.XIC, "I0"))
    rung.outputs.append(Element(Output.OTE, "Q0"))
    net.rungs.append(rung)

    assert net.comment == "Motor Control"
    assert rung.comment == "Start/Stop with seal-in"
    print("  ✓ Rung comment preserved")

    # Serialize/deserialize
    net_dict = net.to_dict()
    net2 = Network.from_dict(net_dict)
    assert net2.rungs[0].comment == "Start/Stop with seal-in"
    print("  ✓ Comment survives serialization")

    print("  ✅ Rung comments PASSED\n")


def test_tonr_retentive():
    """Test TONR (Retentive On-Delay) timer."""
    print("\n" + "=" * 60)
    print("TEST: TONR Retentive Timer")
    print("=" * 60)

    interp = LadderInterpreter(scan_cycle=0.1)
    net = interp.add_network(0)
    rung = Rung(0)
    rung.inline_branch.elements.append(Element(Contact.XIC, "I0"))
    rung.outputs.append(Element(Output.OTE, "T1", 500))  # placeholder
    net.rungs.append(rung)

    interp.set_state(PLCState.RUNNING)

    # Create TONR timer manually
    interp.timers["T1"] = TimerState("T1", 500)
    interp.timers["T1"].type = Timer.TONR
    interp.timers["T1"].en = True

    # Scan 5 times (500ms)
    for _ in range(5):
        interp.timers["T1"].en = True
        interp._update_timers()
    assert interp.timers["T1"].dn, "TONR should be done"
    assert interp.timers["T1"].elapsed >= 500
    print("  ✓ TONR done after 500ms")

    # Stop input - elapsed should be PRESERVED
    interp.timers["T1"].en = False
    interp._update_timers()
    elapsed_before = interp.timers["T1"].elapsed
    assert interp.timers["T1"].elapsed > 0, "TONR should preserve elapsed"
    print(f"  ✓ TONR preserved elapsed={elapsed_before}ms after de-energize")

    # Resume - timer continues from where it left off
    interp.timers["T1"].en = True
    interp._update_timers()
    assert interp.timers["T1"].elapsed >= elapsed_before
    print("  ✓ TONR resumed from preserved elapsed")

    # Reset via RES instruction
    interp.reset_timer("T1")
    assert interp.timers["T1"].elapsed == 0
    assert not interp.timers["T1"].dn
    print("  ✓ TONR reset via RES")

    print("  ✅ TONR timer PASSED\n")


def test_ctud_counter():
    """Test CTUD (Count Up/Down) counter."""
    print("\n" + "=" * 60)
    print("TEST: CTUD Counter")
    print("=" * 60)

    interp = LadderInterpreter(scan_cycle=0.1)
    net = interp.add_network(0)
    rung = Rung(0)
    rung.outputs.append(Element(Output.OTE, "C1", 3))  # CTUD preset=3
    net.rungs.append(rung)

    # Manually set counter type to CTUD
    if "C1" not in interp.counters:
        interp.counters["C1"] = CounterState("C1", 3)
    interp.counters["C1"].type = Counter.CTUD

    interp.set_state(PLCState.RUNNING)

    # Count up 3 times
    for i in range(3):
        interp.counters["C1"].en = True
        interp.counters["C1"]._last_enabled = False
        interp.counters["C1"]._cd_en = False
        interp._update_counters()
    assert interp.counters["C1"].acc == 3
    print(f"  ✓ CTUD count up to {interp.counters['C1'].acc}")

    # Count down 1 time
    interp.counters["C1"].en = False
    interp.counters["C1"]._cd_en = True
    interp.counters["C1"]._last_cd_en = False
    interp._update_counters()
    assert interp.counters["C1"].acc == 2
    print(f"  ✓ CTUD count down to {interp.counters['C1'].acc}")

    print("  ✅ CTUD counter PASSED\n")


def test_state_machine():
    """Test PLC state machine transitions."""
    print("\n" + "=" * 60)
    print("TEST: PLC State Machine")
    print("=" * 60)

    interp = LadderInterpreter()

    # Initial state is INIT
    assert interp.get_state() == PLCState.INIT
    print("  ✓ Initial state is INIT")

    # INIT -> RUNNING
    assert interp.set_state(PLCState.RUNNING)
    assert interp.get_state() == PLCState.RUNNING
    assert interp.running
    print("  ✓ INIT → RUNNING (running=True)")

    # RUNNING -> STOPPED
    assert interp.set_state(PLCState.STOPPED)
    assert interp.get_state() == PLCState.STOPPED
    assert not interp.running
    print("  ✓ RUNNING → STOPPED (running=False)")

    # STOPPED -> RUNNING
    assert interp.set_state(PLCState.RUNNING)
    assert interp.get_state() == PLCState.RUNNING
    print("  ✓ STOPPED → RUNNING")

    # RUNNING -> ERROR
    assert interp.set_state(PLCState.ERROR)
    assert interp.get_state() == PLCState.ERROR
    assert not interp.running
    print("  ✓ RUNNING → ERROR (running=False)")

    # ERROR -> STOPPED (not RUNNING directly)
    assert not interp.set_state(PLCState.RUNNING)  # Invalid!
    assert interp.get_state() == PLCState.ERROR
    print("  ✓ ERROR → RUNNING rejected (must go via STOPPED)")

    assert interp.set_state(PLCState.STOPPED)
    print("  ✓ ERROR → STOPPED")

    # STOPPED -> EMPTY
    assert interp.set_state(PLCState.EMPTY)
    print("  ✓ STOPPED → EMPTY")

    # EMPTY -> RUNNING (direct start)
    interp2 = LadderInterpreter()
    interp2.set_state(PLCState.EMPTY)
    assert interp2.set_state(PLCState.RUNNING)
    print("  ✓ EMPTY → RUNNING (direct start)")

    print("  ✅ State machine PASSED\n")


def test_scan_timing():
    """Test scan timing statistics."""
    print("\n" + "=" * 60)
    print("TEST: Scan Timing")
    print("=" * 60)

    interp = LadderInterpreter(scan_cycle=0.05)
    net = interp.add_network(0)
    rung = Rung(0)
    rung.inline_branch.elements.append(Element(Contact.XIC, "I0"))
    rung.outputs.append(Element(Output.OTE, "Q0"))
    net.rungs.append(rung)

    interp.set_state(PLCState.RUNNING)
    interp.digital_inputs["I0"] = True

    # Run several scans
    for _ in range(10):
        interp.scan()

    stats = interp.get_stats()
    assert stats["scan_count"] == 10
    print(f"  ✓ Scan count: {stats['scan_count']}")
    print(f"  ✓ Avg scan time: {stats['scan_time_ms']['avg']}ms")
    print(f"  ✓ Avg cycle time: {stats['cycle_time_ms']['avg']}ms")

    # Verify timing is reasonable (should be < 1ms for simple program)
    assert stats["scan_time_ms"]["avg"] < 10
    print("  ✓ Scan times reasonable")

    print("  ✅ Scan timing PASSED\n")


def test_new_instructions():
    """Test new AB instructions: OTN, NEG, ABS, INC, DEC, SCD, AND, OR, XOR."""
    print("\n" + "=" * 60)
    print("TEST: New AB Instructions")
    print("=" * 60)

    interp = LadderInterpreter(scan_cycle=0.1)
    interp.memory_int["N0"] = 42
    interp.memory_int["N1"] = -10
    interp.memory_real["R0"] = 0.0
    interp.digital_inputs["I0"] = True
    interp.digital_inputs["I1"] = False

    interp.set_state(PLCState.RUNNING)

    # OTN - inverted output coil
    net = interp.add_network(0)
    rung = Rung(0)
    rung.inline_branch.elements.append(Element(Contact.XIC, "I0"))  # True
    rung.outputs.append(Element(Output.OTN, "Q0"))
    net.rungs.append(rung)
    interp.scan()
    assert not interp.digital_outputs.get("Q0", True)  # Inverted!
    print("  ✓ OTN: True input → False output (inverted)")

    # NEG
    rung2 = Rung(1)
    rung2.outputs.append(Element(Output.OTE, "N2"))
    neg_elem = Element(Math.NEG, "N0")
    neg_elem.expression = "N2"
    rung2.outputs.append(neg_elem)
    net.rungs.append(rung2)
    interp.scan()
    assert interp.memory_int["N2"] == ~42  # One's complement
    print(f"  ✓ NEG(42) = {interp.memory_int['N2']}")

    # ABS
    abs_elem = Element(Math.ABS, "N1")
    abs_elem.expression = "N3"
    rung2 = Rung(2)
    rung2.outputs.append(abs_elem)
    net.rungs.append(rung2)
    interp.scan()
    assert interp.memory_int["N3"] == 10
    print(f"  ✓ ABS(-10) = {interp.memory_int['N3']}")

    # INC
    inc_elem = Element(Math.INC, "N0")
    inc_elem.expression = "N0"
    rung3 = Rung(3)
    rung3.outputs.append(inc_elem)
    net.rungs.append(rung3)
    old_n0 = interp.memory_int["N0"]
    interp.scan()
    assert interp.memory_int["N0"] == old_n0 + 1
    print(f"  ✓ INC({old_n0}) = {interp.memory_int['N0']}")

    # DEC - create fresh interpreter to avoid INC interference
    dec_interp = LadderInterpreter(scan_cycle=0.1)
    dec_interp.memory_int["N0"] = interp.memory_int["N0"]  # Start from INC result
    dec_interp.set_state(PLCState.RUNNING)
    dec_net = dec_interp.add_network(0)
    rung_dec = Rung(0)
    dec_elem = Element(Math.DEC, "N0")
    dec_elem.expression = "N0"
    rung_dec.outputs.append(dec_elem)
    dec_net.rungs.append(rung_dec)
    n0_before_dec = dec_interp.memory_int["N0"]
    dec_interp.scan()
    assert dec_interp.memory_int["N0"] == n0_before_dec - 1
    print(f"  ✓ DEC({n0_before_dec}) = {dec_interp.memory_int['N0']}")

    # SCD
    scd_elem = Element(Data.SCD, "R0")
    scd_elem.expression = "R1"
    scd_elem.value2 = "0,1023,0.0,100.0"  # Scale 0-1023 → 0-100
    rung5 = Rung(5)
    rung5.outputs.append(scd_elem)
    net.rungs.append(rung5)
    interp.scan()
    # R0 = 0.0, scaled should be 0.0
    assert abs(interp.memory_real["R1"] - 0.0) < 0.01
    interp.memory_real["R0"] = 511.5  # ~50%
    interp.scan()
    assert abs(interp.memory_real["R1"] - 50.0) < 1.0
    print(f"  ✓ SCD(511.5, 0-1023→0-100) ≈ {interp.memory_real['R1']}")

    # AND, OR, XOR logic
    and_elem = Element(Logic.AND, "I0")
    and_elem.value = "I1"
    and_elem.value2 = "M0"
    or_elem = Element(Logic.OR, "I0")
    or_elem.value = "I1"
    or_elem.value2 = "M1"
    xor_elem = Element(Logic.XOR, "I0")
    xor_elem.value = "I1"
    xor_elem.value2 = "M2"
    rung6 = Rung(6)
    rung6.outputs.extend([and_elem, or_elem, xor_elem])
    net.rungs.append(rung6)
    interp.scan()
    assert not interp.memory_bool["M0"]  # True AND False = False
    assert interp.memory_bool["M1"]      # True OR False = True
    assert interp.memory_bool["M2"]      # True XOR False = True
    print("  ✓ AND/OR/XOR logic correct")

    print("  ✅ New instructions PASSED\n")


def test_error_flags():
    """Test instruction error flags (ERR/ENO)."""
    print("\n" + "=" * 60)
    print("TEST: Instruction Error Flags")
    print("=" * 60)

    interp = LadderInterpreter(scan_cycle=0.1)
    interp.memory_int["N0"] = 100
    interp.memory_int["N1"] = 0

    interp.set_state(PLCState.RUNNING)

    # DIV by zero → should set error flag
    net = interp.add_network(0)
    rung = Rung(0)
    div_elem = Element(Math.DIV, "N0")
    div_elem.value = "N1"  # Dividing by 0
    div_elem.expression = "N2"
    rung.outputs.append(div_elem)
    net.rungs.append(rung)
    interp.scan()
    assert interp.memory_bool.get("N2_ERR", False), "DIV/0 should set ERR"
    print("  ✓ DIV by zero → ERR flag set")

    # SQRT of negative → should set error flag
    interp.memory_int["N0"] = -5
    rung2 = Rung(1)
    sqrt_elem = Element(Math.SQT, "N0")
    sqrt_elem.expression = "N3"
    rung2.outputs.append(sqrt_elem)
    net.rungs.append(rung2)
    interp.scan()
    assert interp.memory_bool.get("N3_ERR", False), "SQRT(neg) should set ERR"
    print("  ✓ SQRT(negative) → ERR flag set")

    print("  ✅ Error flags PASSED\n")


def run_all_tests():
    print("\n" + "=" * 60)
    print("MechBase PLC — Ladder Logic Engine v4 Test Suite")
    print("=" * 60)

    tests = [
        test_xic_xio_contacts,
        test_latch_unlatch,
        test_parallel_branching,
        test_timers,
        test_counters,
        test_comparisons,
        test_math,
        test_mov,
        test_rung_comment,
        # v4: New tests
        test_tonr_retentive,
        test_ctud_counter,
        test_state_machine,
        test_scan_timing,
        test_new_instructions,
        test_error_flags,
    ]

    passed = 0
    failed = 0
    for test in tests:
        try:
            test()
            passed += 1
        except Exception as e:
            print(f"  ❌ {test.__name__} FAILED: {e}\n")
            failed += 1

    print("=" * 60)
    print(f"Results: {passed} passed, {failed} failed out of {len(tests)}")
    print("=" * 60)
    return failed == 0


if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument("--quick", action="store_true", help="Quick test mode")
    args = parser.parse_args()

    success = run_all_tests()
    exit(0 if success else 1)