connections.py 2.63 KB
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


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)


def get_source_conn() -> pymysql.Connection:
    return pymysql.connect(**SOURCE_DB, cursorclass=pymysql.cursors.DictCursor)


def get_spider_conn() -> pymysql.Connection:
    cfg = SOURCE_DB.copy()
    cfg['database'] = 'hikoon-data-spider'
    return pymysql.connect(**cfg, cursorclass=pymysql.cursors.DictCursor)


def get_hk_songs_conn() -> pymysql.Connection:
    return pymysql.connect(**HK_SONGS_DB, cursorclass=pymysql.cursors.DictCursor)


def get_pg_conn() -> _NulSafePgConnection:
    conn = pg8000.connect(**CRAWLER_DB)
    conn.run("SET TimeZone = 'Asia/Shanghai'")
    return _NulSafePgConnection(conn)


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'])