# shared.py version: 1.6

import asyncio
import asyncpg
import logging
import os
import gc
import random
from datetime import datetime

import config
from pyrogram.errors import (
    UserDeactivated, UserDeactivatedBan, UserInvalid,
    PeerIdInvalid, InputUserDeactivated, SessionRevoked,
    SessionExpired, AuthKeyUnregistered, AuthKeyDuplicated,
    UserBannedInChannel, FloodWait
)

log = logging.getLogger(__name__)

# ← إعدادات PostgreSQL
PG_DSN = "postgresql://mahde:MahdeFab1YssR@localhost/mahde"

BOT_DELAY = 10
BATCH_REST = 300

# ← Pool مشترك بدل db_lock
_pg_pool: asyncpg.Pool = None

clients = {}
flood_map = {}
stop_event = asyncio.Event()
account_join_count = {}
_join_notify_buffer: list = []
_join_notify_lock: asyncio.Lock = None

def get_join_notify_lock() -> asyncio.Lock:
    global _join_notify_lock
    if _join_notify_lock is None:
        _join_notify_lock = asyncio.Lock()
    return _join_notify_lock


dead_sessions: set = set()
account_locks: dict = {}

def get_account_lock(phone: str) -> asyncio.Lock:
    if phone not in account_locks:
        account_locks[phone] = asyncio.Lock()
    log.debug(f"[LOCK] account_lock {phone} ▶")
    return account_locks[phone]

all_bots_task = None
all_bots_on = False
leave_all_task = None
leave_all_on = False
leave_all_paused = False
_background_tasks: set = set()
_admins_cache: list = []

_get_client_locks: dict = {}

def get_client_open_lock(phone: str) -> asyncio.Lock:
    if phone not in _get_client_locks:
        _get_client_locks[phone] = asyncio.Lock()
    log.debug(f"[LOCK] client_open_lock {phone} ▶")
    return _get_client_locks[phone]
_banned_channels: set = set()

def is_channel_banned(target: str) -> bool:
    return norm_ch_key(target) in _banned_channels

async def ban_channel(target: str):
    key = norm_ch_key(target)
    _banned_channels.add(key)
    log.warning(f"[CHANNEL BAN] قناة محظورة: {key}")
    async with _pg_pool.acquire() as con:
        await con.execute(
            "INSERT INTO banned_channels VALUES ($1,$2) ON CONFLICT DO NOTHING",
            key, datetime.now().isoformat()
        )


async def init_db():
    """إنشاء Pool والجداول إذا لم تكن موجودة"""
    global _pg_pool
    _pg_pool = await asyncpg.create_pool(
        PG_DSN,
        min_size=5,
        max_size=20,
        command_timeout=30
    )
    async with _pg_pool.acquire() as con:
        await con.execute("""
            CREATE TABLE IF NOT EXISTS settings (
                key TEXT PRIMARY KEY,
                value TEXT
            );
            CREATE TABLE IF NOT EXISTS admins (
                user_id BIGINT PRIMARY KEY
            );
            CREATE TABLE IF NOT EXISTS accounts (
                phone TEXT PRIMARY KEY,
                banned INTEGER DEFAULT 0
            );
            CREATE TABLE IF NOT EXISTS banned_accounts (
                phone TEXT PRIMARY KEY,
                banned_at TEXT,
                reason TEXT
            );
            CREATE TABLE IF NOT EXISTS bots (
                username TEXT PRIMARY KEY,
                url TEXT
            );
            CREATE TABLE IF NOT EXISTS channels (
                id BIGINT,
                title TEXT,
                username TEXT,
                url TEXT,
                joined_at TEXT,
                PRIMARY KEY (id, url)
            );
            CREATE TABLE IF NOT EXISTS force_sub_channels (
                bot_username TEXT NOT NULL,
                url TEXT NOT NULL,
                channel_id BIGINT,
                PRIMARY KEY (bot_username, url)
            );
        """)
        # إدراج الإعدادات الافتراضية
        defaults = [
            ("gathering_speed", "5"),
            ("concurrent", "20"),
            ("is_active", "0"),
            ("work_minutes", "15"),
            ("rest_minutes", "5"),
            ("active_bot", ""),
        ]
        for k, v in defaults:
            await con.execute(
                "INSERT INTO settings VALUES ($1,$2) ON CONFLICT(key) DO NOTHING",
                k, v
            )
        # إعادة تعيين حالة التجميع عند الإقلاع
        await con.execute(
            "INSERT INTO settings VALUES ($1,$2) "
            "ON CONFLICT(key) DO UPDATE SET value=$2",
            "is_active", "0"
        )
        await con.execute(
            "INSERT INTO settings VALUES ($1,$2) "
            "ON CONFLICT(key) DO UPDATE SET value=$2",
            "active_bot", ""
        )
    rows = await _pg_pool.fetch("SELECT channel_key FROM banned_channels")
    for r in rows:
        _banned_channels.add(r["channel_key"])
    log.info(f"[DB] تم تحميل {len(_banned_channels)} قناة محظورة")
    log.info("[DB] PostgreSQL pool جاهز")


