nex_math/backend/services/question_generator.py

321 lines
12 KiB
Python
Raw Permalink 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.

"""按教材章节调用大模型生成客观题,并做结构化校验后入库。"""
from __future__ import annotations
from models import Chapter, Knowledge, LlmSetting, Question, Textbook
from services import llm_client
SYSTEM_PROMPT = (
"你是一位经验丰富的高中数学教师和命题人,擅长编写规范、严谨、"
"区分度良好的单项选择题。所有题面必须自洽且答案唯一。"
)
def build_user_prompt(
*,
textbook_name: str,
chapter_name: str,
chapter_summary: str,
count: int,
difficulty: int,
knowledge_name: str,
instructions: str,
available_knowledge: list[str],
need_figure: bool = False,
) -> str:
difficulty_text = {1: "基础识记与直接计算", 2: "概念理解与中等计算", 3: "综合推理与易错辨析"}[
difficulty
]
knowledge_hint = (
f"知识点限定为“{knowledge_name}”。"
if knowledge_name
else "知识点由你依据章节内容选择,尽量使用下面给出的知识点名称;"
"若下面列表没有合适名称,可自拟简洁的中文知识点名。"
)
extra = f"额外要求:{instructions}" if instructions else ""
figure_requirement = (
(
"配图要求:每题必须在 image_svg 字段中给出与题意对应的原创 SVG 图像"
"(只画该题需要的函数图像/几何示意图,含坐标轴与必要标注),"
"SVG 用单行字符串输出(内部换行用 \\n),不要用 Markdown 图片或外链;"
"题干中不要再重复插入任何图片标记。"
)
if need_figure
else "配图要求:默认不配图,image_svg 字段填空字符串 \"\"。"
)
figure_hint = f"\n{figure_requirement}\n" if need_figure else f"\n{figure_requirement}"
return f"""请围绕以下教材章节命制 {count} 道高中数学单选题。
教材:{textbook_name}
章节:{chapter_name}
章节内容提示:{chapter_summary or "请结合章节标题自行判断"}
难度:{difficulty_text}
{knowledge_hint}
本系统已有知识点:{", ".join(available_knowledge) or "(暂无)"}
{extra}
{figure_hint}
硬性要求:
1. 每题必须恰好 4 个选项,且只有一个正确选项;
2. 不使用“以上都对”“以上都错”这类选项;
3. 涉及数学符号、公式时必须用 LaTeX:行内公式写 $...$,独立公式写 $$...$$,
不要使用 x²、x³、√、¼ 这类 Unicode 记号,例如 $x^2$、$\\sqrt{{x+1}}$、
$\\frac{{1}}{{x^2-9}}$、$f(-x)=-f(x)$;
4. 图像只通过 JSON 的 image_svg 字段提供:配图要求开启时,每题必须有与题目
语义一致的原创 SVG 图(函数图或几何示意图,含坐标轴/关键标注),
题干中不得再写任何 Markdown 图片;未开启配图时不要输出 image_svg;
5. explanation 用 1-2 句中文解释关键思路,便于学生理解,含公式时同样用 LaTeX;
6. 不要输出多余文字,只输出 JSON:
{{"questions":[
{{"stem":"题干","options":["A选项","B选项","C选项","D选项"],
"correct_index":0,"explanation":"解析","knowledge_name":"知识点","difficulty":{difficulty},
"image_svg":"<svg xmlns=\\"http://www.w3.org/2000/svg\\" ...></svg> 或空字符串"}}
]}}"""
def _extract_json(content: str) -> dict:
from services.json_utils import extract_json_lax
return extract_json_lax(content)
def _normalize_item(raw: dict, fallback_knowledge: str, difficulty: int) -> dict:
stem = str(raw.get("stem", "")).strip()
options = raw.get("options") or []
options = [str(option).strip() for option in options]
explanation = str(raw.get("explanation", "")).strip()
knowledge_name = str(raw.get("knowledge_name", "")).strip() or fallback_knowledge
image_svg = str(raw.get("image_svg") or "").strip()
# 部分兼容网关会对 JSON 内的 SVG 引号/换行做双重转义
image_svg = image_svg.replace("\\\\", "\\")
image_svg = (
image_svg.replace('\\"', '"')
.replace("\\n", "\n")
.replace("\\t", "\t")
.replace("\\r", "\r")
)
if not stem:
raise ValueError("模型返回的题干为空")
if len(options) != 4 or any(not option for option in options):
raise ValueError(f"选项必须是 4 个非空选项,实际为 {len(options)} 个")
if len(set(options)) != 4:
raise ValueError("选项存在重复")
try:
correct_index = int(raw.get("correct_index", -1))
except (TypeError, ValueError) as exc:
raise ValueError("correct_index 必须是数字") from exc
if correct_index < 0 or correct_index >= 4:
raise ValueError("correct_index 超出选项范围")
return {
"stem": stem,
"options": options,
"correct_index": correct_index,
"explanation": explanation or "请依据本章定义与性质重新判断。",
"knowledge_name": knowledge_name,
"difficulty": max(1, min(3, difficulty)),
"image_svg": image_svg,
}
def generate_questions(
*,
db,
setting: LlmSetting,
chapter: Chapter,
textbook: Textbook,
count: int,
difficulty: int,
knowledge_name: str,
instructions: str,
need_figure: bool = False,
) -> list[dict]:
if not setting or not setting.api_key:
raise ValueError("尚未配置大模型 API Key,请先到“模型配置”页面保存并测试")
# 带图模式先按无图生成文本题目,再逐题补 SVG,避免长 JSON 超时/截断
if need_figure:
text_items = generate_questions(
db=db,
setting=setting,
chapter=chapter,
textbook=textbook,
count=count,
difficulty=difficulty,
knowledge_name=knowledge_name,
instructions=instructions,
need_figure=False,
)
for item in text_items:
item["image_svg"] = _request_figure_svg(setting, item["stem"])
return text_items
knowledge_names = {name for (name,) in db.query(Knowledge.name).all()}
knowledge_names.update(
name for (name,) in db.query(Question.knowledge_name).all()
)
available_knowledge = sorted(knowledge_names)
prompt = build_user_prompt(
textbook_name=textbook.name,
chapter_name=chapter.name,
chapter_summary=chapter.summary,
count=count,
difficulty=difficulty,
knowledge_name=knowledge_name,
instructions=instructions,
available_knowledge=available_knowledge,
need_figure=need_figure,
)
payload = {}
last_error = ""
for attempt in range(3):
try:
content = llm_client.chat_completion(
base_url=setting.base_url,
api_key=setting.api_key,
model=setting.model,
temperature=setting.temperature,
max_tokens=_effective_max_tokens(setting, need_figure),
timeout=200.0,
messages=[
{
"role": "system",
"content": (
SYSTEM_PROMPT
+ " JSON 必须完整且闭合:字符串中的双引号写 \\\","
"反斜杠写 \\\\,不要输出多余字符。"
),
},
{"role": "user", "content": prompt},
],
)
except ValueError as exc:
last_error = str(exc)
continue
try:
payload = _extract_json(content)
break
except ValueError as exc:
last_error = str(exc)
else:
raise ValueError(
f"模型连续 3 次未返回合法 JSON,最后一次:{last_error}"
)
raw_items = payload.get("questions")
if not isinstance(raw_items, list) or not raw_items:
raise ValueError("模型返回内容中没有 questions 列表")
generated: list[dict] = []
for raw in raw_items[:count]:
item = _normalize_item(raw, knowledge_name or "章节综合", difficulty)
if need_figure and (
"<svg" not in item["image_svg"].lower()
or "</svg>" not in item["image_svg"].lower()
):
# 补图:只让模型针对该题生成一张 SVG,缩短输出降低截断概率
for _ in range(3):
svg_content = llm_client.chat_completion(
base_url=setting.base_url,
api_key=setting.api_key,
model=setting.model,
temperature=0.1,
max_tokens=max(
int(setting.max_tokens)
if setting.max_tokens
else 3000,
1200,
),
timeout=120.0,
messages=[
{
"role": "system",
"content": "只输出 JSON:{\"image_svg\":\"<svg ...>\"}",
},
{
"role": "user",
"content": (
f"请为下列题干生成与题意一致的原创 SVG 示意图:\n{item['stem']}"
),
},
],
)
try:
svg_payload = _extract_json(svg_content)
svg_value = str(svg_payload.get("image_svg") or "")
svg_value = svg_value.replace("\\\\", "\\")
svg_value = (
svg_value.replace('\\"', '"')
.replace("\\n", "\n")
.replace("\\t", "\t")
.replace("\\r", "\r")
)
if (
"<svg" in svg_value.lower()
and "</svg>" in svg_value.lower()
):
item["image_svg"] = svg_value
break
except ValueError:
continue
if (
"<svg" not in item["image_svg"].lower()
or "</svg>" not in item["image_svg"].lower()
):
raise ValueError(
"模型未能生成本题有效配图,请尝试减少题量或更换模型通道"
)
generated.append(item)
return generated
def _request_figure_svg(setting: LlmSetting, stem: str) -> str:
"""针对单一题干请求一张尽量简单的 SVG 配图。"""
max_tokens = int(setting.max_tokens) if setting.max_tokens else 1600
for _ in range(3):
try:
content = llm_client.chat_completion(
base_url=setting.base_url,
api_key=setting.api_key,
model=setting.model,
temperature=0.1,
max_tokens=max_tokens,
timeout=300.0,
messages=[
{
"role": "system",
"content": (
"你是数学示意图生成器。只输出 JSON:"
'{"image_svg":"<svg ...></svg>"}。SVG 必须非常简洁,'
"建议坐标轴+图形+少量标注,控制在 100 行内。"
),
},
{
"role": "user",
"content": f"请为以下题干生成与题意一致的示意图:\n{stem}",
},
],
)
except ValueError:
continue
try:
payload = _extract_json(content)
svg = str(payload.get("image_svg") or "")
svg = svg.replace("\\\\", "\\")
svg = (
svg.replace('\\"', '"')
.replace("\\n", "\n")
.replace("\\t", "\t")
.replace("\\r", "\r")
)
if "<svg" in svg.lower() and "</svg>" in svg.lower():
return svg
except ValueError:
continue
raise ValueError("模型未能为本题生成有效 SVG,请更换模型通道或关闭配图")
def _effective_max_tokens(setting: LlmSetting, need_figure: bool) -> int:
if getattr(setting, "max_tokens", None):
return int(setting.max_tokens)
return 3000 if need_figure else 1600