#! /usr/bin/python3
"""Backfill draws.jsonl from rotated watter.log history.

Walks watter.log* in chronological order (oldest first), feeds the sensor
records through a fresh FlowRateTracker + DrawLogger, and emits the same
JSONL records that the live daemon would have produced.

Cross-file continuity is implicit: a single tracker/logger pair sees every
record in time order, so a draw that opens in watter.log.2.gz and closes
in watter.log.1 yields one stitched record.

Open draws at end-of-stream are dropped (matches production: only the
True->False flow edge closes a draw).
"""
import argparse
import glob
import gzip
import json
import os
import re
import sys

from watter_collector.config import wtr_cfg
from watter_collector.draw_logger import DrawLogger, DRAWS_FILENAME
from watter_collector.flow_rate import FlowRateTracker


class _StderrLogger:
    """Minimal logger for off-collectd use: prints errors to stderr."""

    def error(self, msg):
        print(msg, file=sys.stderr)

DEFAULT_LOG_DIR = "/var/log/watter-collector"
DEFAULT_LOG_GLOB = "watter.log*"


def sort_log_files(paths):
    """Return paths sorted oldest-first.

    Rotation suffix convention: watter.log (newest) -> watter.log.1 ->
    watter.log.2.gz -> ... numeric, not lexical, so .10 sorts after .2.
    Files without a numeric suffix (the live log) sort last.
    """
    def key(path):
        name = os.path.basename(path)
        m = re.match(r"^watter\.log(?:\.(\d+))?(?:\.gz)?$", name)
        if not m or m.group(1) is None:
            return (sys.maxsize, name)
        return (-int(m.group(1)), name)

    return sorted(paths, key=key)


def open_log(path):
    """Open a plain or gzipped log for line-by-line text reads."""
    if path.endswith(".gz"):
        return gzip.open(path, "rt", encoding="utf-8", errors="replace")
    return open(path, "r", encoding="utf-8", errors="replace")


def load_existing_starts(output_path):
    """Return the set of started_at values already present in draws.jsonl."""
    starts = set()
    if not os.path.exists(output_path):
        return starts
    with open(output_path, "r", encoding="utf-8") as fh:
        for line in fh:
            line = line.strip()
            if not line:
                continue
            try:
                rec = json.loads(line)
            except ValueError:
                continue
            started = rec.get("started_at")
            if started is not None:
                starts.add(started)
    return starts


class _DedupeWriter:
    """DrawLogger writer that drops records whose started_at is already known."""

    def __init__(self, target, known_starts):
        self._target = target
        self._known = known_starts
        self.skipped = 0
        self.written = 0

    def write(self, record):
        if record.get("started_at") in self._known:
            self.skipped += 1
            return
        # pylint: disable=protected-access
        self._target._write(record)
        self._known.add(record.get("started_at"))
        self.written += 1


class _StdoutWriter:
    """DrawLogger writer that prints records to stdout (dry-run mode)."""

    def __init__(self):
        self.written = 0

    def write(self, record):
        print(json.dumps(record))
        self.written += 1


def _patch_logger_write(logger, writer):
    """Redirect DrawLogger._write through a custom writer."""
    logger._write = writer.write  # pylint: disable=protected-access


def process_files(paths, logger, tracker, on_malformed=None):
    """Stream every log file's JSONL records through tracker + logger.

    Returns (records_seen, malformed_lines).
    """
    seen = 0
    malformed = 0
    state = None

    for path in paths:
        with open_log(path) as fh:
            for line in fh:
                line = line.strip()
                if not line:
                    continue
                try:
                    entry = json.loads(line)
                except ValueError:
                    malformed += 1
                    if on_malformed:
                        on_malformed(path, line)
                    continue

                t = entry.get("time")
                if "state" in entry:
                    ev = entry["state"].get("event")
                    if ev is not None:
                        state = ev
                    continue

                if "gpios" in entry:
                    he = entry["gpios"].get("heat-element")
                    if he and he.get("active"):
                        logger.mark_element_fired()
                    continue

                if "sensors" in entry:
                    sensors = dict(entry["sensors"])
                    if t is not None:
                        # watter.log records publish_time=0 inside sensors;
                        # the wrapper "time" field is the real timestamp.
                        sensors["publish_time"] = t
                    tracker.update(sensors)
                    logger.update(sensors, tracker.is_flowing, state)
                    seen += 1

    return seen, malformed


def discover_logs(log_dir):
    """Return rotated log file paths under log_dir, oldest first."""
    paths = glob.glob(os.path.join(log_dir, DEFAULT_LOG_GLOB))
    return sort_log_files(paths)


def main(argv=None):
    parser = argparse.ArgumentParser(
        description="Backfill draws.jsonl from rotated watter.log history.")
    parser.add_argument(
        "paths", nargs="*",
        help="Log files in any order; defaults to /var/log/watter-collector/watter.log*")
    parser.add_argument(
        "--log-dir", default=DEFAULT_LOG_DIR,
        help="Directory to discover logs in when paths are not given.")
    parser.add_argument(
        "--output",
        help="Override draws.jsonl path; defaults to wtr_cfg.status.dir/draws.jsonl.")
    parser.add_argument(
        "--dedupe", action="store_true",
        help="Skip records whose started_at already exists in the output file.")
    parser.add_argument(
        "--dry-run", action="store_true",
        help="Print records to stdout instead of appending to draws.jsonl.")
    args = parser.parse_args(argv)

    if args.paths:
        paths = sort_log_files(args.paths)
    else:
        paths = discover_logs(args.log_dir)

    if not paths:
        print("No log files found.", file=sys.stderr)
        return 1

    output_path = args.output or os.path.join(wtr_cfg.status.dir, DRAWS_FILENAME)

    logger = DrawLogger(_StderrLogger(), sigaction=False)
    tracker = FlowRateTracker()

    writer = None
    if args.dry_run:
        writer = _StdoutWriter()
        _patch_logger_write(logger, writer)
    elif args.dedupe:
        known = load_existing_starts(output_path)
        os.makedirs(os.path.dirname(output_path) or ".", exist_ok=True)
        wtr_cfg.status.dir = os.path.dirname(output_path) or wtr_cfg.status.dir
        writer = _DedupeWriter(logger, known)
        _patch_logger_write(logger, writer)
    else:
        os.makedirs(os.path.dirname(output_path) or ".", exist_ok=True)
        wtr_cfg.status.dir = os.path.dirname(output_path) or wtr_cfg.status.dir

    seen, malformed = process_files(paths, logger, tracker)
    logger.close()

    print(f"Read {seen} sensor records from {len(paths)} file(s)"
          + (f", skipped {malformed} malformed line(s)" if malformed else ""),
          file=sys.stderr)
    if isinstance(writer, _DedupeWriter):
        print(f"Wrote {writer.written} new draw(s); skipped {writer.skipped} duplicate(s).",
              file=sys.stderr)
    elif isinstance(writer, _StdoutWriter):
        print(f"Emitted {writer.written} draw record(s) (dry-run).",
              file=sys.stderr)

    return 0


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