221 lines
7.9 KiB
Python
221 lines
7.9 KiB
Python
"""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()
|