nex_math/backend/database.py

221 lines
7.9 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

"""SQLite 数据库连接与会话管理。
配置来源(优先级从高到低):进程环境变量 → 仓库根 ``.env`` →
``backend/.env`` → 代码里的默认值。``.env`` 在本模块导入时即被读入,
因此 ``DATABASE_URL`` / ``CORS_ORIGINS`` / ``JWT_SECRET`` 等都在此生效。
"""
from __future__ import annotations
import os
import sys
from pathlib import Path
from sqlalchemy import create_engine, event
from sqlalchemy import inspect, text
from sqlalchemy.orm import DeclarativeBase, sessionmaker
REPO_ROOT = Path(__file__).resolve().parent.parent
DEFAULT_DB_PATH = REPO_ROOT / "data" / "math.db"
BACKEND_DIR = Path(__file__).resolve().parent
def _load_env_files() -> None:
"""读取 ``.env``:仓库根优先,``backend/.env`` 只补前者没有的键。
已经存在的环境变量不会被覆盖,所以容器/脚本用 ``DATABASE_URL=…``
直接启动时仍然以环境变量为准。
"""
try:
from dotenv import load_dotenv
except ImportError:
# 不静默:否则 .env 会被“看起来没生效”地忽略掉
print(
"[config] 未安装 python-dotenv,.env 不会生效。"
"请执行:backend/.venv/bin/pip install -r backend/requirements.txt",
file=sys.stderr,
)
return
for path in (REPO_ROOT / ".env", BACKEND_DIR / ".env"):
if path.is_file():
load_dotenv(path, override=False)
_load_env_files()
class Base(DeclarativeBase):
pass
def _engine_url() -> str:
url = os.getenv("DATABASE_URL", "").strip()
if url:
return url
DEFAULT_DB_PATH.parent.mkdir(parents=True, exist_ok=True)
return f"sqlite:///{DEFAULT_DB_PATH}"
DATABASE_URL = _engine_url()
connect_args = {"check_same_thread": False} if DATABASE_URL.startswith("sqlite") else {}
engine = create_engine(DATABASE_URL, connect_args=connect_args)
SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
if DATABASE_URL.startswith("sqlite"):
@event.listens_for(engine, "connect")
def _enable_sqlite_foreign_keys(dbapi_connection, _connection_record):
cursor = dbapi_connection.cursor()
cursor.execute("PRAGMA foreign_keys=ON")
cursor.close()
def ensure_schema() -> None:
"""建表并为旧库补充新增列(SQLite 不会由 create_all 自动加列)。
只做**表与字段**层面的对齐,不写任何业务数据:后端每次启动都可以安全
重放。种子数据与历史数据迁移只在初始化时执行 ——
``cd backend && .venv/bin/python -m seed``。
"""
Base.metadata.create_all(bind=engine)
inspector = inspect(engine)
additions: dict[str, list[tuple[str, str]]] = {
"questions": [
("chapter_id", "INTEGER"),
("difficulty", "INTEGER NOT NULL DEFAULT 1"),
("is_generated", "BOOLEAN NOT NULL DEFAULT 0"),
("used_count", "INTEGER NOT NULL DEFAULT 0"),
],
"llm_settings": [
("name", "VARCHAR(128) NOT NULL DEFAULT ''"),
("is_default", "BOOLEAN NOT NULL DEFAULT 0"),
("created_at", "DATETIME"),
("max_tokens", "INTEGER"),
],
"attempt_items": [
("knowledge_names", "TEXT NOT NULL DEFAULT '[]'"),
],
"knowledge": [
("domain", "VARCHAR(16) NOT NULL DEFAULT '初等'"),
("category", "VARCHAR(64) NOT NULL DEFAULT ''"),
],
"error_entries": [
("attempt_id", "INTEGER"),
("question_id", "INTEGER"),
("stem", "TEXT NOT NULL DEFAULT ''"),
("options", "TEXT NOT NULL DEFAULT ''"),
("selected", "INTEGER"),
("correct_index", "INTEGER"),
("explanation", "TEXT NOT NULL DEFAULT ''"),
("knowledge_names", "TEXT NOT NULL DEFAULT '[]'"),
],
"attempt_sessions": [
("chapter_id", "INTEGER"),
],
"chapters": [
("ebook_page", "INTEGER NOT NULL DEFAULT 0"),
],
"textbooks": [
("ebook_file", "VARCHAR(256) NOT NULL DEFAULT ''"),
("ebook_name", "VARCHAR(256) NOT NULL DEFAULT ''"),
("ebook_format", "VARCHAR(8) NOT NULL DEFAULT ''"),
("ebook_size", "INTEGER NOT NULL DEFAULT 0"),
("ebook_pages", "INTEGER NOT NULL DEFAULT 0"),
("ebook_uploaded_at", "DATETIME"),
],
"knowledge_resources": [
("course_id", "INTEGER"),
],
"daily_completions": [
("chapter_id", "INTEGER"),
("rating", "VARCHAR(16) NOT NULL DEFAULT ''"),
],
}
for table_name, columns in additions.items():
if table_name not in inspector.get_table_names():
continue
existing = {
column["name"] for column in inspector.get_columns(table_name)
}
for name, definition in columns:
if name in existing:
continue
with engine.begin() as connection:
connection.execute(
text(
f"ALTER TABLE {table_name} "
f"ADD COLUMN {name} {definition}"
)
)
def migrate_legacy_courses() -> None:
"""把旧 textbooks 表里的“在线课程”行搬到 courses 表,并去掉 kind 列。
教材与课程从此各用一张表:教材带章节目录与电子书,课程只有外链。
属于数据迁移,只在 ``python -m seed`` 初始化时执行;调用前需先跑
``ensure_schema()``(要用到新补的 ``knowledge_resources.course_id``)。
幂等:kind 列不存在(已迁移过)时直接返回。
"""
inspector = inspect(engine)
tables = inspector.get_table_names()
if "textbooks" not in tables or "courses" not in tables:
return
columns = {column["name"] for column in inspector.get_columns("textbooks")}
if "kind" not in columns:
return
with engine.begin() as connection:
rows = connection.execute(
text(
"SELECT id, name, publisher, link, grade, description, position "
"FROM textbooks WHERE kind = 'course'"
)
).fetchall()
for row in rows:
exists = connection.execute(
text("SELECT id FROM courses WHERE name = :name"),
{"name": row[1]},
).first()
if exists is None:
connection.execute(
text(
"INSERT INTO courses "
"(name, provider, url, grade, description, position) "
"VALUES (:name, :provider, :url, :grade, :description, :position)"
),
{
"name": row[1],
"provider": row[2] or "",
"url": row[3] or "",
"grade": row[4] or "",
"description": row[5] or "",
"position": row[6] or 0,
},
)
exists = connection.execute(
text("SELECT id FROM courses WHERE name = :name"),
{"name": row[1]},
).first()
# 知识点上的视频资源改挂课程,删教材行时不会被级联删掉
connection.execute(
text(
"UPDATE knowledge_resources "
"SET course_id = :course_id, textbook_id = NULL "
"WHERE textbook_id = :textbook_id"
),
{"course_id": exists[0], "textbook_id": row[0]},
)
connection.execute(text("DELETE FROM textbooks WHERE kind = 'course'"))
connection.execute(text("ALTER TABLE textbooks DROP COLUMN kind"))
def get_db():
db = SessionLocal()
try:
yield db
finally:
db.close()