#!/usr/bin/env python3
# SPDX-License-Identifier: GPL-2.0-only
"""Checkmk special agent for Citrix Virtual Apps and Desktops (CVAD) via REST API.

Polls a Citrix Delivery Controller (DDC) over the CVAD REST API
(``/cvad/manage/...``, on-premises, requires CVAD >= 2212 with the
Orchestration Service active) and emits classic Checkmk piggyback sections
that are BYTE-FOR-BYTE compatible with the sections produced by Checkmk's own
Windows agent plugin ``citrix_farm.ps1`` for the per-machine/per-hypervisor
data:

    <<<<<hosted-machine-name>>>>>
    <<<citrix_state>>>
    <<<citrix_serverload>>>
    <<<citrix_sessions>>>

    <<<<<hosting-server-name>>>>>
    <<<citrix_hostsystem>>>

Because these are the exact sections Checkmk's built-in ``citrix_state``,
``citrix_serverload``, ``citrix_sessions`` and ``citrix_hostsystem`` check
plugins already parse, those are reused unchanged for the piggyback hosts --
there is no collision at the check-plugin level with Checkmk's built-in
Citrix monitoring there. The package does ship one own check plugin, on the
host running this special agent itself (the DDC host): it reports the
``citrix_ddc_rest_source`` section as a "Citrix DDC Statistics" service, i.e. the
special agent's own run health (auth/site-id lookup ok, machines/sessions
retrieved, query duration), since piggyback sections alone give no visible
service if the agent run itself fails.

Known limitation: the Citrix CVAD REST API has no documented equivalent for
``Get-BrokerController`` (controller health/licensing state). The
``citrix_controller`` section/service (and with it ``citrix_controller_*``
checks, e.g. license state) can therefore NOT be populated via REST and
still requires ``citrix_farm.ps1`` running locally on the DDC with Citrix
Admin rights if that information is needed.

Do not run this special agent AND ``citrix_farm.ps1`` for the same Citrix
machines/hosts at the same time -- both would write to the same piggyback
hosts/sections and whichever runs later in a check cycle wins, which can
blur the data source for troubleshooting.
"""

from __future__ import annotations

import argparse
import base64
import json
import ssl
import sys
from datetime import datetime, timezone
from typing import Any
from urllib.error import HTTPError, URLError
from urllib.request import Request, urlopen

VERSION = "1.0.0"


class AgentError(RuntimeError):
    pass


def _hide_process_title() -> None:
    try:
        import setproctitle  # type: ignore

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


def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="Checkmk special agent for Citrix CVAD REST API")
    parser.add_argument("--version", action="version", version=VERSION)
    parser.add_argument("--ddc", required=True, help="Delivery Controller address or FQDN")
    parser.add_argument("--port", type=int, default=443, help="HTTPS port of the DDC (default 443)")
    parser.add_argument("--user", required=True, help="Citrix admin user, e.g. DOMAIN\\\\svc_checkmk")
    parser.add_argument("--password", required=True, help="Password for --user")
    parser.add_argument("--no-cert-check", action="store_true", help="Disable TLS certificate verification")
    parser.add_argument("--timeout", type=int, default=20, help="Per-request timeout in seconds")
    parser.add_argument("--max-machines", type=int, default=2000, help="Safety limit for fetched machines")
    parser.add_argument("--max-sessions", type=int, default=10000, help="Safety limit for fetched sessions")
    parser.add_argument(
        "--no-hostsystem",
        action="store_true",
        help="Do not emit citrix_hostsystem piggyback sections for the hypervisor hosts",
    )
    return parser.parse_args(argv)


def _ssl_context(no_cert_check: bool) -> ssl.SSLContext:
    ctx = ssl.create_default_context()
    if no_cert_check:
        ctx.check_hostname = False
        ctx.verify_mode = ssl.CERT_NONE
    return ctx


