#!/usr/bin/env python3
# SPDX-License-Identifier: GPL-2.0-only
"""Checkmk special agent for Home Assistant.

Fetches entity states from the Home Assistant REST API and device/entity metadata
from the Home Assistant WebSocket API. It emits a source section plus piggyback
sections grouped by Home Assistant area/room.
"""

from __future__ import annotations

import argparse
import base64
import hashlib
import json
import os
import re
import socket
import ssl
import struct
import sys
import unicodedata
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any, Iterable
from urllib.parse import urlparse

import requests

VERSION = "0.2.4"


class AgentError(RuntimeError):
    pass


def _hide_process_title() -> None:
    """Best effort: remove credentials from the visible process title quickly."""
    try:
        import setproctitle  # type: ignore

        setproctitle.setproctitle("agent_homeassistant")
    except Exception:
        pass


def _json_line(data: Any) -> str:
    return json.dumps(data, ensure_ascii=False, separators=(",", ":"))


def _safe_label(value: Any, max_len: int = 200) -> str:
    text = str(value or "").replace(":", "-").replace("\n", " ").replace("\r", " ")
    return text[:max_len]


def _domain(entity_id: str) -> str:
    return entity_id.split(".", 1)[0] if "." in entity_id else ""


def _compile_optional_regex(pattern: str, field: str) -> re.Pattern[str] | None:
    if not pattern:
        return None
    try:
        return re.compile(pattern, re.IGNORECASE)
    except re.error as exc:
        raise AgentError(f"Invalid {field} regular expression: {exc}") from exc


def _parse_iso8601(value: str | None) -> datetime | None:
    if not value:
        return None
    try:
        return datetime.fromisoformat(value.replace("Z", "+00:00"))
    except ValueError:
        return None


def _age_seconds(value: str | None, now: datetime | None = None) -> float | None:
    parsed = _parse_iso8601(value)
    if parsed is None:
        return None
    if parsed.tzinfo is None:
        parsed = parsed.replace(tzinfo=timezone.utc)
    now = now or datetime.now(timezone.utc)
    return max(0.0, (now - parsed.astimezone(timezone.utc)).total_seconds())


def _display_name(state: dict[str, Any]) -> str:
    attrs = state.get("attributes") or {}
    return str(attrs.get("friendly_name") or state.get("entity_id") or "unknown")


def filter_states(
    states: Iterable[dict[str, Any]],
    domains: set[str],
    include_regex: re.Pattern[str] | None,
    exclude_regex: re.Pattern[str] | None,
    ignore_unavailable: bool = False,
) -> list[dict[str, Any]]:
    selected: list[dict[str, Any]] = []
    for state in states:
        entity_id = str(state.get("entity_id") or "")
        if not entity_id or _domain(entity_id) not in domains:
            continue
        if ignore_unavailable and str(state.get("state") or "").strip().lower() == "unavailable":
            continue
        haystack = f"{entity_id} {_display_name(state)}"
        if include_regex and not include_regex.search(haystack):
            continue
        if exclude_regex and exclude_regex.search(haystack):
            continue
        selected.append(state)
    return selected


@dataclass
class RegistryData:
    entity_to_device: dict[str, str]
    entity_area: dict[str, str]
    devices: dict[str, dict[str, Any]]
    areas: dict[str, str]
    warnings: list[str]


