connections.py
2.63 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
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'])