class CvadClient:
    def __init__(self, ddc: str, port: int, ctx: ssl.SSLContext, timeout: int) -> None:
        host = ddc
        if ":" in host and not host.startswith("["):
            host = f"[{host}]"
        self.base = f"https://{host}:{port}"
        self.ctx = ctx
        self.timeout = timeout
        self.bearer = ""
        self.customer_id = "CitrixOnPremises"
        self.instance_id = ""

    def _request(self, method: str, path: str, headers: dict[str, str], body: bytes | None = None) -> Any:
        url = f"{self.base}{path}"
        req = Request(url, data=body, method=method)
        for key, value in headers.items():
            req.add_header(key, value)
        try:
            with urlopen(req, context=self.ctx, timeout=self.timeout) as resp:
                raw = resp.read().decode("utf-8")
        except HTTPError as exc:
            detail = exc.read().decode("utf-8", "replace")[:500]
            raise AgentError(f"HTTP {exc.code} on {path}: {detail}") from exc
        except URLError as exc:
            raise AgentError(f"Connection to {self.base} failed: {exc.reason}") from exc
        if not raw:
            return None
        try:
            return json.loads(raw)
        except ValueError as exc:
            raise AgentError(f"Invalid JSON response from {path}: {exc}") from exc

    def authenticate(self, user: str, password: str) -> None:
        basic = base64.b64encode(f"{user}:{password}".encode("utf-8")).decode("ascii")
        resp = self._request(
            "POST",
            "/cvad/manage/Tokens",
            headers={"Authorization": f"Basic {basic}", "Accept": "application/json"},
        )
        if not isinstance(resp, dict):
            raise AgentError("Token response was not a JSON object")
        # The CVAD REST API returns PascalCase field names ("Token",
        # "CustomerId") -- look them up case-insensitively so a future
        # casing change does not silently break authentication.
        lower_map = {k.lower(): k for k in resp}
        bearer = resp.get(lower_map["token"]) if "token" in lower_map else None
        if not bearer:
            keys = sorted(resp.keys())
            raise AgentError(f"Token response contained no recognizable token field. Available fields: {keys}")
        self.bearer = str(bearer)
        if "customerid" in lower_map and resp.get(lower_map["customerid"]):
            self.customer_id = str(resp[lower_map["customerid"]])

    def fetch_site_id(self) -> None:
        """Resolve the Citrix-InstanceId (site id) via /cvad/manage/Me.

        The CVAD REST API requires this header on data endpoints such as
        Machines/Sessions; Tokens alone is not enough, so this must run
        right after authenticate() and before any other request.
        """
        resp = self._request(
            "GET",
            "/cvad/manage/Me",
            headers={
                "Authorization": f"CWSAuth Bearer={self.bearer}",
                "Citrix-CustomerId": self.customer_id,
                "Accept": "application/json",
            },
        )
        if not isinstance(resp, dict):
            raise AgentError("Me response was not a JSON object")
        lower_map = {k.lower(): k for k in resp}
        customers = resp.get(lower_map["customers"]) if "customers" in lower_map else None
        if not isinstance(customers, list) or not customers:
            raise AgentError("Me response contained no customers array, could not determine site ID")
        first_customer = customers[0]
        if not isinstance(first_customer, dict):
            raise AgentError("Me response: first customer entry was not a JSON object")
        customer_lower = {k.lower(): k for k in first_customer}
        sites = first_customer.get(customer_lower["sites"]) if "sites" in customer_lower else None
        if not isinstance(sites, list) or not sites:
            raise AgentError("Me response contained no sites, could not determine site ID")
        first_site = sites[0]
        if not isinstance(first_site, dict):
            raise AgentError("Me response: first site entry was not a JSON object")
        site_lower = {k.lower(): k for k in first_site}
        site_id = first_site.get(site_lower["id"]) if "id" in site_lower else None
        if not site_id:
            raise AgentError("Me response: site entry contained no id")
        self.instance_id = str(site_id)

    def _auth_headers(self) -> dict[str, str]:
        headers = {
            "Authorization": f"CWSAuth Bearer={self.bearer}",
            "Citrix-CustomerId": self.customer_id,
            "Accept": "application/json",
        }
        if self.instance_id:
            headers["Citrix-InstanceId"] = self.instance_id
        return headers

    def get_items(self, path: str, max_items: int) -> list[dict[str, Any]]:
        """GET a CVAD collection endpoint, following ContinuationToken paging."""
        items: list[dict[str, Any]] = []
        next_path: str | None = path
        guard = 0
        while next_path and len(items) < max_items:
            guard += 1
            if guard > 200:
                break
            resp = self._request("GET", next_path, headers=self._auth_headers())
            if not isinstance(resp, dict):
                break
            page_items = resp.get("Items")
            if isinstance(page_items, list):
                items.extend(item for item in page_items if isinstance(item, dict))
            token = resp.get("ContinuationToken")
            if token:
                sep = "&" if "?" in path else "?"
                next_path = f"{path}{sep}continuationToken={token}"
            else:
                next_path = None
        return items[:max_items]


