Spaces:
Running
Running
| import threading | |
| from sqlalchemy import create_engine, Column, TEXT, Boolean, Numeric, BigInteger, Integer | |
| from sqlalchemy.ext.declarative import declarative_base | |
| from sqlalchemy.orm import sessionmaker, scoped_session | |
| from sqlalchemy.pool import QueuePool | |
| from sqlalchemy.orm.exc import NoResultFound | |
| from mfinder import DB_URL, LOGGER | |
| BASE = declarative_base() | |
| def merge_env_shortener_into_list(shorteners_json, env_site, env_api, env_download): | |
| import json | |
| try: | |
| shorteners = json.loads(shorteners_json) if shorteners_json else [] | |
| except Exception: | |
| shorteners = [] | |
| parsed_site = None | |
| if env_site: | |
| ss = env_site.strip() | |
| parsed_site = None if ss.lower() in ("off", "disable", "disabled", "none") else ss | |
| parsed_api = env_api.strip() if env_api else None | |
| parsed_hd = None | |
| if env_download: | |
| hd = env_download.strip() | |
| parsed_hd = None if hd.lower() in ("off", "disable", "disabled", "none") else hd | |
| if not parsed_site or not parsed_api: | |
| return shorteners_json, False | |
| exists = False | |
| updated = False | |
| for s in shorteners: | |
| if s.get("site") == parsed_site: | |
| if s.get("api") != parsed_api: | |
| s["api"] = parsed_api | |
| updated = True | |
| if parsed_hd is not None and s.get("how_to_download") != parsed_hd: | |
| s["how_to_download"] = parsed_hd | |
| updated = True | |
| exists = True | |
| break | |
| if not exists: | |
| shorteners.append({ | |
| "site": parsed_site, | |
| "api": parsed_api, | |
| "how_to_download": parsed_hd, | |
| "status": True | |
| }) | |
| updated = True | |
| return json.dumps(shorteners), updated | |
| class AdminSettings(BASE): | |
| __tablename__ = "admin_settings" | |
| setting_name = Column(TEXT, primary_key=True) | |
| auto_delete = Column(Numeric) | |
| custom_caption = Column(TEXT) | |
| fsub_channel = Column(Numeric) | |
| channel_link = Column(TEXT) | |
| caption_uname = Column(TEXT) | |
| repair_mode = Column(Boolean) | |
| shortener_status = Column(Boolean) | |
| admins_list = Column(TEXT) | |
| db_channels_list = Column(TEXT) | |
| token_shortener_enabled = Column(Boolean) | |
| token_timeout = Column(Numeric) | |
| log_channel = Column(Numeric) | |
| log_channel_access_hash = Column(Numeric) | |
| fsub_channel_access_hash = Column(Numeric) | |
| fsub_channels_list = Column(TEXT) | |
| fsub_channels_access_hashes = Column(TEXT) | |
| shorteners_list = Column(TEXT) | |
| smart_rotator = Column(Boolean, default=False) | |
| newsletter_days = Column(TEXT) | |
| newsletter_time = Column(TEXT) | |
| newsletter_count = Column(Integer) | |
| newsletter_enabled = Column(Boolean, default=False) | |
| newsletter_target = Column(TEXT) | |
| announcement_enabled = Column(Boolean, default=False) | |
| announcement_channel = Column(Numeric) | |
| announcement_channel_access_hash = Column(Numeric) | |
| def __init__(self, setting_name="default"): | |
| self.setting_name = setting_name | |
| self.auto_delete = 0 | |
| self.custom_caption = None | |
| self.fsub_channel = None | |
| self.channel_link = None | |
| self.caption_uname = None | |
| self.repair_mode = False | |
| self.shortener_site = None | |
| self.shortener_api = None | |
| self.shortener_status = True | |
| self.admins_list = None | |
| self.db_channels_list = None | |
| self.token_shortener_enabled = False | |
| self.token_timeout = 3600 | |
| self.log_channel = None | |
| self.how_to_download = None | |
| self.log_channel_access_hash = None | |
| self.fsub_channel_access_hash = None | |
| self.fsub_channels_list = None | |
| self.fsub_channels_access_hashes = "{}" | |
| self.shorteners_list = "[]" | |
| self.smart_rotator = False | |
| self.newsletter_days = "Friday" | |
| self.newsletter_time = "20:00" | |
| self.newsletter_count = 10 | |
| self.newsletter_enabled = False | |
| self.newsletter_target = "subscribed" | |
| self.newsletter_days = "Friday" | |
| self.newsletter_time = "20:00" | |
| self.newsletter_count = 10 | |
| self.newsletter_enabled = False | |
| class PendingFsubRequests(BASE): | |
| __tablename__ = "pending_fsub_requests" | |
| user_id = Column(Numeric, primary_key=True) | |
| search_query = Column(TEXT) | |
| target_file_id = Column(TEXT) | |
| created_at = Column(Numeric) | |
| def __init__(self, user_id, search_query=None, target_file_id=None): | |
| self.user_id = user_id | |
| self.search_query = search_query | |
| self.target_file_id = target_file_id | |
| import time | |
| self.created_at = time.time() | |
| class Settings(BASE): | |
| __tablename__ = "settings" | |
| user_id = Column(BigInteger, primary_key=True) | |
| precise_mode = Column(Boolean) | |
| button_mode = Column(Boolean) | |
| link_mode = Column(Boolean) | |
| list_mode = Column(Boolean) | |
| def __init__(self, user_id, precise_mode, button_mode, link_mode, list_mode): | |
| self.user_id = user_id | |
| self.precise_mode = precise_mode | |
| self.button_mode = button_mode | |
| self.link_mode = link_mode | |
| self.list_mode = list_mode | |
| class UserTokens(BASE): | |
| __tablename__ = "user_tokens" | |
| user_id = Column(BigInteger, primary_key=True) | |
| token = Column(TEXT) | |
| created_at = Column(Numeric) | |
| expires_at = Column(Numeric) | |
| status = Column(TEXT) | |
| search_query = Column(TEXT) | |
| target_file_id = Column(TEXT) | |
| req_msg_id = Column(Numeric, nullable=True) | |
| def __init__(self, user_id, token, created_at, expires_at=0, status="pending", search_query=None, target_file_id=None, req_msg_id=None): | |
| self.user_id = user_id | |
| self.token = token | |
| self.created_at = created_at | |
| self.expires_at = expires_at | |
| self.status = status | |
| self.search_query = search_query | |
| self.target_file_id = target_file_id | |
| self.req_msg_id = req_msg_id | |
| class UserShortenerUsage(BASE): | |
| __tablename__ = "user_shortener_usage" | |
| user_id = Column(BigInteger, primary_key=True) | |
| last_index = Column(Integer, default=0) | |
| last_date = Column(TEXT) # Format: YYYY-MM-DD | |
| def __init__(self, user_id, last_index=0, last_date=""): | |
| self.user_id = user_id | |
| self.last_index = last_index | |
| self.last_date = last_date | |
| # SQLite configuration for auto-delete queue to save Neon compute hours | |
| SQLITE_DB_URL = "sqlite:///mfinder_autodelete.db" | |
| sqlite_engine = create_engine(SQLITE_DB_URL, pool_pre_ping=True) | |
| sqlite_session_maker = sessionmaker(bind=sqlite_engine, autoflush=False) | |
| SQLITE_SESSION = scoped_session(sqlite_session_maker) | |
| SQLITE_BASE = declarative_base() | |
| class AutoDeleteMessage(SQLITE_BASE): | |
| __tablename__ = "auto_delete_messages" | |
| id = Column(Integer, primary_key=True, autoincrement=True) | |
| chat_id = Column(Numeric) | |
| message_id = Column(Integer) | |
| delete_at = Column(Numeric) # Unix timestamp | |
| def __init__(self, chat_id, message_id, delete_at): | |
| self.chat_id = chat_id | |
| self.message_id = message_id | |
| self.delete_at = delete_at | |
| def start() -> scoped_session: | |
| connect_args = {} | |
| if DB_URL and DB_URL.startswith(("postgres://", "postgresql://")): | |
| connect_args["sslmode"] = "require" | |
| engine = create_engine( | |
| DB_URL, | |
| connect_args=connect_args, | |
| poolclass=QueuePool, | |
| pool_size=10, | |
| max_overflow=20, | |
| pool_pre_ping=True, | |
| pool_recycle=1800 # Recycle connections every 1800 seconds (30 minutes) | |
| ) | |
| BASE.metadata.bind = engine | |
| BASE.metadata.create_all(engine) | |
| SQLITE_BASE.metadata.bind = sqlite_engine | |
| SQLITE_BASE.metadata.create_all(sqlite_engine) | |
| # Migrate and drop legacy columns | |
| from sqlalchemy import text | |
| try: | |
| from sqlalchemy import inspect | |
| inspector = inspect(engine) | |
| columns = [col["name"] for col in inspector.get_columns("admin_settings")] | |
| # Only migrate and drop if obsolete columns actually exist in the table | |
| has_obsolete = any(c in columns for c in ["shortener_site", "shortener_api", "how_to_download"]) | |
| if has_obsolete: | |
| # 1. Attempt to migrate old settings | |
| old_site = None | |
| old_api = None | |
| old_download = None | |
| shorteners_list_json = None | |
| col_selectors = [] | |
| if "shortener_site" in columns: | |
| col_selectors.append("shortener_site") | |
| if "shortener_api" in columns: | |
| col_selectors.append("shortener_api") | |
| if "how_to_download" in columns: | |
| col_selectors.append("how_to_download") | |
| if len(col_selectors) > 0 and "shorteners_list" in columns: | |
| col_selectors.append("shorteners_list") | |
| try: | |
| with engine.connect() as conn: | |
| query_str = f"SELECT {', '.join(col_selectors)} FROM admin_settings LIMIT 1" | |
| res = conn.execute(text(query_str)).fetchone() | |
| if res: | |
| val_map = dict(zip(col_selectors, res)) | |
| old_site = val_map.get("shortener_site") | |
| old_api = val_map.get("shortener_api") | |
| old_download = val_map.get("how_to_download") | |
| shorteners_list_json = val_map.get("shorteners_list") | |
| if old_site and old_api: | |
| import json | |
| try: | |
| shorteners = json.loads(shorteners_list_json) if shorteners_list_json else [] | |
| except Exception: | |
| shorteners = [] | |
| exists = any(s.get("site") == old_site and s.get("api") == old_api for s in shorteners) | |
| if not exists: | |
| shorteners.append({ | |
| "site": old_site, | |
| "api": old_api, | |
| "how_to_download": old_download, | |
| "status": True | |
| }) | |
| with engine.begin() as t_conn: | |
| t_conn.execute( | |
| text("UPDATE admin_settings SET shorteners_list = :sl"), | |
| {"sl": json.dumps(shorteners)} | |
| ) | |
| LOGGER.info(f"Successfully migrated old shortener {old_site} to shorteners_list.") | |
| except Exception as e: | |
| LOGGER.warning(f"Error migrating old settings: {e}") | |
| # 2. Drop the obsolete columns since they exist (each column drop in its own transaction block) | |
| for col_to_drop in ["shortener_site", "shortener_api", "how_to_download"]: | |
| if col_to_drop in columns: | |
| try: | |
| with engine.begin() as t_conn: | |
| t_conn.execute(text(f"ALTER TABLE admin_settings DROP COLUMN IF EXISTS {col_to_drop}")) | |
| LOGGER.info(f"Obsolete column {col_to_drop} successfully dropped from admin_settings table.") | |
| except Exception as e: | |
| LOGGER.warning(f"Error dropping column {col_to_drop}: {e}") | |
| except Exception as e: | |
| LOGGER.warning(f"Error checking/dropping obsolete columns: {e}") | |
| from sqlalchemy import text | |
| try: | |
| with engine.connect() as conn: | |
| with conn.begin(): | |
| # conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS shortener_site TEXT")) | |
| # conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS shortener_api TEXT")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS shortener_status BOOLEAN DEFAULT TRUE")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS admins_list TEXT")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS db_channels_list TEXT")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS token_shortener_enabled BOOLEAN DEFAULT FALSE")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS token_timeout NUMERIC DEFAULT 3600")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS log_channel NUMERIC")) | |
| # conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS how_to_download TEXT")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS log_channel_access_hash NUMERIC")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS fsub_channel_access_hash NUMERIC")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS fsub_channels_list TEXT")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS fsub_channels_access_hashes TEXT")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS shorteners_list TEXT")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS smart_rotator BOOLEAN DEFAULT FALSE")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS newsletter_days TEXT DEFAULT 'Friday'")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS newsletter_time TEXT DEFAULT '20:00'")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS newsletter_count INTEGER DEFAULT 10")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS newsletter_enabled BOOLEAN DEFAULT FALSE")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS newsletter_target TEXT DEFAULT 'subscribed'")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS announcement_enabled BOOLEAN DEFAULT FALSE")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS announcement_channel NUMERIC")) | |
| conn.execute(text("ALTER TABLE admin_settings ADD COLUMN IF NOT EXISTS announcement_channel_access_hash NUMERIC")) | |
| conn.execute(text(""" | |
| CREATE TABLE IF NOT EXISTS user_shortener_usage ( | |
| user_id BIGINT PRIMARY KEY, | |
| last_index INTEGER DEFAULT 0, | |
| last_date TEXT | |
| ) | |
| """)) | |
| conn.execute(text("ALTER TABLE user_tokens ADD COLUMN IF NOT EXISTS req_msg_id NUMERIC")) | |
| conn.execute(text(""" | |
| CREATE TABLE IF NOT EXISTS pending_fsub_requests ( | |
| user_id NUMERIC PRIMARY KEY, | |
| search_query TEXT, | |
| target_file_id TEXT, | |
| created_at NUMERIC | |
| ) | |
| """)) | |
| try: | |
| conn.execute(text("ALTER TABLE pending_fsub_requests ALTER COLUMN user_id TYPE NUMERIC")) | |
| except Exception: | |
| pass | |
| except Exception as e: | |
| LOGGER.warning("Could not execute DDL migrations for admin_settings: %s", str(e)) | |
| return scoped_session(sessionmaker(bind=engine, autoflush=False)) | |
| SESSION = start() | |
| import asyncio | |
| import time | |
| import functools | |
| from sqlalchemy.exc import OperationalError, PendingRollbackError | |
| def db_retry(func): | |
| if asyncio.iscoroutinefunction(func): | |
| async def wrapper(*args, **kwargs): | |
| for attempt in range(3): | |
| try: | |
| return await func(*args, **kwargs) | |
| except (OperationalError, PendingRollbackError) as e: | |
| LOGGER.warning("Database connection error in %s (attempt %s): %s. Reconnecting...", func.__name__, attempt + 1, str(e)) | |
| try: | |
| SESSION.rollback() | |
| SESSION.remove() | |
| except Exception: | |
| pass | |
| if attempt == 2: | |
| raise | |
| await asyncio.sleep(0.5) | |
| return wrapper | |
| else: | |
| def wrapper(*args, **kwargs): | |
| for attempt in range(3): | |
| try: | |
| return func(*args, **kwargs) | |
| except (OperationalError, PendingRollbackError) as e: | |
| LOGGER.warning("Database connection error in %s (attempt %s): %s. Reconnecting...", func.__name__, attempt + 1, str(e)) | |
| try: | |
| SESSION.rollback() | |
| SESSION.remove() | |
| except Exception: | |
| pass | |
| if attempt == 2: | |
| raise | |
| time.sleep(0.5) | |
| return wrapper | |
| # ----------------- IN-MEMORY CACHE & MUTEX LOCK ----------------- | |
| class CachedAdminSettings: | |
| def __init__(self, admin_settings_db=None): | |
| self.setting_name = "default" | |
| self.auto_delete = 0 | |
| self.custom_caption = None | |
| self.fsub_channel = None | |
| self.channel_link = None | |
| self.caption_uname = None | |
| self.repair_mode = False | |
| self.shortener_status = True | |
| self.admins_list = None | |
| self.db_channels_list = None | |
| self.token_shortener_enabled = False | |
| self.token_timeout = 3600 | |
| self.log_channel = None | |
| self.log_channel_access_hash = None | |
| self.fsub_channel_access_hash = None | |
| self.fsub_channels_list = None | |
| self.fsub_channels_access_hashes = "{}" | |
| self.shorteners_list = "[]" | |
| self.smart_rotator = False | |
| self.newsletter_days = "Friday" | |
| self.newsletter_time = "20:00" | |
| self.newsletter_count = 10 | |
| self.newsletter_enabled = False | |
| self.newsletter_target = "subscribed" | |
| self.announcement_enabled = False | |
| self.announcement_channel = None | |
| self.announcement_channel_access_hash = None | |
| self.shortener_site = None | |
| self.shortener_api = None | |
| self.how_to_download = None | |
| if admin_settings_db: | |
| def to_numeric(val): | |
| if val is None: | |
| return None | |
| try: | |
| val_str = str(val) | |
| if "." not in val_str: | |
| return int(val) | |
| f_val = float(val) | |
| if f_val.is_integer(): | |
| return int(f_val) | |
| return f_val | |
| except (ValueError, TypeError): | |
| return val | |
| self.setting_name = admin_settings_db.setting_name | |
| self.auto_delete = to_numeric(admin_settings_db.auto_delete) or 0 | |
| self.custom_caption = admin_settings_db.custom_caption | |
| self.fsub_channel = to_numeric(admin_settings_db.fsub_channel) | |
| self.channel_link = admin_settings_db.channel_link | |
| self.caption_uname = admin_settings_db.caption_uname | |
| self.repair_mode = bool(admin_settings_db.repair_mode) | |
| self.shortener_status = bool(admin_settings_db.shortener_status) if admin_settings_db.shortener_status is not None else True | |
| self.admins_list = admin_settings_db.admins_list | |
| self.db_channels_list = admin_settings_db.db_channels_list | |
| self.token_shortener_enabled = bool(admin_settings_db.token_shortener_enabled) if admin_settings_db.token_shortener_enabled is not None else False | |
| self.token_timeout = to_numeric(admin_settings_db.token_timeout) or 3600 | |
| self.log_channel = to_numeric(admin_settings_db.log_channel) | |
| self.log_channel_access_hash = to_numeric(admin_settings_db.log_channel_access_hash) | |
| self.fsub_channel_access_hash = to_numeric(admin_settings_db.fsub_channel_access_hash) | |
| self.fsub_channels_list = admin_settings_db.fsub_channels_list | |
| self.fsub_channels_access_hashes = admin_settings_db.fsub_channels_access_hashes or "{}" | |
| self.shorteners_list = admin_settings_db.shorteners_list or "[]" | |
| self.smart_rotator = bool(admin_settings_db.smart_rotator) if admin_settings_db.smart_rotator is not None else False | |
| self.newsletter_days = admin_settings_db.newsletter_days or "Friday" | |
| self.newsletter_time = admin_settings_db.newsletter_time or "20:00" | |
| self.newsletter_count = to_numeric(admin_settings_db.newsletter_count) or 10 | |
| self.newsletter_enabled = bool(admin_settings_db.newsletter_enabled) if admin_settings_db.newsletter_enabled is not None else False | |
| self.newsletter_target = admin_settings_db.newsletter_target or "subscribed" | |
| self.announcement_enabled = bool(admin_settings_db.announcement_enabled) if admin_settings_db.announcement_enabled is not None else False | |
| self.announcement_channel = to_numeric(admin_settings_db.announcement_channel) | |
| self.announcement_channel_access_hash = to_numeric(admin_settings_db.announcement_channel_access_hash) | |
| # Dynamically resolve legacy columns from shorteners_list JSON | |
| import json | |
| try: | |
| sh_list = json.loads(self.shorteners_list) | |
| except Exception: | |
| sh_list = [] | |
| if sh_list: | |
| first_sh = sh_list[0] | |
| self.shortener_site = first_sh.get("site") | |
| self.shortener_api = first_sh.get("api") | |
| self.how_to_download = first_sh.get("how_to_download") | |
| else: | |
| self.shortener_site = None | |
| self.shortener_api = None | |
| self.how_to_download = None | |
| _cached_admin_settings = None | |
| _cache_lock = threading.Lock() | |
| def _update_cache_from_db(admin_setting): | |
| global _cached_admin_settings | |
| with _cache_lock: | |
| _cached_admin_settings = CachedAdminSettings(admin_setting) | |
| def get_admin_settings_sync(): | |
| global _cached_admin_settings | |
| with _cache_lock: | |
| if _cached_admin_settings is not None: | |
| return _cached_admin_settings | |
| for attempt in range(3): | |
| try: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| session.commit() | |
| _cached_admin_settings = CachedAdminSettings(admin_setting) | |
| return _cached_admin_settings | |
| except (OperationalError, PendingRollbackError) as e: | |
| LOGGER.warning("Database connection error in get_admin_settings_sync (attempt %d): %s. Reconnecting...", attempt + 1, str(e)) | |
| try: | |
| SESSION.rollback() | |
| except Exception: | |
| pass | |
| try: | |
| SESSION.remove() | |
| except Exception: | |
| pass | |
| if attempt == 2: | |
| return CachedAdminSettings() | |
| time.sleep(0.5) | |
| except Exception as e: | |
| LOGGER.warning("Error in get_admin_settings_sync: %s", str(e)) | |
| try: | |
| SESSION.rollback() | |
| except Exception: | |
| pass | |
| try: | |
| SESSION.remove() | |
| except Exception: | |
| pass | |
| return CachedAdminSettings() | |
| finally: | |
| SESSION.close() | |
| def heal_admin_settings(): | |
| import os | |
| try: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| env_admins = os.environ.get("ADMINS", "") | |
| owner_id = os.environ.get("OWNER_ID", "") | |
| admins_set = set() | |
| for x in env_admins.split(): | |
| x = x.strip() | |
| if x: | |
| admins_set.add(x) | |
| if owner_id: | |
| admins_set.add(owner_id.strip()) | |
| default_admins = ",".join(sorted(list(admins_set))) if admins_set else None | |
| env_channels = os.environ.get("DB_CHANNELS", "") | |
| channels_list = [] | |
| for x in env_channels.split(): | |
| x = x.strip() | |
| if x: | |
| channels_list.append(x) | |
| default_channels = ",".join(channels_list) if channels_list else None | |
| # Parse helpers | |
| def parse_bool(env_val, default=None): | |
| if env_val is None: | |
| return default | |
| val = env_val.strip().lower() | |
| if val in ("true", "yes", "on", "1"): | |
| return True | |
| if val in ("false", "no", "off", "0"): | |
| return False | |
| return default | |
| def parse_int(env_val, default=None): | |
| if env_val is None: | |
| return default | |
| try: | |
| return int(env_val.strip()) | |
| except ValueError: | |
| return default | |
| def parse_time(env_val, default=None): | |
| if env_val is None: | |
| return default | |
| val = env_val.strip().lower() | |
| if val.endswith("h"): | |
| try: | |
| return int(float(val[:-1]) * 3600) | |
| except ValueError: | |
| pass | |
| elif val.endswith("m"): | |
| try: | |
| return int(float(val[:-1]) * 60) | |
| except ValueError: | |
| pass | |
| elif val.endswith("d"): | |
| try: | |
| return int(float(val[:-1]) * 86400) | |
| except ValueError: | |
| pass | |
| elif val.endswith("s"): | |
| try: | |
| return int(float(val[:-1])) | |
| except ValueError: | |
| pass | |
| else: | |
| try: | |
| return int(val) | |
| except ValueError: | |
| pass | |
| return default | |
| env_auto_delete = os.environ.get("AUTO_DELETE") | |
| env_custom_caption = os.environ.get("CUSTOM_CAPTION") | |
| env_fsub_channel = os.environ.get("FSUB_CHANNEL") | |
| env_channel_link = os.environ.get("CHANNEL_LINK") | |
| env_caption_uname = os.environ.get("CAPTION_UNAME") | |
| env_repair_mode = os.environ.get("REPAIR_MODE") | |
| env_shortener_site = os.environ.get("SHORTENER_SITE") | |
| env_shortener_api = os.environ.get("SHORTENER_API") | |
| env_shortener_status = os.environ.get("SHORTENER_STATUS") | |
| env_token_shortener_enabled = os.environ.get("TOKEN_SHORTENER_ENABLED") | |
| env_token_timeout = os.environ.get("TOKEN_TIMEOUT") | |
| env_log_channel = os.environ.get("LOG_CHANNEL") | |
| env_how_to_download = os.environ.get("HOW_TO_DOWNLOAD") | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| admin_setting.admins_list = default_admins | |
| admin_setting.db_channels_list = default_channels | |
| if env_auto_delete is not None: | |
| admin_setting.auto_delete = parse_int(env_auto_delete, 0) | |
| if env_custom_caption is not None: | |
| cc = env_custom_caption.strip() | |
| admin_setting.custom_caption = None if cc.lower() in ("off", "disable", "disabled", "none") else cc | |
| if env_fsub_channel is not None: | |
| admin_setting.fsub_channel = parse_int(env_fsub_channel) | |
| if env_channel_link is not None: | |
| cl = env_channel_link.strip() | |
| admin_setting.channel_link = None if cl.lower() in ("off", "disable", "disabled", "none") else cl | |
| if env_caption_uname is not None: | |
| cu = env_caption_uname.strip() | |
| admin_setting.caption_uname = None if cu.lower() in ("off", "disable", "disabled", "none") else cu | |
| if env_repair_mode is not None: | |
| admin_setting.repair_mode = parse_bool(env_repair_mode, False) | |
| if env_shortener_site and env_shortener_api: | |
| new_list, _ = merge_env_shortener_into_list( | |
| admin_setting.shorteners_list, | |
| env_shortener_site, | |
| env_shortener_api, | |
| env_how_to_download | |
| ) | |
| admin_setting.shorteners_list = new_list | |
| if env_shortener_status is not None: | |
| admin_setting.shortener_status = parse_bool(env_shortener_status, True) | |
| if env_token_shortener_enabled is not None: | |
| admin_setting.token_shortener_enabled = parse_bool(env_token_shortener_enabled, False) | |
| if env_token_timeout is not None: | |
| admin_setting.token_timeout = parse_time(env_token_timeout, 3600) | |
| if env_log_channel is not None: | |
| admin_setting.log_channel = parse_int(env_log_channel) | |
| # Enforce mutual exclusion | |
| if admin_setting.token_shortener_enabled and admin_setting.shortener_status: | |
| admin_setting.shortener_status = False | |
| session.add(admin_setting) | |
| session.commit() | |
| LOGGER.info("Created default AdminSettings with env values.") | |
| else: | |
| updated = False | |
| # Check if admins_list is None, empty, or has the test dummy values | |
| if (not admin_setting.admins_list or | |
| admin_setting.admins_list == "12345,67890" or | |
| (owner_id and owner_id.strip() not in admin_setting.admins_list.split(','))): | |
| admin_setting.admins_list = default_admins | |
| updated = True | |
| # Check if db_channels_list is None, empty, or has the test dummy values | |
| if (not admin_setting.db_channels_list or | |
| admin_setting.db_channels_list == "-100123,-100456"): | |
| admin_setting.db_channels_list = default_channels | |
| updated = True | |
| # Heal other environment variables | |
| if env_auto_delete is not None: | |
| parsed_auto_delete = parse_int(env_auto_delete) | |
| if parsed_auto_delete is not None and admin_setting.auto_delete != parsed_auto_delete: | |
| admin_setting.auto_delete = parsed_auto_delete | |
| updated = True | |
| if env_custom_caption is not None: | |
| parsed_custom_caption = env_custom_caption.strip() | |
| if parsed_custom_caption.lower() in ("off", "disable", "disabled", "none"): | |
| parsed_custom_caption = None | |
| if admin_setting.custom_caption != parsed_custom_caption: | |
| admin_setting.custom_caption = parsed_custom_caption | |
| updated = True | |
| if env_fsub_channel is not None: | |
| parsed_fsub_channel = parse_int(env_fsub_channel) | |
| if parsed_fsub_channel is not None and admin_setting.fsub_channel != parsed_fsub_channel: | |
| admin_setting.fsub_channel = parsed_fsub_channel | |
| updated = True | |
| if env_channel_link is not None: | |
| parsed_channel_link = env_channel_link.strip() | |
| if parsed_channel_link.lower() in ("off", "disable", "disabled", "none"): | |
| parsed_channel_link = None | |
| if admin_setting.channel_link != parsed_channel_link: | |
| admin_setting.channel_link = parsed_channel_link | |
| updated = True | |
| if env_caption_uname is not None: | |
| parsed_caption_uname = env_caption_uname.strip() | |
| if parsed_caption_uname.lower() in ("off", "disable", "disabled", "none"): | |
| parsed_caption_uname = None | |
| if admin_setting.caption_uname != parsed_caption_uname: | |
| admin_setting.caption_uname = parsed_caption_uname | |
| updated = True | |
| if env_repair_mode is not None: | |
| parsed_repair_mode = parse_bool(env_repair_mode) | |
| if parsed_repair_mode is not None and admin_setting.repair_mode != parsed_repair_mode: | |
| admin_setting.repair_mode = parsed_repair_mode | |
| updated = True | |
| if env_shortener_site and env_shortener_api: | |
| new_list, list_updated = merge_env_shortener_into_list( | |
| admin_setting.shorteners_list, | |
| env_shortener_site, | |
| env_shortener_api, | |
| env_how_to_download | |
| ) | |
| if list_updated: | |
| admin_setting.shorteners_list = new_list | |
| updated = True | |
| if env_shortener_status is not None: | |
| parsed_shortener_status = parse_bool(env_shortener_status) | |
| if parsed_shortener_status is not None and admin_setting.shortener_status != parsed_shortener_status: | |
| admin_setting.shortener_status = parsed_shortener_status | |
| updated = True | |
| if env_token_shortener_enabled is not None: | |
| parsed_token_shortener_enabled = parse_bool(env_token_shortener_enabled) | |
| if parsed_token_shortener_enabled is not None and admin_setting.token_shortener_enabled != parsed_token_shortener_enabled: | |
| admin_setting.token_shortener_enabled = parsed_token_shortener_enabled | |
| updated = True | |
| if env_token_timeout is not None: | |
| parsed_token_timeout = parse_time(env_token_timeout) | |
| if parsed_token_timeout is not None and admin_setting.token_timeout != parsed_token_timeout: | |
| admin_setting.token_timeout = parsed_token_timeout | |
| updated = True | |
| if env_log_channel is not None: | |
| parsed_log_channel = parse_int(env_log_channel) | |
| if parsed_log_channel is not None and admin_setting.log_channel != parsed_log_channel: | |
| admin_setting.log_channel = parsed_log_channel | |
| updated = True | |
| # Enforce mutual exclusion | |
| if admin_setting.token_shortener_enabled and admin_setting.shortener_status: | |
| admin_setting.shortener_status = False | |
| updated = True | |
| if updated: | |
| session.commit() | |
| LOGGER.info("Healed/restored AdminSettings from env values.") | |
| # Prime the cache on boot | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Could not heal admin settings: %s", str(e)) | |
| finally: | |
| SESSION.close() | |
| heal_admin_settings() | |
| INSERTION_LOCK = threading.RLock() | |
| async def get_search_settings(user_id): | |
| for attempt in range(3): | |
| try: | |
| with INSERTION_LOCK: | |
| settings = SESSION.query(Settings).filter_by(user_id=user_id).first() | |
| return settings | |
| except (OperationalError, PendingRollbackError) as e: | |
| LOGGER.warning("Database connection error in get_search_settings (attempt %d): %s. Reconnecting...", attempt + 1, str(e)) | |
| try: | |
| SESSION.rollback() | |
| except Exception: | |
| pass | |
| try: | |
| SESSION.remove() | |
| except Exception: | |
| pass | |
| if attempt == 2: | |
| return None | |
| await asyncio.sleep(0.5) | |
| except Exception as e: | |
| LOGGER.warning("Error getting search settings: %s", str(e)) | |
| try: | |
| SESSION.rollback() | |
| except Exception: | |
| pass | |
| try: | |
| SESSION.remove() | |
| except Exception: | |
| pass | |
| return None | |
| finally: | |
| SESSION.close() | |
| async def change_search_settings(user_id, precise_mode=None, button_mode=None, link_mode=None, list_mode=None): | |
| try: | |
| with INSERTION_LOCK: | |
| settings = SESSION.query(Settings).filter_by(user_id=user_id).first() | |
| if settings: | |
| if precise_mode is not None: | |
| settings.precise_mode = precise_mode | |
| if button_mode is not None: | |
| settings.button_mode = button_mode | |
| if link_mode is not None: | |
| settings.link_mode = link_mode | |
| if list_mode is not None: | |
| settings.list_mode = list_mode | |
| else: | |
| new_settings = Settings( | |
| user_id=user_id, | |
| precise_mode=precise_mode if precise_mode is not None else False, | |
| button_mode=button_mode if button_mode is not None else False, | |
| link_mode=link_mode if link_mode is not None else True, | |
| list_mode=list_mode if list_mode is not None else False | |
| ) | |
| SESSION.add(new_settings) | |
| SESSION.commit() | |
| return True | |
| except Exception as e: | |
| LOGGER.warning("Error changing search settings: %s", str(e)) | |
| try: | |
| SESSION.rollback() | |
| except Exception: | |
| pass | |
| try: | |
| SESSION.remove() | |
| except Exception: | |
| pass | |
| return False | |
| finally: | |
| SESSION.close() | |
| async def set_repair_mode(repair_mode): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| session.commit() | |
| admin_setting.repair_mode = repair_mode | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Error setting repair mode: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def set_auto_delete(dur): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| session.commit() | |
| admin_setting.auto_delete = dur | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Error setting auto delete: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def get_admin_settings(): | |
| return get_admin_settings_sync() | |
| async def set_custom_caption(caption): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| session.commit() | |
| admin_setting.custom_caption = caption | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Error setting custom caption: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def set_force_sub(channel, access_hash=None): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| session.commit() | |
| admin_setting.fsub_channel = channel | |
| if access_hash is not None: | |
| admin_setting.fsub_channel_access_hash = access_hash | |
| # Keep fsub_channels_list in sync for backwards compatibility | |
| if channel: | |
| admin_setting.fsub_channels_list = str(channel) | |
| if access_hash is not None: | |
| import json | |
| admin_setting.fsub_channels_access_hashes = json.dumps({str(channel): access_hash}) | |
| else: | |
| admin_setting.fsub_channels_list = None | |
| admin_setting.fsub_channels_access_hashes = "{}" | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Error setting Force Sub channel: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def set_force_sub_channels(channels_list, hashes_map): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| session.commit() | |
| admin_setting.fsub_channels_list = channels_list | |
| import json | |
| admin_setting.fsub_channels_access_hashes = json.dumps(hashes_map) | |
| # Sync first channel to fsub_channel legacy columns for safety | |
| if channels_list: | |
| first_ch = channels_list.split(',')[0].strip() | |
| try: | |
| admin_setting.fsub_channel = int(first_ch) | |
| except ValueError: | |
| admin_setting.fsub_channel = None | |
| admin_setting.fsub_channel_access_hash = hashes_map.get(first_ch) | |
| else: | |
| admin_setting.fsub_channel = None | |
| admin_setting.fsub_channel_access_hash = None | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Error setting Force Sub channels list: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def save_pending_fsub_request(user_id, search_query=None, target_file_id=None): | |
| import time | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| req = session.query(PendingFsubRequests).filter(PendingFsubRequests.user_id == user_id).first() | |
| if req: | |
| req.search_query = search_query | |
| req.target_file_id = target_file_id | |
| req.created_at = time.time() | |
| else: | |
| req = PendingFsubRequests(user_id=user_id, search_query=search_query, target_file_id=target_file_id) | |
| req.created_at = time.time() | |
| session.add(req) | |
| session.commit() | |
| except Exception as e: | |
| LOGGER.warning("Error saving pending fsub request: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def get_pending_fsub_request(user_id): | |
| for attempt in range(3): | |
| try: | |
| session = SESSION() | |
| req = session.query(PendingFsubRequests).filter(PendingFsubRequests.user_id == user_id).first() | |
| return req | |
| except (OperationalError, PendingRollbackError) as e: | |
| LOGGER.warning("Database connection error in get_pending_fsub_request (attempt %d): %s. Reconnecting...", attempt + 1, str(e)) | |
| try: | |
| SESSION.rollback() | |
| except Exception: | |
| pass | |
| try: | |
| SESSION.remove() | |
| except Exception: | |
| pass | |
| if attempt == 2: | |
| return None | |
| await asyncio.sleep(0.5) | |
| except Exception as e: | |
| LOGGER.warning("Error getting pending fsub request: %s", str(e)) | |
| try: | |
| SESSION.rollback() | |
| except Exception: | |
| pass | |
| try: | |
| SESSION.remove() | |
| except Exception: | |
| pass | |
| return None | |
| finally: | |
| session.close() | |
| async def delete_pending_fsub_request(user_id): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| session.query(PendingFsubRequests).filter(PendingFsubRequests.user_id == user_id).delete() | |
| session.commit() | |
| except Exception as e: | |
| LOGGER.warning("Error deleting pending fsub request: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def set_channel_link(link): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| session.commit() | |
| admin_setting.channel_link = link | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Error adding Force Sub channel link: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def get_channel(): | |
| settings = get_admin_settings_sync() | |
| return settings.fsub_channel if settings.fsub_channel else False | |
| async def get_link(): | |
| settings = get_admin_settings_sync() | |
| return settings.channel_link if settings.channel_link else False | |
| async def set_username(username): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| session.commit() | |
| admin_setting.caption_uname = username | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Error adding username: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def set_shortener_settings(site=None, api=None, status=None): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| session.commit() | |
| import json | |
| try: | |
| shorteners = json.loads(admin_setting.shorteners_list) if admin_setting.shorteners_list else [] | |
| except Exception: | |
| shorteners = [] | |
| if site is not None or api is not None: | |
| if shorteners: | |
| if site is not None: | |
| shorteners[0]["site"] = site | |
| if api is not None: | |
| shorteners[0]["api"] = api | |
| else: | |
| shorteners.append({ | |
| "site": site or "", | |
| "api": api or "", | |
| "how_to_download": None, | |
| "status": True | |
| }) | |
| admin_setting.shorteners_list = json.dumps(shorteners) | |
| if status is not None: | |
| admin_setting.shortener_status = status | |
| if status: | |
| admin_setting.token_shortener_enabled = False | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Error setting shortener settings: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def delete_shortener_settings(): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if admin_setting: | |
| import json | |
| try: | |
| shorteners = json.loads(admin_setting.shorteners_list) if admin_setting.shorteners_list else [] | |
| except Exception: | |
| shorteners = [] | |
| if shorteners: | |
| shorteners.pop(0) | |
| admin_setting.shorteners_list = json.dumps(shorteners) | |
| admin_setting.shortener_status = False | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Error deleting shortener settings: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def update_admins_list(admins_str): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| session.commit() | |
| admin_setting.admins_list = admins_str | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Error updating admins list: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def update_db_channels_list(channels_str): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| session.commit() | |
| admin_setting.db_channels_list = channels_str | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Error updating db channels list: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def get_user_token(user_id): | |
| for attempt in range(3): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| return session.query(UserTokens).filter_by(user_id=user_id).first() | |
| except (OperationalError, PendingRollbackError) as e: | |
| LOGGER.warning("Database connection error in get_user_token (attempt %d): %s. Reconnecting...", attempt + 1, str(e)) | |
| try: | |
| SESSION.rollback() | |
| except Exception: | |
| pass | |
| try: | |
| SESSION.remove() | |
| except Exception: | |
| pass | |
| if attempt == 2: | |
| return None | |
| await asyncio.sleep(0.5) | |
| except Exception as e: | |
| LOGGER.warning("Error getting user token: %s", str(e)) | |
| try: | |
| SESSION.rollback() | |
| except Exception: | |
| pass | |
| try: | |
| SESSION.remove() | |
| except Exception: | |
| pass | |
| return None | |
| finally: | |
| SESSION.close() | |
| async def create_user_token(user_id, token, search_query=None, target_file_id=None): | |
| import time | |
| try: | |
| await prune_all_expired_tokens() | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| tok_rec = session.query(UserTokens).filter_by(user_id=user_id).first() | |
| if tok_rec: | |
| tok_rec.token = token | |
| tok_rec.created_at = time.time() | |
| tok_rec.expires_at = 0 | |
| tok_rec.status = "pending" | |
| tok_rec.search_query = search_query | |
| tok_rec.target_file_id = target_file_id | |
| else: | |
| tok_rec = UserTokens( | |
| user_id=user_id, | |
| token=token, | |
| created_at=time.time(), | |
| search_query=search_query, | |
| target_file_id=target_file_id | |
| ) | |
| session.add(tok_rec) | |
| session.commit() | |
| return tok_rec | |
| except Exception as e: | |
| LOGGER.warning("Error creating user token: %s", str(e)) | |
| SESSION.rollback() | |
| return None | |
| finally: | |
| SESSION.close() | |
| async def activate_user_token(user_id, expires_at): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| tok_rec = session.query(UserTokens).filter_by(user_id=user_id).first() | |
| if tok_rec: | |
| tok_rec.status = "active" | |
| tok_rec.expires_at = expires_at | |
| session.commit() | |
| return True | |
| return False | |
| except Exception as e: | |
| LOGGER.warning("Error activating user token: %s", str(e)) | |
| SESSION.rollback() | |
| return False | |
| finally: | |
| SESSION.close() | |
| async def delete_user_token(user_id): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| tok_rec = session.query(UserTokens).filter_by(user_id=user_id).first() | |
| if tok_rec: | |
| session.delete(tok_rec) | |
| session.commit() | |
| return True | |
| return False | |
| except Exception as e: | |
| LOGGER.warning("Error deleting user token: %s", str(e)) | |
| SESSION.rollback() | |
| return False | |
| finally: | |
| SESSION.close() | |
| async def prune_all_expired_tokens(): | |
| import time | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| current_time = time.time() | |
| # Delete active tokens that have expired | |
| expired_active = session.query(UserTokens).filter( | |
| UserTokens.status == "active", | |
| UserTokens.expires_at < current_time | |
| ) | |
| # Delete pending tokens that are older than 30 minutes (1800 seconds) | |
| stale_pending = session.query(UserTokens).filter( | |
| UserTokens.status == "pending", | |
| (current_time - UserTokens.created_at) > 1800 | |
| ) | |
| deleted_active = expired_active.delete(synchronize_session=False) | |
| deleted_pending = stale_pending.delete(synchronize_session=False) | |
| if deleted_active or deleted_pending: | |
| session.commit() | |
| LOGGER.info("Database row protection: Pruned %d expired active and %d stale pending tokens.", deleted_active, deleted_pending) | |
| return True | |
| except Exception as e: | |
| LOGGER.warning("Error pruning expired tokens from database: %s", str(e)) | |
| SESSION.rollback() | |
| return False | |
| finally: | |
| SESSION.close() | |
| async def set_token_shortener_state(enabled): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| admin_setting.token_shortener_enabled = enabled | |
| if enabled: | |
| admin_setting.shortener_status = False | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| return True | |
| except Exception as e: | |
| LOGGER.warning("Error setting token shortener state: %s", str(e)) | |
| SESSION.rollback() | |
| return False | |
| finally: | |
| SESSION.close() | |
| async def set_token_timeout(seconds): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| admin_setting.token_timeout = seconds | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| return True | |
| except Exception as e: | |
| LOGGER.warning("Error setting token timeout: %s", str(e)) | |
| SESSION.rollback() | |
| return False | |
| finally: | |
| SESSION.close() | |
| async def set_log_channel(channel, access_hash=None): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| admin_setting.log_channel = channel | |
| if access_hash is not None: | |
| admin_setting.log_channel_access_hash = access_hash | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| return True | |
| except Exception as e: | |
| LOGGER.warning("Error setting log channel: %s", str(e)) | |
| SESSION.rollback() | |
| return False | |
| finally: | |
| SESSION.close() | |
| async def get_log_channel(): | |
| settings = get_admin_settings_sync() | |
| return settings.log_channel if settings.log_channel else False | |
| async def set_how_to_download(url): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| admin_setting.how_to_download = url | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| return True | |
| except Exception as e: | |
| LOGGER.warning("Error setting how to download URL: %s", str(e)) | |
| SESSION.rollback() | |
| return False | |
| finally: | |
| SESSION.close() | |
| async def update_user_token_msg_id(user_id, msg_id): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| tok_rec = session.query(UserTokens).filter_by(user_id=user_id).first() | |
| if tok_rec: | |
| tok_rec.req_msg_id = msg_id | |
| session.commit() | |
| return True | |
| return False | |
| except Exception as e: | |
| LOGGER.warning("Error updating user token message ID: %s", str(e)) | |
| SESSION.rollback() | |
| return False | |
| finally: | |
| SESSION.close() | |
| async def set_smart_rotator_state(state: bool): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| session.commit() | |
| admin_setting.smart_rotator = state | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Error setting smart rotator state: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def update_shorteners_list(shorteners_json_str): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| session.commit() | |
| admin_setting.shorteners_list = shorteners_json_str | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Error updating shorteners list: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def get_user_shortener_index(user_id: int, today_str: str) -> int: | |
| for attempt in range(3): | |
| try: | |
| session = SESSION() | |
| record = session.query(UserShortenerUsage).filter_by(user_id=user_id).first() | |
| if record and record.last_date == today_str: | |
| return record.last_index or 0 | |
| return 0 | |
| except (OperationalError, PendingRollbackError) as e: | |
| LOGGER.warning("Database connection error in get_user_shortener_index: %s. Reconnecting...", str(e)) | |
| try: | |
| SESSION.rollback() | |
| except Exception: | |
| pass | |
| try: | |
| SESSION.remove() | |
| except Exception: | |
| pass | |
| if attempt == 2: | |
| return -1 | |
| await asyncio.sleep(0.5) | |
| except Exception as e: | |
| LOGGER.warning("Error in get_user_shortener_index: %s", str(e)) | |
| return -1 | |
| finally: | |
| SESSION.close() | |
| async def update_user_shortener_index(user_id: int, index: int, today_str: str): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| record = session.query(UserShortenerUsage).filter_by(user_id=user_id).first() | |
| if record: | |
| record.last_index = index | |
| record.last_date = today_str | |
| else: | |
| record = UserShortenerUsage(user_id=user_id, last_index=index, last_date=today_str) | |
| session.add(record) | |
| session.commit() | |
| except Exception as e: | |
| LOGGER.warning("Error in update_user_shortener_index: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def advance_user_shortener_index(user_id: int): | |
| try: | |
| import datetime | |
| local_now = datetime.datetime.utcnow() + datetime.timedelta(hours=5, minutes=30) | |
| today_str = local_now.strftime("%Y-%m-%d") | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| record = session.query(UserShortenerUsage).filter_by(user_id=user_id).first() | |
| if record: | |
| if record.last_date == today_str: | |
| record.last_index = (record.last_index or 0) + 1 | |
| else: | |
| record.last_index = 1 | |
| record.last_date = today_str | |
| else: | |
| record = UserShortenerUsage(user_id=user_id, last_index=1, last_date=today_str) | |
| session.add(record) | |
| session.commit() | |
| except Exception as e: | |
| LOGGER.warning("Error in advance_user_shortener_index: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def clear_all_user_shortener_usages(): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| session.query(UserShortenerUsage).delete() | |
| session.commit() | |
| except Exception as e: | |
| LOGGER.warning("Error clearing UserShortenerUsage: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| async def set_newsletter_settings(enabled=None, days=None, time_str=None, count=None, target=None): | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| admin_setting = session.query(AdminSettings).first() | |
| if not admin_setting: | |
| admin_setting = AdminSettings(setting_name="default") | |
| session.add(admin_setting) | |
| session.commit() | |
| if enabled is not None: | |
| admin_setting.newsletter_enabled = enabled | |
| if days is not None: | |
| admin_setting.newsletter_days = days | |
| if time_str is not None: | |
| admin_setting.newsletter_time = time_str | |
| if count is not None: | |
| admin_setting.newsletter_count = count | |
| if target is not None: | |
| admin_setting.newsletter_target = target | |
| session.commit() | |
| _update_cache_from_db(admin_setting) | |
| except Exception as e: | |
| LOGGER.warning("Error setting newsletter settings: %s", str(e)) | |
| SESSION.rollback() | |
| finally: | |
| SESSION.close() | |
| def add_auto_delete_msg_sync(chat_id, message_id, delay): | |
| delete_at = time.time() + delay | |
| session = SQLITE_SESSION() | |
| try: | |
| msg = AutoDeleteMessage(chat_id=chat_id, message_id=message_id, delete_at=delete_at) | |
| session.add(msg) | |
| session.commit() | |
| except Exception as e: | |
| session.rollback() | |
| LOGGER.warning(f"Error in add_auto_delete_msg_sync: {e}") | |
| finally: | |
| SQLITE_SESSION.remove() | |
| def remove_auto_delete_msg_sync(chat_id, message_id): | |
| session = SQLITE_SESSION() | |
| try: | |
| session.query(AutoDeleteMessage).filter_by(chat_id=chat_id, message_id=message_id).delete() | |
| session.commit() | |
| except Exception as e: | |
| session.rollback() | |
| LOGGER.warning(f"Error in remove_auto_delete_msg_sync: {e}") | |
| finally: | |
| SQLITE_SESSION.remove() | |
| def get_all_auto_delete_msgs_sync(): | |
| session = SQLITE_SESSION() | |
| try: | |
| all_msgs = session.query(AutoDeleteMessage).all() | |
| return [(msg.chat_id, msg.message_id, msg.delete_at) for msg in all_msgs] | |
| except Exception as e: | |
| LOGGER.warning(f"Error in get_all_auto_delete_msgs_sync: {e}") | |
| return [] | |
| finally: | |
| SQLITE_SESSION.remove() | |
| async def set_announcement_settings(enabled: bool = None, channel = None, access_hash = None) -> None: | |
| retries = 3 | |
| while retries > 0: | |
| try: | |
| with INSERTION_LOCK: | |
| session = SESSION() | |
| settings = session.query(AdminSettings).first() | |
| if not settings: | |
| settings = AdminSettings(setting_name="default") | |
| session.add(settings) | |
| if enabled is not None: | |
| settings.announcement_enabled = enabled | |
| if channel is not None: | |
| if str(channel).lower() == "off": | |
| settings.announcement_channel = None | |
| settings.announcement_channel_access_hash = None | |
| else: | |
| settings.announcement_channel = int(channel) | |
| if access_hash is not None: | |
| settings.announcement_channel_access_hash = int(access_hash) | |
| session.commit() | |
| _update_cache_from_db(settings) | |
| return True | |
| except OperationalError: | |
| reconnect_session() | |
| retries -= 1 | |
| except Exception as e: | |
| LOGGER.warning(f"Error setting announcement settings: {e}") | |
| session.rollback() | |
| retries -= 1 | |
| finally: | |
| try: | |
| session.close() | |
| except Exception: | |
| pass | |
| return False | |