connections.py 9.58 KB
import time
import pymysql
import pymysql.cursors
import pg8000
import oss2
from .config import SOURCE_DB, HK_SONGS_DB, CRAWLER_DB, OSS_CONFIG, OSS_CONNECTION_POOL_SIZE


CONN_POOL_SIZE = 4
CONN_RECYCLE_SECONDS = 600  # 连接最长存活 10 分钟,超时自动换新


def _strip_nul(value):
    """PostgreSQL text/json 参数不接受 NUL 字符,递归清理全部绑定参数。"""
    if isinstance(value, str):
        return value.replace('\x00', '')
    if isinstance(value, tuple):
        return tuple(_strip_nul(item) for item in value)
    if isinstance(value, list):
        return [_strip_nul(item) for item in value]
    if isinstance(value, dict):
        return {key: _strip_nul(item) for key, item in value.items()}
    return value


class _NulSafePgCursor:
    """在所有 PostgreSQL 参数化查询前统一移除文本中的 NUL 字符。"""

    def __init__(self, cursor):
        self._cursor = cursor

    def __enter__(self):
        self._cursor.__enter__()
        return self

    def __exit__(self, exc_type, exc_value, traceback):
        return self._cursor.__exit__(exc_type, exc_value, traceback)

    def execute(self, operation, parameters=None):
        if parameters is None:
            return self._cursor.execute(operation)
        return self._cursor.execute(operation, _strip_nul(parameters))

    def executemany(self, operation, parameter_sets):
        return self._cursor.executemany(
            operation,
            [_strip_nul(parameters) for parameters in parameter_sets],
        )

    def __getattr__(self, name):
        return getattr(self._cursor, name)


class _NulSafePgConnection:
    def __init__(self, connection):
        self._connection = connection

    def cursor(self, *args, **kwargs):
        return _NulSafePgCursor(self._connection.cursor(*args, **kwargs))

    def __getattr__(self, name):
        return getattr(self._connection, name)


class _MySqlConnPool:
    """MySQL 连接池:按需创建连接,定期回收,获取时自动重连。"""

    def __init__(self, connect_kwargs, size=CONN_POOL_SIZE, recycle=CONN_RECYCLE_SECONDS):
        self._connect_kwargs = connect_kwargs
        self._size = size
        self._recycle = recycle
        self._pool = []
        self._created_at = []

    def _new_conn(self):
        conn = pymysql.connect(**self._connect_kwargs, cursorclass=pymysql.cursors.DictCursor)
        self._pool.append(conn)
        self._created_at.append(time.monotonic())
        return conn

    def get(self) -> pymysql.Connection:
        for i in range(len(self._pool)):
            conn = self._pool[i]
            try:
                conn.ping(reconnect=False)
            except Exception:
                pass
            else:
                if time.monotonic() - self._created_at[i] < self._recycle:
                    return conn
                # 连接超时回收
                try:
                    conn.close()
                except Exception:
                    pass
                self._pool.pop(i)
                self._created_at.pop(i)
                return self._new_conn()
            # ping 失败,尝试重连
            try:
                conn.ping(reconnect=True)
                self._created_at[i] = time.monotonic()
                return conn
            except Exception:
                try:
                    conn.close()
                except Exception:
                    pass
                self._pool.pop(i)
                self._created_at.pop(i)
                return self._new_conn()
        return self._new_conn()

    def check(self, conn) -> pymysql.Connection:
        """确保连接可用;如断开则自动替换为新连接。"""
        try:
            conn.ping(reconnect=False)
        except Exception:
            pass
        else:
            if id(conn) in {id(c) for c in self._pool}:
                return conn
        try:
            conn.ping(reconnect=True)
            return conn
        except Exception:
            pass
        try:
            idx = None
            for i, c in enumerate(self._pool):
                if id(c) == id(conn):
                    idx = i
                    break
            if idx is not None:
                try:
                    conn.close()
                except Exception:
                    pass
                self._pool.pop(idx)
                self._created_at.pop(idx)
        except Exception:
            pass
        return self._new_conn()

    def close_all(self):
        for conn in self._pool:
            try:
                conn.close()
            except Exception:
                pass
        self._pool.clear()
        self._created_at.clear()


