"""gRPC servicer 测试(AiServicer 8 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.proto_gen import ai_pb2 class TestAiServicer: """AiServicer 8 RPC 测试.""" def setup_method(self) -> None: self.chat_svc = AsyncMock() self.question_svc = AsyncMock() self.expr_svc = AsyncMock() self.workflow_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, ) 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 # ------------------------------------------------------------------ # # 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"