Files
2026-07-28 15:04:17 +08:00

311 lines
11 KiB
Python

"""数据库连接配置:支持 SQLite / MySQL / PostgreSQL。"""
from __future__ import annotations
import os
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from urllib.parse import quote_plus
from dotenv import load_dotenv
from sqlalchemy import event, text
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlalchemy.engine import make_url
from sqlalchemy.pool import AsyncAdaptedQueuePool, StaticPool
BACKEND_DIR = Path(__file__).resolve().parent.parent
PROJECT_ROOT = BACKEND_DIR.parent
ENV_FILE = PROJECT_ROOT / ".env"
load_dotenv(ENV_FILE)
SUPPORTED_DB_TYPES = ("sqlite", "mysql", "postgresql")
PASSWORD_PLACEHOLDER = "******"
DEFAULT_SQLITE_PATH = BACKEND_DIR / "kefu.db"
def _env(name: str, default: str = "") -> str:
return (os.getenv(name) or default).strip()
def _env_int(name: str, default: int) -> int:
raw = _env(name)
if not raw:
return default
try:
return int(raw)
except ValueError:
return default
def _bounded_env_int(name: str, default: int, minimum: int, maximum: int) -> int:
return max(minimum, min(maximum, _env_int(name, default)))
def normalize_db_type(value: str | None) -> str:
raw = (value or "sqlite").strip().lower()
if raw in ("postgres", "pgsql"):
return "postgresql"
if raw in SUPPORTED_DB_TYPES:
return raw
return "sqlite"
@dataclass
class DatabaseConfig:
db_type: str = "sqlite"
db_host: str = "127.0.0.1"
db_port: int = 0
db_user: str = ""
db_password: str = ""
db_name: str = "kefu"
db_path: str = ""
database_url: str = ""
def normalized_type(self) -> str:
return normalize_db_type(self.db_type)
def resolved_port(self) -> int:
if self.db_port:
return self.db_port
if self.normalized_type() == "mysql":
return 3306
if self.normalized_type() == "postgresql":
return 5432
return 0
def sqlite_path(self) -> Path:
if self.db_path:
path = Path(self.db_path)
if not path.is_absolute():
path = BACKEND_DIR / path
return path
if self.database_url and self.database_url.startswith("sqlite"):
# sqlite+aiosqlite:///path
raw = self.database_url.split("///", 1)[-1]
return Path(raw)
return DEFAULT_SQLITE_PATH
def read_database_config() -> DatabaseConfig:
explicit_url = _env("KEFU_DATABASE_URL")
db_type = normalize_db_type(_env("KEFU_DB_TYPE", "sqlite"))
db_path = _env("KEFU_DB_PATH")
if explicit_url and not _env("KEFU_DB_TYPE"):
lowered = explicit_url.lower()
if lowered.startswith("mysql"):
db_type = "mysql"
elif lowered.startswith("postgresql") or lowered.startswith("postgres"):
db_type = "postgresql"
elif lowered.startswith("sqlite"):
db_type = "sqlite"
return DatabaseConfig(
db_type=db_type,
db_host=_env("KEFU_DB_HOST", "127.0.0.1"),
db_port=_env_int("KEFU_DB_PORT", 0),
db_user=_env("KEFU_DB_USER"),
db_password=_env("KEFU_DB_PASSWORD"),
db_name=_env("KEFU_DB_NAME", "kefu"),
db_path=db_path,
database_url=explicit_url,
)
def build_database_url(config: DatabaseConfig | None = None) -> str:
cfg = config or read_database_config()
if cfg.database_url:
return cfg.database_url
db_type = cfg.normalized_type()
if db_type == "sqlite":
path = cfg.sqlite_path()
path.parent.mkdir(parents=True, exist_ok=True)
return f"sqlite+aiosqlite:///{path.as_posix()}"
user = quote_plus(cfg.db_user or "")
password = quote_plus(cfg.db_password or "")
host = cfg.db_host or "127.0.0.1"
port = cfg.resolved_port()
db_name = cfg.db_name or "kefu"
if db_type == "mysql":
auth = f"{user}:{password}@" if user else ""
return (
f"mysql+asyncmy://{auth}{host}:{port}/{db_name}"
"?charset=utf8mb4"
)
auth = f"{user}:{password}@" if user else ""
return f"postgresql+asyncpg://{auth}{host}:{port}/{db_name}"
def engine_kwargs_for_url(url: str) -> dict[str, Any]:
kwargs: dict[str, Any] = {"echo": False}
if url.startswith("sqlite"):
# SQLAlchemy 2.0.30 defaults file-backed aiosqlite to NullPool. At
# hundreds of hosted accounts that creates and tears down an aiosqlite
# worker thread for every short query. Reuse a small bounded pool;
# SQLite still serializes writers, so a large pool only adds lock
# contention and does not improve throughput.
busy_timeout_ms = _bounded_env_int(
"KEFU_SQLITE_BUSY_TIMEOUT_MS", 30_000, 1_000, 120_000
)
kwargs["connect_args"] = {"timeout": busy_timeout_ms / 1000.0}
sqlite_database = make_url(url).database
is_memory_database = (
not sqlite_database
or sqlite_database == ":memory:"
or "mode=memory" in url.lower()
)
if is_memory_database:
# Every connection to an in-memory SQLite URL otherwise receives a
# different database. StaticPool preserves the single shared
# connection expected by tests and utility callers.
kwargs["poolclass"] = StaticPool
else:
kwargs["poolclass"] = AsyncAdaptedQueuePool
kwargs["pool_size"] = _bounded_env_int(
"KEFU_SQLITE_POOL_SIZE", 5, 1, 10
)
kwargs["max_overflow"] = 0
kwargs["pool_timeout"] = max(5.0, busy_timeout_ms / 1000.0)
else:
kwargs["pool_pre_ping"] = True
kwargs["pool_recycle"] = 3600
# 默认连接池仅 pool_size=5 + max_overflow=10。多账号托管 + 前端并发请求时
# 容易耗尽连接导致请求阻塞/失败,这里放大连接池(可用环境变量覆盖)。
try:
_pool = int(os.getenv("KEFU_DB_POOL_SIZE", "20") or 20)
except ValueError:
_pool = 20
try:
_overflow = int(os.getenv("KEFU_DB_MAX_OVERFLOW", "40") or 40)
except ValueError:
_overflow = 40
kwargs["pool_size"] = max(5, _pool)
kwargs["max_overflow"] = max(0, _overflow)
kwargs["pool_timeout"] = 30
return kwargs
def _configure_sqlite_connection(dbapi_connection, _connection_record) -> None:
"""Apply process-wide SQLite settings to every pooled connection.
WAL lets readers continue while the single writer commits. NORMAL avoids
a full disk sync for every small log/status transaction while retaining
WAL crash consistency. busy_timeout turns transient writer contention
into bounded waiting instead of immediate ``database is locked`` errors.
"""
busy_timeout_ms = _bounded_env_int(
"KEFU_SQLITE_BUSY_TIMEOUT_MS", 30_000, 1_000, 120_000
)
cursor = dbapi_connection.cursor()
try:
cursor.execute(f"PRAGMA busy_timeout={busy_timeout_ms}")
cursor.execute("PRAGMA journal_mode=WAL")
cursor.execute("PRAGMA synchronous=NORMAL")
cursor.execute("PRAGMA foreign_keys=ON")
finally:
cursor.close()
def create_database_engine(config: DatabaseConfig | None = None) -> AsyncEngine:
url = build_database_url(config)
engine = create_async_engine(url, **engine_kwargs_for_url(url))
if url.startswith("sqlite"):
event.listen(engine.sync_engine, "connect", _configure_sqlite_connection)
return engine
def mask_database_url(url: str) -> str:
if "@" not in url or "://" not in url:
return url
scheme, rest = url.split("://", 1)
if "@" not in rest:
return url
creds, host_part = rest.rsplit("@", 1)
if ":" in creds:
user = creds.split(":", 1)[0]
return f"{scheme}://{user}:{PASSWORD_PLACEHOLDER}@{host_part}"
return f"{scheme}://{PASSWORD_PLACEHOLDER}@{host_part}"
def database_config_to_response(cfg: DatabaseConfig | None = None) -> dict[str, Any]:
cfg = cfg or read_database_config()
url = build_database_url(cfg)
return {
"db_type": cfg.normalized_type(),
"db_host": cfg.db_host,
"db_port": cfg.resolved_port(),
"db_user": cfg.db_user,
"db_name": cfg.db_name,
"db_path": str(cfg.sqlite_path()) if cfg.normalized_type() == "sqlite" else "",
"db_password": PASSWORD_PLACEHOLDER if cfg.db_password else "",
"db_password_configured": bool(cfg.db_password),
"database_url_display": mask_database_url(url),
"supported_types": list(SUPPORTED_DB_TYPES),
}
def persist_database_config(payload: dict[str, Any]) -> None:
"""将数据库配置写入项目根 .env 文件。"""
from dotenv import set_key
if not ENV_FILE.exists():
ENV_FILE.write_text("", encoding="utf-8")
db_type = normalize_db_type(payload.get("db_type"))
env_map = {
"KEFU_DB_TYPE": db_type,
"KEFU_DB_HOST": (payload.get("db_host") or "127.0.0.1").strip(),
"KEFU_DB_PORT": str(payload.get("db_port") or ""),
"KEFU_DB_USER": (payload.get("db_user") or "").strip(),
"KEFU_DB_NAME": (payload.get("db_name") or "kefu").strip(),
"KEFU_DB_PATH": (payload.get("db_path") or "").strip(),
}
password = payload.get("db_password")
if password and str(password).strip() not in ("", PASSWORD_PLACEHOLDER):
env_map["KEFU_DB_PASSWORD"] = str(password).strip()
# 清除完整 URL,避免与分项配置冲突
set_key(str(ENV_FILE), "KEFU_DATABASE_URL", "")
for key, value in env_map.items():
set_key(str(ENV_FILE), key, value or "")
# 让当前进程也能读到新值(重启后才会重建 engine)
for key, value in env_map.items():
os.environ[key] = value or ""
os.environ.pop("KEFU_DATABASE_URL", None)
async def test_database_connection(payload: dict[str, Any]) -> None:
cfg = DatabaseConfig(
db_type=normalize_db_type(payload.get("db_type")),
db_host=(payload.get("db_host") or "127.0.0.1").strip(),
db_port=int(payload.get("db_port") or 0),
db_user=(payload.get("db_user") or "").strip(),
db_password=(payload.get("db_password") or "").strip(),
db_name=(payload.get("db_name") or "kefu").strip(),
db_path=(payload.get("db_path") or "").strip(),
)
current = read_database_config()
if cfg.db_password in ("", PASSWORD_PLACEHOLDER):
cfg.db_password = current.db_password
url = build_database_url(cfg)
engine = create_async_engine(url, **engine_kwargs_for_url(url))
try:
async with engine.connect() as conn:
await conn.execute(text("SELECT 1"))
except ModuleNotFoundError as exc:
driver = "asyncmy" if cfg.normalized_type() == "mysql" else "asyncpg"
raise RuntimeError(
f"缺少数据库驱动,请执行: pip install {driver}"
) from exc
finally:
await engine.dispose()