#!/usr/bin/env python3
"""Watch Bambu print state through OctoEverywhere's local MQTT relays.

Designed for the Creality Sonic Pad's Python 3.7 environment.  Multiple
printers share one Python process, while each printer still has its own MQTT
client because every OctoEverywhere connector exposes a separate relay port.
"""

import argparse
import configparser
import json
import logging
import logging.handlers
import os
import queue
import signal
import ssl
import stat
import threading
import time
import urllib.request

import certifi
import paho.mqtt.client as mqtt


PRINT_KEYS = {
    "gcode_file",
    "gcode_state",
    "mc_print_stage",
    "mc_remaining_time",
    "project_id",
    "subtask_id",
    "subtask_name",
}
ACTIVE_STATES = {"PREPARE", "RUNNING", "PAUSE", "PAUSED"}
TERMINAL_STATES = {"IDLE", "FAILED", "FINISH", "FINISHED", "COMPLETED"}
TERMINATE_COMMAND = '{"print":{"sequence_id":"0","command":"stop"}}'
PUSH_ALL_COMMAND = '{"pushing":{"sequence_id":"0","command":"pushall"}}'


def update_print_status(original, new_values):
    """Merge only the print fields used by this watcher."""
    for key, value in new_values.items():
        if key not in PRINT_KEYS:
            continue
        if isinstance(value, dict) and isinstance(original.get(key), dict):
            update_print_status(original[key], value)
        else:
            original[key] = value


def strip_print_suffix(filename):
    value = (filename or "").strip().lower()
    for suffix in (".3mf", ".gcode"):
        if value.endswith(suffix):
            return value[:-len(suffix)]
    return value


def as_positive_minutes(value):
    try:
        result = float(value)
    except (TypeError, ValueError):
        return 0.0
    return result if result > 0 else 0.0


def read_octoeverywhere_bambu_config(path):
    """Read the Bambu identity and local relay settings owned by OctoEverywhere."""
    parser = configparser.ConfigParser(interpolation=None)
    with open(path, "r") as config_file:
        parser.read_file(config_file)

    try:
        serial = parser.get("bambu", "printer_serial_number").strip()
        access_code = parser.get("bambu", "access_token").strip()
        relay_port = parser.getint("mqtt", "port", fallback=1883)
        relay_enabled = parser.getboolean("mqtt", "enable", fallback=True)
        require_upstream_auth = parser.getboolean(
            "mqtt", "require_upstream_auth", fallback=True
        )
    except (configparser.Error, ValueError) as exc:
        raise ValueError("invalid OctoEverywhere config %s: %s" % (path, exc))

    if not serial:
        raise ValueError("printer_serial_number is empty in %s" % path)
    if require_upstream_auth and not access_code:
        raise ValueError("access_token is empty in %s" % path)
    if not relay_enabled:
        raise ValueError("the [mqtt] relay is disabled in %s" % path)
    if relay_port < 1 or relay_port > 65535:
        raise ValueError("invalid MQTT relay port in %s" % path)

    return {
        "serial": serial,
        "access_code": access_code,
        "relay_port": relay_port,
        "require_upstream_auth": require_upstream_auth,
        "static_username": parser.get("mqtt", "username", fallback="").strip(),
        "static_password": parser.get("mqtt", "password", fallback="").strip(),
    }


class StateStore:
    """Small atomic state file that prevents duplicate events after a restart."""

    def __init__(self, path, logger):
        self.path = path
        self.logger = logger
        self.lock = threading.Lock()
        self.data = {"printers": {}}
        self._load()

    def _load(self):
        try:
            with open(self.path, "r") as state_file:
                loaded = json.load(state_file)
            if isinstance(loaded, dict) and isinstance(loaded.get("printers"), dict):
                self.data = loaded
        except FileNotFoundError:
            return
        except Exception as exc:
            self.logger.warning("Could not read state file %s: %s", self.path, exc)

    def get_printer(self, serial_number):
        with self.lock:
            return dict(self.data["printers"].get(serial_number, {}))

    def update_printer(self, serial_number, values):
        with self.lock:
            current = self.data["printers"].setdefault(serial_number, {})
            current.update(values)
            self._write_locked()

    def _write_locked(self):
        parent = os.path.dirname(self.path)
        if parent:
            os.makedirs(parent, exist_ok=True)
        temporary = self.path + ".tmp"
        with open(temporary, "w") as state_file:
            json.dump(self.data, state_file, indent=2, sort_keys=True)
            state_file.write("\n")
        os.chmod(temporary, 0o600)
        os.replace(temporary, self.path)


