186 lines
6.5 KiB
Python
186 lines
6.5 KiB
Python
"""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()
|