nex_math/backend/database.py

186 lines
6.5 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 数据库连接与会话管理。"""
from __future__ import annotations
import os
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"
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 自动加列)。"""
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}"
)
)
# 补列之后再搬数据:迁移要写 knowledge_resources.course_id
_split_courses()
def _split_courses() -> None:
"""把旧 textbooks 表里的“在线课程”行搬到 courses 表,并去掉 kind 列。
教材与课程从此各用一张表:教材带章节目录与电子书,课程只有外链。
幂等: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()