def get_admins_sync():
    """للاستخدام عند الإقلاع قبل بدء event loop"""
    import psycopg2
    try:
        con = psycopg2.connect(PG_DSN)
        cur = con.cursor()
        cur.execute("SELECT user_id FROM admins")
        rows = cur.fetchall()
        con.close()
        return [r[0] for r in rows]
    except Exception:
        return []


async def get_setting(key, default=None):
    async with _pg_pool.acquire() as con:
        row = await con.fetchrow(
            "SELECT value FROM settings WHERE key=$1", key
        )
        return row["value"] if row else default


async def set_setting(key, value):
    async with _pg_pool.acquire() as con:
        await con.execute(
            "INSERT INTO settings VALUES ($1,$2) "
            "ON CONFLICT(key) DO UPDATE SET value=$2",
            key, str(value)
        )


async def get_speed():
    try:
        return max(int(await get_setting("gathering_speed", 5)), 0)
    except Exception:
        return 5


async def get_concurrent():
    try:
        return max(int(await get_setting("concurrent", 20)), 1)
    except Exception:
        return 20

async def get_concurrent_all_bots() -> int:
    try:
        return max(int(await get_setting("concurrent_all_bots", 60)), 1)
    except Exception:
        return 60


async def get_work_minutes():
    try:
        return max(int(await get_setting("work_minutes", 15)), 1)
    except Exception:
        return 15


async def get_rest_minutes():
    try:
        return max(int(await get_setting("rest_minutes", 5)), 1)
    except Exception:
        return 5


async def get_accounts(active_only=False):
    async with _pg_pool.acquire() as con:
        if active_only:
            rows = await con.fetch(
                "SELECT phone FROM accounts WHERE banned=0"
            )
            return [{"phone": r["phone"]} for r in rows]
        rows = await con.fetch(
            "SELECT phone, banned FROM accounts WHERE banned=0"
        )
        return [{"phone": r["phone"], "banned": bool(r["banned"])} for r in rows]


async def get_active_count():
    async with _pg_pool.acquire() as con:
        row = await con.fetchrow(
            "SELECT COUNT(*) FROM accounts WHERE banned=0"
        )
        return row["count"] if row else 0


async def add_account(phone):
    async with _pg_pool.acquire() as con:
        await con.execute(
            "INSERT INTO accounts VALUES ($1,0) ON CONFLICT DO NOTHING",
            phone
        )


async def del_account(phone):
    async with _pg_pool.acquire() as con:
        await con.execute(
            "DELETE FROM accounts WHERE phone=$1", phone
        )


