"""
Field output.

Two formats, deliberately:

npz   Compressed NumPy archives, one per frame. This is what
      visualize_waterfall.py consumes. Self-describing, no external
      dependency, and cheap enough to write mid-run without stalling the
      solver.
vti   VTK ImageData, written as raw binary appended data. Only produced on
      request. This exists so results can be opened in ParaView -- not to
      replace the built-in visualisation, but because the moment a result
      looks surprising the first thing anyone wants is to poke at the actual
      field in a real post-processor.

Every frame carries the grid metadata needed to reconstruct coordinates, so a
frame directory is interpretable without the config that produced it.
"""

import base64
import json
import os

import numpy as np


def frame_path(outdir, index, prefix="frame"):
    return os.path.join(outdir, f"{prefix}_{index:05d}.npz")


def save_frame(outdir, index, solver, fields=("u", "p", "vorticity"),
               prefix="frame"):
    """
    Write one frame. `fields` selects which derived quantities to compute --
    vorticity and Q each cost a full pass over the grid, so they are opt-in.
    """
    os.makedirs(outdir, exist_ok=True)
    g = solver.grid
    uc, vc, wc, pc = solver.cell_fields()

    data = {
        "u": uc.astype(np.float32),
        "v": vc.astype(np.float32),
        "w": wc.astype(np.float32),
        "p": pc.astype(np.float32),
        "eta": solver.eta_c[1:-1, 1:-1, 1:-1].astype(np.float32),
    }
    if "vorticity" in fields:
        ox, oy, oz = solver.vorticity_field()
        data["omega_x"] = ox.astype(np.float32)
        data["omega_y"] = oy.astype(np.float32)
        data["omega_z"] = oz.astype(np.float32)
    if "q" in fields:
        data["q"] = solver.q_field().astype(np.float32)

    meta = {
        "index": int(index),
        "time": float(solver.t),
        "step": int(solver.step_count),
        "nx": g.nx, "ny": g.ny, "nz": g.nz,
        "dx": g.dx, "dy": g.dy, "dz": g.dz,
        "origin": [float(v) for v in g.origin],
        "u_inf": [float(v) for v in solver.bc.u_inf],
        "rho": solver.rho,
        "nu": solver.nu,
        "force": [float(v) for v in solver.force],
    }
    data["_meta"] = np.frombuffer(
        json.dumps(meta).encode("utf-8"), dtype=np.uint8
    )
    np.savez_compressed(frame_path(outdir, index, prefix), **data)


def load_frame(path):
    """Returns (dict of arrays, metadata dict)."""
    z = np.load(path)
    meta = json.loads(bytes(z["_meta"]).decode("utf-8"))
    arrays = {k: z[k] for k in z.files if k != "_meta"}
    return arrays, meta


def list_frames(outdir, prefix="frame"):
    if not os.path.isdir(outdir):
        return []
    names = [
        n for n in os.listdir(outdir)
        if n.startswith(prefix + "_") and n.endswith(".npz")
    ]
    return [os.path.join(outdir, n) for n in sorted(names)]


# ----------------------------------------------------------------------
# VTK ImageData
# ----------------------------------------------------------------------

def save_vti(path, solver, include_q=False):
    """
    Write a VTK ImageData file for ParaView.

    Point data is written at cell centres and declared as CellData, so
    ParaView shows the same discrete values the solver holds rather than an
    interpolation of them.
    """
    g = solver.grid
    uc, vc, wc, pc = solver.cell_fields()
    eta = solver.eta_c[1:-1, 1:-1, 1:-1]

    arrays = [
        ("velocity", np.stack([uc, vc, wc], axis=-1).astype(np.float32), 3),
        ("pressure", pc.astype(np.float32), 1),
        ("solid_fraction", eta.astype(np.float32), 1),
    ]
    if include_q:
        arrays.append(("q_criterion", solver.q_field().astype(np.float32), 1))

    ox, oy, oz = g.origin
    ext = f"0 {g.nx} 0 {g.ny} 0 {g.nz}"

    header = [
        '<?xml version="1.0"?>',
        '<VTKFile type="ImageData" version="1.0" byte_order="LittleEndian" '
        'header_type="UInt64">',
        f'  <ImageData WholeExtent="{ext}" Origin="{ox} {oy} {oz}" '
        f'Spacing="{g.dx} {g.dy} {g.dz}">',
        f'    <Piece Extent="{ext}">',
        '      <CellData>',
    ]
    payload = b""
    offset = 0
    for name, arr, ncomp in arrays:
        header.append(
            f'        <DataArray type="Float32" Name="{name}" '
            f'NumberOfComponents="{ncomp}" format="appended" offset="{offset}"/>'
        )
        # VTK expects Fortran-ordered cell data (x fastest).
        flat = np.ascontiguousarray(
            arr.transpose(2, 1, 0) if ncomp == 1
            else arr.transpose(2, 1, 0, 3)
        ).astype(np.float32).tobytes()
        block = np.array([len(flat)], dtype=np.uint64).tobytes() + flat
        payload += block
        offset += len(block)

    header += [
        '      </CellData>',
        '    </Piece>',
        '  </ImageData>',
        '  <AppendedData encoding="raw">',
        '   _',
    ]
    with open(path, "wb") as fh:
        fh.write("\n".join(header).encode("ascii"))
        fh.write(payload)
        fh.write(b"\n  </AppendedData>\n</VTKFile>\n")
    _ = base64
