Files
Edu/services/ai/tests/test_services.py
SpecialX aac26c7c6f feat(ai): v2 新增 GenerateReport RPC + ReportService
第 9 个 RPC GenerateReport(学情报告生成):data-ana 学情数据 → LLM 生成 → 结构化提取

新增 ReportService 业务编排层 + GenerateReportRequest/GeneratedReport 模型

gRPC servicer + HTTP POST /v1/ai/generate/report(权限 ai:report:generate)

proto_gen 重新生成 + 测试覆盖(servicer/service/HTTP/模型/权限 共 26 用例)

402 测试通过,覆盖率 88.5%
2026-07-14 22:57:57 +08:00

331 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""服务层测试ChatService / QuestionService / ExpressionService / ReportService."""
import json
from src.ai.clients.data_ana_client import DataAnaClientMock
from src.ai.models.question import GenerateQuestionRequest
from src.ai.providers import ProviderFailoverChain
from src.ai.providers.circuit_breaker import CircuitBreaker
from src.ai.services.chat_service import ChatService
from src.ai.services.evaluation import QualityGate, RuleValidator
from src.ai.services.expression_service import ExpressionService
from src.ai.services.question_service import QuestionService
from src.ai.services.report_service import ReportService
from .conftest import MockProvider
def _make_chain(provider: MockProvider) -> ProviderFailoverChain:
return ProviderFailoverChain([provider], CircuitBreaker())
class TestChatService:
"""ChatService 测试."""
async def test_chat_success(self) -> None:
"""聊天成功."""
chain = _make_chain(MockProvider(response_content="你好"))
svc = ChatService(failover_chain=chain, default_model="gpt-4o")
data = await svc.chat(messages=[{"role": "user", "content": "hi"}])
assert data.content == "你好"
assert data.degraded is False
assert data.usage.total_tokens == 30
async def test_chat_degraded(self) -> None:
"""LLM 不可用时降级."""
chain = _make_chain(MockProvider(fail=True))
svc = ChatService(failover_chain=chain)
data = await svc.chat(messages=[{"role": "user", "content": "hi"}])
assert data.degraded is True
assert "degraded" in data.content
async def test_stream_chat(self) -> None:
"""流式聊天."""
chain = _make_chain(MockProvider(stream_chunks=["hello", " world"]))
svc = ChatService(failover_chain=chain)
chunks = []
async for chunk in svc.stream_chat(messages=[{"role": "user", "content": "hi"}]):
chunks.append(chunk)
assert len(chunks) == 2
assert chunks[-1].done is True
async def test_stream_chat_degraded(self) -> None:
"""流式降级."""
chain = _make_chain(MockProvider(fail=True))
svc = ChatService(failover_chain=chain)
chunks = []
async for chunk in svc.stream_chat(messages=[{"role": "user", "content": "hi"}]):
chunks.append(chunk)
assert chunks[-1].done is True
assert "degraded" in chunks[-1].content
def test_build_system_prompt_no_template(self) -> None:
"""无模板服务时返回默认 prompt."""
chain = _make_chain(MockProvider())
svc = ChatService(failover_chain=chain)
prompt = svc._build_system_prompt({})
assert "educational assistant" in prompt.lower()
class TestQuestionService:
"""QuestionService 测试."""
async def test_generate_success(self) -> None:
"""生成题目成功."""
output = json.dumps(
{
"question": "1+1等于几",
"answer": "2",
"explanation": "基础加法",
"difficulty": "easy",
"question_type": "short_answer",
}
)
chain = _make_chain(MockProvider(response_content=output))
gate = QualityGate(rule_validator=RuleValidator())
svc = QuestionService(
failover_chain=chain,
quality_gate=gate,
default_model="gpt-4o",
)
request = GenerateQuestionRequest(
prompt="生成加法题",
subject="数学",
difficulty="easy",
)
data = await svc.generate(request)
assert data.question == "1+1等于几"
assert data.answer == "2"
assert data.degraded is False
async def test_generate_degraded_llm_fail(self) -> None:
"""LLM 失败降级."""
chain = _make_chain(MockProvider(fail=True))
svc = QuestionService(failover_chain=chain)
request = GenerateQuestionRequest(prompt="生成题目", subject="数学")
data = await svc.generate(request)
assert data.degraded is True
assert "degraded" in data.explanation
async def test_generate_invalid_json(self) -> None:
"""LLM 输出非 JSON 降级."""
chain = _make_chain(MockProvider(response_content="这不是JSON"))
svc = QuestionService(failover_chain=chain)
request = GenerateQuestionRequest(prompt="生成题目", subject="数学")
data = await svc.generate(request)
# 非JSON → 规则校验失败 → degraded
assert data.degraded is True
async def test_stream_generate(self) -> None:
"""流式生成题目."""
output = json.dumps({"question": "", "answer": ""})
chain = _make_chain(MockProvider(stream_chunks=[output]))
svc = QuestionService(failover_chain=chain)
request = GenerateQuestionRequest(prompt="生成题目", subject="数学")
chunks = []
async for chunk in svc.stream_generate(request):
chunks.append(chunk)
assert chunks[-1].done is True
assert chunks[-1].complete_question is not None
async def test_stream_generate_degraded(self) -> None:
"""流式生成降级."""
chain = _make_chain(MockProvider(fail=True))
svc = QuestionService(failover_chain=chain)
request = GenerateQuestionRequest(prompt="生成题目", subject="数学")
chunks = []
async for chunk in svc.stream_generate(request):
chunks.append(chunk)
assert chunks[-1].done is True
assert chunks[-1].complete_question is not None
assert chunks[-1].complete_question.degraded is True
def test_fallback_prompt(self) -> None:
"""降级 prompt."""
chain = _make_chain(MockProvider())
svc = QuestionService(failover_chain=chain)
request = GenerateQuestionRequest(prompt="测试", subject="语文")
prompt = svc._fallback_prompt(request)
assert "语文" in prompt
class TestExpressionService:
"""ExpressionService 测试."""
async def test_optimize_success(self) -> None:
"""优化成功."""
output = json.dumps(
{
"optimized": "优化后的文字",
"suggestions": ["建议1"],
}
)
chain = _make_chain(MockProvider(response_content=output))
svc = ExpressionService(failover_chain=chain)
data = await svc.optimize(text="原始文字")
assert data.optimized == "优化后的文字"
assert data.suggestions == ["建议1"]
assert data.degraded is False
async def test_optimize_degraded(self) -> None:
"""LLM 失败降级."""
chain = _make_chain(MockProvider(fail=True))
svc = ExpressionService(failover_chain=chain)
data = await svc.optimize(text="原始文字")
assert data.degraded is True
assert data.optimized == "原始文字"
async def test_optimize_invalid_json(self) -> None:
"""非 JSON 输出降级."""
chain = _make_chain(MockProvider(response_content="纯文本"))
svc = ExpressionService(failover_chain=chain)
data = await svc.optimize(text="原始文字")
assert data.degraded is True
assert "json" in data.degraded_reason.lower()
async def test_optimize_markdown_json(self) -> None:
"""markdown 包裹 JSON 能解析."""
output = '```json\n{"optimized": "ok", "suggestions": []}\n```'
chain = _make_chain(MockProvider(response_content=output))
svc = ExpressionService(failover_chain=chain)
data = await svc.optimize(text="原始文字")
assert data.optimized == "ok"
def test_fallback_prompt(self) -> None:
"""降级 prompt."""
chain = _make_chain(MockProvider())
svc = ExpressionService(failover_chain=chain)
prompt = svc._fallback_prompt("文字", "上下文")
assert "文字" in prompt
assert "上下文" in prompt
class TestReportService:
"""ReportService 测试."""
async def test_generate_class_summary_success(self) -> None:
"""班级学情总结报告生成成功."""
chain = _make_chain(
MockProvider(
response_content=(
"# 班级学情报告\n\n"
"## 摘要\n班级平均分 78.5,及格率 85%\n\n"
"## 详细分析\n整体表现良好。\n\n"
"## 教学建议\n- 加强函数概念\n- 增加练习题\n"
),
)
)
svc = ReportService(
failover_chain=chain,
data_ana_client=DataAnaClientMock(),
default_model="gpt-4o-mini",
)
data = await svc.generate(
class_id="c-1",
report_type="class_summary",
)
assert data.degraded is False
assert "班级学情报告" in data.content
assert data.summary # 非空
assert len(data.recommendations) == 2
assert "函数概念" in data.recommendations[0]
async def test_generate_student_detail_with_student_id(self) -> None:
"""学生详情报告(带 student_id."""
chain = _make_chain(
MockProvider(
response_content=(
"# 学生学情详情\n\n"
"## 摘要\n该生在函数概念上较薄弱。\n\n"
"## 教学建议\n- 针对性练习\n"
),
)
)
svc = ReportService(
failover_chain=chain,
data_ana_client=DataAnaClientMock(),
)
data = await svc.generate(
class_id="c-1",
report_type="student_detail",
student_id="s-001",
)
assert data.degraded is False
assert "学生学情详情" in data.content
async def test_generate_degraded_llm_fail(self) -> None:
"""LLM 不可用时降级."""
chain = _make_chain(MockProvider(fail=True))
svc = ReportService(
failover_chain=chain,
data_ana_client=DataAnaClientMock(),
)
data = await svc.generate(
class_id="c-1",
report_type="exam_analysis",
)
assert data.degraded is True
assert "LLM unavailable" in data.degraded_reason
assert data.content == ""
async def test_generate_degraded_no_data_ana_client(self) -> None:
"""无 data-ana 客户端时上下文降级(但 LLM 仍可生成)."""
chain = _make_chain(MockProvider(response_content="# 报告\n\n## 摘要\n无数据。"))
svc = ReportService(
failover_chain=chain,
data_ana_client=None,
)
data = await svc.generate(
class_id="c-1",
report_type="class_summary",
)
# LLM 可用 → 报告生成成功(上下文降级但不影响 LLM 调用)
assert data.degraded is False
assert "报告" in data.content
def test_extract_summary_from_explicit_section(self) -> None:
"""从「摘要」段落提取摘要."""
chain = _make_chain(MockProvider())
svc = ReportService(failover_chain=chain)
content = "# 报告\n\n## 摘要\n这是摘要内容。\n\n## 详细\n详情"
summary = svc._extract_summary(content)
assert "这是摘要内容" in summary
def test_extract_summary_fallback_first_200_chars(self) -> None:
"""无「摘要」段落时取前 200 字."""
chain = _make_chain(MockProvider())
svc = ReportService(failover_chain=chain)
content = "这是一段没有摘要标题的报告内容。"
summary = svc._extract_summary(content)
assert summary == content
def test_extract_recommendations_from_dash_list(self) -> None:
"""从「-」列表提取建议."""
chain = _make_chain(MockProvider())
svc = ReportService(failover_chain=chain)
content = "## 教学建议\n- 建议一\n- 建议二\n## 其他\n"
recs = svc._extract_recommendations(content)
assert recs == ["建议一", "建议二"]
def test_extract_recommendations_from_numbered_list(self) -> None:
"""从数字列表提取建议."""
chain = _make_chain(MockProvider())
svc = ReportService(failover_chain=chain)
content = "## 建议\n1. 第一条\n2. 第二条\n"
recs = svc._extract_recommendations(content)
assert recs == ["第一条", "第二条"]
def test_build_prompt_with_context(self) -> None:
"""构建 prompt 包含上下文数据."""
chain = _make_chain(MockProvider())
svc = ReportService(failover_chain=chain)
context = {
"class_id": "c-1",
"report_type": "class_summary",
"average_score": 78.5,
"pass_rate": 0.85,
}
prompt = svc._build_prompt("class_summary", context)
assert "class_summary" in prompt or "班级学情总结" in prompt
assert "78.5" in prompt