def _short_name(dns_or_name: str) -> str:
    """Reduce a DNS/NetBIOS machine identity to a bare, piggyback-safe host name."""
    name = dns_or_name or ""
    if "\\" in name:
        name = name.rsplit("\\", 1)[-1]
    return name.strip()


def _machine_identity(machine: dict[str, Any]) -> str:
    """Pick the best available piggyback host name for a CVAD machine.

    Preference order mirrors citrix_farm.ps1 (Hosting.HostedMachineName) but
    falls back to the DNS/machine name instead of silently skipping the
    machine, so physical/Remote-PC/manually provisioned machines (which have
    no HostedMachineName) still get monitored.
    """
    hosting = machine.get("Hosting") or {}
    hosted_name = hosting.get("HostedMachineName")
    if hosted_name:
        return _short_name(str(hosted_name))
    dns_name = machine.get("DnsName")
    if dns_name:
        return _short_name(str(dns_name))
    return _short_name(str(machine.get("Name") or ""))


def _bool_str(value: Any) -> str | None:
    if value is None:
        return None
    return "True" if value else "False"


def build_session_counts(sessions: list[dict[str, Any]]) -> dict[str, dict[str, int]]:
    counts: dict[str, dict[str, int]] = {}
    for session in sessions:
        machine = session.get("Machine") or {}
        machine_id = machine.get("Id")
        if not machine_id:
            continue
        state = str(session.get("State") or "")
        entry = counts.setdefault(machine_id, {"active": 0, "inactive": 0})
        if state == "Active":
            entry["active"] += 1
        elif state == "Disconnected":
            entry["inactive"] += 1
    return counts


def emit_state_section(machine: dict[str, Any]) -> list[str]:
    lines = ["<<<citrix_state>>>"]
    hosting = machine.get("Hosting") or {}
    catalog = (machine.get("MachineCatalog") or {}).get("Name")
    if catalog:
        lines.append(f"Catalog {catalog}")
    controller = machine.get("ControllerDnsName")
    if controller:
        lines.append(f"Controller {controller}")
    desktop_group = (machine.get("DeliveryGroup") or {}).get("Name")
    if desktop_group:
        lines.append(f"DesktopGroupName {desktop_group}")
    fault_state = machine.get("FaultState")
    if fault_state:
        lines.append(f"FaultState {fault_state}")
    hosting_server = hosting.get("HostingServerName")
    if hosting_server:
        lines.append(f"HostingServer {hosting_server}")
    maintenance = _bool_str(machine.get("InMaintenanceMode"))
    if maintenance is not None:
        lines.append(f"MaintenanceMode {maintenance}")
    power_state = machine.get("PowerState")
    if power_state:
        lines.append(f"PowerState {power_state}")
    registration_state = machine.get("RegistrationState")
    if registration_state:
        lines.append(f"RegistrationState {registration_state}")
    vm_tools_state = machine.get("VMToolsState")
    if vm_tools_state:
        lines.append(f"VMToolsState {vm_tools_state}")
    agent_version = machine.get("AgentVersion")
    if agent_version:
        lines.append(f"AgentVersion {agent_version}")
    return lines


