Files
dy/backend/models/db_config.py
T
2026-07-23 17:56:25 +08:00

253 lines
8.3 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 text
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
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 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 not url.startswith("sqlite"):
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 create_database_engine(config: DatabaseConfig | None = None) -> AsyncEngine:
url = build_database_url(config)
return create_async_engine(url, **engine_kwargs_for_url(url))
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()