class _PgConnPool:
    """PostgreSQL 连接池:按需创建连接,定期回收,获取时自动重连。"""

    def __init__(self, connect_kwargs, size=CONN_POOL_SIZE, recycle=CONN_RECYCLE_SECONDS):
        self._connect_kwargs = connect_kwargs
        self._size = size
        self._recycle = recycle
        self._pool = []
        self._created_at = []

    def _new_conn(self):
        conn = pg8000.connect(**self._connect_kwargs)
        conn.run("SET TimeZone = 'Asia/Shanghai'")
        wrapped = _NulSafePgConnection(conn)
        self._pool.append(conn)
        self._created_at.append(time.monotonic())
        return wrapped

    def get(self) -> _NulSafePgConnection:
        for i in range(len(self._pool)):
            conn = self._pool[i]
            try:
                conn.run("SELECT 1")
            except Exception:
                try:
                    conn.close()
                except Exception:
                    pass
                self._pool.pop(i)
                self._created_at.pop(i)
                return self._new_conn()
            if time.monotonic() - self._created_at[i] < self._recycle:
                return _NulSafePgConnection(conn)
            # 回收
            try:
                conn.close()
            except Exception:
                pass
            self._pool.pop(i)
            self._created_at.pop(i)
            return self._new_conn()
        return self._new_conn()

    def check(self, conn) -> _NulSafePgConnection:
        """确保连接可用;如断开则自动替换为新连接。"""
        try:
            conn.run("SELECT 1")
        except Exception:
            pass
        else:
            raw = conn._connection if hasattr(conn, '_connection') else conn
            if id(raw) in {id(c) for c in self._pool}:
                return conn
        # 创建新连接替换当前连接
        try:
            idx = None
            for i, c in enumerate(self._pool):
                if id(c) == id(raw):
                    idx = i
                    break
            if idx is not None:
                try:
                    raw.close()
                except Exception:
                    pass
                self._pool.pop(idx)
                self._created_at.pop(idx)
        except Exception:
            pass
        return self._new_conn()

    def close_all(self):
        for conn in self._pool:
            try:
                conn.close()
            except Exception:
                pass
        self._pool.clear()
        self._created_at.clear()


# ── 全局连接池实例 ──

_source_pool = None
_hk_songs_pool = None
_spider_pool = None
_pg_pool = None


def _get_source_pool() -> _MySqlConnPool:
    global _source_pool
    if _source_pool is None:
        _source_pool = _MySqlConnPool(SOURCE_DB)
    return _source_pool


def _get_hk_songs_pool() -> _MySqlConnPool:
    global _hk_songs_pool
    if _hk_songs_pool is None:
        _hk_songs_pool = _MySqlConnPool(HK_SONGS_DB)
    return _hk_songs_pool


def _get_spider_pool() -> _MySqlConnPool:
    global _spider_pool
    if _spider_pool is None:
        cfg = SOURCE_DB.copy()
        cfg['database'] = 'hikoon-data-spider'
        _spider_pool = _MySqlConnPool(cfg)
    return _spider_pool


def _get_pg_pool() -> _PgConnPool:
    global _pg_pool
    if _pg_pool is None:
        _pg_pool = _PgConnPool(CRAWLER_DB)
    return _pg_pool


def get_source_conn() -> pymysql.Connection:
    return _get_source_pool().get()


def get_spider_conn() -> pymysql.Connection:
    return _get_spider_pool().get()


def get_hk_songs_conn() -> pymysql.Connection:
    return _get_hk_songs_pool().get()


def get_pg_conn() -> _NulSafePgConnection:
    return _get_pg_pool().get()


def get_oss_bucket() -> oss2.Bucket:
    oss2.defaults.connection_pool_size = OSS_CONNECTION_POOL_SIZE
    auth = oss2.Auth(OSS_CONFIG['access_key_id'], OSS_CONFIG['access_key_secret'])
    return oss2.Bucket(auth, OSS_CONFIG['endpoint'], OSS_CONFIG['bucket_name'])


def close_all_pools():
    """关闭所有连接池中的连接。应在 ETL 退出时调用。"""
    for pool in (_source_pool, _hk_songs_pool, _spider_pool):
        if pool is not None:
            pool.close_all()
    if _pg_pool is not None:
        _pg_pool.close_all()


def refresh_conn(conn, pool_type: str):
    """验证连接健康度,断开时返回新连接,正常时原样返回。

    用法:在批次循环开头调用,如 ``conn = refresh_conn(conn, 'pg')``。
    pool_type 取 'source' | 'hk_songs' | 'spider' | 'pg'。
    """
    try:
        if pool_type == 'pg':
            conn.run("SELECT 1")
            return conn
        # MySQL
        conn.ping(reconnect=False)
        return conn
    except Exception:
        pass

    pools = {
        'source': _get_source_pool,
        'hk_songs': _get_hk_songs_pool,
        'spider': _get_spider_pool,
        'pg': _get_pg_pool,
    }
    pool = pools.get(pool_type)
    if pool is None:
        raise ValueError(f"unknown pool_type: {pool_type}")

    return pool().check(conn)