nex_math/backend/services/question_engine.py

90 lines
2.9 KiB
Python

"""题目引擎:按章节动态组卷(掌握度缺口 + 错题 + 轮换扰动)。"""
from __future__ import annotations
import json
from sqlalchemy.orm import Session
from sqlalchemy import func
from models import (
AttemptItem,
Chapter,
ErrorEntry,
Knowledge,
Question,
UserKnowledge,
)
from services.knowledge_service import question_knowledge_names
def load_questions(db: Session, ids: list[int]) -> list[Question]:
rows = db.query(Question).filter(Question.id.in_(ids)).all()
by_id = {row.id: row for row in rows}
return [by_id[qid] for qid in ids if qid in by_id]
def chapter_question_ids(
db: Session,
user_id: int,
chapter_id: int,
limit: int = 5,
variant: int = 0,
) -> list[int]:
"""按章节动态组卷:只取该章题库,不足则返回实际题量,不从其他章节补题。
权重 = 掌握度缺口 + 错题加成 + 上次已做惩罚 + 轮换扰动,避免每轮完全相同。
"""
chapter = db.get(Chapter, chapter_id)
if chapter is None:
return []
knowledge_map = {
knowledge.name: user_knowledge.mastery
for user_knowledge, knowledge in (
db.query(UserKnowledge, Knowledge)
.join(Knowledge, Knowledge.id == UserKnowledge.knowledge_id)
.filter(UserKnowledge.user_id == user_id)
.all()
)
}
error_counts: dict[str, int] = {}
for primary, raw_names in (
db.query(ErrorEntry.knowledge_name, ErrorEntry.knowledge_names)
.filter(ErrorEntry.user_id == user_id)
.all()
):
try:
names = json.loads(raw_names or "[]")
if not isinstance(names, list) or not names:
names = [primary]
except (TypeError, ValueError):
names = [primary]
for name in {str(name) for name in names}:
error_counts[name] = error_counts.get(name, 0) + 1
attempt_counts = dict(
db.query(AttemptItem.question_id, func.count(AttemptItem.id))
.filter(AttemptItem.user_id == user_id, AttemptItem.question_id.isnot(None))
.group_by(AttemptItem.question_id)
.all()
)
chapter_questions = (
db.query(Question).filter(Question.chapter_id == chapter_id).all()
)
ranked = []
for question in chapter_questions:
tags = question_knowledge_names(db, question)
knowledge_deficit = max(
(110 - knowledge_map.get(tag, 60) for tag in tags),
default=50,
)
errors = sum(error_counts.get(tag, 0) for tag in tags)
repeated = attempt_counts.get(question.id, 0) * 22
jitter = (question.id * 7 + variant * 13) % 19
weight = knowledge_deficit + errors * 12 - repeated + jitter
ranked.append((weight, question))
ranked.sort(key=lambda pair: (-pair[0], pair[1].id))
return [q.id for _, q in ranked[:limit]]