更新
This commit is contained in:
@@ -0,0 +1,252 @@
|
||||
"""数据库连接配置:支持 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()
|
||||
Reference in New Issue
Block a user