#!/usr/bin/env python3
"""
Render causal DAGs as SVG figures for the Causality Atlas.
Uses Graphviz for main diagrams and outputs Mermaid markdown blocks.

Usage:
    python3 scripts/atlas_rendering/render_dags.py
"""
import subprocess, json, os

OUTPUT_DIR = os.path.expanduser("~/Documents/Causality/outputs/figures")

dags = {
    "confounding_dag": {
        "graph": "digraph Confounding {\n  rankdir=LR;\n  C -> X;\n  C -> Y;\n  X -> Y;\n  labelloc=t;\n  label=\"Basic Confounding Structure\\n(C confounds X→Y)\";\n}",
        "mermaid": "graph TD\n  subgraph Basic Confounding\n    C --> X\n    C --> Y\n    X --> Y\n  end\n  style C fill:#ffcccc\n  style X fill:#ccffcc\n  style Y fill:#ccccff"
    },
    "back_door": {
        "graph": "digraph BackDoor {\n  rankdir=LR;\n  C -> X; C -> Y; X -> Y;\n  U [shape=box, style=dashed]; U -> X; U -> Y;\n  labelloc=t; label=\"Back-Door Path\\nX ← C → Y and X ← U → Y\";\n}",
        "mermaid": "graph LR\n  U[Unmeasured<br/>Confounder] -.-> X\n  U -.-> Y\n  C --> X\n  C --> Y\n  X --> Y\n  style U fill:#ffcccc,stroke-dasharray: 5 5"
    },
    "front_door": {
        "graph": "digraph FrontDoor {\n  rankdir=LR;\n  X -> M; M -> Y;\n  U [shape=box, style=dashed]; U -> X; U -> Y;\n  labelloc=t; label=\"Front-Door Path: X→M→Y\";\n}",
        "mermaid": "graph LR\n  X --> M\n  M --> Y\n  U[Unmeasured<br/>Confounder] -.-> X\n  U -.-> Y\n  style U fill:#ffcccc,stroke-dasharray:5 5"
    },
    "collider_dag": {
        "graph": "digraph Collider {\n  rankdir=LR;\n  X -> C; Y -> C;\n  labelloc=t; label=\"Collider Structure: X→C←Y\";\n}",
        "mermaid": "graph LR\n  X --> C\n  Y --> C\n  style C fill:#ffffcc"
    },
    "iv_dag": {
        "graph": "digraph IV {\n  rankdir=LR;\n  Z -> A; A -> Y;\n  U [shape=box, style=dashed]; U -> A; U -> Y;\n  labelloc=t; label=\"Instrumental Variable DAG\\nZ→A→Y (Exclusion restriction: no Z→Y)\";\n}",
        "mermaid": "graph LR\n  Z --> A\n  A --> Y\n  U[Confounder] -.-> A\n  U -.-> Y\n  style Z fill:#ccffcc\n  style U fill:#ffcccc,stroke-dasharray:5 5"
    },
    "mediation_dag": {
        "graph": "digraph Mediation {\n  rankdir=LR;\n  X -> M; X -> Y; M -> Y;\n  C [shape=box, style=dashed]; C -> M; C -> Y;\n  labelloc=t; label=\"Mediation: X→M→Y\\nDirect: X→Y, Indirect: X→M→Y\";\n}",
        "mermaid": "graph LR\n  X --> M\n  X --> Y\n  M --> Y\n  C[Confounder] -.-> M\n  C -.-> Y\n  style C fill:#ffcccc,stroke-dasharray:5 5"
    },
    "time_varying_dag": {
        "graph": "digraph TimeVarying {\n  rankdir=TB;\n  L0 -> A1; L0 -> L1; A1 -> L1;\n  L1 -> A2; L1 -> Y; A1 -> Y; A2 -> Y;\n  labelloc=t; label=\"Time-Varying Confounding\\n(Affected by prior treatment)\";\n}",
        "mermaid": "graph TD\n  L0[L₀] --> A1[A₁]\n  L0 --> L1[L₁]\n  A1 --> L1\n  L1 --> A2[A₂]\n  L1 --> Y\n  A1 --> Y\n  A2 --> Y\n  subgraph t=1\n    L0\n  end\n  subgraph t=2\n    A1; L1\n  end\n  subgraph t=3\n    A2; Y\n  end"
    },
    "missing_data_mcar": {
        "graph": "digraph MCAR {\n  rankdir=LR;\n  Y -> R [style=dashed, constraint=false];\n  labelloc=t; label=\"MCAR: Missingness (R) independent of Y\";\n}",
        "mermaid": "graph LR\n  Y[True Y] -.->|observed if| R[Response R=1]\n  style R fill:#ffffcc"
    },
    "causal_hierarchy": {
        "graph": "digraph CausalHierarchy {\n  rankdir=TB;\n  node [shape=box, style=rounded];\n  Counterfactuals [label=\"Rung 3\\nCounterfactuals\\n\\\"What if?\\\"\\nP(Y_x | X=x', Y=y)\"];\n  Intervention [label=\"Rung 2\\nIntervention\\n\\\"What if I do?\\\"\\nP(Y | do(X))\\n\"];\n  Association [label=\"Rung 1\\nAssociation\\n\\\"What if I see?\\\"\\nP(Y | X)\\n\"];\n  Association -> Intervention -> Counterfactuals;\n  labelloc=t; label=\"Pearl's Causal Hierarchy\";\n}",
        "mermaid": "graph TD\n  CF[<b>Rung 3: Counterfactuals</b><br/>P(Y_x | X=x', Y=y)<br/>'What if I had...?']\n  IN[<b>Rung 2: Intervention</b><br/>P(Y | do(X))<br/>'What if I do?']\n  AS[<b>Rung 1: Association</b><br/>P(Y | X)<br/>'What if I see?']\n  AS --> IN --> CF"
    },
    "scm_structure": {
        "graph": "digraph SCM {\n  rankdir=TB;\n  U1 [shape=diamond, label=\"U₁\"];\n  U2 [shape=diamond, label=\"U₂\"];\n  U1 -> X; U2 -> Y;\n  X -> M; X -> Y; M -> Y;\n  labelloc=t; label=\"Structural Causal Model\\nU = exogenous (diamonds)\";\n}",
        "mermaid": "graph TD\n  U1((U₁)) --> X\n  U2((U₂)) --> Y\n  X --> M\n  X --> Y\n  M --> Y\n  style U1 fill:#ccffcc,shape:circle\n  style U2 fill:#ccffcc,shape:circle"
    }
}

