Files
Edu/services/ai/tests/test_grpc_servicer.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

487 lines
20 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.
"""gRPC servicer 测试AiServicer 9 RPC + interceptors 辅助函数)."""
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import grpc
import pytest
from src.ai.errors import AIError, AILLMUnavailableError, ErrorCode
from src.ai.grpc_server.interceptors import _grpc_status, get_user_context
from src.ai.grpc_server.servicer import (
AiServicer,
_degraded_chat_response,
_degraded_question_response,
)
from src.ai.middleware.auth import UserContext
from src.ai.models.chat import ChatData, Usage
from src.ai.models.expression import OptimizedExpressionData
from src.ai.models.question import GeneratedQuestionData
from src.ai.models.report import GeneratedReportData
from src.ai.proto_gen import ai_pb2
class TestAiServicer:
"""AiServicer 9 RPC 测试."""
def setup_method(self) -> None:
self.chat_svc = AsyncMock()
self.question_svc = AsyncMock()
self.expr_svc = AsyncMock()
self.workflow_svc = AsyncMock()
self.report_svc = AsyncMock()
self.servicer = AiServicer(
chat_service=self.chat_svc,
question_service=self.question_svc,
expression_service=self.expr_svc,
workflow_service=self.workflow_svc,
report_service=self.report_svc,
)
self.context = MagicMock()
self.context.user_context = UserContext(user_id="u-1", role="teacher")
# ------------------------------------------------------------------ #
# Chat
# ------------------------------------------------------------------ #
async def test_chat_success(self) -> None:
self.chat_svc.chat.return_value = ChatData(
content="hi",
model="gpt-4o",
usage=Usage(),
)
request = ai_pb2.ChatRequest(model="gpt-4o")
request.messages.add(role="user", content="hello")
result = await self.servicer.Chat(request, self.context)
assert result.content == "hi"
assert result.model == "gpt-4o"
assert result.degraded is False
async def test_chat_no_service_degraded(self) -> None:
servicer = AiServicer(chat_service=None)
request = ai_pb2.ChatRequest(model="gpt-4o")
request.messages.add(role="user", content="hello")
result = await servicer.Chat(request, self.context)
assert result.degraded is True
assert "chat_service not initialized" in result.degraded_reason
async def test_chat_llm_unavailable_degraded(self) -> None:
self.chat_svc.chat.side_effect = AILLMUnavailableError("all providers failed")
request = ai_pb2.ChatRequest(model="gpt-4o")
request.messages.add(role="user", content="hello")
result = await self.servicer.Chat(request, self.context)
assert result.degraded is True
assert "all providers failed" in result.degraded_reason
async def test_chat_internal_error_raises(self) -> None:
self.chat_svc.chat.side_effect = RuntimeError("boom")
request = ai_pb2.ChatRequest(model="gpt-4o")
request.messages.add(role="user", content="hello")
with pytest.raises(AIError) as exc_info:
await self.servicer.Chat(request, self.context)
assert exc_info.value.code == ErrorCode.AI_INTERNAL_ERROR
# ------------------------------------------------------------------ #
# StreamChat
# ------------------------------------------------------------------ #
async def test_stream_chat_success(self) -> None:
async def mock_stream(**kwargs: object) -> None:
yield SimpleNamespace(content="chunk1", done=False)
yield SimpleNamespace(content="chunk2", done=True)
self.chat_svc.stream_chat = mock_stream
request = ai_pb2.ChatRequest(model="gpt-4o")
request.messages.add(role="user", content="hello")
chunks = []
async for chunk in self.servicer.StreamChat(request, self.context):
chunks.append(chunk)
assert len(chunks) == 2
assert chunks[0].content == "chunk1"
assert chunks[1].done is True
async def test_stream_chat_no_service_degraded(self) -> None:
servicer = AiServicer(chat_service=None)
request = ai_pb2.ChatRequest(model="gpt-4o")
request.messages.add(role="user", content="hello")
chunks = []
async for chunk in servicer.StreamChat(request, self.context):
chunks.append(chunk)
assert len(chunks) == 1
assert chunks[0].done is True
assert "chat_service not initialized" in chunks[0].content
async def test_stream_chat_error_yields_error_chunk(self) -> None:
async def mock_stream_error(**kwargs: object) -> None:
raise RuntimeError("stream boom")
yield # noqa -- makes this an async generator function
self.chat_svc.stream_chat = mock_stream_error
request = ai_pb2.ChatRequest(model="gpt-4o")
request.messages.add(role="user", content="hello")
chunks = []
async for chunk in self.servicer.StreamChat(request, self.context):
chunks.append(chunk)
assert len(chunks) == 1
assert chunks[0].done is True
assert "error" in chunks[0].content
# ------------------------------------------------------------------ #
# GenerateQuestion
# ------------------------------------------------------------------ #
async def test_generate_question_success(self) -> None:
self.question_svc.generate.return_value = GeneratedQuestionData(
question="1+1=?",
answer="2",
explanation="addition",
question_type="short_answer",
difficulty="easy",
knowledge_point_ids=["kp_1"],
evaluation_score=0.9,
)
request = ai_pb2.GenerateQuestionRequest(
prompt="生成加法题",
subject="数学",
difficulty="easy",
)
result = await self.servicer.GenerateQuestion(request, self.context)
assert result.question == "1+1=?"
assert result.answer == "2"
assert result.degraded is False
async def test_generate_question_no_service_degraded(self) -> None:
servicer = AiServicer(question_service=None)
request = ai_pb2.GenerateQuestionRequest(
prompt="生成加法题",
subject="数学",
difficulty="easy",
)
result = await servicer.GenerateQuestion(request, self.context)
assert result.degraded is True
assert "question_service not initialized" in result.degraded_reason
async def test_generate_question_llm_unavailable_degraded(self) -> None:
self.question_svc.generate.side_effect = AILLMUnavailableError("llm down")
request = ai_pb2.GenerateQuestionRequest(
prompt="生成加法题",
subject="数学",
difficulty="easy",
)
result = await self.servicer.GenerateQuestion(request, self.context)
assert result.degraded is True
assert "llm down" in result.degraded_reason
# ------------------------------------------------------------------ #
# StreamGenerateQuestion
# ------------------------------------------------------------------ #
async def test_stream_generate_question_success(self) -> None:
complete = ai_pb2.GeneratedQuestion(
question="q",
answer="a",
question_type="short_answer",
)
async def mock_stream_gen(request: object) -> None:
yield SimpleNamespace(content="chunk1", done=False, complete_question=None)
yield SimpleNamespace(content="", done=True, complete_question=complete)
self.question_svc.stream_generate = mock_stream_gen
request = ai_pb2.GenerateQuestionRequest(
prompt="生成加法题",
subject="数学",
difficulty="easy",
)
chunks = []
async for chunk in self.servicer.StreamGenerateQuestion(request, self.context):
chunks.append(chunk)
assert len(chunks) == 2
assert chunks[0].content == "chunk1"
assert chunks[1].done is True
assert chunks[1].HasField("complete_question")
assert chunks[1].complete_question.question == "q"
async def test_stream_generate_question_no_service_degraded(self) -> None:
servicer = AiServicer(question_service=None)
request = ai_pb2.GenerateQuestionRequest(
prompt="生成加法题",
subject="数学",
difficulty="easy",
)
chunks = []
async for chunk in servicer.StreamGenerateQuestion(request, self.context):
chunks.append(chunk)
assert len(chunks) == 1
assert chunks[0].done is True
assert "question_service not initialized" in chunks[0].content
# ------------------------------------------------------------------ #
# OptimizeExpression
# ------------------------------------------------------------------ #
async def test_optimize_expression_success(self) -> None:
self.expr_svc.optimize.return_value = OptimizedExpressionData(
optimized="优化后",
suggestions=["建议1"],
)
request = ai_pb2.OptimizeExpressionRequest(text="原始", context="")
result = await self.servicer.OptimizeExpression(request, self.context)
assert result.optimized == "优化后"
assert list(result.suggestions) == ["建议1"]
assert result.degraded is False
async def test_optimize_expression_no_service_degraded(self) -> None:
servicer = AiServicer(expression_service=None)
request = ai_pb2.OptimizeExpressionRequest(text="原始", context="")
result = await servicer.OptimizeExpression(request, self.context)
assert result.degraded is True
assert "expression_service not initialized" in result.degraded_reason
async def test_optimize_expression_llm_unavailable_degraded(self) -> None:
self.expr_svc.optimize.side_effect = AILLMUnavailableError("llm down")
request = ai_pb2.OptimizeExpressionRequest(text="原始", context="")
result = await self.servicer.OptimizeExpression(request, self.context)
assert result.degraded is True
assert "llm down" in result.degraded_reason
# ------------------------------------------------------------------ #
# GenerateLessonPlan
# ------------------------------------------------------------------ #
async def test_generate_lesson_plan_success(self) -> None:
self.workflow_svc.start.return_value = SimpleNamespace(
workflow_id="wf-1",
status="pending",
estimated_completion_seconds=60,
degraded=False,
degraded_reason="",
)
request = ai_pb2.GenerateLessonPlanRequest(
class_id="c-1",
subject_id="math",
topic="函数",
target_difficulty="medium",
question_count=3,
)
result = await self.servicer.GenerateLessonPlan(request, self.context)
assert result.workflow_id == "wf-1"
assert result.status == "pending"
assert result.estimated_completion_seconds == 60
assert result.degraded is False
async def test_generate_lesson_plan_no_service_degraded(self) -> None:
servicer = AiServicer(workflow_service=None)
request = ai_pb2.GenerateLessonPlanRequest(
class_id="c-1",
subject_id="math",
topic="函数",
)
result = await servicer.GenerateLessonPlan(request, self.context)
assert result.degraded is True
assert result.status == "failed"
assert "workflow_service not initialized" in result.degraded_reason
# ------------------------------------------------------------------ #
# GetLessonPlanStatus
# ------------------------------------------------------------------ #
async def test_get_lesson_plan_status_success(self) -> None:
question = GeneratedQuestionData(
question="q1",
answer="a1",
explanation="e1",
question_type="short_answer",
difficulty="easy",
knowledge_point_ids=["kp_1"],
evaluation_score=0.8,
)
self.workflow_svc.get_status.return_value = SimpleNamespace(
workflow_id="wf-1",
status="pending_review",
questions=[question],
error=None,
degraded=False,
degraded_reason="",
)
request = ai_pb2.GetLessonPlanStatusRequest(workflow_id="wf-1")
result = await self.servicer.GetLessonPlanStatus(request, self.context)
assert result.workflow_id == "wf-1"
assert result.status == "pending_review"
assert len(result.questions) == 1
assert result.questions[0].question == "q1"
assert result.questions[0].answer == "a1"
async def test_get_lesson_plan_status_no_service_degraded(self) -> None:
servicer = AiServicer(workflow_service=None)
request = ai_pb2.GetLessonPlanStatusRequest(workflow_id="wf-1")
result = await servicer.GetLessonPlanStatus(request, self.context)
assert result.degraded is True
assert result.status == "failed"
assert "workflow_service not initialized" in result.degraded_reason
# ------------------------------------------------------------------ #
# ConfirmLessonPlan
# ------------------------------------------------------------------ #
async def test_confirm_lesson_plan_success(self) -> None:
self.workflow_svc.confirm.return_value = SimpleNamespace(
success=True,
persisted_question_ids=["q_1", "q_2"],
error=None,
)
request = ai_pb2.ConfirmLessonPlanRequest(workflow_id="wf-1")
result = await self.servicer.ConfirmLessonPlan(request, self.context)
assert result.success is True
assert list(result.persisted_question_ids) == ["q_1", "q_2"]
async def test_confirm_lesson_plan_no_service_degraded(self) -> None:
servicer = AiServicer(workflow_service=None)
request = ai_pb2.ConfirmLessonPlanRequest(workflow_id="wf-1")
result = await servicer.ConfirmLessonPlan(request, self.context)
assert result.success is False
assert "workflow_service not initialized" in result.error
# ------------------------------------------------------------------ #
# GenerateReport
# ------------------------------------------------------------------ #
async def test_generate_report_success(self) -> None:
"""生成学情报告成功."""
self.report_svc.generate.return_value = GeneratedReportData(
id="report-1",
content="# 班级学情报告\n\n摘要\n班级整体表现良好。",
summary="班级整体表现良好。",
recommendations=["加强函数概念教学", "增加练习题量"],
degraded=False,
degraded_reason="",
)
request = ai_pb2.GenerateReportRequest(
class_id="c-1",
report_type="class_summary",
user_id="u-1",
)
result = await self.servicer.GenerateReport(request, self.context)
assert result.id == "report-1"
assert "班级学情报告" in result.content
assert result.summary == "班级整体表现良好。"
assert list(result.recommendations) == ["加强函数概念教学", "增加练习题量"]
assert result.degraded is False
async def test_generate_report_student_detail_with_student_id(self) -> None:
"""学生详情报告(带 student_id."""
self.report_svc.generate.return_value = GeneratedReportData(
id="report-2",
content="学生个人报告",
summary="学生薄弱点分析",
recommendations=["针对性练习"],
degraded=False,
degraded_reason="",
)
request = ai_pb2.GenerateReportRequest(
class_id="c-1",
report_type="student_detail",
student_id="s-001",
user_id="u-1",
)
result = await self.servicer.GenerateReport(request, self.context)
# 验证 service 被调用时 student_id 正确传递
self.report_svc.generate.assert_awaited_once()
call_kwargs = self.report_svc.generate.call_args.kwargs
assert call_kwargs["student_id"] == "s-001"
assert call_kwargs["report_type"] == "student_detail"
assert result.id == "report-2"
async def test_generate_report_no_service_degraded(self) -> None:
"""report_service 未初始化时返回降级响应."""
servicer = AiServicer(report_service=None)
request = ai_pb2.GenerateReportRequest(
class_id="c-1",
report_type="class_summary",
)
result = await servicer.GenerateReport(request, self.context)
assert result.degraded is True
assert "report_service not initialized" in result.degraded_reason
assert result.id == ""
assert result.content == ""
async def test_generate_report_llm_unavailable_degraded(self) -> None:
"""LLM 不可用时 ReportService 返回降级数据."""
self.report_svc.generate.return_value = GeneratedReportData(
id="report-3",
content="",
summary="",
recommendations=[],
degraded=True,
degraded_reason="LLM unavailable: all providers failed",
)
request = ai_pb2.GenerateReportRequest(
class_id="c-1",
report_type="exam_analysis",
)
result = await self.servicer.GenerateReport(request, self.context)
assert result.degraded is True
assert "LLM unavailable" in result.degraded_reason
async def test_generate_report_internal_error_raises(self) -> None:
"""未知异常转 AI_INTERNAL_ERROR."""
self.report_svc.generate.side_effect = RuntimeError("boom")
request = ai_pb2.GenerateReportRequest(
class_id="c-1",
report_type="class_summary",
)
with pytest.raises(AIError) as exc_info:
await self.servicer.GenerateReport(request, self.context)
assert exc_info.value.code == ErrorCode.AI_INTERNAL_ERROR
# ------------------------------------------------------------------ #
# degraded response helpers
# ------------------------------------------------------------------ #
def test_degraded_chat_response_helper(self) -> None:
result = _degraded_chat_response("gpt-4o", "service down")
assert result.degraded is True
assert result.model == "gpt-4o"
assert "service down" in result.content
assert result.degraded_reason == "service down"
assert result.usage.prompt_tokens == 0
def test_degraded_question_response_helper(self) -> None:
result = _degraded_question_response("llm unavailable")
assert result.degraded is True
assert "llm unavailable" in result.question
assert result.degraded_reason == "llm unavailable"
assert result.answer == ""
class TestInterceptorsHelpers:
"""interceptors 辅助函数测试."""
def test_grpc_status_mapping(self) -> None:
assert _grpc_status(0) == grpc.StatusCode.OK
assert _grpc_status(3) == grpc.StatusCode.INVALID_ARGUMENT
assert _grpc_status(5) == grpc.StatusCode.NOT_FOUND
assert _grpc_status(7) == grpc.StatusCode.PERMISSION_DENIED
assert _grpc_status(8) == grpc.StatusCode.UNAUTHENTICATED
assert _grpc_status(13) == grpc.StatusCode.INTERNAL
assert _grpc_status(14) == grpc.StatusCode.UNAVAILABLE
def test_grpc_status_unknown_code_defaults_to_unknown(self) -> None:
assert _grpc_status(999) == grpc.StatusCode.UNKNOWN
def test_get_user_context_default(self) -> None:
"""无 user_context 属性时返回默认 UserContext."""
bare_context = SimpleNamespace()
ctx = get_user_context(bare_context)
assert isinstance(ctx, UserContext)
assert ctx.user_id == ""
assert ctx.is_empty is True
def test_get_user_context_with_value(self) -> None:
"""有 user_context 属性时返回注入的 UserContext."""
context = MagicMock()
context.user_context = UserContext(user_id="u-1", role="teacher")
ctx = get_user_context(context)
assert ctx.user_id == "u-1"
assert ctx.role == "teacher"