Files
2026-06-08 18:14:59 +08:00

298 lines
8.3 KiB
Python

#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""外部数据库连接池(按连接目标复用)"""
import asyncio
import hashlib
import logging
from typing import Any, Dict, Optional
logger = logging.getLogger(__name__)
_POOL_MAX_SIZE = 5
_pool_lock = asyncio.Lock()
def _pool_identity(user: str, password: str) -> str:
"""生成连接池身份标识(含密码摘要,避免换密后复用旧池)"""
digest = hashlib.sha256(f"{user or ''}\0{password or ''}".encode()).hexdigest()[:12]
return f"{user or '-'}#{digest}"
_pg_pool_cache: Dict[str, Any] = {}
_mysql_pool_cache: Dict[str, Any] = {}
_mssql_pool_cache: Dict[str, Any] = {}
_oracle_pool_cache: Dict[str, Any] = {}
async def get_asyncpg_pool(
host: str,
port: int,
user: str,
password: str,
database: str,
):
import asyncpg
db_name = database or "postgres"
key = f"pg://{_pool_identity(user, password)}@{host}:{port}/{db_name}"
async with _pool_lock:
pool = _pg_pool_cache.get(key)
if pool is None or getattr(pool, "_closed", False):
pool = await asyncpg.create_pool(
host=host,
port=port,
user=user,
password=password,
database=db_name,
min_size=0,
max_size=_POOL_MAX_SIZE,
timeout=10,
command_timeout=60,
)
_pg_pool_cache[key] = pool
logger.debug("Created asyncpg pool: %s (max=%s)", key, _POOL_MAX_SIZE)
return pool
async def get_aiomysql_pool(
host: str,
port: int,
user: str,
password: str,
database: str,
):
import aiomysql
db_name = database or ""
key = f"mysql://{_pool_identity(user, password)}@{host}:{port}/{db_name}"
async with _pool_lock:
pool = _mysql_pool_cache.get(key)
if pool is None or pool.closed:
pool = await aiomysql.create_pool(
host=host,
port=port,
user=user,
password=password,
db=db_name,
charset="utf8mb4",
autocommit=True,
minsize=0,
maxsize=_POOL_MAX_SIZE,
connect_timeout=10,
)
_mysql_pool_cache[key] = pool
logger.debug("Created aiomysql pool: %s (max=%s)", key, _POOL_MAX_SIZE)
return pool
def build_mssql_dsn(
host: str,
port: int,
user: str,
password: str,
database: str,
extra_options: Optional[Dict[str, Any]] = None,
) -> str:
opts = extra_options or {}
driver = opts.get("odbc_driver", "ODBC Driver 18 for SQL Server")
db_name = database or "master"
encrypt = opts.get("encrypt", "yes")
trust = opts.get("trust_server_certificate", "yes")
return (
f"DRIVER={{{driver}}};"
f"SERVER={host},{port};"
f"DATABASE={db_name};"
f"UID={user};"
f"PWD={password};"
f"Encrypt={encrypt};"
f"TrustServerCertificate={trust};"
)
async def get_aioodbc_pool(
host: str,
port: int,
user: str,
password: str,
database: str,
extra_options: Optional[Dict[str, Any]] = None,
):
import aioodbc
db_name = database or "master"
key = f"mssql://{_pool_identity(user, password)}@{host}:{port}/{db_name}"
async with _pool_lock:
pool = _mssql_pool_cache.get(key)
if pool is None or pool.closed:
dsn = build_mssql_dsn(host, port, user, password, db_name, extra_options)
pool = await aioodbc.create_pool(
dsn=dsn,
minsize=0,
maxsize=_POOL_MAX_SIZE,
autocommit=True,
)
_mssql_pool_cache[key] = pool
logger.debug("Created aioodbc pool: %s (max=%s)", key, _POOL_MAX_SIZE)
return pool
def build_oracle_dsn(host: str, port: int, service_name: str) -> str:
import oracledb
service = service_name or "ORCL"
return oracledb.make_dsn(host, port, service_name=service)
async def get_oracledb_pool(
host: str,
port: int,
user: str,
password: str,
database: str,
extra_options: Optional[Dict[str, Any]] = None,
):
import oracledb
service = database or (extra_options or {}).get("service_name") or "ORCL"
key = f"oracle://{_pool_identity(user, password)}@{host}:{port}/{service}"
async with _pool_lock:
pool = _oracle_pool_cache.get(key)
if pool is None or (hasattr(pool, "opened") and not pool.opened):
dsn = build_oracle_dsn(host, port, service)
pool = oracledb.create_pool_async(
user=user,
password=password,
dsn=dsn,
min=0,
max=_POOL_MAX_SIZE,
)
await pool.open()
_oracle_pool_cache[key] = pool
logger.debug("Created oracledb pool: %s (max=%s)", key, _POOL_MAX_SIZE)
return pool
async def probe_asyncpg_connection(
host: str,
port: int,
user: str,
password: str,
database: str,
) -> None:
"""单次 PostgreSQL 探测(不走连接池,避免多 worker 下池缓存导致测试结果抖动)"""
import asyncpg
conn = await asyncpg.connect(
host=host,
port=port,
user=user,
password=password,
database=database or "postgres",
timeout=10,
)
try:
await conn.execute("SELECT 1")
finally:
await conn.close()
async def probe_aiomysql_connection(
host: str,
port: int,
user: str,
password: str,
database: str,
) -> None:
"""单次 MySQL 探测"""
import aiomysql
conn = await aiomysql.connect(
host=host,
port=port,
user=user,
password=password,
db=database or "",
charset="utf8mb4",
connect_timeout=10,
)
try:
async with conn.cursor() as cursor:
await cursor.execute("SELECT 1")
finally:
conn.close()
async def probe_aioodbc_connection(
host: str,
port: int,
user: str,
password: str,
database: str,
extra_options: Optional[Dict[str, Any]] = None,
) -> None:
"""单次 SQL Server 探测"""
import aioodbc
dsn = build_mssql_dsn(host, port, user, password, database or "master", extra_options)
conn = await aioodbc.connect(dsn=dsn, timeout=10)
try:
async with conn.cursor() as cursor:
await cursor.execute("SELECT 1")
finally:
await conn.close()
async def probe_oracledb_connection(
host: str,
port: int,
user: str,
password: str,
database: str,
extra_options: Optional[Dict[str, Any]] = None,
) -> None:
"""单次 Oracle 探测"""
import oracledb
service = database or (extra_options or {}).get("service_name") or "ORCL"
dsn = build_oracle_dsn(host, port, service)
conn = await oracledb.connect_async(user=user, password=password, dsn=dsn)
try:
async with conn.cursor() as cursor:
await cursor.execute("SELECT 1 FROM DUAL")
finally:
await conn.close()
async def close_all_manager_pools() -> None:
"""关闭所有外部数据库连接池(应用 shutdown 时调用)"""
async with _pool_lock:
for key, pool in list(_pg_pool_cache.items()):
try:
await pool.close()
except Exception as exc:
logger.warning("Failed to close asyncpg pool %s: %s", key, exc)
_pg_pool_cache.clear()
for key, pool in list(_mysql_pool_cache.items()):
try:
pool.close()
await pool.wait_closed()
except Exception as exc:
logger.warning("Failed to close aiomysql pool %s: %s", key, exc)
_mysql_pool_cache.clear()
for key, pool in list(_mssql_pool_cache.items()):
try:
pool.close()
await pool.wait_closed()
except Exception as exc:
logger.warning("Failed to close aioodbc pool %s: %s", key, exc)
_mssql_pool_cache.clear()
for key, pool in list(_oracle_pool_cache.items()):
try:
await pool.close()
except Exception as exc:
logger.warning("Failed to close oracledb pool %s: %s", key, exc)
_oracle_pool_cache.clear()