Commit 3a86c294 3a86c294702648355efd1ecd7ab296dce68cbe4b by 沈秋雨

refactor(connections): 实现数据库连接池和连接管理机制

- 新增 MySQL 和 PostgreSQL 连接池类,支持按需创建和自动回收连接
- 增加连接池全局实例及其获取函数,实现连接复用
- 修改原有连接获取函数,改用连接池管理连接
- 添加 refresh_conn 函数用于检测和刷新断开的连接
- 在 runner 模块中批处理循环前调用 refresh_conn 保证长时间任务连接可用
- 用 close_all_pools 统一关闭所有连接池连接,替代原单独关闭调用
- 提升长期运行任务的数据库连接稳定性与效率
1 parent cdf0d9e2
import time
import pymysql
import pymysql.cursors
import pg8000
......@@ -5,6 +6,10 @@ 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):
......@@ -57,27 +62,272 @@ class _NulSafePgConnection:
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 pymysql.connect(**SOURCE_DB, cursorclass=pymysql.cursors.DictCursor)
return _get_source_pool().get()
def get_spider_conn() -> pymysql.Connection:
cfg = SOURCE_DB.copy()
cfg['database'] = 'hikoon-data-spider'
return pymysql.connect(**cfg, cursorclass=pymysql.cursors.DictCursor)
return _get_spider_pool().get()
def get_hk_songs_conn() -> pymysql.Connection:
return pymysql.connect(**HK_SONGS_DB, cursorclass=pymysql.cursors.DictCursor)
return _get_hk_songs_pool().get()
def get_pg_conn() -> _NulSafePgConnection:
conn = pg8000.connect(**CRAWLER_DB)
conn.run("SET TimeZone = 'Asia/Shanghai'")
return _NulSafePgConnection(conn)
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}")
if pool_type == 'pg':
return pool().check(conn)
return pool().check(conn)
......
......@@ -5,7 +5,7 @@ from concurrent.futures import ThreadPoolExecutor, as_completed
from tqdm import tqdm
from .config import PLATFORM_QQ, PLATFORM_KUGOU, PLATFORM_NETEASE, BATCH_SIZE, BACKFILL_BATCH_SIZE, OSS_CONFIG
from .connections import get_hk_songs_conn, get_source_conn, get_spider_conn, get_pg_conn, get_oss_bucket
from .connections import get_hk_songs_conn, get_source_conn, get_spider_conn, get_pg_conn, get_oss_bucket, refresh_conn, close_all_pools
from .reader import (
iter_hk_songs_batches,
iter_hk_songs_records2_batches,
......@@ -1185,6 +1185,11 @@ def initialize_yinyan_song_records(platforms: list[str], max_batches: int | None
skipped_batches = 0
try:
for i, batch in enumerate(tqdm(iter_hk_songs_batches(hk_conn, BATCH_SIZE), desc='init-yinyan')):
# 长任务连接可能断开,每批次开始前刷新
hk_conn = refresh_conn(hk_conn, 'hk_songs')
src_conn = refresh_conn(src_conn, 'source')
pg_conn = refresh_conn(pg_conn, 'pg')
if max_batches is not None and i >= max_batches:
break
# 过滤掉已存在的 song_id
......@@ -1236,9 +1241,7 @@ def initialize_yinyan_song_records(platforms: list[str], max_batches: int | None
hk_conn.commit()
log.info("Init: marked %d unmatched songs as deleted", len(unmatched_ids))
finally:
hk_conn.close()
src_conn.close()
pg_conn.close()
close_all_pools()
log.info("Initialized yinyan_song_records candidates=%d, skipped_batches=%d", total, skipped_batches)
......@@ -1254,6 +1257,11 @@ def initialize_yinyan_song_records2(platforms: list[str], max_batches: int | Non
for index, batch in enumerate(
tqdm(iter_hk_songs_records2_batches(hk_conn, BATCH_SIZE), desc='init-yinyan-records2')
):
# 长任务连接可能断开,每批次开始前刷新
hk_conn = refresh_conn(hk_conn, 'hk_songs')
src_conn = refresh_conn(src_conn, 'source')
pg_conn = refresh_conn(pg_conn, 'pg')
if max_batches is not None and index >= max_batches:
break
......@@ -1281,9 +1289,7 @@ def initialize_yinyan_song_records2(platforms: list[str], max_batches: int | Non
deleted = delete_yinyan_song_records2_existing_relations(pg_cur)
pg_conn.commit()
finally:
hk_conn.close()
src_conn.close()
pg_conn.close()
close_all_pools()
log.info(
"Initialized yinyan_song_records2 candidates=%d, removed_existing_records1_relations=%d",
......@@ -1302,6 +1308,10 @@ def backfill_yinyan_record_platforms(max_batches: int | None = None) -> None:
batch_index = 0
pbar = tqdm(desc='backfill-yinyan-platforms')
while max_batches is None or batch_index < max_batches:
# 长任务连接可能断开,每批次开始前刷新
src_conn = refresh_conn(src_conn, 'source')
pg_conn = refresh_conn(pg_conn, 'pg')
with pg_conn.cursor() as pg_cur:
rows = fetch_yinyan_records_missing_platform(pg_cur, BACKFILL_BATCH_SIZE)
if not rows:
......@@ -1334,8 +1344,7 @@ def backfill_yinyan_record_platforms(max_batches: int | None = None) -> None:
pbar.update(1)
pbar.close()
finally:
src_conn.close()
pg_conn.close()
close_all_pools()
log.info("Backfilled yinyan_song_records platform rows=%d, skipped=%d", total, skipped)
......@@ -1345,6 +1354,10 @@ def _backfill_qq_singers(spider_conn, pg_conn, bucket, base_url, max_batches):
total_ok = total_skip = 0
pbar = tqdm(desc='backfill-qq-singers')
while max_batches is None or batch_index < max_batches:
# 长任务连接可能断开,每批次开始前刷新
spider_conn = refresh_conn(spider_conn, 'spider')
pg_conn = refresh_conn(pg_conn, 'pg')
with pg_conn.cursor() as pg_cur:
rows = fetch_qq_songs_missing_singers(pg_cur, BATCH_SIZE)
if not rows:
......@@ -1421,6 +1434,10 @@ def _backfill_kugou_singers(spider_conn, pg_conn, bucket, base_url, max_batches)
total_ok = total_skip = 0
pbar = tqdm(desc='backfill-kugou-singers')
while max_batches is None or batch_index < max_batches:
# 长任务连接可能断开,每批次开始前刷新
spider_conn = refresh_conn(spider_conn, 'spider')
pg_conn = refresh_conn(pg_conn, 'pg')
with pg_conn.cursor() as pg_cur:
rows = fetch_kugou_songs_missing_singers(pg_cur, BATCH_SIZE)
if not rows:
......@@ -1491,6 +1508,10 @@ def _backfill_netease_singers(spider_conn, pg_conn, bucket, base_url, max_batche
total_ok = total_skip = 0
pbar = tqdm(desc='backfill-netease-singers')
while max_batches is None or batch_index < max_batches:
# 长任务连接可能断开,每批次开始前刷新
spider_conn = refresh_conn(spider_conn, 'spider')
pg_conn = refresh_conn(pg_conn, 'pg')
with pg_conn.cursor() as pg_cur:
rows = fetch_netease_songs_missing_singers(pg_cur, BATCH_SIZE)
if not rows:
......@@ -1569,8 +1590,7 @@ def backfill_empty_singers(platforms: list[str], max_batches: int | None = None)
if PLATFORM_NETEASE in platforms:
_backfill_netease_singers(spider_conn, pg_conn, bucket, base_url, max_batches)
finally:
spider_conn.close()
pg_conn.close()
close_all_pools()
def _is_target_oss_url(url: str | None, base_url: str) -> bool:
......@@ -1589,6 +1609,12 @@ def backfill_qq_invalid_covers(max_batches: int | None = None) -> None:
pbar = tqdm(desc='backfill-qq-invalid-covers')
try:
while max_batches is None or batch_index < max_batches:
# 长任务连接可能断开,每批次开始前刷新
hk_conn = refresh_conn(hk_conn, 'hk_songs')
src_conn = refresh_conn(src_conn, 'source')
spider_conn = refresh_conn(spider_conn, 'spider')
pg_conn = refresh_conn(pg_conn, 'pg')
with pg_conn.cursor() as pg_cur:
invalid_rows = fetch_qq_songs_with_invalid_covers(
pg_cur, QQ_MISSING_ALBUM_COVER, BATCH_SIZE,
......@@ -1674,10 +1700,7 @@ def backfill_qq_invalid_covers(max_batches: int | None = None) -> None:
pbar.update(1)
finally:
pbar.close()
hk_conn.close()
src_conn.close()
spider_conn.close()
pg_conn.close()
close_all_pools()
log.info('QQ invalid cover backfill: replaced=%d filtered=%d retry_later=%d', replaced, filtered, retry_later)
......@@ -1698,6 +1721,12 @@ def run(
batch_index = 0
pbar = tqdm(desc='batches')
while max_batches is None or batch_index < max_batches:
# 长任务连接可能断开,每批次开始前刷新
hk_conn = refresh_conn(hk_conn, 'hk_songs')
src_conn = refresh_conn(src_conn, 'source')
spider_conn = refresh_conn(spider_conn, 'spider')
pg_conn = refresh_conn(pg_conn, 'pg')
with pg_conn.cursor() as pg_cur:
pending_records = fetch_pending_yinyan_song_records(pg_cur, BATCH_SIZE)
if not pending_records:
......@@ -1880,10 +1909,7 @@ def run(
pbar.close()
finally:
hk_conn.close()
src_conn.close()
spider_conn.close()
pg_conn.close()
close_all_pools()
log.info("Done. OK=%d ERR=%d", total_ok, total_err)
if imported:
......