#!/usr/bin/python3
"""Keep every ECZOS optical LIO target exclusive to one initiator."""

import fcntl
import json
import logging
import logging.handlers
import os
import signal
import tempfile
import time
from pathlib import Path

from rtslib_fb import FabricModule, Target

STATE_DIR = Path("/var/lib/eczos/network-optical")
STATE = STATE_DIR / "state.json"
LOCK = STATE_DIR / "state.lock"
RUNNING = True


def stop(_signum, _frame):
    global RUNNING
    RUNNING = False


def log_setup():
    log = logging.getLogger("eczos-network-optical-guard")
    log.setLevel(logging.INFO)
    try:
        log.addHandler(logging.handlers.SysLogHandler(address="/dev/log"))
    except OSError:
        log.addHandler(logging.StreamHandler())
    return log


LOG = log_setup()


def configuration():
    timeout = 30
    try:
        for line in Path("/etc/eczos/network-optical.conf").read_text().splitlines():
            if line.startswith("ReleaseTimeoutSec="):
                timeout = int(line.partition("=")[2])
    except (OSError, ValueError):
        pass
    return max(10, min(timeout, 600))


def update_once(timeout):
    STATE_DIR.mkdir(mode=0o755, parents=True, exist_ok=True)
    os.chmod(STATE_DIR, 0o755)
    with LOCK.open("a+") as lock:
        os.chmod(LOCK, 0o600)
        fcntl.flock(lock, fcntl.LOCK_EX)
        try:
            state = json.loads(STATE.read_text(encoding="utf-8"))
        except (OSError, ValueError):
            return
        changed = False
        now = int(time.time())
        for record in state.get("shares", {}).values():
            try:
                target = Target(FabricModule("iscsi"), record["iqn"], mode="lookup")
                tpg = next(iter(target.tpgs))
                active = [acl for acl in tpg.node_acls if acl.session]
            except Exception as error:
                LOG.warning("cannot inspect %s: %s", record.get("iqn", "unknown"), error)
                continue
            claim = record.get("claim", "")
            active_names = [acl.node_wwn for acl in active]
            if not claim and active_names:
                claim = active_names[0]
                record["claim"] = claim
                record["claimMissingSince"] = 0
                tpg.set_attribute("generate_node_acls", "0")
                changed = True
                LOG.info("claimed %s for initiator %s", record["iqn"], claim)
            if claim:
                tpg.set_attribute("generate_node_acls", "0")
                for acl in list(tpg.node_acls):
                    if acl.node_wwn != claim:
                        try:
                            acl.delete()
                            LOG.warning("rejected competing initiator %s for %s", acl.node_wwn, record["iqn"])
                        except Exception as error:
                            LOG.error("could not reject competing initiator: %s", error)
                if claim in active_names:
                    if record.get("claimMissingSince"):
                        record["claimMissingSince"] = 0
                        changed = True
                else:
                    missing_since = int(record.get("claimMissingSince") or 0)
                    if not missing_since:
                        record["claimMissingSince"] = now
                        changed = True
                    elif now - missing_since >= timeout:
                        for acl in list(tpg.node_acls):
                            if acl.node_wwn == claim:
                                try:
                                    acl.delete()
                                except Exception as error:
                                    LOG.error("could not release ACL %s: %s", claim, error)
                                    break
                        else:
                            record["claim"] = ""
                            record["claimMissingSince"] = 0
                            tpg.set_attribute("generate_node_acls", "1")
                            changed = True
                            LOG.info("released %s after disconnect timeout", record["iqn"])
            elif not active_names:
                tpg.set_attribute("generate_node_acls", "1")
        if changed:
            fd, temporary = tempfile.mkstemp(prefix="state.", dir=STATE_DIR)
            try:
                with os.fdopen(fd, "w", encoding="utf-8") as stream:
                    json.dump(state, stream, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
                    stream.write("\n")
                    stream.flush()
                    os.fsync(stream.fileno())
                os.chmod(temporary, 0o644)
                os.replace(temporary, STATE)
            finally:
                if os.path.exists(temporary):
                    os.unlink(temporary)


def main():
    signal.signal(signal.SIGTERM, stop)
    signal.signal(signal.SIGINT, stop)
    timeout = configuration()
    while RUNNING:
        try:
            update_once(timeout)
        except Exception as error:
            LOG.exception("exclusive-access guard failure: %s", error)
        time.sleep(0.25)


if __name__ == "__main__":
    main()
