126 lines
3.3 KiB
Python
126 lines
3.3 KiB
Python
"""
|
|
db.py — Пул соединений с keepalive, проверкой живости, переподключением.
|
|
"""
|
|
|
|
import os
|
|
import logging
|
|
from psycopg2 import pool as _pgpool
|
|
from psycopg2 import OperationalError, InterfaceError
|
|
import psycopg2
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
_pool = None
|
|
|
|
|
|
def _get_pool():
|
|
global _pool
|
|
if _pool is None:
|
|
_pool = _pgpool.ThreadedConnectionPool(
|
|
minconn=2,
|
|
maxconn=10,
|
|
host=os.getenv("DB_HOST"),
|
|
port=os.getenv("DB_PORT", 5432),
|
|
dbname=os.getenv("DB_NAME"),
|
|
user=os.getenv("DB_USER"),
|
|
password=os.getenv("DB_PASS"),
|
|
sslmode=os.getenv("DB_SSLMODE", "disable"),
|
|
connect_timeout=10,
|
|
keepalives=1,
|
|
keepalives_idle=30,
|
|
keepalives_interval=10,
|
|
keepalives_count=3,
|
|
)
|
|
return _pool
|
|
|
|
|
|
def _pg_connect(dbname):
|
|
"""Сырое подключение к ЛЮБОЙ базе (для /test createdb)."""
|
|
return psycopg2.connect(
|
|
host=os.getenv("DB_HOST"),
|
|
port=os.getenv("DB_PORT", 5432),
|
|
dbname=dbname,
|
|
user=os.getenv("DB_USER"),
|
|
password=os.getenv("DB_PASS"),
|
|
sslmode=os.getenv("DB_SSLMODE", "disable"),
|
|
)
|
|
|
|
|
|
def connect():
|
|
"""Подключение из пула с проверкой живости."""
|
|
pool = _get_pool()
|
|
for attempt in range(3):
|
|
try:
|
|
conn = pool.getconn()
|
|
# Проверить живость
|
|
cur = conn.cursor()
|
|
cur.execute("SELECT 1")
|
|
cur.close()
|
|
return conn, None
|
|
except (OperationalError, InterfaceError) as e:
|
|
if conn:
|
|
try: pool.putconn(conn, close=True)
|
|
except: pass
|
|
if attempt == 2:
|
|
return None, str(e)
|
|
except Exception as e:
|
|
if conn:
|
|
try: pool.putconn(conn)
|
|
except: pass
|
|
return None, str(e)
|
|
return None, "pool exhausted"
|
|
|
|
|
|
def put_conn(conn):
|
|
"""Вернуть соединение в пул."""
|
|
try:
|
|
_get_pool().putconn(conn)
|
|
except Exception:
|
|
pass
|
|
|
|
|
|
def query(sql_text, params=None):
|
|
conn, err = connect()
|
|
if err:
|
|
return None, f"connect: {err}"
|
|
try:
|
|
cur = conn.cursor()
|
|
cur.execute(sql_text, params)
|
|
rows = cur.fetchall()
|
|
cols = [desc[0] for desc in cur.description] if cur.description else []
|
|
cur.close()
|
|
return {"columns": cols, "rows": [list(r) for r in rows]}, None
|
|
except Exception as e:
|
|
return None, str(e)
|
|
finally:
|
|
put_conn(conn)
|
|
|
|
|
|
def execute(sql_text, params=None):
|
|
conn, err = connect()
|
|
if err:
|
|
return None, f"connect: {err}"
|
|
try:
|
|
cur = conn.cursor()
|
|
cur.execute(sql_text, params)
|
|
conn.commit()
|
|
rowcount = cur.rowcount
|
|
cur.close()
|
|
return rowcount, None
|
|
except Exception as e:
|
|
try: conn.rollback()
|
|
except: pass
|
|
return None, str(e)
|
|
finally:
|
|
put_conn(conn)
|
|
|
|
|
|
def query_one(sql_text, params=None):
|
|
result, err = query(sql_text, params)
|
|
if err:
|
|
return None, err
|
|
rows = result["rows"]
|
|
if not rows:
|
|
return None, None
|
|
return dict(zip(result["columns"], rows[0])), None
|