311 lines
11 KiB
Python
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()
|