#!/usr/bin/env python3
"""Show / optionally terminate Postgres sessions blocking channel_listing_path.

Usage:
  python -m app.jobs.db_lock_inspect
  python -m app.jobs.db_lock_inspect --relation channel_listing_path
  python -m app.jobs.db_lock_inspect --terminate-blockers
"""

from __future__ import annotations

import argparse
import json
import sys

from sqlalchemy import text

from db.session import get_session_factory


BLOCKERS_SQL = text(
    """
    SELECT
      blocked.pid AS blocked_pid,
      blocked.usename AS blocked_user,
      blocked.application_name AS blocked_app,
      left(blocked.query, 160) AS blocked_query,
      now() - blocked.xact_start AS blocked_xact_age,
      blocking.pid AS blocking_pid,
      blocking.usename AS blocking_user,
      blocking.application_name AS blocking_app,
      blocking.state AS blocking_state,
      now() - blocking.xact_start AS blocking_xact_age,
      left(blocking.query, 160) AS blocking_query
    FROM pg_stat_activity blocked
    JOIN pg_locks bl ON bl.pid = blocked.pid AND NOT bl.granted
    JOIN pg_locks kl
      ON kl.locktype = bl.locktype
     AND kl.database IS NOT DISTINCT FROM bl.database
     AND kl.relation IS NOT DISTINCT FROM bl.relation
     AND kl.page IS NOT DISTINCT FROM bl.page
     AND kl.tuple IS NOT DISTINCT FROM bl.tuple
     AND kl.virtualxid IS NOT DISTINCT FROM bl.virtualxid
     AND kl.transactionid IS NOT DISTINCT FROM bl.transactionid
     AND kl.classid IS NOT DISTINCT FROM bl.classid
     AND kl.objid IS NOT DISTINCT FROM bl.objid
     AND kl.objsubid IS NOT DISTINCT FROM bl.objsubid
     AND kl.pid <> bl.pid
     AND kl.granted
    JOIN pg_stat_activity blocking ON blocking.pid = kl.pid
    WHERE blocked.datname = current_database()
    ORDER BY blocking.xact_start NULLS LAST
    """
)

RELATION_LOCKS_SQL = text(
    """
    SELECT
      a.pid,
      a.application_name,
      a.state,
      a.wait_event_type,
      a.wait_event,
      now() - a.xact_start AS xact_age,
      now() - a.query_start AS query_age,
      l.mode,
      l.granted,
      left(a.query, 180) AS query
    FROM pg_locks l
    JOIN pg_class c ON c.oid = l.relation
    JOIN pg_stat_activity a ON a.pid = l.pid
    WHERE c.relname = :relation
      AND a.datname = current_database()
    ORDER BY l.granted DESC, a.xact_start NULLS LAST
    """
)

LONG_TX_SQL = text(
    """
    SELECT
      pid,
      application_name,
      state,
      wait_event_type,
      wait_event,
      now() - xact_start AS xact_age,
      now() - query_start AS query_age,
      left(query, 180) AS query
    FROM pg_stat_activity
    WHERE datname = current_database()
      AND pid <> pg_backend_pid()
      AND xact_start IS NOT NULL
    ORDER BY xact_start
    """
)


def main() -> int:
    parser = argparse.ArgumentParser(description="Inspect Postgres lock blockers for taxonomy import.")
    parser.add_argument("--relation", default="channel_listing_path")
    parser.add_argument(
        "--terminate-blockers",
        action="store_true",
        help="pg_terminate_backend() on distinct blocking_pids from the blocker query.",
    )
    parser.add_argument(
        "--cancel-blockers",
        action="store_true",
        help="pg_cancel_backend() on distinct blocking_pids (gentler than terminate).",
    )
    args = parser.parse_args()

    session = get_session_factory()()
    try:
        blockers = [dict(row) for row in session.execute(BLOCKERS_SQL).mappings().all()]
        relation_locks = [
            dict(row) for row in session.execute(RELATION_LOCKS_SQL, {"relation": args.relation}).mappings().all()
        ]
        long_tx = [dict(row) for row in session.execute(LONG_TX_SQL).mappings().all()]

        blocking_pids = sorted({int(row["blocking_pid"]) for row in blockers if row.get("blocking_pid")})
        actions = []
        if args.cancel_blockers or args.terminate_blockers:
            for pid in blocking_pids:
                if args.cancel_blockers:
                    ok = session.execute(text("SELECT pg_cancel_backend(:pid)"), {"pid": pid}).scalar()
                    actions.append({"pid": pid, "action": "cancel", "ok": bool(ok)})
                if args.terminate_blockers:
                    ok = session.execute(text("SELECT pg_terminate_backend(:pid)"), {"pid": pid}).scalar()
                    actions.append({"pid": pid, "action": "terminate", "ok": bool(ok)})
            session.commit()

        print(
            json.dumps(
                {
                    "relation": args.relation,
                    "blocker_count": len(blockers),
                    "blocking_pids": blocking_pids,
                    "blockers": blockers,
                    "relation_locks": relation_locks,
                    "long_transactions": long_tx,
                    "actions": actions,
                },
                indent=2,
                default=str,
            )
        )
        return 0
    finally:
        session.close()


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