"""数据库连接配置:支持 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()