class SimpleWebSocket:
    """Small RFC6455 client sufficient for Home Assistant JSON commands.

    This deliberately uses only the Python standard library so the special agent
    does not depend on an additional WebSocket package on the Checkmk server.
    """

    def __init__(self, url: str, timeout: float, verify_tls: bool) -> None:
        self.url = url
        self.timeout = timeout
        self.verify_tls = verify_tls
        self.sock: socket.socket | ssl.SSLSocket | None = None
        self._buffer = b""

    def __enter__(self) -> "SimpleWebSocket":
        self.connect()
        return self

    def __exit__(self, exc_type, exc, tb) -> None:  # type: ignore[no-untyped-def]
        self.close()

    def connect(self) -> None:
        parsed = urlparse(self.url)
        if parsed.scheme not in ("ws", "wss"):
            raise AgentError(f"Unsupported WebSocket URL scheme: {parsed.scheme}")
        host = parsed.hostname
        if not host:
            raise AgentError("WebSocket URL has no hostname")
        port = parsed.port or (443 if parsed.scheme == "wss" else 80)
        raw = socket.create_connection((host, port), timeout=self.timeout)
        raw.settimeout(self.timeout)
        if parsed.scheme == "wss":
            context = ssl.create_default_context() if self.verify_tls else ssl._create_unverified_context()
            sock: socket.socket | ssl.SSLSocket = context.wrap_socket(raw, server_hostname=host)
        else:
            sock = raw

        path = parsed.path or "/"
        if parsed.query:
            path += "?" + parsed.query
        key = base64.b64encode(os.urandom(16)).decode("ascii")
        host_header = host
        default_port = 443 if parsed.scheme == "wss" else 80
        if port != default_port:
            host_header = f"{host}:{port}"
        request = (
            f"GET {path} HTTP/1.1\r\n"
            f"Host: {host_header}\r\n"
            "Upgrade: websocket\r\n"
            "Connection: Upgrade\r\n"
            f"Sec-WebSocket-Key: {key}\r\n"
            "Sec-WebSocket-Version: 13\r\n"
            "User-Agent: checkmk-homeassistant/0.2.4\r\n"
            "\r\n"
        ).encode("ascii")
        sock.sendall(request)
        response = b""
        while b"\r\n\r\n" not in response:
            chunk = sock.recv(4096)
            if not chunk:
                raise AgentError("Home Assistant closed the WebSocket during handshake")
            response += chunk
            if len(response) > 65536:
                raise AgentError("Oversized WebSocket handshake response")
        header, self._buffer = response.split(b"\r\n\r\n", 1)
        status_line = header.split(b"\r\n", 1)[0]
        if b" 101 " not in status_line:
            raise AgentError(f"WebSocket handshake failed: {status_line.decode('latin1', 'replace')}")

        headers: dict[str, str] = {}
        for line in header.split(b"\r\n")[1:]:
            if b":" in line:
                k, v = line.split(b":", 1)
                headers[k.decode("ascii", "ignore").lower()] = v.decode("latin1").strip()
        expected = base64.b64encode(
            hashlib.sha1((key + "258EAFA5-E914-47DA-95CA-C5AB0DC85B11").encode("ascii")).digest()
        ).decode("ascii")
        if headers.get("sec-websocket-accept") != expected:
            raise AgentError("Invalid Sec-WebSocket-Accept header")
        self.sock = sock

    def _recv_exact(self, length: int) -> bytes:
        if self.sock is None:
            raise AgentError("WebSocket is not connected")
        while len(self._buffer) < length:
            chunk = self.sock.recv(max(4096, length - len(self._buffer)))
            if not chunk:
                raise AgentError("Home Assistant closed the WebSocket")
            self._buffer += chunk
        data, self._buffer = self._buffer[:length], self._buffer[length:]
        return data

    def _send_frame(self, opcode: int, payload: bytes) -> None:
        if self.sock is None:
            raise AgentError("WebSocket is not connected")
        first = 0x80 | (opcode & 0x0F)
        mask = os.urandom(4)
        length = len(payload)
        if length < 126:
            header = bytes([first, 0x80 | length])
        elif length <= 0xFFFF:
            header = bytes([first, 0x80 | 126]) + struct.pack("!H", length)
        else:
            header = bytes([first, 0x80 | 127]) + struct.pack("!Q", length)
        masked = bytes(byte ^ mask[i % 4] for i, byte in enumerate(payload))
        self.sock.sendall(header + mask + masked)

    def send_json(self, data: dict[str, Any]) -> None:
        self._send_frame(0x1, _json_line(data).encode("utf-8"))

    def recv_text(self) -> str:
        fragments: list[bytes] = []
        text_started = False
        while True:
            b1, b2 = self._recv_exact(2)
            fin = bool(b1 & 0x80)
            opcode = b1 & 0x0F
            masked = bool(b2 & 0x80)
            length = b2 & 0x7F
            if length == 126:
                length = struct.unpack("!H", self._recv_exact(2))[0]
            elif length == 127:
                length = struct.unpack("!Q", self._recv_exact(8))[0]
            mask = self._recv_exact(4) if masked else b""
            payload = self._recv_exact(length)
            if masked:
                payload = bytes(byte ^ mask[i % 4] for i, byte in enumerate(payload))

            if opcode == 0x8:
                raise AgentError("Home Assistant closed the WebSocket")
            if opcode == 0x9:
                self._send_frame(0xA, payload)
                continue
            if opcode == 0xA:
                continue
            if opcode == 0x1:
                fragments = [payload]
                text_started = True
            elif opcode == 0x0 and text_started:
                fragments.append(payload)
            elif opcode == 0x2:
                raise AgentError("Unexpected binary WebSocket message from Home Assistant")
            else:
                continue
            if fin:
                return b"".join(fragments).decode("utf-8")

    def recv_json(self) -> dict[str, Any]:
        try:
            return json.loads(self.recv_text())
        except json.JSONDecodeError as exc:
            raise AgentError(f"Invalid JSON from Home Assistant WebSocket: {exc}") from exc

    def command(self, command_id: int, command_type: str) -> Any:
        self.send_json({"id": command_id, "type": command_type})
        while True:
            response = self.recv_json()
            if response.get("id") != command_id:
                continue
            if response.get("type") != "result" or not response.get("success"):
                error = response.get("error") or response
                raise AgentError(f"Home Assistant WebSocket command {command_type!r} failed: {error}")
            return response.get("result")

    def close(self) -> None:
        if self.sock is not None:
            try:
                self._send_frame(0x8, b"")
            except Exception:
                pass
            try:
                self.sock.close()
            except Exception:
                pass
            self.sock = None