async def ban_account(phone, reason="محظور") -> bool:
    """
    يعيد True إذا كان الحظر جديداً
    يعيد False إذا كان محظوراً مسبقاً
    """
    log.warning(f"[BAN] {phone} ▶ حظر: {reason}")
    async with _pg_pool.acquire() as con:
        already = await con.fetchrow(
            "SELECT 1 FROM banned_accounts WHERE phone=$1", phone
        )
        await con.execute(
            "DELETE FROM accounts WHERE phone=$1", phone
        )
        await con.execute(
            "INSERT INTO banned_accounts VALUES ($1,$2,$3) "
            "ON CONFLICT(phone) DO UPDATE SET reason=$3",
            phone, datetime.now().isoformat(), reason
        )
    if already:
        log.warning(f"[BAN] {phone} ■ كان محظوراً مسبقاً")
        return False
    await cleanup_client(phone)
    # لا نحذف ملف الجلسة نهائياً بل ننقله إلى حجر صحي —
    # سبب "الحظر" قد يكون خطأً عابراً أو تشخيصاً خاطئاً، وحذف الملف
    # يجعل الخسارة نهائية وغير قابلة للتراجع.
    qdir = os.path.join("sessions", "quarantine", phone)
    os.makedirs(qdir, exist_ok=True)
    moved = 0
    for name in [f"{phone}.session", f"{phone}.session-journal"]:
        path = os.path.join("sessions", name)
        try:
            if os.path.exists(path):
                os.rename(path, os.path.join(qdir, name))
                moved += 1
        except Exception as e:
            log.error(f"[BAN] {phone} ✗ نقل لل quarantين فشل: {e}")
    log.info(f"[BAN] {phone} ✓ جلسة نُقلت لل quarantين ({moved} ملف)")
    return True


async def get_banned():
    async with _pg_pool.acquire() as con:
        rows = await con.fetch(
            "SELECT phone, banned_at, reason FROM banned_accounts"
        )
        return [
            {
                "phone": r["phone"],
                "banned_at": r["banned_at"],
                "reason": r["reason"] or "محظور"
            }
            for r in rows
        ]


async def get_bots():
    async with _pg_pool.acquire() as con:
        rows = await con.fetch("SELECT username, url FROM bots")
        return [{"username": r["username"], "url": r["url"]} for r in rows]


async def add_bot(username, url):
    async with _pg_pool.acquire() as con:
        await con.execute(
            "INSERT INTO bots VALUES ($1,$2) ON CONFLICT DO NOTHING",
            username, url
        )


async def del_bot(username):
    async with _pg_pool.acquire() as con:
        await con.execute(
            "DELETE FROM bots WHERE username=$1", username
        )


async def get_admins():
    async with _pg_pool.acquire() as con:
        rows = await con.fetch("SELECT user_id FROM admins")
        return [r["user_id"] for r in rows]


async def add_admin(uid):
    global _admins_cache
    async with _pg_pool.acquire() as con:
        await con.execute(
            "INSERT INTO admins VALUES ($1) ON CONFLICT DO NOTHING",
            uid
        )
    _admins_cache = await get_admins()


async def del_admin(uid):
    global _admins_cache
    async with _pg_pool.acquire() as con:
        await con.execute(
            "DELETE FROM admins WHERE user_id=$1", uid
        )
    _admins_cache = await get_admins()


async def get_channels():
    async with _pg_pool.acquire() as con:
        rows = await con.fetch(
            "SELECT id, title, username, url, joined_at FROM channels"
        )
        return [
            {
                "id": r["id"],
                "title": r["title"],
                "username": r["username"],
                "url": r["url"],
                "joined_at": r["joined_at"]
            }
            for r in rows
        ]


async def add_channel(ch: dict):
    ch_id = ch.get("id") or 0
    url = ch.get("url") or ""
    async with _pg_pool.acquire() as con:
        exists = await con.fetchrow(
            "SELECT 1 FROM channels WHERE (id=$1 AND id!=0) OR (url=$2 AND url!='')",
            ch_id, url
        )
        if not exists:
            await con.execute(
                "INSERT INTO channels VALUES ($1,$2,$3,$4,$5) ON CONFLICT DO NOTHING",
                ch_id,
                ch.get("title"),
                ch.get("username"),
                url,
                ch.get("joined_at", datetime.now().isoformat())
            )


async def del_channel(ch_id=None, url=None):
    async with _pg_pool.acquire() as con:
        if ch_id:
            await con.execute(
                "DELETE FROM channels WHERE id=$1", ch_id
            )
        elif url:
            await con.execute(
                "DELETE FROM channels WHERE url=$1", url
            )


async def add_force_sub_channel(bot_username: str, url: str, channel_id: int = None):
    async with _pg_pool.acquire() as con:
        await con.execute(
            "INSERT INTO force_sub_channels (bot_username, url, channel_id) "
            "VALUES ($1,$2,$3) ON CONFLICT (bot_username, url) "
            "DO UPDATE SET channel_id = COALESCE(EXCLUDED.channel_id, force_sub_channels.channel_id)",
            bot_username.lower().strip('@'), url.strip(), channel_id
        )


