# -*- coding: utf-8 -*- # Copyright (c) 2025 relakkes@gmail.com # # This file is part of MediaCrawler project. # Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/monitor/db.py # GitHub: https://github.com/NanmiCoder # Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 # # 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: # 1. 不得用于任何商业用途。 # 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 # 3. 不得进行大规模爬取或对平台造成运营干扰。 # 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 # 5. 不得用于任何非法或不当的用途。 # # 详细许可条款请参阅项目根目录下的LICENSE文件。 # 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 """Database engine for the monitoring layer. **MySQL** by default (see ``config/db_config.py`` and ``.env``), with SQLite kept as an option so the test suite can run without a reachable server. Three things here exist because of specific MySQL 5.7 behaviour: * **utf8mb4 is forced per table.** This instance's server *and* the target schema default to ``latin1``; relying on either would mangle or reject Chinese text. The charset is set on every table rather than on the database, so it holds no matter what the schema default is. * **Connections are recycled.** The monitor runs for weeks, and MySQL drops idle connections after ``wait_timeout`` (8h by default). Without ``pool_recycle`` and ``pool_pre_ping`` the first query after a quiet night fails with "server has gone away". * **The connected schema is asserted at startup.** A misconfigured database name is caught immediately instead of silently writing to the wrong schema. Only the configured schema is ever touched: no ``CREATE DATABASE``, no ``USE``, no cross-schema query. """ import os import sys from contextlib import asynccontextmanager from pathlib import Path from typing import AsyncIterator, Optional from sqlalchemy import event, text from sqlalchemy.ext.asyncio import ( AsyncEngine, AsyncSession, async_sessionmaker, create_async_engine, ) from .models import MonitorBase PROJECT_ROOT = Path(__file__).parent.parent.parent DATA_DIR = PROJECT_ROOT / "data" DEFAULT_SQLITE_PATH = DATA_DIR / "monitor.db" # Load .env here as well as in api/main.py: this module is imported directly by # scripts and tests, and a configuration that only applies when the server is the # entry point is a trap. load_dotenv does not override real environment variables. from dotenv import load_dotenv load_dotenv(PROJECT_ROOT / ".env") # Kept identical to config/db_config.py's defaults so one .env drives both the # monitor database and the crawler's own DB output. MYSQL_HOST = lambda: os.getenv("MYSQL_DB_HOST", "localhost") # noqa: E731 MYSQL_PORT = lambda: int(os.getenv("MYSQL_DB_PORT", "3306")) # noqa: E731 MYSQL_USER = lambda: os.getenv("MYSQL_DB_USER", "root") # noqa: E731 MYSQL_PWD = lambda: os.getenv("MYSQL_DB_PWD", "") # noqa: E731 MYSQL_DB_NAME = lambda: os.getenv("MYSQL_DB_NAME", "mediacrawler") # noqa: E731 _engine: Optional[AsyncEngine] = None _session_factory: Optional[async_sessionmaker[AsyncSession]] = None # None means "resolve from the environment" (MySQL). Tests set a SQLite URL. _db_url: Optional[str] = None _expected_schema: Optional[str] = None def resolve_db_url() -> str: """Build the connection URL. MySQL unless overridden.""" if _db_url is not None: return _db_url from urllib.parse import quote_plus user = quote_plus(MYSQL_USER()) password = quote_plus(MYSQL_PWD()) host = MYSQL_HOST() port = MYSQL_PORT() name = MYSQL_DB_NAME() return f"mysql+aiomysql://{user}:{password}@{host}:{port}/{name}?charset=utf8mb4" def is_mysql() -> bool: return resolve_db_url().startswith("mysql") def set_sqlite_path(path: Path) -> None: """Point the layer at SQLite. Used by the test suite only.""" global _db_url, _engine, _session_factory, _expected_schema _db_url = f"sqlite+aiosqlite:///{Path(path)}" _engine = None _session_factory = None _expected_schema = None def set_db_url(url: str, expected_schema: Optional[str] = None) -> None: """Point the layer at an explicit URL. ``expected_schema`` enables the guard.""" global _db_url, _engine, _session_factory, _expected_schema _db_url = url _engine = None _session_factory = None _expected_schema = expected_schema def expected_schema() -> Optional[str]: """The schema the connection must be using, if the guard applies.""" if _expected_schema is not None: return _expected_schema return MYSQL_DB_NAME() if is_mysql() else None def get_engine() -> AsyncEngine: global _engine if _engine is None: url = resolve_db_url() kwargs: dict = {"future": True} if url.startswith("mysql"): # Recycle well inside MySQL's default 8h wait_timeout, and verify a # pooled connection before handing it out. kwargs.update(pool_recycle=3600, pool_pre_ping=True, pool_size=5, max_overflow=5) kwargs["connect_args"] = {"charset": "utf8mb4"} else: Path(url.split("///", 1)[-1]).parent.mkdir(parents=True, exist_ok=True) _engine = create_async_engine(url, **kwargs) if url.startswith("sqlite"): @event.listens_for(_engine.sync_engine, "connect") def _set_sqlite_pragmas(dbapi_connection, _connection_record): # pragma: no cover cursor = dbapi_connection.cursor() cursor.execute("PRAGMA journal_mode=WAL") cursor.execute("PRAGMA foreign_keys=ON") cursor.close() return _engine def get_session_factory() -> async_sessionmaker[AsyncSession]: global _session_factory if _session_factory is None: _session_factory = async_sessionmaker( bind=get_engine(), class_=AsyncSession, expire_on_commit=False, ) return _session_factory @asynccontextmanager async def get_session() -> AsyncIterator[AsyncSession]: """Transactional session. Commits on success, rolls back on error.""" factory = get_session_factory() async with factory() as session: try: yield session await session.commit() except Exception: await session.rollback() raise async def _assert_correct_schema(conn) -> None: """Refuse to run against anything but the configured schema. A guard, not the guarantee: the real protection is a MySQL account scoped to this one schema (see UPSTREAM.md). This catches the ordinary mistake of a wrong database name in configuration, before a single row is written. """ if not is_mysql(): return expected = expected_schema() if not expected: return current = (await conn.execute(text("SELECT DATABASE()"))).scalar() if current is None: raise RuntimeError( f"数据库连接未选定 schema,期望 {expected!r}。请检查 MYSQL_DB_NAME。" ) # lower_case_table_names=1 makes names case-insensitive server-side. if current.lower() != expected.lower(): raise RuntimeError( f"连接的库是 {current!r},但配置要求 {expected!r}。" f"为避免误写其它库,已拒绝启动。" ) print(f"[monitor.db] 已连接 MySQL schema: {current}", flush=True) async def init_db() -> None: """Create missing tables, then run the small in-place migrations.""" engine = get_engine() async with engine.begin() as conn: await _assert_correct_schema(conn) await conn.run_sync(MonitorBase.metadata.create_all) await _ensure_columns(conn) await _migrate_setting_keys(conn) # Columns added to a table after it may already exist. ``create_all`` only # creates missing *tables*, so new columns need an explicit ALTER TABLE. _ADDED_COLUMNS: dict[str, list[tuple[str, str]]] = { "monitor_task": [ ("notify_enabled", "BOOLEAN NOT NULL DEFAULT 0"), ("last_notified_at", "BIGINT NULL"), ], } async def _existing_columns(conn, table: str) -> set[str]: if is_mysql(): rows = await conn.execute( text( "SELECT COLUMN_NAME FROM information_schema.COLUMNS " "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = :t" ), {"t": table}, ) return {row[0] for row in rows} rows = await conn.execute(text(f"PRAGMA table_info({table})")) return {row[1] for row in rows} async def _ensure_columns(conn) -> None: for table, columns in _ADDED_COLUMNS.items(): existing = await _existing_columns(conn, table) if not existing: # Table did not exist before this run; create_all built it complete. continue for name, ddl in columns: if name not in existing: await conn.execute(text(f"ALTER TABLE {table} ADD COLUMN {name} {ddl}")) async def _migrate_setting_keys(conn) -> None: """Move pre-namespacing setting keys to their scoped names. Idempotent: the legacy row is only renamed when the new key is absent, so an operator's later value is never overwritten. """ from .models import LEGACY_SETTING_KEY_RENAMES for legacy, scoped in LEGACY_SETTING_KEY_RENAMES.items(): exists = ( await conn.execute( text("SELECT 1 FROM monitor_setting WHERE `key` = :k"), {"k": legacy} ) ).first() if not exists: continue already = ( await conn.execute( text("SELECT 1 FROM monitor_setting WHERE `key` = :k"), {"k": scoped} ) ).first() if already: # Both present: the scoped one is authoritative; drop the stale row. await conn.execute( text("DELETE FROM monitor_setting WHERE `key` = :k"), {"k": legacy} ) continue await conn.execute( text("UPDATE monitor_setting SET `key` = :new WHERE `key` = :old"), {"new": scoped, "old": legacy}, ) async def dispose_engine() -> None: global _engine, _session_factory if _engine is not None: await _engine.dispose() _engine = None _session_factory = None