def _websocket_url(base_url: str) -> str:
    parsed = urlparse(base_url)
    scheme = "wss" if parsed.scheme == "https" else "ws"
    path = parsed.path.rstrip("/") + "/api/websocket"
    netloc = parsed.netloc
    return f"{scheme}://{netloc}{path}"


def fetch_registry(base_url: str, token: str, timeout: float, verify_tls: bool) -> RegistryData:
    warnings: list[str] = []
    ws_url = _websocket_url(base_url)
    with SimpleWebSocket(ws_url, timeout, verify_tls) as ws:
        auth_required = ws.recv_json()
        if auth_required.get("type") != "auth_required":
            raise AgentError(f"Unexpected Home Assistant WebSocket greeting: {auth_required.get('type')}")
        ws.send_json({"type": "auth", "access_token": token})
        auth_result = ws.recv_json()
        if auth_result.get("type") != "auth_ok":
            raise AgentError(f"Home Assistant WebSocket authentication failed: {auth_result.get('message', auth_result)}")

        entity_result = ws.command(1, "config/entity_registry/list_for_display")

        device_result = ws.command(2, "config/device_registry/list")
        if not isinstance(device_result, list):
            raise AgentError("Home Assistant device registry did not return a list")
        devices: dict[str, dict[str, Any]] = {
            str(d.get("id")): d for d in device_result if isinstance(d, dict) and d.get("id")
        }

        area_result = ws.command(3, "config/area_registry/list")
        if not isinstance(area_result, list):
            raise AgentError("Home Assistant area registry did not return a list")
        areas: dict[str, str] = {
            str(a.get("area_id") or a.get("id")): str(a.get("name") or "")
            for a in area_result
            if isinstance(a, dict) and (a.get("area_id") or a.get("id"))
        }

    entities_raw = entity_result.get("entities", []) if isinstance(entity_result, dict) else []
    entity_to_device: dict[str, str] = {}
    entity_area: dict[str, str] = {}
    for entry in entities_raw:
        if not isinstance(entry, dict):
            continue
        entity_id = str(entry.get("ei") or "")
        if not entity_id:
            continue
        device_id = entry.get("di")
        area_id = entry.get("ai")
        if device_id:
            entity_to_device[entity_id] = str(device_id)
        if area_id:
            entity_area[entity_id] = str(area_id)
    return RegistryData(entity_to_device, entity_area, devices, areas, warnings)


