第一次提交
This commit is contained in:
@@ -0,0 +1,111 @@
|
||||
"""AI 任务:判分、讲解、知识点总结。统一记录到 ai_generations。"""
|
||||
import json
|
||||
import logging
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models.practice import AiGeneration
|
||||
from app.models.question import Question
|
||||
from app.models.user import UserSettings
|
||||
from app.services.llm.client import LLMClient
|
||||
from app.services.llm.prompts import (
|
||||
EXPLAIN_SYSTEM,
|
||||
EXPLAIN_USER_TEMPLATE,
|
||||
JUDGE_SYSTEM,
|
||||
JUDGE_USER_TEMPLATE,
|
||||
TAGS_SYSTEM,
|
||||
TAGS_USER_TEMPLATE,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def load_llm(db: AsyncSession, user_id: str) -> LLMClient:
|
||||
"""加载用户的 LLMClient(未配置会抛 LLMConfigError)。"""
|
||||
row = await db.scalar(select(UserSettings).where(UserSettings.user_id == user_id))
|
||||
return LLMClient(row)
|
||||
|
||||
|
||||
def _parse_json(raw: str) -> dict:
|
||||
text = raw.strip()
|
||||
if text.startswith("```"):
|
||||
text = text.split("```", 2)[1]
|
||||
if text.startswith("json"):
|
||||
text = text[4:]
|
||||
return json.loads(text.strip())
|
||||
|
||||
|
||||
async def _record(
|
||||
db: AsyncSession, user_id: str, question_id: str | None, task: str, result, content: str
|
||||
):
|
||||
db.add(
|
||||
AiGeneration(
|
||||
user_id=user_id,
|
||||
question_id=question_id,
|
||||
task=task,
|
||||
model=result.model,
|
||||
request_tokens=result.prompt_tokens,
|
||||
response_tokens=result.completion_tokens,
|
||||
content_markdown=content,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def judge_answer(
|
||||
db: AsyncSession, llm: LLMClient, user_id: str, q: Question, user_answer: list[str]
|
||||
) -> dict:
|
||||
"""主观题判分。返回 {is_correct, feedback_markdown, key_points_missed}。"""
|
||||
answer = "、".join(q.correct_answer) if q.correct_answer else "(未提供标准答案)"
|
||||
result = await llm.chat_text(
|
||||
JUDGE_SYSTEM,
|
||||
JUDGE_USER_TEMPLATE.format(
|
||||
stem=q.stem_markdown, answer=answer, user_answer="\n".join(user_answer)
|
||||
),
|
||||
json_mode=True,
|
||||
)
|
||||
await _record(db, user_id, q.id, "judge", result, result.content)
|
||||
try:
|
||||
data = _parse_json(result.content)
|
||||
return {
|
||||
"is_correct": bool(data.get("is_correct")),
|
||||
"feedback_markdown": data.get("feedback_markdown", ""),
|
||||
"key_points_missed": data.get("key_points_missed", []),
|
||||
}
|
||||
except (json.JSONDecodeError, KeyError) as e:
|
||||
logger.warning("judge 解析失败:%s", e)
|
||||
return {
|
||||
"is_correct": False,
|
||||
"feedback_markdown": result.content,
|
||||
"key_points_missed": [],
|
||||
}
|
||||
|
||||
|
||||
async def explain_question(
|
||||
db: AsyncSession, llm: LLMClient, user_id: str, q: Question
|
||||
) -> str:
|
||||
answer_hint = ""
|
||||
if q.correct_answer:
|
||||
answer_hint = f"参考答案:{'、'.join(q.correct_answer)}\n\n"
|
||||
result = await llm.chat_text(
|
||||
EXPLAIN_SYSTEM,
|
||||
EXPLAIN_USER_TEMPLATE.format(stem=q.stem_markdown, answer_hint=answer_hint),
|
||||
)
|
||||
await _record(db, user_id, q.id, "explain", result, result.content)
|
||||
return result.content
|
||||
|
||||
|
||||
async def summarize_tags(
|
||||
db: AsyncSession, llm: LLMClient, user_id: str, q: Question
|
||||
) -> list[str]:
|
||||
result = await llm.chat_text(
|
||||
TAGS_SYSTEM, TAGS_USER_TEMPLATE.format(stem=q.stem_markdown), json_mode=True
|
||||
)
|
||||
await _record(db, user_id, q.id, "summarize", result, result.content)
|
||||
try:
|
||||
data = _parse_json(result.content)
|
||||
tags = data.get("tags", [])
|
||||
return [str(t).strip() for t in tags if str(t).strip()][:4]
|
||||
except (json.JSONDecodeError, KeyError) as e:
|
||||
logger.warning("tags 解析失败:%s", e)
|
||||
return []
|
||||
Reference in New Issue
Block a user