298 lines
8.3 KiB
Python
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()
|