class CalendarQueue:
    """Runs HTTP requests away from Paho MQTT callback threads."""

    def __init__(self, config, logger):
        self.endpoint = config["calendar_endpoint"]
        self.timeout = float(config.get("http_timeout_seconds", 10))
        self.retries = max(1, int(config.get("http_retries", 3)))
        self.retry_seconds = max(1, int(config.get("http_retry_seconds", 5)))
        self.logger = logger
        self.tasks = queue.Queue()
        self.threads = []
        worker_count = max(1, min(4, int(config.get("calendar_workers", 2))))
        self.ssl_context = ssl.create_default_context(cafile=certifi.where())
        for number in range(worker_count):
            thread = threading.Thread(
                target=self._worker,
                name="CalendarWorker-%d" % (number + 1),
                daemon=True,
            )
            thread.start()
            self.threads.append(thread)

    def submit(self, watcher, job_key, payload):
        self.tasks.put((watcher, job_key, payload))

    def stop(self):
        for _thread in self.threads:
            self.tasks.put(None)

    def _worker(self):
        while True:
            task = self.tasks.get()
            try:
                if task is None:
                    return
                watcher, job_key, payload = task
                self._post_with_retries(watcher, job_key, payload)
            finally:
                self.tasks.task_done()

    def _post_with_retries(self, watcher, job_key, payload):
        last_error = None
        for attempt in range(1, self.retries + 1):
            try:
                body = json.dumps(payload).encode("utf-8")
                request = urllib.request.Request(
                    self.endpoint,
                    data=body,
                    headers={"Content-Type": "application/json"},
                    method="POST",
                )
                with urllib.request.urlopen(
                    request,
                    timeout=self.timeout,
                    context=self.ssl_context,
                ) as response:
                    response_body = response.read().decode("utf-8").strip()
                result = json.loads(response_body) if response_body else {}
                if not isinstance(result, dict):
                    raise ValueError("calendar endpoint returned a non-object JSON response")
                watcher.calendar_succeeded(job_key, result)
                return
            except Exception as exc:
                last_error = exc
                self.logger.warning(
                    "[%s] Calendar request %d/%d failed: %s",
                    watcher.label,
                    attempt,
                    self.retries,
                    exc,
                )
                if attempt < self.retries:
                    time.sleep(self.retry_seconds * attempt)
        watcher.calendar_failed(job_key, last_error)