async def get_force_sub_channels(bot_username: str = None) -> list:
    async with _pg_pool.acquire() as con:
        if bot_username:
            rows = await con.fetch(
                "SELECT bot_username, url, channel_id FROM force_sub_channels WHERE bot_username=$1",
                bot_username.lower().strip('@')
            )
        else:
            rows = await con.fetch(
                "SELECT bot_username, url, channel_id FROM force_sub_channels"
            )
        return [dict(r) for r in rows]


async def get_all_force_sub_urls() -> set:
    async with _pg_pool.acquire() as con:
        rows = await con.fetch("SELECT url FROM force_sub_channels")
        return {r['url'].strip() for r in rows}


async def get_all_force_sub_keys() -> tuple:
    """يعيد (مجموعة مفاتيح موحّدة عبر norm_ch_key، مجموعة channel_id) لقنوات
    الاشتراك الإجباري. يُستخدم للمقارنة الدقيقة بدل تطابق نصي خام على الرابط،
    لأن الرابط المخزَّن هنا خام كما ورد من الرسالة (قد يختلف بالحالة/الصيغة
    عن الرابط الطبيعي المخزَّن في جدول channels)."""
    async with _pg_pool.acquire() as con:
        rows = await con.fetch("SELECT url, channel_id FROM force_sub_channels")
        keys = {norm_ch_key(r['url']) for r in rows if r['url']}
        ids = {int(r['channel_id']) for r in rows if r['channel_id'] is not None}
        return keys, ids


async def del_force_sub_channels(bot_username: str):
    async with _pg_pool.acquire() as con:
        await con.execute(
            "DELETE FROM force_sub_channels WHERE bot_username=$1",
            bot_username.lower().strip('@')
        )


async def cleanup_client(phone: str, client=None):
    log.info(f"[CLEANUP] {phone} ▶ بدء التنظيف")
    lock = get_client_open_lock(phone)
    async with lock:
        current = clients.get(phone)
        if current is None:
            log.info(f"[CLEANUP] {phone} ■ لا يوجد عميل في الكاش")
            return
        # حارس الهوية: لا نُغلق عميلاً حلّ غيره في الكاش منذ استدعاء هذا الإغلاق
        # (أخطر مصدر لفتح مزدوج داخلي: worker أنشأ عميلاً جديداً بينما يُغلق القديم).
        if client is not None and current is not client:
            log.warning(f"[CLEANUP] {phone} ■ حارس الهوية: عميل مختلف، تجاهل")
            return
        clients.pop(phone, None)
        try:
            await asyncio.wait_for(client.stop(), timeout=10)
            log.info(f"[CLEANUP] {phone} ✓ stop() نجح")
        except Exception as e:
            log.warning(f"[CLEANUP] {phone} ⚠️ stop() فشل: {e}")
            # لو فشل stop أو تجاوز المهلة، نضمن قطع الاتصال يدوياً
            # حتى لا يبقى socket/جلسة عالقة في حلقة إعادة محاولة.
            try:
                await asyncio.wait_for(client.disconnect(), timeout=10)
                log.info(f"[CLEANUP] {phone} ✓ disconnect() نجح")
            except Exception as e2:
                log.warning(f"[CLEANUP] {phone} ✗ disconnect() فشل: {e2}")


def create_tracked_task(coro):
    task = asyncio.create_task(coro)
    _background_tasks.add(task)
    task.add_done_callback(_background_tasks.discard)
    return task


def rand_delay(base: float) -> float:
    low = max(1.0, base * 0.6)
    val = random.uniform(low, base)
    log.debug(f"[RAND_DELAY] base={base} → {val:.2f}s")
    return val


async def safe_sleep(seconds: float):
    log.debug(f"[SAFE_SLEEP] {seconds:.1f}s ▶")
    elapsed = 0
    interval = 1.0
    while elapsed < seconds:
        if stop_event.is_set():
            log.debug(f"[SAFE_SLEEP] ■ stop_event بعد {elapsed:.1f}s")
            return
        await asyncio.sleep(min(interval, seconds - elapsed))
        elapsed += interval
    log.debug(f"[SAFE_SLEEP] {seconds:.1f}s ✓")


def ch_link(ch):
    return ch.get("url") or (
        f"https://t.me/{ch['username']}" if ch.get("username") else None
    )


