570 lines
19 KiB
Python
570 lines
19 KiB
Python
import sqlite3
|
||
import os
|
||
import logging
|
||
from datetime import datetime
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
_active_db = None
|
||
_mysql_pool = None
|
||
|
||
|
||
def _get_project_root():
|
||
return os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||
|
||
|
||
def _get_sqlite_db_path():
|
||
from lib.config import get_sqlite_config
|
||
cfg = get_sqlite_config()
|
||
path = cfg['path']
|
||
if not os.path.isabs(path):
|
||
path = os.path.join(_get_project_root(), path)
|
||
return path
|
||
|
||
|
||
def _try_mysql_connect():
|
||
try:
|
||
import pymysql
|
||
except ImportError:
|
||
logger.warning("pymysql 未安装, 无法使用 MySQL")
|
||
return None
|
||
|
||
from lib.config import get_mysql_config
|
||
cfg = get_mysql_config()
|
||
|
||
try:
|
||
conn = pymysql.connect(
|
||
host=cfg['host'],
|
||
port=cfg['port'],
|
||
user=cfg['user'],
|
||
password=cfg['password'],
|
||
database=cfg['database'],
|
||
charset='utf8mb4',
|
||
cursorclass=pymysql.cursors.DictCursor,
|
||
connect_timeout=5,
|
||
)
|
||
return conn
|
||
except Exception as e:
|
||
logger.warning("MySQL 连接失败: %s", str(e))
|
||
return None
|
||
|
||
|
||
def _ensure_mysql_database():
|
||
try:
|
||
import pymysql
|
||
except ImportError:
|
||
return False
|
||
|
||
from lib.config import get_mysql_config
|
||
cfg = get_mysql_config()
|
||
|
||
try:
|
||
conn = pymysql.connect(
|
||
host=cfg['host'],
|
||
port=cfg['port'],
|
||
user=cfg['user'],
|
||
password=cfg['password'],
|
||
charset='utf8mb4',
|
||
cursorclass=pymysql.cursors.DictCursor,
|
||
connect_timeout=5,
|
||
)
|
||
with conn.cursor() as cursor:
|
||
cursor.execute(
|
||
"SELECT SCHEMA_NAME FROM INFORMATION_SCHEMA.SCHEMATA WHERE SCHEMA_NAME = %s",
|
||
(cfg['database'],)
|
||
)
|
||
if not cursor.fetchone():
|
||
cursor.execute(
|
||
"CREATE DATABASE `%s` CHARACTER SET utf8mb4 COLLATE utf8mb4_general_ci" % cfg['database']
|
||
)
|
||
conn.commit()
|
||
logger.info("MySQL 数据库 '%s' 创建成功", cfg['database'])
|
||
conn.close()
|
||
return True
|
||
except Exception as e:
|
||
logger.warning("MySQL 确保数据库存在失败: %s", str(e))
|
||
return False
|
||
|
||
|
||
def _get_mysql_connection():
|
||
global _mysql_pool
|
||
if _mysql_pool is not None:
|
||
try:
|
||
_mysql_pool.ping(reconnect=True)
|
||
return _mysql_pool
|
||
except Exception:
|
||
_mysql_pool = None
|
||
|
||
conn = _try_mysql_connect()
|
||
if conn:
|
||
_mysql_pool = conn
|
||
return conn
|
||
|
||
|
||
def _get_sqlite_connection():
|
||
db_path = _get_sqlite_db_path()
|
||
os.makedirs(os.path.dirname(db_path), exist_ok=True)
|
||
conn = sqlite3.connect(db_path)
|
||
conn.row_factory = sqlite3.Row
|
||
conn.execute("PRAGMA journal_mode=WAL")
|
||
conn.execute("PRAGMA foreign_keys=ON")
|
||
return conn
|
||
|
||
|
||
def get_connection():
|
||
if _active_db == 'mysql':
|
||
conn = _get_mysql_connection()
|
||
if conn:
|
||
return conn
|
||
if _active_db == 'sqlite':
|
||
return _get_sqlite_connection()
|
||
return None
|
||
|
||
|
||
def init_database(strict_mysql=False):
|
||
"""初始化数据库。
|
||
|
||
参数:
|
||
- ``strict_mysql``: True 表示严格按 SRS 15.3.4 — 配置了 mysql 但不可用时返回失败,
|
||
不静默回退 SQLite。False(启动期默认)允许回退以便开发模式可用。
|
||
"""
|
||
global _active_db
|
||
|
||
from lib.config import get_db_type
|
||
db_type = get_db_type()
|
||
|
||
if db_type == 'mysql':
|
||
ok = _ensure_mysql_database()
|
||
conn = _try_mysql_connect() if ok else None
|
||
if conn:
|
||
_active_db = 'mysql'
|
||
_init_mysql_tables(conn)
|
||
conn.close()
|
||
logger.info("数据库初始化完成 (MySQL)")
|
||
return True
|
||
else:
|
||
if strict_mysql:
|
||
logger.error("MySQL 不可用且 strict_mysql=True,拒绝服务")
|
||
return False
|
||
logger.warning("MySQL 不可用, 回退到 SQLite")
|
||
_active_db = 'sqlite'
|
||
elif db_type == 'sqlite':
|
||
_active_db = 'sqlite'
|
||
else:
|
||
conn = _try_mysql_connect()
|
||
if conn:
|
||
_active_db = 'mysql'
|
||
_ensure_mysql_database()
|
||
_init_mysql_tables(conn)
|
||
conn.close()
|
||
logger.info("数据库初始化完成 (MySQL, 自动检测)")
|
||
return True
|
||
else:
|
||
if strict_mysql:
|
||
logger.error("MySQL 不可用且 strict_mysql=True,拒绝服务")
|
||
return False
|
||
logger.warning("MySQL 不可用, 回退到 SQLite")
|
||
_active_db = 'sqlite'
|
||
|
||
if _active_db == 'sqlite':
|
||
_init_sqlite_tables()
|
||
logger.info("数据库初始化完成 (SQLite): %s", _get_sqlite_db_path())
|
||
return True
|
||
|
||
|
||
def ensure_mysql_available():
|
||
"""SRS 15.3.4:API 路径在配置 type=mysql 时必须保证 MySQL 可用。
|
||
|
||
返回 True 表示可用;False 表示应返回 5xx。
|
||
"""
|
||
from lib.config import get_db_type
|
||
db_type = get_db_type()
|
||
if db_type == 'sqlite':
|
||
return True
|
||
global _active_db
|
||
if _active_db == 'mysql':
|
||
conn = _get_mysql_connection()
|
||
return conn is not None
|
||
return False
|
||
|
||
|
||
def _init_mysql_tables(conn):
|
||
try:
|
||
with conn.cursor() as cursor:
|
||
cursor.execute("""
|
||
CREATE TABLE IF NOT EXISTS user_sign (
|
||
user_id VARCHAR(50) PRIMARY KEY,
|
||
user_name VARCHAR(100) NOT NULL,
|
||
sign_image LONGBLOB NOT NULL,
|
||
is_signature TINYINT(1) NOT NULL DEFAULT 0,
|
||
create_time DATETIME DEFAULT CURRENT_TIMESTAMP,
|
||
update_time DATETIME DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP
|
||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4
|
||
""")
|
||
conn.commit()
|
||
logger.info("MySQL 数据表初始化完成")
|
||
except Exception as e:
|
||
logger.error("MySQL 数据表初始化失败: %s", str(e))
|
||
raise
|
||
|
||
|
||
def _init_sqlite_tables():
|
||
conn = _get_sqlite_connection()
|
||
try:
|
||
conn.execute("""
|
||
CREATE TABLE IF NOT EXISTS user_sign (
|
||
user_id TEXT PRIMARY KEY,
|
||
user_name TEXT NOT NULL,
|
||
sign_image BLOB NOT NULL,
|
||
is_signature INTEGER NOT NULL DEFAULT 0,
|
||
create_time TEXT DEFAULT (datetime('now','localtime')),
|
||
update_time TEXT DEFAULT (datetime('now','localtime'))
|
||
)
|
||
""")
|
||
conn.commit()
|
||
except Exception as e:
|
||
conn.rollback()
|
||
logger.error("SQLite 数据表初始化失败: %s", str(e))
|
||
raise
|
||
finally:
|
||
conn.close()
|
||
|
||
|
||
def query_by_user_id(user_id):
|
||
if _active_db == 'mysql':
|
||
return _mysql_query_by_user_id(user_id)
|
||
return _sqlite_query_by_user_id(user_id)
|
||
|
||
|
||
def _mysql_query_by_user_id(user_id):
|
||
conn = _get_mysql_connection()
|
||
if not conn:
|
||
return _sqlite_query_by_user_id(user_id)
|
||
try:
|
||
with conn.cursor() as cursor:
|
||
cursor.execute("SELECT * FROM user_sign WHERE user_id = %s", (user_id,))
|
||
row = cursor.fetchone()
|
||
return row if row else None
|
||
except Exception as e:
|
||
logger.error("MySQL 查询失败: %s", str(e))
|
||
return _sqlite_query_by_user_id(user_id)
|
||
|
||
|
||
def _sqlite_query_by_user_id(user_id):
|
||
conn = _get_sqlite_connection()
|
||
try:
|
||
cursor = conn.execute("SELECT * FROM user_sign WHERE user_id = ?", (user_id,))
|
||
row = cursor.fetchone()
|
||
return dict(row) if row else None
|
||
finally:
|
||
conn.close()
|
||
|
||
|
||
def query_by_user_name(user_name):
|
||
if _active_db == 'mysql':
|
||
return _mysql_query_by_user_name(user_name)
|
||
return _sqlite_query_by_user_name(user_name)
|
||
|
||
|
||
def _mysql_query_by_user_name(user_name):
|
||
conn = _get_mysql_connection()
|
||
if not conn:
|
||
return _sqlite_query_by_user_name(user_name)
|
||
try:
|
||
with conn.cursor() as cursor:
|
||
cursor.execute("SELECT * FROM user_sign WHERE user_name = %s", (user_name,))
|
||
return cursor.fetchall()
|
||
except Exception as e:
|
||
logger.error("MySQL 查询失败: %s", str(e))
|
||
return _sqlite_query_by_user_name(user_name)
|
||
|
||
|
||
def _sqlite_query_by_user_name(user_name):
|
||
conn = _get_sqlite_connection()
|
||
try:
|
||
cursor = conn.execute("SELECT * FROM user_sign WHERE user_name = ?", (user_name,))
|
||
return [dict(r) for r in cursor.fetchall()]
|
||
finally:
|
||
conn.close()
|
||
|
||
|
||
def query_all(page=1, page_size=20, keyword=None):
|
||
if _active_db == 'mysql':
|
||
return _mysql_query_all(page, page_size, keyword)
|
||
return _sqlite_query_all(page, page_size, keyword)
|
||
|
||
|
||
def _mysql_query_all(page=1, page_size=20, keyword=None):
|
||
conn = _get_mysql_connection()
|
||
if not conn:
|
||
return _sqlite_query_all(page, page_size, keyword)
|
||
try:
|
||
offset = (page - 1) * page_size
|
||
with conn.cursor() as cursor:
|
||
if keyword:
|
||
cursor.execute(
|
||
"SELECT COUNT(*) as total FROM user_sign WHERE user_id LIKE %s OR user_name LIKE %s",
|
||
(f'%{keyword}%', f'%{keyword}%')
|
||
)
|
||
total = cursor.fetchone()['total']
|
||
cursor.execute(
|
||
"SELECT user_id, user_name, is_signature, create_time, update_time FROM user_sign WHERE user_id LIKE %s OR user_name LIKE %s ORDER BY create_time DESC LIMIT %s OFFSET %s",
|
||
(f'%{keyword}%', f'%{keyword}%', page_size, offset)
|
||
)
|
||
else:
|
||
cursor.execute("SELECT COUNT(*) as total FROM user_sign")
|
||
total = cursor.fetchone()['total']
|
||
cursor.execute(
|
||
"SELECT user_id, user_name, is_signature, create_time, update_time FROM user_sign ORDER BY create_time DESC LIMIT %s OFFSET %s",
|
||
(page_size, offset)
|
||
)
|
||
rows = cursor.fetchall()
|
||
return {'total': total, 'page': page, 'page_size': page_size, 'data': rows}
|
||
except Exception as e:
|
||
logger.error("MySQL 查询失败: %s", str(e))
|
||
return _sqlite_query_all(page, page_size, keyword)
|
||
|
||
|
||
def _sqlite_query_all(page=1, page_size=20, keyword=None):
|
||
conn = _get_sqlite_connection()
|
||
try:
|
||
offset = (page - 1) * page_size
|
||
if keyword:
|
||
count_row = conn.execute(
|
||
"SELECT COUNT(*) as total FROM user_sign WHERE user_id LIKE ? OR user_name LIKE ?",
|
||
(f'%{keyword}%', f'%{keyword}%')
|
||
).fetchone()
|
||
total = count_row['total']
|
||
cursor = conn.execute(
|
||
"SELECT user_id, user_name, is_signature, create_time, update_time FROM user_sign WHERE user_id LIKE ? OR user_name LIKE ? ORDER BY create_time DESC LIMIT ? OFFSET ?",
|
||
(f'%{keyword}%', f'%{keyword}%', page_size, offset)
|
||
)
|
||
else:
|
||
count_row = conn.execute("SELECT COUNT(*) as total FROM user_sign").fetchone()
|
||
total = count_row['total']
|
||
cursor = conn.execute(
|
||
"SELECT user_id, user_name, is_signature, create_time, update_time FROM user_sign ORDER BY create_time DESC LIMIT ? OFFSET ?",
|
||
(page_size, offset)
|
||
)
|
||
rows = [dict(r) for r in cursor.fetchall()]
|
||
return {'total': total, 'page': page, 'page_size': page_size, 'data': rows}
|
||
finally:
|
||
conn.close()
|
||
|
||
|
||
def insert_record(user_id, user_name, sign_image, is_signature=0):
|
||
if _active_db == 'mysql':
|
||
return _mysql_insert_record(user_id, user_name, sign_image, is_signature)
|
||
return _sqlite_insert_record(user_id, user_name, sign_image, is_signature)
|
||
|
||
|
||
def _mysql_insert_record(user_id, user_name, sign_image, is_signature=0):
|
||
conn = _get_mysql_connection()
|
||
if not conn:
|
||
return _sqlite_insert_record(user_id, user_name, sign_image, is_signature)
|
||
try:
|
||
with conn.cursor() as cursor:
|
||
cursor.execute(
|
||
"INSERT INTO user_sign (user_id, user_name, sign_image, is_signature) VALUES (%s, %s, %s, %s)",
|
||
(user_id, user_name, sign_image, is_signature)
|
||
)
|
||
conn.commit()
|
||
logger.info("新增记录(MySQL): user_id=%s, user_name=%s", user_id, user_name)
|
||
return True
|
||
except Exception as e:
|
||
conn.rollback()
|
||
if 'Duplicate entry' in str(e) or 'PRIMARY' in str(e):
|
||
logger.error("新增记录失败(主键冲突): user_id=%s, %s", user_id, str(e))
|
||
raise ValueError(f"人员ID '{user_id}' 已存在")
|
||
logger.error("MySQL 新增记录失败: %s", str(e))
|
||
raise
|
||
except:
|
||
conn.rollback()
|
||
raise
|
||
|
||
|
||
def _sqlite_insert_record(user_id, user_name, sign_image, is_signature=0):
|
||
conn = _get_sqlite_connection()
|
||
try:
|
||
conn.execute(
|
||
"INSERT INTO user_sign (user_id, user_name, sign_image, is_signature) VALUES (?, ?, ?, ?)",
|
||
(user_id, user_name, sign_image, is_signature)
|
||
)
|
||
conn.commit()
|
||
logger.info("新增记录(SQLite): user_id=%s, user_name=%s", user_id, user_name)
|
||
return True
|
||
except sqlite3.IntegrityError as e:
|
||
conn.rollback()
|
||
logger.error("新增记录失败(主键冲突): user_id=%s, %s", user_id, str(e))
|
||
raise ValueError(f"人员ID '{user_id}' 已存在")
|
||
except Exception as e:
|
||
conn.rollback()
|
||
logger.error("新增记录失败: %s", str(e))
|
||
raise
|
||
finally:
|
||
conn.close()
|
||
|
||
|
||
def update_record(user_id, user_name=None, sign_image=None, is_signature=None):
|
||
if _active_db == 'mysql':
|
||
return _mysql_update_record(user_id, user_name, sign_image, is_signature)
|
||
return _sqlite_update_record(user_id, user_name, sign_image, is_signature)
|
||
|
||
|
||
def _mysql_update_record(user_id, user_name=None, sign_image=None, is_signature=None):
|
||
conn = _get_mysql_connection()
|
||
if not conn:
|
||
return _sqlite_update_record(user_id, user_name, sign_image, is_signature)
|
||
try:
|
||
sets = []
|
||
params = []
|
||
if user_name is not None:
|
||
sets.append("user_name = %s")
|
||
params.append(user_name)
|
||
if sign_image is not None:
|
||
sets.append("sign_image = %s")
|
||
params.append(sign_image)
|
||
if is_signature is not None:
|
||
sets.append("is_signature = %s")
|
||
params.append(is_signature)
|
||
if not sets:
|
||
return True
|
||
params.append(user_id)
|
||
with conn.cursor() as cursor:
|
||
cursor.execute(f"UPDATE user_sign SET {', '.join(sets)} WHERE user_id = %s", params)
|
||
conn.commit()
|
||
logger.info("更新记录(MySQL): user_id=%s", user_id)
|
||
return True
|
||
except Exception as e:
|
||
conn.rollback()
|
||
logger.error("MySQL 更新记录失败: %s", str(e))
|
||
raise
|
||
|
||
|
||
def _sqlite_update_record(user_id, user_name=None, sign_image=None, is_signature=None):
|
||
conn = _get_sqlite_connection()
|
||
try:
|
||
sets = []
|
||
params = []
|
||
if user_name is not None:
|
||
sets.append("user_name = ?")
|
||
params.append(user_name)
|
||
if sign_image is not None:
|
||
sets.append("sign_image = ?")
|
||
params.append(sign_image)
|
||
if is_signature is not None:
|
||
sets.append("is_signature = ?")
|
||
params.append(is_signature)
|
||
if not sets:
|
||
return True
|
||
sets.append("update_time = datetime('now','localtime')")
|
||
params.append(user_id)
|
||
conn.execute(f"UPDATE user_sign SET {', '.join(sets)} WHERE user_id = ?", params)
|
||
conn.commit()
|
||
logger.info("更新记录(SQLite): user_id=%s", user_id)
|
||
return True
|
||
except Exception as e:
|
||
conn.rollback()
|
||
logger.error("更新记录失败: %s", str(e))
|
||
raise
|
||
finally:
|
||
conn.close()
|
||
|
||
|
||
def delete_record(user_id):
|
||
if _active_db == 'mysql':
|
||
return _mysql_delete_record(user_id)
|
||
return _sqlite_delete_record(user_id)
|
||
|
||
|
||
def _mysql_delete_record(user_id):
|
||
conn = _get_mysql_connection()
|
||
if not conn:
|
||
return _sqlite_delete_record(user_id)
|
||
try:
|
||
with conn.cursor() as cursor:
|
||
cursor.execute("DELETE FROM user_sign WHERE user_id = %s", (user_id,))
|
||
conn.commit()
|
||
logger.info("删除记录(MySQL): user_id=%s", user_id)
|
||
return True
|
||
except Exception as e:
|
||
conn.rollback()
|
||
logger.error("MySQL 删除记录失败: %s", str(e))
|
||
raise
|
||
|
||
|
||
def _sqlite_delete_record(user_id):
|
||
conn = _get_sqlite_connection()
|
||
try:
|
||
conn.execute("DELETE FROM user_sign WHERE user_id = ?", (user_id,))
|
||
conn.commit()
|
||
logger.info("删除记录(SQLite): user_id=%s", user_id)
|
||
return True
|
||
except Exception as e:
|
||
conn.rollback()
|
||
logger.error("删除记录失败: %s", str(e))
|
||
raise
|
||
finally:
|
||
conn.close()
|
||
|
||
|
||
def get_image_by_user_id(user_id):
|
||
if _active_db == 'mysql':
|
||
return _mysql_get_image_by_user_id(user_id)
|
||
return _sqlite_get_image_by_user_id(user_id)
|
||
|
||
|
||
def _mysql_get_image_by_user_id(user_id):
|
||
conn = _get_mysql_connection()
|
||
if not conn:
|
||
return _sqlite_get_image_by_user_id(user_id)
|
||
try:
|
||
with conn.cursor() as cursor:
|
||
cursor.execute("SELECT sign_image, is_signature FROM user_sign WHERE user_id = %s", (user_id,))
|
||
row = cursor.fetchone()
|
||
return row if row else None
|
||
except Exception as e:
|
||
logger.error("MySQL 查询图片失败: %s", str(e))
|
||
return _sqlite_get_image_by_user_id(user_id)
|
||
|
||
|
||
def _sqlite_get_image_by_user_id(user_id):
|
||
conn = _get_sqlite_connection()
|
||
try:
|
||
cursor = conn.execute("SELECT sign_image, is_signature FROM user_sign WHERE user_id = ?", (user_id,))
|
||
row = cursor.fetchone()
|
||
return dict(row) if row else None
|
||
finally:
|
||
conn.close()
|
||
|
||
|
||
def get_image_by_user_name(user_name):
|
||
if _active_db == 'mysql':
|
||
return _mysql_get_image_by_user_name(user_name)
|
||
return _sqlite_get_image_by_user_name(user_name)
|
||
|
||
|
||
def _mysql_get_image_by_user_name(user_name):
|
||
conn = _get_mysql_connection()
|
||
if not conn:
|
||
return _sqlite_get_image_by_user_name(user_name)
|
||
try:
|
||
with conn.cursor() as cursor:
|
||
cursor.execute("SELECT sign_image, is_signature, user_id, user_name FROM user_sign WHERE user_name = %s", (user_name,))
|
||
return cursor.fetchall()
|
||
except Exception as e:
|
||
logger.error("MySQL 查询图片失败: %s", str(e))
|
||
return _sqlite_get_image_by_user_name(user_name)
|
||
|
||
|
||
def _sqlite_get_image_by_user_name(user_name):
|
||
conn = _get_sqlite_connection()
|
||
try:
|
||
cursor = conn.execute("SELECT sign_image, is_signature, user_id, user_name FROM user_sign WHERE user_name = ?", (user_name,))
|
||
return [dict(r) for r in cursor.fetchall()]
|
||
finally:
|
||
conn.close()
|
||
|
||
|
||
def get_active_db_type():
|
||
return _active_db or 'unknown' |