feat: auto committed
This commit is contained in:
367
services/ai/tests/test_grpc_servicer.py
Normal file
367
services/ai/tests/test_grpc_servicer.py
Normal file
@@ -0,0 +1,367 @@
|
||||
"""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"
|
||||
Reference in New Issue
Block a user