class PrinterWatcher:
    def __init__(self, printer, config, state_store, calendar_queue, logger):
        self.label = str(printer["label"])
        self.oe_config_path = os.path.abspath(str(printer["octoeverywhere_config"]))
        self.relay_host = str(printer.get("relay_host", "127.0.0.1"))
        self.config = config
        self.state_store = state_store
        self.calendar_queue = calendar_queue
        self.logger = logger
        self.status = {}
        self.lock = threading.Lock()
        self.client_lock = threading.Lock()
        self.stop_event = threading.Event()
        self.source_thread = None
        self.next_submit_at = 0.0

        source = read_octoeverywhere_bambu_config(self.oe_config_path)
        self.serial = source["serial"]
        self.access_code = source["access_code"]
        self.relay_port = source["relay_port"]
        self.require_upstream_auth = source["require_upstream_auth"]
        self.static_username = source["static_username"]
        self.static_password = source["static_password"]

        saved = state_store.get_printer(self.serial)
        self.job_open = bool(saved.get("job_open", False))
        self.job_key = saved.get("job_key")
        self.job_filename = saved.get("job_filename")
        self.job_identifier = saved.get("job_identifier")
        self.calendar_submitted = bool(saved.get("calendar_submitted", False))
        self.calendar_pending = False

        self.client = self._create_client()

    def _create_client(self):
        safe_label = "".join(c for c in self.label if c.isalnum())[:8] or "printer"
        client_id = "cal-%s-%s" % (safe_label, self.serial[-6:])
        client = mqtt.Client(
            callback_api_version=mqtt.CallbackAPIVersion.VERSION2,
            client_id=client_id,
            clean_session=True,
            protocol=mqtt.MQTTv311,
            transport="tcp",
        )
        if self.require_upstream_auth:
            client.username_pw_set("bblp", self.access_code)
        elif self.static_username or self.static_password:
            client.username_pw_set(self.static_username, self.static_password)
        # The OctoEverywhere local relay is ordinary MQTT/TCP, not MQTT/TLS.
        client.reconnect_delay_set(min_delay=1, max_delay=60)
        client.on_connect = self.on_connect
        client.on_disconnect = self.on_disconnect
        client.on_message = self.on_message
        return client

    @property
    def report_topic(self):
        return "device/%s/report" % self.serial

    @property
    def request_topic(self):
        return "device/%s/request" % self.serial

    def start(self):
        self.logger.info(
            "[%s] Starting relay client on %s:%d",
            self.label,
            self.relay_host,
            self.relay_port,
        )
        with self.client_lock:
            self.client.connect_async(self.relay_host, self.relay_port, keepalive=60)
            self.client.loop_start()
        self.source_thread = threading.Thread(
            target=self._source_monitor,
            name="OctoEverywhereConfig-%s" % self.label,
            daemon=True,
        )
        self.source_thread.start()

    def stop(self):
        self.stop_event.set()
        if self.source_thread is not None:
            self.source_thread.join(timeout=2)
        with self.client_lock:
            try:
                self.client.disconnect()
            finally:
                self.client.loop_stop()

    def _source_monitor(self):
        interval = max(5, int(self.config.get("source_reload_seconds", 10)))
        while not self.stop_event.wait(interval):
            try:
                source = read_octoeverywhere_bambu_config(self.oe_config_path)
                current = (
                    self.serial,
                    self.access_code,
                    self.relay_port,
                    self.require_upstream_auth,
                    self.static_username,
                    self.static_password,
                )
                updated = (
                    source["serial"],
                    source["access_code"],
                    source["relay_port"],
                    source["require_upstream_auth"],
                    source["static_username"],
                    source["static_password"],
                )
                if current != updated:
                    self._apply_source_change(source)
            except Exception as exc:
                self.logger.error(
                    "[%s] Could not reload %s: %s",
                    self.label,
                    self.oe_config_path,
                    exc,
                )

    def _apply_source_change(self, source):
        old_serial = self.serial
        old_port = self.relay_port
        with self.client_lock:
            old_client = self.client
            try:
                old_client.disconnect()
            finally:
                old_client.loop_stop()

            self.serial = source["serial"]
            self.access_code = source["access_code"]
            self.relay_port = source["relay_port"]
            self.require_upstream_auth = source["require_upstream_auth"]
            self.static_username = source["static_username"]
            self.static_password = source["static_password"]

            if self.serial != old_serial:
                with self.lock:
                    self.status = {}
                    saved = self.state_store.get_printer(self.serial)
                    self.job_open = bool(saved.get("job_open", False))
                    self.job_key = saved.get("job_key")
                    self.job_filename = saved.get("job_filename")
                    self.job_identifier = saved.get("job_identifier")
                    self.calendar_submitted = bool(saved.get("calendar_submitted", False))
                    self.calendar_pending = False

            self.client = self._create_client()
            self.client.connect_async(self.relay_host, self.relay_port, keepalive=60)
            self.client.loop_start()

        self.logger.info(
            "[%s] Reloaded OctoEverywhere settings (serial %s, relay port %d -> %d)",
            self.label,
            self.serial,
            old_port,
            self.relay_port,
        )

    def on_connect(self, client, userdata, flags, reason_code, properties=None):
        if reason_code != 0:
            self.logger.warning("[%s] MQTT connection rejected: %s", self.label, reason_code)
            return
        self.logger.info("[%s] Connected to local MQTT relay", self.label)
        client.subscribe(self.report_topic, qos=0)
        if bool(self.config.get("push_full_state_on_connect", True)):
            client.publish(self.request_topic, PUSH_ALL_COMMAND, qos=0)

    def on_disconnect(self, client, userdata, disconnect_flags, reason_code, properties=None):
        if reason_code != 0:
            self.logger.warning(
                "[%s] MQTT disconnected (%s); Paho will reconnect",
                self.label,
                reason_code,
            )
        else:
            self.logger.info("[%s] MQTT disconnected", self.label)

    def on_message(self, client, userdata, message):
        try:
            decoded = json.loads(message.payload.decode("utf-8"))
            print_update = decoded.get("print")
            if not isinstance(print_update, dict):
                return
            with self.lock:
                update_print_status(self.status, print_update)
                self._evaluate_status_locked()
        except Exception:
            self.logger.exception("[%s] Failed to process MQTT message", self.label)

    def _evaluate_status_locked(self):
        state = str(self.status.get("gcode_state", "")).upper()
        stage = str(self.status.get("mc_print_stage", ""))

        if state in TERMINAL_STATES or stage == "1":
            if not (
                self.job_open
                or self.job_key
                or self.job_filename
                or self.job_identifier
                or self.calendar_submitted
                or self.calendar_pending
            ):
                return
            self.logger.info("[%s] Print is no longer active", self.label)
            self.job_open = False
            self.job_key = None
            self.job_filename = None
            self.job_identifier = None
            self.calendar_submitted = False
            self.calendar_pending = False
            self._persist_locked()
            return

        is_active = state in ACTIVE_STATES or stage == "2"
        remaining_minutes = as_positive_minutes(self.status.get("mc_remaining_time"))
        filename = strip_print_suffix(self.status.get("gcode_file", ""))
        if not is_active or remaining_minutes <= 0 or not filename:
            return

        new_identifier = self._get_job_identifier()
        is_new_job = not self.job_open or self.job_filename != filename
        if (
            not is_new_job
            and self.job_identifier
            and new_identifier
            and self.job_identifier != new_identifier
        ):
            is_new_job = True

        if is_new_job:
            self.job_open = True
            self.job_key = "%s:%s:%d" % (self.serial, filename, int(time.time()))
            self.job_filename = filename
            self.job_identifier = new_identifier
            self.calendar_submitted = False
            self.calendar_pending = False
            self.next_submit_at = 0.0
            self._persist_locked()
            self.logger.info(
                "[%s] New print: %s [%.0f min remaining]",
                self.label,
                filename,
                remaining_minutes,
            )
        elif not self.job_identifier and new_identifier:
            # A partial report can be followed by a full report. Remember the
            # stronger identifier without treating the same print as new.
            self.job_identifier = new_identifier
            self._persist_locked()

        if self.calendar_submitted or self.calendar_pending:
            return
        if time.monotonic() < self.next_submit_at:
            return

        now = int(time.time())
        payload = {
            "printer": self.label,
            "file": filename,
            "start": str(now),
            "end": str(now + int(round(remaining_minutes * 60))),
        }
        self.calendar_pending = True
        self.calendar_queue.submit(self, self.job_key, payload)

    def _get_job_identifier(self):
        for key_name in ("subtask_id", "project_id"):
            value = str(self.status.get(key_name, "")).strip()
            if value and value != "0":
                return "%s:%s" % (key_name, value)
        return None

    def calendar_succeeded(self, job_key, response):
        should_stop = False
        with self.lock:
            if not self.job_open or self.job_key != job_key:
                return
            self.calendar_pending = False
            self.calendar_submitted = True
            self._persist_locked()
            should_stop = (
                bool(self.config.get("stop_print_if_response_contains_end", True))
                and "end" in response
            )

        self.logger.info("[%s] Calendar booking confirmed", self.label)
        if should_stop:
            self.logger.warning("[%s] Stopping print as requested by calendar server", self.label)
            result = self.client.publish(self.request_topic, TERMINATE_COMMAND, qos=0)
            if result.rc != mqtt.MQTT_ERR_SUCCESS:
                self.logger.error("[%s] Stop command publish failed: %s", self.label, result.rc)

    def calendar_failed(self, job_key, error):
        with self.lock:
            if not self.job_open or self.job_key != job_key:
                return
            self.calendar_pending = False
            retry_delay = max(10, int(self.config.get("retry_after_failure_seconds", 60)))
            self.next_submit_at = time.monotonic() + retry_delay
        self.logger.error(
            "[%s] Calendar booking failed; it will be attempted again: %s",
            self.label,
            error,
        )

    def _persist_locked(self):
        self.state_store.update_printer(
            self.serial,
            {
                "job_open": self.job_open,
                "job_key": self.job_key,
                "job_filename": self.job_filename,
                "job_identifier": self.job_identifier,
                "calendar_submitted": self.calendar_submitted,
            },
        )