def emit_serverload_section(machine: dict[str, Any]) -> list[str]:
    load = machine.get("LoadIndex")
    if load is None:
        return []
    return ["<<<citrix_serverload>>>", str(load)]


def emit_sessions_section(machine: dict[str, Any], session_counts: dict[str, dict[str, int]]) -> list[str]:
    machine_id = machine.get("Id")
    counts = session_counts.get(machine_id, {"active": 0, "inactive": 0})
    total = machine.get("SessionCount")
    if total is None:
        total = counts["active"] + counts["inactive"]
    return [
        "<<<citrix_sessions>>>",
        f"sessions {total}",
        f"active_sessions {counts['active']}",
        f"inactive_sessions {counts['inactive']}",
    ]


def emit_piggyback(
    machines: list[dict[str, Any]],
    session_counts: dict[str, dict[str, int]],
    emit_hostsystem: bool,
) -> None:
    for machine in machines:
        identity = _machine_identity(machine)
        if not identity:
            continue
        print(f"<<<<{identity}>>>>")
        print("\n".join(emit_state_section(machine)))
        serverload_lines = emit_serverload_section(machine)
        if serverload_lines:
            print("\n".join(serverload_lines))
        print("\n".join(emit_sessions_section(machine, session_counts)))
        print("<<<<>>>>")

        if not emit_hostsystem:
            continue
        hosting = machine.get("Hosting") or {}
        hosting_server = hosting.get("HostingServerName")
        if not hosting_server:
            continue
        hosting_server_name = _short_name(str(hosting_server))
        pool_name = (hosting.get("HypervisorConnection") or {}).get("Name") or ""
        print(f"<<<<{hosting_server_name}>>>>")
        print("<<<citrix_hostsystem>>>")
        print(f"VMName {identity}")
        if pool_name:
            print(f"CitrixPoolName {pool_name}")
        print("<<<<>>>>")


def emit_source(payload: dict[str, Any]) -> None:
    print("<<<citrix_ddc_rest_source:sep(0)>>>")
    print(json.dumps(payload, ensure_ascii=False, separators=(",", ":")))


def main(argv: list[str] | None = None) -> int:
    args = parse_args(argv)
    _hide_process_title()
    started = datetime.now(timezone.utc)
    ctx = _ssl_context(args.no_cert_check)
    client = CvadClient(args.ddc, args.port, ctx, args.timeout)

    try:
        client.authenticate(args.user, args.password)
        client.fetch_site_id()
        machines = client.get_items("/cvad/manage/Machines", args.max_machines)
        sessions = client.get_items("/cvad/manage/Sessions", args.max_sessions)
        session_counts = build_session_counts(sessions)
        emit_piggyback(machines, session_counts, not args.no_hostsystem)
        duration = (datetime.now(timezone.utc) - started).total_seconds()
        emit_source(
            {
                "ok": True,
                "version": VERSION,
                "ddc": args.ddc,
                "customer_id": client.customer_id,
                "machines": len(machines),
                "sessions": len(sessions),
                "duration_seconds": round(duration, 3),
                "timestamp": datetime.now(timezone.utc).isoformat(),
            }
        )
        return 0
    except Exception as exc:
        duration = (datetime.now(timezone.utc) - started).total_seconds()
        emit_source(
            {
                "ok": False,
                "version": VERSION,
                "ddc": args.ddc,
                "error": str(exc),
                "duration_seconds": round(duration, 3),
                "timestamp": datetime.now(timezone.utc).isoformat(),
            }
        )
        sys.stderr.write(f"citrix_ddc_rest: {exc}\n")
        return 1


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