def fetch_states(base_url: str, token: str, timeout: float, verify_tls: bool) -> list[dict[str, Any]]:
    url = base_url.rstrip("/") + "/api/states"
    headers = {"Authorization": f"Bearer {token}", "Accept": "application/json"}
    try:
        response = requests.get(url, headers=headers, timeout=timeout, verify=verify_tls)
    except requests.RequestException as exc:
        raise AgentError(f"Home Assistant REST request failed: {exc}") from exc
    if response.status_code != 200:
        body = response.text.replace("\n", " ")[:300]
        raise AgentError(f"Home Assistant REST API returned HTTP {response.status_code}: {body}")
    try:
        data = response.json()
    except ValueError as exc:
        raise AgentError(f"Home Assistant REST API returned invalid JSON: {exc}") from exc
    if not isinstance(data, list):
        raise AgentError("Home Assistant /api/states did not return a list")
    return [item for item in data if isinstance(item, dict)]


def _area_for_state(state: dict[str, Any], registry: RegistryData) -> str:
    """Return the effective Home Assistant area for an entity.

    An entity-level area assignment wins. If none exists, inherit the area
    from the entity's device. This mirrors how Home Assistant presents
    entities in areas.
    """
    entity_id = str(state.get("entity_id") or "")
    entity_area = registry.entity_area.get(entity_id)
    if entity_area:
        return entity_area

    device_id = registry.entity_to_device.get(entity_id)
    if not device_id:
        return ""
    device = registry.devices.get(device_id, {})
    return str(device.get("area_id") or "")


def _slugify_area_name(name: str) -> str:
    """Create a human-readable Checkmk-safe hostname component."""
    text = str(name or "").strip().lower()
    # Preserve common German transliterations before Unicode decomposition.
    text = (
        text.replace("ä", "ae")
        .replace("ö", "oe")
        .replace("ü", "ue")
        .replace("ß", "ss")
    )
    text = unicodedata.normalize("NFKD", text)
    text = "".join(ch for ch in text if not unicodedata.combining(ch))
    text = re.sub(r"[^a-z0-9]+", "-", text).strip("-")
    text = re.sub(r"-+", "-", text)
    return text or "area"


def _area_hostnames(area_ids: Iterable[str], registry: RegistryData, host_prefix: str) -> dict[str, str]:
    """Build readable host names and resolve slug collisions deterministically."""
    base_by_area: dict[str, str] = {}
    areas_by_base: dict[str, list[str]] = {}
    for area_id in sorted(set(area_ids)):
        area_name = registry.areas.get(area_id) or area_id
        base = f"{host_prefix}{_slugify_area_name(area_name)}"
        base_by_area[area_id] = base
        areas_by_base.setdefault(base, []).append(area_id)

    result: dict[str, str] = {}
    for area_id, base in base_by_area.items():
        if len(areas_by_base[base]) == 1:
            result[area_id] = base
            continue
        suffix = hashlib.sha1(area_id.encode("utf-8")).hexdigest()[:6]
        result[area_id] = f"{base}-{suffix}"
    return result