def load_config(path):
    with open(path, "r") as config_file:
        config = json.load(config_file)
    if not isinstance(config, dict):
        raise ValueError("configuration must be a JSON object")
    endpoint = str(config.get("calendar_endpoint", ""))
    if not endpoint.startswith("https://"):
        raise ValueError("calendar_endpoint must be an https:// URL")
    printers = config.get("printers")
    if not isinstance(printers, list) or not printers:
        raise ValueError("printers must be a non-empty JSON array")

    required = ("label", "octoeverywhere_config")
    ports = set()
    for index, printer in enumerate(printers):
        if not isinstance(printer, dict):
            raise ValueError("printer %d must be a JSON object" % (index + 1))
        missing = [key for key in required if not str(printer.get(key, "")).strip()]
        if missing:
            raise ValueError("printer %d is missing: %s" % (index + 1, ", ".join(missing)))
        source_path = os.path.abspath(str(printer["octoeverywhere_config"]))
        source = read_octoeverywhere_bambu_config(source_path)
        port = source["relay_port"]
        if port in ports:
            raise ValueError("relay_port %d is assigned to more than one printer" % port)
        ports.add(port)
    return config


def setup_logging(log_path):
    parent = os.path.dirname(log_path)
    if parent:
        os.makedirs(parent, exist_ok=True)
    logger = logging.getLogger("print-watcher")
    logger.setLevel(logging.INFO)
    formatter = logging.Formatter("%(asctime)s - %(levelname)s - %(message)s")
    stream_handler = logging.StreamHandler()
    stream_handler.setFormatter(formatter)
    file_handler = logging.handlers.RotatingFileHandler(
        log_path,
        maxBytes=2 * 1024 * 1024,
        backupCount=2,
    )
    file_handler.setFormatter(formatter)
    logger.addHandler(stream_handler)
    logger.addHandler(file_handler)
    return logger


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--config", required=True, help="Path to config.json")
    args = parser.parse_args()

    config_path = os.path.abspath(args.config)
    config = load_config(config_path)
    base_dir = os.path.dirname(config_path)
    log_path = os.path.abspath(config.get("log_file", os.path.join(base_dir, "print_watcher.log")))
    state_path = os.path.abspath(config.get("state_file", os.path.join(base_dir, "state.json")))
    logger = setup_logging(log_path)

    try:
        config_mode = stat.S_IMODE(os.stat(config_path).st_mode)
        if config_mode & 0o077:
            logger.warning("Config file is readable by other users; run chmod 600 %s", config_path)
    except OSError:
        pass

    stop_event = threading.Event()

    def request_stop(signum, frame):
        logger.info("Shutdown requested")
        stop_event.set()

    signal.signal(signal.SIGINT, request_stop)
    signal.signal(signal.SIGTERM, request_stop)

    state_store = StateStore(state_path, logger)
    calendar_queue = CalendarQueue(config, logger)
    watchers = [
        PrinterWatcher(printer, config, state_store, calendar_queue, logger)
        for printer in config["printers"]
    ]

    try:
        for watcher in watchers:
            watcher.start()
        logger.info("Print watcher is running for %d printers", len(watchers))
        while not stop_event.wait(1):
            pass
    finally:
        for watcher in watchers:
            try:
                watcher.stop()
            except Exception:
                logger.exception("[%s] Error while stopping MQTT client", watcher.label)
        calendar_queue.stop()
        logger.info("Print watcher stopped")


if __name__ == "__main__":
    main()
