#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
第 17 章 · ClickHouse 健康巡检脚本
====================================

干什么:
    一次性查 5 张关键 system 视图, 把「副本延迟 / Part 数 / Merge 队列 /
    Mutation 队列 / 磁盘水位 / 内存 / 失败查询」的现状汇总成一张 Markdown 风格
    的健康报告. 适合放到 cron / systemd timer 里每 5 分钟跑一次.

前置:
    python -m pip install clickhouse-connect

用法:
    python ops_check.py
    python ops_check.py --host 127.0.0.1 --port 8123
    python ops_check.py --json   # 输出 JSON, 给 Alertmanager / 飞书机器人吃

退出码:
    0 = 全部 OK
    1 = WARNING (有指标超阈值, 但不致命)
    2 = CRITICAL (必须立刻处理)
"""
from __future__ import annotations

import argparse
import json
import sys
from dataclasses import dataclass, field
from typing import Any, Dict, List

import clickhouse_connect


# ============== 阈值表(可改) ==============
THRESH = {
    "replica_delay_warn":   30,     # 秒
    "replica_delay_crit":   120,
    "queue_size_warn":      50,
    "queue_size_crit":      300,
    "parts_per_partition_warn": 150,
    "parts_per_partition_crit": 300,
    "mutation_pending_warn":  10,
    "mutation_pending_crit":  50,
    "disk_pct_warn":          0.80,
    "disk_pct_crit":          0.90,
    "mem_pct_warn":           0.75,
    "mem_pct_crit":           0.90,
    "mark_cache_hit_warn":    0.90,
}

LEVEL_OK   = 0
LEVEL_WARN = 1
LEVEL_CRIT = 2
LEVEL_NAME = {0: "OK", 1: "WARN", 2: "CRIT"}


@dataclass
class Section:
    name: str
    level: int = LEVEL_OK
    rows:  List[Dict[str, Any]] = field(default_factory=list)
    notes: List[str] = field(default_factory=list)


def get_client(host: str, port: int, user: str, password: str):
    return clickhouse_connect.get_client(
        host=host, port=port, username=user, password=password,
    )


# ---------------------------------------------------------------------------
# 1) 副本延迟 + 队列
# ---------------------------------------------------------------------------
def check_replicas(client) -> Section:
    s = Section("1) 副本同步状态 (system.replicas)")
    rows = client.query("""
        SELECT
            database, table,
            is_readonly, is_session_expired,
            absolute_delay,
            queue_size, inserts_in_queue, merges_in_queue
        FROM system.replicas
        ORDER BY absolute_delay DESC
        LIMIT 20
    """).result_rows
    cols = ["db", "table", "ro", "expired", "delay_s",
            "queue", "ins_q", "mrg_q"]
    for r in rows:
        d = dict(zip(cols, r))
        if d["delay_s"] >= THRESH["replica_delay_crit"] \
                or d["queue"] >= THRESH["queue_size_crit"] \
                or d["expired"] == 1:
            d["__level"] = LEVEL_CRIT
            s.level = max(s.level, LEVEL_CRIT)
        elif d["delay_s"] >= THRESH["replica_delay_warn"] \
                or d["queue"] >= THRESH["queue_size_warn"] \
                or d["ro"] == 1:
            d["__level"] = LEVEL_WARN
            s.level = max(s.level, LEVEL_WARN)
        else:
            d["__level"] = LEVEL_OK
        s.rows.append(d)
    if not rows:
        s.notes.append("(无 Replicated 表, 跳过)")
    return s


# ---------------------------------------------------------------------------
# 2) Part 数量 (找出单分区 Part 数过多)
# ---------------------------------------------------------------------------
def check_parts(client) -> Section:
    s = Section("2) 单分区 Part 数 (system.parts active=1)")
    rows = client.query("""
        SELECT
            database, table, partition,
            count() AS parts,
            sum(rows) AS rows,
            formatReadableSize(sum(bytes_on_disk)) AS size
        FROM system.parts
        WHERE active AND database NOT IN ('system','INFORMATION_SCHEMA')
        GROUP BY database, table, partition
        HAVING parts >= 50
        ORDER BY parts DESC
        LIMIT 20
    """).result_rows
    cols = ["db", "table", "partition", "parts", "rows", "size"]
    for r in rows:
        d = dict(zip(cols, r))
        if d["parts"] >= THRESH["parts_per_partition_crit"]:
            d["__level"] = LEVEL_CRIT
            s.level = max(s.level, LEVEL_CRIT)
        elif d["parts"] >= THRESH["parts_per_partition_warn"]:
            d["__level"] = LEVEL_WARN
            s.level = max(s.level, LEVEL_WARN)
        else:
            d["__level"] = LEVEL_OK
        s.rows.append(d)
    if not rows:
        s.notes.append("所有分区 part 数 < 50, 健康.")
    return s


# ---------------------------------------------------------------------------
# 3) 后台 Merge / Mutation 队列
# ---------------------------------------------------------------------------
def check_merges(client) -> Section:
    s = Section("3) 后台任务 (system.merges + system.mutations)")
    merges = client.query("""
        SELECT database, table, elapsed, progress, num_parts,
               formatReadableSize(total_size_bytes_compressed) AS size,
               is_mutation
        FROM system.merges
        ORDER BY elapsed DESC LIMIT 10
    """).result_rows
    cols = ["db", "table", "elapsed_s", "progress",
            "num_parts", "size", "is_mut"]
    for r in merges:
        d = dict(zip(cols, r))
        d["__level"] = LEVEL_OK
        s.rows.append(d)

    pending = client.query("""
        SELECT database, table, mutation_id,
               command, parts_to_do, latest_fail_reason
        FROM system.mutations
        WHERE is_done = 0
        ORDER BY create_time
    """).result_rows
    cols2 = ["db", "table", "mut_id", "cmd", "parts_left", "fail_reason"]
    for r in pending:
        d = dict(zip(cols2, r))
        d["__level"] = LEVEL_WARN
        s.level = max(s.level, LEVEL_WARN)
        s.rows.append(d)

    if len(pending) >= THRESH["mutation_pending_crit"]:
        s.level = LEVEL_CRIT
    elif len(pending) >= THRESH["mutation_pending_warn"]:
        s.level = max(s.level, LEVEL_WARN)
    if not merges and not pending:
        s.notes.append("无活跃 merge, 无待处理 mutation, 健康.")
    return s


# ---------------------------------------------------------------------------
# 4) 磁盘水位
# ---------------------------------------------------------------------------
def check_disks(client) -> Section:
    s = Section("4) 磁盘水位 (system.disks)")
    rows = client.query("""
        SELECT name, path,
               formatReadableSize(total_space) AS total,
               formatReadableSize(free_space) AS free,
               total_space, free_space, type
        FROM system.disks
    """).result_rows
    cols = ["name", "path", "total", "free", "total_b", "free_b", "type"]
    for r in rows:
        d = dict(zip(cols, r))
        used_pct = 1.0 - (d["free_b"] / max(d["total_b"], 1))
        d["used_pct"] = round(used_pct * 100, 1)
        if used_pct >= THRESH["disk_pct_crit"]:
            d["__level"] = LEVEL_CRIT
            s.level = max(s.level, LEVEL_CRIT)
        elif used_pct >= THRESH["disk_pct_warn"]:
            d["__level"] = LEVEL_WARN
            s.level = max(s.level, LEVEL_WARN)
        else:
            d["__level"] = LEVEL_OK
        d.pop("total_b"); d.pop("free_b")
        s.rows.append(d)
    return s


# ---------------------------------------------------------------------------
# 5) 内存 + Mark Cache 命中
# ---------------------------------------------------------------------------
def check_memory(client) -> Section:
    s = Section("5) 内存与缓存 (system.metrics + system.events)")

    mem_track = client.query("""
        SELECT value FROM system.metrics WHERE metric = 'MemoryTracking'
    """).result_rows[0][0]

    try:
        mem_total_row = client.query("""
            SELECT value FROM system.asynchronous_metrics
            WHERE metric = 'OSMemoryTotal'
        """).result_rows
        mem_total = mem_total_row[0][0] if mem_total_row else None
    except Exception:
        mem_total = None

    mem_pct = None
    if mem_total and mem_total > 0:
        mem_pct = mem_track / mem_total

    s.rows.append({
        "item": "MemoryTracking",
        "value": f"{mem_track / (1024 ** 3):.2f} GiB",
        "of_total_pct": (
            f"{mem_pct * 100:.1f}%" if mem_pct is not None else "n/a"
        ),
        "__level": (
            LEVEL_CRIT if mem_pct and mem_pct >= THRESH["mem_pct_crit"]
            else LEVEL_WARN if mem_pct and mem_pct >= THRESH["mem_pct_warn"]
            else LEVEL_OK
        ),
    })
    if mem_pct and mem_pct >= THRESH["mem_pct_crit"]:
        s.level = max(s.level, LEVEL_CRIT)
    elif mem_pct and mem_pct >= THRESH["mem_pct_warn"]:
        s.level = max(s.level, LEVEL_WARN)

    cache = client.query("""
        SELECT
            sumIf(value, event = 'MarkCacheHits')   AS hits,
            sumIf(value, event = 'MarkCacheMisses') AS misses
        FROM system.events
        WHERE event LIKE 'MarkCache%'
    """).result_rows[0]
    hits, misses = cache
    total = hits + misses
    hit_ratio = hits / total if total else 1.0
    lvl = LEVEL_OK if hit_ratio >= THRESH["mark_cache_hit_warn"] else LEVEL_WARN
    s.level = max(s.level, lvl)
    s.rows.append({
        "item": "MarkCacheHitRatio",
        "value": f"{hit_ratio * 100:.2f}%",
        "of_total_pct": f"hits={hits:,} misses={misses:,}",
        "__level": lvl,
    })
    return s


# ---------------------------------------------------------------------------
# 6) 最近 5 分钟失败查询 Top
# ---------------------------------------------------------------------------
def check_failed(client) -> Section:
    s = Section("6) 最近 5 分钟失败查询 (system.query_log)")
    rows = client.query("""
        SELECT
            count() AS cnt,
            substring(any(exception), 1, 100) AS exception_sample,
            substring(normalizeQuery(any(query)), 1, 100) AS query_sample
        FROM system.query_log
        WHERE event_time > now() - INTERVAL 5 MINUTE
          AND type IN ('ExceptionBeforeStart', 'ExceptionWhileProcessing')
        GROUP BY normalizeQueryHash(query)
        ORDER BY cnt DESC
        LIMIT 10
    """).result_rows
    if not rows:
        s.notes.append("最近 5 分钟无失败查询, 健康.")
    cols = ["cnt", "exception", "query_template"]
    for r in rows:
        d = dict(zip(cols, r))
        d["__level"] = LEVEL_WARN if d["cnt"] >= 5 else LEVEL_OK
        s.level = max(s.level, d["__level"])
        s.rows.append(d)
    return s


# ---------------------------------------------------------------------------
# 报告渲染
# ---------------------------------------------------------------------------
SYMBOL = {LEVEL_OK: "[OK]", LEVEL_WARN: "[WARN]", LEVEL_CRIT: "[CRIT]"}


def render_text(sections: List[Section]) -> str:
    out = []
    overall = max((s.level for s in sections), default=LEVEL_OK)
    out.append("=" * 64)
    out.append(f"ClickHouse 健康巡检报告  总体状态: {SYMBOL[overall]}")
    out.append("=" * 64)
    for s in sections:
        out.append("")
        out.append(f"## {SYMBOL[s.level]} {s.name}")
        for n in s.notes:
            out.append(f"  - {n}")
        if not s.rows:
            continue
        keys = [k for k in s.rows[0].keys() if not k.startswith("__")]
        widths = {k: max(len(k), max(len(str(r.get(k, ""))) for r in s.rows))
                  for k in keys}
        header = "  " + " | ".join(k.ljust(widths[k]) for k in keys)
        sep    = "  " + "-+-".join("-" * widths[k] for k in keys)
        out.append(header)
        out.append(sep)
        for r in s.rows:
            tag = SYMBOL[r.get("__level", LEVEL_OK)]
            line = "  " + " | ".join(
                str(r.get(k, "")).ljust(widths[k]) for k in keys
            )
            out.append(f"{line}   {tag}")
    out.append("")
    out.append("=" * 64)
    return "\n".join(out)


def render_json(sections: List[Section]) -> str:
    overall = max((s.level for s in sections), default=LEVEL_OK)
    body = {
        "overall": LEVEL_NAME[overall],
        "sections": [
            {
                "name": s.name,
                "level": LEVEL_NAME[s.level],
                "notes": s.notes,
                "rows": [
                    {**{k: v for k, v in r.items() if not k.startswith("__")},
                     "level": LEVEL_NAME[r.get("__level", LEVEL_OK)]}
                    for r in s.rows
                ],
            } for s in sections
        ],
    }
    return json.dumps(body, ensure_ascii=False, indent=2)


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--host", default="127.0.0.1")
    ap.add_argument("--port", type=int, default=8123)
    ap.add_argument("--user", default="default")
    ap.add_argument("--password", default="")
    ap.add_argument("--json", action="store_true",
                    help="以 JSON 格式输出, 适合喂给告警系统")
    args = ap.parse_args()

    client = get_client(args.host, args.port, args.user, args.password)
    sections = [
        check_replicas(client),
        check_parts(client),
        check_merges(client),
        check_disks(client),
        check_memory(client),
        check_failed(client),
    ]

    if args.json:
        print(render_json(sections))
    else:
        print(render_text(sections))

    overall = max((s.level for s in sections), default=LEVEL_OK)
    sys.exit(overall)


if __name__ == "__main__":
    main()