def build_groups(
    states: list[dict[str, Any]], registry: RegistryData, host_prefix: str
) -> list[tuple[str, dict[str, Any], list[dict[str, Any]]]]:
    by_area: dict[str, list[dict[str, Any]]] = {}
    unassigned: list[dict[str, Any]] = []

    for state in states:
        area_id = _area_for_state(state, registry)
        if area_id:
            by_area.setdefault(area_id, []).append(state)
        else:
            unassigned.append(state)

    hostnames = _area_hostnames(by_area.keys(), registry, host_prefix)
    groups: list[tuple[str, dict[str, Any], list[dict[str, Any]]]] = []
    for area_id, area_states in by_area.items():
        area_name = registry.areas.get(area_id) or area_id
        metadata = {
            "object_type": "area",
            "device_id": "",
            "name": area_name,
            "area_id": area_id,
            "area": area_name,
            "manufacturer": "",
            "model": "",
            "sw_version": "",
        }
        groups.append(
            (
                hostnames[area_id],
                metadata,
                sorted(area_states, key=lambda x: str(x.get("entity_id"))),
            )
        )

    if unassigned:
        metadata = {
            "object_type": "unassigned",
            "device_id": "",
            "name": "Home Assistant unassigned entities",
            "area_id": "",
            "area": "",
            "manufacturer": "",
            "model": "",
            "sw_version": "",
        }
        groups.append(
            (
                f"{host_prefix}unassigned",
                metadata,
                sorted(unassigned, key=lambda x: str(x.get("entity_id"))),
            )
        )

    return sorted(groups, key=lambda item: item[0])

def apply_limits(
    groups: list[tuple[str, dict[str, Any], list[dict[str, Any]]]], max_hosts: int, max_entities: int
) -> tuple[list[tuple[str, dict[str, Any], list[dict[str, Any]]]], list[str]]:
    warnings: list[str] = []
    limited_groups = groups
    if max_hosts > 0 and len(limited_groups) > max_hosts:
        warnings.append(f"Host limit reached: {len(limited_groups)} groups found, only {max_hosts} emitted")
        limited_groups = limited_groups[:max_hosts]

    if max_entities <= 0:
        return limited_groups, warnings

    remaining = max_entities
    final: list[tuple[str, dict[str, Any], list[dict[str, Any]]]] = []
    total = sum(len(states) for _, _, states in limited_groups)
    for hostname, metadata, states in limited_groups:
        if remaining <= 0:
            break
        taken = states[:remaining]
        if taken:
            final.append((hostname, metadata, taken))
            remaining -= len(taken)
    emitted = sum(len(states) for _, _, states in final)
    if total > max_entities:
        warnings.append(f"Entity limit reached: {total} selected entities, only {emitted} emitted")
    return final, warnings


def _entity_payload(state: dict[str, Any], stale_after: float) -> dict[str, Any]:
    attrs = state.get("attributes") or {}
    return {
        "kind": "entity",
        "entity_id": str(state.get("entity_id") or ""),
        "friendly_name": _display_name(state),
        "domain": _domain(str(state.get("entity_id") or "")),
        "state": str(state.get("state") if state.get("state") is not None else ""),
        "unit": str(attrs.get("unit_of_measurement") or ""),
        "device_class": str(attrs.get("device_class") or ""),
        "state_class": str(attrs.get("state_class") or ""),
        "last_changed": str(state.get("last_changed") or ""),
        "last_updated": str(state.get("last_updated") or ""),
        "age_seconds": _age_seconds(str(state.get("last_updated") or "")),
        "stale_after": stale_after,
    }


def emit_source(payload: dict[str, Any]) -> None:
    print("<<<homeassistant_source:sep(0)>>>")
    print(_json_line(payload))


def emit_piggyback(
    groups: list[tuple[str, dict[str, Any], list[dict[str, Any]]]], stale_after: float
) -> None:
    for hostname, metadata, states in groups:
        print(f"<<<<{hostname}>>>>")
        labels = {
            "homeassistant/object": _safe_label(metadata.get("object_type")),
            "homeassistant/name": _safe_label(metadata.get("name")),
        }
        if metadata.get("device_id"):
            labels["homeassistant/device_id"] = _safe_label(metadata["device_id"])
        if metadata.get("area_id"):
            labels["homeassistant/area_id"] = _safe_label(metadata["area_id"])
        if metadata.get("area"):
            labels["homeassistant/area"] = _safe_label(metadata["area"])
        if metadata.get("manufacturer"):
            labels["homeassistant/manufacturer"] = _safe_label(metadata["manufacturer"])
        if metadata.get("model"):
            labels["homeassistant/model"] = _safe_label(metadata["model"])
        print("<<<labels:sep(0)>>>")
        print(_json_line(labels))
        print("<<<homeassistant:sep(0)>>>")
        print(_json_line({"kind": "meta", **metadata, "entity_count": len(states)}))
        for state in states:
            print(_json_line(_entity_payload(state, stale_after)))
        print("<<<<>>>>")