def render_svg(dot_source, filename):
    """Render a Graphviz DOT source to SVG."""
    path = os.path.join(OUTPUT_DIR, filename)
    try:
        proc = subprocess.run(
            ["dot", "-Tsvg"],
            input=dot_source,
            capture_output=True,
            text=True,
            check=True
        )
        with open(path, 'w') as f:
            f.write(proc.stdout)
        print(f"  ✓ {filename} ({len(proc.stdout)} bytes)")
        return True
    except subprocess.CalledProcessError as e:
        print(f"  ✗ {filename}: {e.stderr}")
        return False

def generate_mermaid_block(name, dag):
    """Generate a Mermaid markdown block for inclusion in docs."""
    print(f"\n### {name}")
    print("```mermaid")
    print(dag["mermaid"])
    print("```")

if __name__ == "__main__":
    os.makedirs(OUTPUT_DIR, exist_ok=True)
    print(f"Rendering {len(dags)} DAGs to {OUTPUT_DIR}...\n")

    for name, dag in dags.items():
        svg_fn = f"{name}.svg"
        render_svg(dag["graph"], svg_fn)

    print(f"\nDone! Generated {len(dags)} SVG diagrams.")

    # Generate Mermaid markdown blocks for inclusion
    print("\n\n--- Mermaid Blocks (for embedding in docs) ---")
    for name, dag in dags.items():
        generate_mermaid_block(name, dag)