def norm_target(value):
    if not value:
        return value
    value = value.strip()
    for p in ["https://", "http://", "t.me/", "telegram.me/"]:
        if value.startswith(p):
            value = value[len(p):]
    if value.startswith("@"):
        value = value[1:]
    if "?" in value:
        value = value.split("?", 1)[0]
    if value.startswith("+") or value.startswith("joinchat/"):
        return f"https://t.me/{value}"
    if "/" in value:
        value = value.split("/", 1)[0]
    return value


def norm_ch_key(target: str) -> str:
    if not target:
        return ""
    target = target.strip().lower()
    for p in ["https://t.me/", "http://t.me/", "t.me/", "@"]:
        if target.startswith(p):
            target = target[len(p):]
    return target.split("?")[0].rstrip("/")


def parse_bot_input(value: str) -> dict:
    value = value.strip()
    username = None
    if value.startswith("https://") or value.startswith("http://"):
        raw = value.replace("https://", "").replace("http://", "")
        for p in ["t.me/", "telegram.me/"]:
            if raw.startswith(p):
                raw = raw[len(p):]
        raw = raw.split("?")[0].split("/")[0].lstrip("@")
        username = raw
    elif value.startswith("@"):
        username = value[1:]
    else:
        username = value
    return {
        "username": username,
        "url": f"https://t.me/{username}" if username else None
    }


BAN_REASONS = {
    "UserDeactivatedBan": "محظور نهائياً",
    "UserDeactivated": "تم تجميده",
    "UserInvalid": "حساب غير صالح (محذوف)",
    "InputUserDeactivated": "مستخدم معطل",
    "PeerIdInvalid": "معرف غير صالح (حساب محذوف)",
    "AccountBanned": "محظور",
    "AuthKeyUnregistered": "الجلسة منتهية",
    "SessionRevoked": "تم إلغاء الجلسة",
    "SessionExpired": "انتهت صلاحية الجلسة",
    "AuthKeyDuplicated": "جلسة مكررة",
    "UserBannedInChannel": "محظور في القناة",
    "AccountInvalidated": "الحساب غير صالح",
    "DEACTIVATED": "تم تجميده",
    "BANNED": "محظور",
    "AUTH_KEY_UNREGISTERED": "الجلسة منتهية",
    "SESSION_REVOKED": "تم إلغاء الجلسة",
    "SESSION_EXPIRED": "انتهت صلاحية الجلسة",
    "AUTH_KEY_DUPLICATED": "جلسة مكررة",
    "USER_INVALID": "حساب غير صالح",
    "INPUT_USER_DEACTIVATED": "مستخدم معطل",
    "PEER_ID_INVALID": "معرف غير صالح",
}


def get_ban_reason(e) -> str:
    type_name = type(e).__name__
    if type_name in BAN_REASONS:
        return BAN_REASONS[type_name]
    err = str(e).upper()
    for key, reason in BAN_REASONS.items():
        if key.upper() in err:
            return reason
    return "محظور"


def is_ban_error(e) -> bool:
    type_name = type(e).__name__
    if type_name in BAN_REASONS:
        if type_name in ("UserBannedInChannel",):
            return False
        return True
    err = str(e).upper()
    ban_keywords = [
        "DEACTIVATED", "AUTH_KEY_UNREGISTERED",
        "SESSION_REVOKED", "SESSION_EXPIRED", "AUTH_KEY_DUPLICATED",
        "ACCOUNT_INVALID", "USER_DEACTIVATED", "USER_INVALID",
        "INPUT_USER_DEACTIVATED", "USER_IS_BLOCKED",
    ]
    return any(k in err for k in ban_keywords)


async def verify_account_validity(phone: str) -> tuple[bool, str]:
    try:
        from pyrogram import Client
        c = Client(
            f"sessions/{phone}",
            api_id=config.API_ID,
            api_hash=config.API_HASH,
            no_updates=True
        )
        await c.start()
        me = await c.get_me()
        await c.stop()
        if not me:
            await ban_account(phone, "حساب غير موجود")
            return False, "حساب غير موجود"
        return True, ""
    except Exception as e:
        if is_ban_error(e):
            reason = get_ban_reason(e)
            await ban_account(phone, reason)
            return False, reason
        return False, str(e)

