第 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%
331 lines
13 KiB
Python
331 lines
13 KiB
Python
"""服务层测试(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
|