def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="Checkmk special agent for Home Assistant")
    parser.add_argument("--version", action="version", version=VERSION)
    parser.add_argument("--url", required=True, help="Home Assistant base URL")
    parser.add_argument("--token", required=True, help="Home Assistant long-lived access token")
    parser.add_argument("--timeout", type=float, default=10.0)
    parser.add_argument("--no-verify-tls", action="store_true")
    parser.add_argument("--domains", default="sensor,binary_sensor")
    parser.add_argument("--include-regex", default="")
    parser.add_argument("--exclude-regex", default="")
    parser.add_argument("--ignore-unavailable", action="store_true", help="Do not emit entities whose Home Assistant state is unavailable")
    parser.add_argument("--host-prefix", default="ha-")
    parser.add_argument("--max-hosts", type=int, default=150)
    parser.add_argument("--max-entities", type=int, default=450)
    parser.add_argument("--stale-after", type=float, default=0.0)
    return parser.parse_args(argv)


def main(argv: list[str] | None = None) -> int:
    args = parse_args(argv)
    _hide_process_title()
    started = datetime.now(timezone.utc)
    domains = {part.strip() for part in args.domains.split(",") if part.strip()}
    if not domains:
        emit_source({"ok": False, "error": "No Home Assistant domains configured", "version": VERSION})
        return 1
    if not re.fullmatch(r"[a-z0-9][a-z0-9-]*", args.host_prefix):
        emit_source({"ok": False, "error": "Host prefix may only contain lowercase letters, digits and hyphens", "version": VERSION})
        return 1
    if args.timeout <= 0:
        emit_source({"ok": False, "error": "Timeout must be greater than zero", "version": VERSION})
        return 1
    if args.stale_after < 0 or args.max_hosts < 0 or args.max_entities < 0:
        emit_source({"ok": False, "error": "Stale threshold and safety limits must not be negative", "version": VERSION})
        return 1

    try:
        include_regex = _compile_optional_regex(args.include_regex, "include")
        exclude_regex = _compile_optional_regex(args.exclude_regex, "exclude")
        states = fetch_states(args.url, args.token, args.timeout, not args.no_verify_tls)
        registry = fetch_registry(args.url, args.token, args.timeout, not args.no_verify_tls)
        selected = filter_states(states, domains, include_regex, exclude_regex, args.ignore_unavailable)
        groups = build_groups(selected, registry, args.host_prefix)
        emitted_groups, limit_warnings = apply_limits(groups, args.max_hosts, args.max_entities)
        warnings = list(registry.warnings) + limit_warnings
        emitted_entities = sum(len(group_states) for _, _, group_states in emitted_groups)
        duration = (datetime.now(timezone.utc) - started).total_seconds()
        source = {
            "ok": True,
            "version": VERSION,
            "ha_url": args.url,
            "total_states": len(states),
            "selected_entities": len(selected),
            "emitted_entities": emitted_entities,
            "generated_hosts": len(emitted_groups),
            "domains": sorted(domains),
            "ignore_unavailable": bool(args.ignore_unavailable),
            "warnings": warnings,
            "duration_seconds": round(duration, 3),
            "timestamp": datetime.now(timezone.utc).isoformat(),
        }
        emit_source(source)
        emit_piggyback(emitted_groups, args.stale_after)
        return 0
    except Exception as exc:
        duration = (datetime.now(timezone.utc) - started).total_seconds()
        emit_source({
            "ok": False,
            "version": VERSION,
            "error": str(exc),
            "duration_seconds": round(duration, 3),
            "timestamp": datetime.now(timezone.utc).isoformat(),
        })
        return 1


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