#!/usr/bin/env python3
"""
05_replication_lag.py
=====================
持续监控复制延迟，输出主库各从库的 pending_send / pending_write /
pending_flush / pending_replay 字节数与时间延迟。

用法：
    python3 05_replication_lag.py                # 监控主库 5432
    python3 05_replication_lag.py --interval 0.5 # 0.5s 一次
    python3 05_replication_lag.py --once         # 只跑一次
    PG_DSN='host=... port=...' python3 05_replication_lag.py

依赖：psycopg[binary]>=3.1
"""
import argparse
import os
import sys
import time
from datetime import datetime

import psycopg

DEFAULT_DSN = "host=127.0.0.1 port=5432 dbname=learn_pg user=postgres password=postgres"


SQL = """
SELECT
    application_name,
    client_addr::text,
    state,
    sync_state,
    pg_wal_lsn_diff(pg_current_wal_lsn(), sent_lsn)    AS pending_send_bytes,
    pg_wal_lsn_diff(sent_lsn, write_lsn)               AS pending_write_bytes,
    pg_wal_lsn_diff(write_lsn, flush_lsn)              AS pending_flush_bytes,
    pg_wal_lsn_diff(flush_lsn, replay_lsn)             AS pending_replay_bytes,
    pg_wal_lsn_diff(pg_current_wal_lsn(), replay_lsn)  AS total_lag_bytes,
    EXTRACT(EPOCH FROM write_lag)::float  AS write_lag_s,
    EXTRACT(EPOCH FROM flush_lag)::float  AS flush_lag_s,
    EXTRACT(EPOCH FROM replay_lag)::float AS replay_lag_s
FROM pg_stat_replication
ORDER BY application_name;
"""


def human_bytes(n) -> str:
    if n is None:
        return "-"
    n = float(n)
    for unit in ["B", "KB", "MB", "GB", "TB"]:
        if abs(n) < 1024.0:
            return f"{n:6.1f}{unit}"
        n /= 1024.0
    return f"{n:6.1f}PB"


def human_secs(s) -> str:
    if s is None:
        return "-"
    if s < 1:
        return f"{int(s*1000)}ms"
    if s < 60:
        return f"{s:.1f}s"
    return f"{s/60:.1f}m"


def render(rows) -> None:
    ts = datetime.now().strftime("%H:%M:%S")
    print(f"\n[{ts}]  {len(rows)} 个从库连接")
    if not rows:
        print("  (无从库连接)")
        return
    print(f"  {'app':<20} {'addr':<16} {'state':<10} {'sync':<10}"
          f" {'send':>9} {'write':>9} {'flush':>9} {'replay':>9} {'total':>9}"
          f" | {'wlag':>6} {'flag':>6} {'rlag':>6}")
    print("  " + "-" * 130)
    for r in rows:
        (app, addr, state, sync,
         ps, pw, pf, pr, total,
         wlag, flag, rlag) = r
        print(f"  {(app or '-'):<20} {(addr or '-'):<16} {(state or '-'):<10} {(sync or '-'):<10}"
              f" {human_bytes(ps):>9} {human_bytes(pw):>9} {human_bytes(pf):>9}"
              f" {human_bytes(pr):>9} {human_bytes(total):>9}"
              f" | {human_secs(wlag):>6} {human_secs(flag):>6} {human_secs(rlag):>6}")


def main() -> int:
    parser = argparse.ArgumentParser()
    parser.add_argument("--interval", type=float, default=2.0, help="刷新间隔秒数")
    parser.add_argument("--once", action="store_true", help="只跑一次")
    parser.add_argument("--dsn", default=os.environ.get("PG_DSN", DEFAULT_DSN))
    args = parser.parse_args()

    try:
        conn = psycopg.connect(args.dsn)
    except psycopg.Error as e:
        print(f"[FATAL] 连接失败: {e}", file=sys.stderr)
        return 1

    try:
        while True:
            with conn.cursor() as cur:
                cur.execute(SQL)
                rows = cur.fetchall()
            conn.commit()  # 释放快照
            render(rows)
            if args.once:
                break
            time.sleep(args.interval)
    except KeyboardInterrupt:
        print("\n[bye]")
    finally:
        conn.close()
    return 0


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