feat: auto committed
This commit is contained in:
742
services/ai/tests/test_coverage_gaps.py
Normal file
742
services/ai/tests/test_coverage_gaps.py
Normal file
@@ -0,0 +1,742 @@
|
||||
"""Coverage gap tests for base_client, gRPC clients, usage, rate_limiter, etc.
|
||||
|
||||
Fills coverage gaps identified in:
|
||||
- clients/base_client.py (BaseGrpcClient, LoggingInterceptor, TracingInterceptor)
|
||||
- clients/content_client.py (ContentClientGrpc impl)
|
||||
- clients/data_ana_client.py (DataAnaClientGrpc impl)
|
||||
- clients/iam_client.py (IamClientGrpc impl)
|
||||
- usage/usage_recorder.py (Redis paths)
|
||||
- usage/kafka_producer.py (start/stop/publish paths)
|
||||
- rate_limiter.py (Redis paths)
|
||||
- middleware/permission.py (require_permission decorator)
|
||||
- middleware/error_handler.py (grpc_error_mapper)
|
||||
- workflow/state_store.py (Redis paths)
|
||||
- config.py (Settings properties)
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import grpc
|
||||
import pytest
|
||||
from redis.exceptions import RedisError
|
||||
|
||||
from src.ai.clients.base_client import BaseGrpcClient, LoggingInterceptor, TracingInterceptor
|
||||
from src.ai.clients.content_client import ContentClientGrpc, QuestionInput
|
||||
from src.ai.clients.data_ana_client import DataAnaClientGrpc
|
||||
from src.ai.clients.iam_client import IamClientGrpc
|
||||
from src.ai.config import Settings
|
||||
from src.ai.errors import AIError, AIRateLimitedError, ErrorCode
|
||||
from src.ai.middleware.auth import UserContext
|
||||
from src.ai.middleware.error_handler import grpc_error_mapper
|
||||
from src.ai.middleware.permission import (
|
||||
PERMISSION_AI_CHAT,
|
||||
PERMISSION_AI_LESSON_GENERATE,
|
||||
PermissionGuard,
|
||||
require_permission,
|
||||
)
|
||||
from src.ai.rate_limiter import RateLimiter
|
||||
from src.ai.usage.kafka_producer import KafkaProducer, UsageEvent
|
||||
from src.ai.usage.usage_recorder import UsageRecord, UsageRecorder
|
||||
from src.ai.workflow.state_store import WorkflowState, WorkflowStateStore
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _ConcreteClient(BaseGrpcClient):
|
||||
"""Concrete subclass for testing abstract BaseGrpcClient."""
|
||||
|
||||
def is_available(self) -> bool:
|
||||
return self._channel is not None
|
||||
|
||||
|
||||
class _AsyncCM:
|
||||
"""Minimal async context manager for mocking Kafka transaction()."""
|
||||
|
||||
async def __aenter__(self) -> "_AsyncCM":
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *args: object) -> None:
|
||||
pass
|
||||
|
||||
|
||||
def _make_redis_pipeline_mock() -> tuple[AsyncMock, MagicMock]:
|
||||
"""Build a (redis, pipeline) mock pair for UsageRecorder pipeline tests."""
|
||||
mock_redis: AsyncMock = AsyncMock()
|
||||
mock_pipe: MagicMock = MagicMock()
|
||||
mock_pipe.incrby = MagicMock(return_value=mock_pipe)
|
||||
mock_pipe.expire = MagicMock(return_value=mock_pipe)
|
||||
mock_pipe.execute = AsyncMock(return_value=[])
|
||||
mock_redis.pipeline = MagicMock(return_value=mock_pipe)
|
||||
return mock_redis, mock_pipe
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# base_client.py — BaseGrpcClient
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_base_client_connect() -> None:
|
||||
with patch("grpc.aio.insecure_channel") as mock_fn:
|
||||
mock_channel = AsyncMock()
|
||||
mock_fn.return_value = mock_channel
|
||||
client = _ConcreteClient("localhost:50054")
|
||||
await client.connect()
|
||||
mock_fn.assert_called_once()
|
||||
assert client._channel is mock_channel
|
||||
|
||||
|
||||
async def test_base_client_connect_idempotent() -> None:
|
||||
with patch("grpc.aio.insecure_channel") as mock_fn:
|
||||
mock_fn.return_value = AsyncMock()
|
||||
client = _ConcreteClient("localhost:50054")
|
||||
await client.connect()
|
||||
await client.connect()
|
||||
mock_fn.assert_called_once()
|
||||
|
||||
|
||||
async def test_base_client_close() -> None:
|
||||
with patch("grpc.aio.insecure_channel") as mock_fn:
|
||||
mock_channel = AsyncMock()
|
||||
mock_channel.close = AsyncMock()
|
||||
mock_fn.return_value = mock_channel
|
||||
client = _ConcreteClient("localhost:50054")
|
||||
await client.connect()
|
||||
await client.close()
|
||||
assert client._channel is None
|
||||
mock_channel.close.assert_called_once()
|
||||
|
||||
|
||||
async def test_base_client_close_not_connected() -> None:
|
||||
client = _ConcreteClient("localhost:50054")
|
||||
await client.close()
|
||||
assert client._channel is None
|
||||
|
||||
|
||||
def test_base_client_channel_not_connected_raises() -> None:
|
||||
client = _ConcreteClient("localhost:50054")
|
||||
with pytest.raises(RuntimeError, match="not connected"):
|
||||
_ = client.channel
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# base_client.py — TracingInterceptor
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_tracing_interceptor_inject_metadata() -> None:
|
||||
interceptor = TracingInterceptor(request_id="req-123")
|
||||
result = interceptor._inject_metadata([])
|
||||
assert ("x-request-id", "req-123") in result
|
||||
assert any(k == "traceparent" for k, _ in result)
|
||||
|
||||
|
||||
def test_tracing_interceptor_no_request_id() -> None:
|
||||
interceptor = TracingInterceptor(request_id="")
|
||||
result = interceptor._inject_metadata([])
|
||||
assert result == []
|
||||
|
||||
|
||||
def test_tracing_interceptor_with_existing_metadata() -> None:
|
||||
interceptor = TracingInterceptor(request_id="req-456")
|
||||
result = interceptor._inject_metadata([("existing", "value")])
|
||||
assert ("existing", "value") in result
|
||||
assert ("x-request-id", "req-456") in result
|
||||
assert any(k == "traceparent" for k, _ in result)
|
||||
|
||||
|
||||
def test_tracing_interceptor_with_none_metadata() -> None:
|
||||
interceptor = TracingInterceptor(request_id="req-789")
|
||||
result = interceptor._inject_metadata(None)
|
||||
assert isinstance(result, list)
|
||||
assert ("x-request-id", "req-789") in result
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# base_client.py — LoggingInterceptor
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_logging_interceptor_unary_success() -> None:
|
||||
interceptor = LoggingInterceptor()
|
||||
continuation = AsyncMock(return_value="response")
|
||||
call_details = MagicMock()
|
||||
call_details.method = "/svc/method"
|
||||
result = await interceptor.intercept_unary_unary(continuation, call_details, "req")
|
||||
assert result == "response"
|
||||
continuation.assert_called_once()
|
||||
|
||||
|
||||
async def test_logging_interceptor_unary_error_reraises() -> None:
|
||||
interceptor = LoggingInterceptor()
|
||||
rpc_error = grpc.aio.AioRpcError(
|
||||
code=grpc.StatusCode.UNAVAILABLE,
|
||||
initial_metadata=[],
|
||||
trailing_metadata=[],
|
||||
details="service unavailable",
|
||||
)
|
||||
continuation = AsyncMock(side_effect=rpc_error)
|
||||
call_details = MagicMock()
|
||||
call_details.method = "/svc/method"
|
||||
with pytest.raises(grpc.aio.AioRpcError):
|
||||
await interceptor.intercept_unary_unary(continuation, call_details, "req")
|
||||
|
||||
|
||||
async def test_logging_interceptor_unary_stream_success() -> None:
|
||||
interceptor = LoggingInterceptor()
|
||||
|
||||
async def _gen() -> Any:
|
||||
yield "chunk1"
|
||||
yield "chunk2"
|
||||
|
||||
async def continuation(cd: Any, req: Any) -> Any:
|
||||
return _gen()
|
||||
|
||||
call_details = MagicMock()
|
||||
call_details.method = "/svc/stream"
|
||||
agen = interceptor.intercept_unary_stream(continuation, call_details, "req")
|
||||
chunks = [c async for c in agen]
|
||||
assert chunks == ["chunk1", "chunk2"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# UsageRecorder — Redis paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_record_with_redis() -> None:
|
||||
mock_redis, mock_pipe = _make_redis_pipeline_mock()
|
||||
recorder = UsageRecorder(redis=mock_redis)
|
||||
record = UsageRecord(
|
||||
user_id="u-1",
|
||||
school_id="s-1",
|
||||
provider="openai",
|
||||
model="gpt-4o",
|
||||
operation="chat",
|
||||
total_tokens=100,
|
||||
)
|
||||
await recorder.record(record)
|
||||
assert mock_pipe.incrby.call_count == 2
|
||||
assert mock_pipe.expire.call_count == 2
|
||||
mock_pipe.execute.assert_called_once()
|
||||
|
||||
|
||||
async def test_record_redis_error_degraded() -> None:
|
||||
mock_redis, mock_pipe = _make_redis_pipeline_mock()
|
||||
mock_pipe.execute = AsyncMock(side_effect=RedisError("conn refused"))
|
||||
recorder = UsageRecorder(redis=mock_redis)
|
||||
record = UsageRecord(
|
||||
user_id="u-1",
|
||||
school_id="s-1",
|
||||
provider="openai",
|
||||
model="gpt-4o",
|
||||
operation="chat",
|
||||
total_tokens=100,
|
||||
)
|
||||
await recorder.record(record)
|
||||
|
||||
|
||||
async def test_get_user_usage_with_redis() -> None:
|
||||
mock_redis = AsyncMock()
|
||||
mock_redis.get = AsyncMock(return_value=b"500")
|
||||
recorder = UsageRecorder(redis=mock_redis)
|
||||
usage = await recorder.get_user_usage("u-1")
|
||||
assert usage == 500
|
||||
|
||||
|
||||
async def test_get_user_usage_redis_error_returns_0() -> None:
|
||||
mock_redis = AsyncMock()
|
||||
mock_redis.get = AsyncMock(side_effect=RedisError("conn"))
|
||||
recorder = UsageRecorder(redis=mock_redis)
|
||||
usage = await recorder.get_user_usage("u-1")
|
||||
assert usage == 0
|
||||
|
||||
|
||||
async def test_get_school_usage_with_redis() -> None:
|
||||
mock_redis = AsyncMock()
|
||||
mock_redis.get = AsyncMock(return_value=b"2000")
|
||||
recorder = UsageRecorder(redis=mock_redis)
|
||||
usage = await recorder.get_school_usage("s-1")
|
||||
assert usage == 2000
|
||||
|
||||
|
||||
async def test_get_school_usage_no_redis_returns_0() -> None:
|
||||
recorder = UsageRecorder(redis=None)
|
||||
usage = await recorder.get_school_usage("s-1")
|
||||
assert usage == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# RateLimiter — Redis paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_check_with_redis_allows() -> None:
|
||||
mock_redis = AsyncMock()
|
||||
mock_redis.script_load = AsyncMock(return_value="sha-abc")
|
||||
mock_redis.evalsha = AsyncMock(return_value=[1, 9])
|
||||
limiter = RateLimiter(redis=mock_redis)
|
||||
results = await limiter.check(user_id="u-1")
|
||||
assert len(results) == 1
|
||||
assert results[0].allowed is True
|
||||
assert results[0].remaining == 9
|
||||
|
||||
|
||||
async def test_check_with_redis_blocks() -> None:
|
||||
mock_redis = AsyncMock()
|
||||
mock_redis.script_load = AsyncMock(return_value="sha-abc")
|
||||
mock_redis.evalsha = AsyncMock(return_value=[0, 0])
|
||||
limiter = RateLimiter(redis=mock_redis)
|
||||
with pytest.raises(AIRateLimitedError):
|
||||
await limiter.check(user_id="u-1")
|
||||
|
||||
|
||||
async def test_check_redis_error_degraded_allows() -> None:
|
||||
mock_redis = AsyncMock()
|
||||
mock_redis.script_load = AsyncMock(return_value="sha-abc")
|
||||
mock_redis.evalsha = AsyncMock(side_effect=RedisError("eval failed"))
|
||||
limiter = RateLimiter(redis=mock_redis)
|
||||
results = await limiter.check(user_id="u-1")
|
||||
assert len(results) == 1
|
||||
assert results[0].allowed is True
|
||||
assert results[0].remaining == 10
|
||||
|
||||
|
||||
async def test_lua_load_failure_degrades() -> None:
|
||||
mock_redis = AsyncMock()
|
||||
mock_redis.script_load = AsyncMock(side_effect=RedisError("load failed"))
|
||||
limiter = RateLimiter(redis=mock_redis)
|
||||
results = await limiter.check(user_id="u-1", ip="1.2.3.4")
|
||||
assert len(results) == 2
|
||||
assert all(r.allowed for r in results)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# KafkaProducer — start / stop / publish paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_start_success() -> None:
|
||||
mock_module = MagicMock()
|
||||
mock_producer_cls = MagicMock()
|
||||
mock_instance = AsyncMock()
|
||||
mock_instance.start = AsyncMock()
|
||||
mock_producer_cls.return_value = mock_instance
|
||||
mock_module.AIOKafkaProducer = mock_producer_cls
|
||||
with patch.dict("sys.modules", {"aiokafka": mock_module}):
|
||||
producer = KafkaProducer()
|
||||
await producer.start()
|
||||
assert producer.is_started is True
|
||||
mock_producer_cls.assert_called_once()
|
||||
|
||||
|
||||
async def test_start_failure_degraded() -> None:
|
||||
with patch.dict("sys.modules", {"aiokafka": None}):
|
||||
producer = KafkaProducer()
|
||||
await producer.start()
|
||||
assert producer.is_started is False
|
||||
assert producer._producer is None
|
||||
|
||||
|
||||
async def test_stop_when_started() -> None:
|
||||
mock_prod = AsyncMock()
|
||||
mock_prod.stop = AsyncMock()
|
||||
producer = KafkaProducer()
|
||||
producer._producer = mock_prod
|
||||
producer._started = True
|
||||
await producer.stop()
|
||||
assert producer.is_started is False
|
||||
assert producer._producer is None
|
||||
mock_prod.stop.assert_called_once()
|
||||
|
||||
|
||||
async def test_stop_when_not_started() -> None:
|
||||
producer = KafkaProducer()
|
||||
await producer.stop()
|
||||
assert producer.is_started is False
|
||||
|
||||
|
||||
async def test_stop_failure_degraded() -> None:
|
||||
mock_prod = AsyncMock()
|
||||
mock_prod.stop = AsyncMock(side_effect=RuntimeError("stop failed"))
|
||||
producer = KafkaProducer()
|
||||
producer._producer = mock_prod
|
||||
producer._started = True
|
||||
await producer.stop()
|
||||
assert producer.is_started is False
|
||||
assert producer._producer is None
|
||||
|
||||
|
||||
async def test_publish_not_started_skipped() -> None:
|
||||
producer = KafkaProducer()
|
||||
event = UsageEvent(user_id="u-1")
|
||||
await producer.publish(event)
|
||||
|
||||
|
||||
async def test_publish_success() -> None:
|
||||
mock_prod = AsyncMock()
|
||||
mock_prod.transaction = MagicMock(return_value=_AsyncCM())
|
||||
mock_prod.send_and_wait = AsyncMock()
|
||||
producer = KafkaProducer()
|
||||
producer._producer = mock_prod
|
||||
producer._started = True
|
||||
event = UsageEvent(user_id="u-1", operation="chat", event_id="e-1")
|
||||
await producer.publish(event)
|
||||
mock_prod.send_and_wait.assert_called_once()
|
||||
|
||||
|
||||
async def test_publish_failure_degraded() -> None:
|
||||
mock_prod = AsyncMock()
|
||||
mock_prod.transaction = MagicMock(return_value=_AsyncCM())
|
||||
mock_prod.send_and_wait = AsyncMock(side_effect=RuntimeError("send failed"))
|
||||
producer = KafkaProducer()
|
||||
producer._producer = mock_prod
|
||||
producer._started = True
|
||||
event = UsageEvent(user_id="u-1", event_id="e-2")
|
||||
await producer.publish(event)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ContentClientGrpc
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_content_grpc_connect_close() -> None:
|
||||
with patch("grpc.aio.insecure_channel") as mock_fn:
|
||||
mock_channel = AsyncMock()
|
||||
mock_channel.close = AsyncMock()
|
||||
mock_fn.return_value = mock_channel
|
||||
client = ContentClientGrpc()
|
||||
assert client.is_available() is False
|
||||
await client.connect()
|
||||
assert client.is_available() is True
|
||||
await client.close()
|
||||
assert client.is_available() is False
|
||||
mock_channel.close.assert_called_once()
|
||||
|
||||
|
||||
async def test_content_grpc_methods_not_connected() -> None:
|
||||
client = ContentClientGrpc()
|
||||
assert client.is_available() is False
|
||||
pre = await client.get_prerequisites("kp-1")
|
||||
assert len(pre) == 1
|
||||
path = await client.get_learning_path("s-1", "math")
|
||||
assert len(path) == 3
|
||||
questions = [
|
||||
QuestionInput(
|
||||
question="q1",
|
||||
answer="a1",
|
||||
explanation="e1",
|
||||
question_type="short_answer",
|
||||
difficulty="easy",
|
||||
knowledge_point_ids=["kp-1"],
|
||||
),
|
||||
]
|
||||
result = await client.create_questions(questions, user_id="u-1")
|
||||
assert len(result) == 1
|
||||
|
||||
|
||||
async def test_content_grpc_methods_connected() -> None:
|
||||
client = ContentClientGrpc()
|
||||
client._channel = MagicMock()
|
||||
assert client.is_available() is True
|
||||
pre = await client.get_prerequisites("kp-1")
|
||||
assert len(pre) == 1
|
||||
path = await client.get_learning_path("s-1", "math")
|
||||
assert len(path) == 3
|
||||
questions = [
|
||||
QuestionInput(
|
||||
question="q1",
|
||||
answer="a1",
|
||||
explanation="e1",
|
||||
question_type="short_answer",
|
||||
difficulty="easy",
|
||||
knowledge_point_ids=["kp-1"],
|
||||
),
|
||||
]
|
||||
result = await client.create_questions(questions, user_id="u-1")
|
||||
assert len(result) == 1
|
||||
|
||||
|
||||
async def test_content_grpc_create_questions_exception_raises() -> None:
|
||||
client = ContentClientGrpc()
|
||||
client._channel = MagicMock()
|
||||
client._mock = MagicMock()
|
||||
client._mock.create_questions = AsyncMock(side_effect=RuntimeError("boom"))
|
||||
with pytest.raises(AIError) as exc_info:
|
||||
await client.create_questions([], user_id="u-1")
|
||||
assert exc_info.value.code == ErrorCode.AI_DOWNSTREAM_UNAVAILABLE
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DataAnaClientGrpc
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_data_ana_grpc_connect_close() -> None:
|
||||
with patch("grpc.aio.insecure_channel") as mock_fn:
|
||||
mock_channel = AsyncMock()
|
||||
mock_channel.close = AsyncMock()
|
||||
mock_fn.return_value = mock_channel
|
||||
client = DataAnaClientGrpc()
|
||||
assert client.is_available() is False
|
||||
await client.connect()
|
||||
assert client.is_available() is True
|
||||
await client.close()
|
||||
assert client.is_available() is False
|
||||
|
||||
|
||||
async def test_data_ana_grpc_methods_not_connected() -> None:
|
||||
client = DataAnaClientGrpc()
|
||||
assert client.is_available() is False
|
||||
perf = await client.get_class_performance("c-1", "math")
|
||||
assert perf.class_id == "c-1"
|
||||
assert perf.average_score > 0
|
||||
weak = await client.get_student_weakness("s-1", "math")
|
||||
assert weak.student_id == "s-1"
|
||||
assert len(weak.weak_points) > 0
|
||||
trend = await client.get_learning_trend("s-1")
|
||||
assert trend.student_id == "s-1"
|
||||
assert len(trend.points) > 0
|
||||
|
||||
|
||||
async def test_data_ana_grpc_methods_connected() -> None:
|
||||
client = DataAnaClientGrpc()
|
||||
client._channel = MagicMock()
|
||||
assert client.is_available() is True
|
||||
perf = await client.get_class_performance("c-1", "math")
|
||||
assert perf.class_id == "c-1"
|
||||
weak = await client.get_student_weakness("s-1", "math")
|
||||
assert weak.student_id == "s-1"
|
||||
trend = await client.get_learning_trend("s-1")
|
||||
assert trend.student_id == "s-1"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# IamClientGrpc
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_iam_grpc_connect_close() -> None:
|
||||
with patch("grpc.aio.insecure_channel") as mock_fn:
|
||||
mock_channel = AsyncMock()
|
||||
mock_channel.close = AsyncMock()
|
||||
mock_fn.return_value = mock_channel
|
||||
client = IamClientGrpc()
|
||||
assert client.is_available() is False
|
||||
await client.connect()
|
||||
assert client.is_available() is True
|
||||
await client.close()
|
||||
assert client.is_available() is False
|
||||
|
||||
|
||||
async def test_iam_grpc_get_effective_data_scope() -> None:
|
||||
client = IamClientGrpc()
|
||||
assert client.is_available() is False
|
||||
scope = await client.get_effective_data_scope("u-1")
|
||||
assert scope.user_id == "u-1"
|
||||
assert scope.school_id == "school_mock_001"
|
||||
assert len(scope.class_ids) > 0
|
||||
client._channel = MagicMock()
|
||||
assert client.is_available() is True
|
||||
scope2 = await client.get_effective_data_scope("u-2")
|
||||
assert scope2.user_id == "u-2"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# require_permission decorator
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_require_permission_passes() -> None:
|
||||
guard = PermissionGuard(dev_mode=False)
|
||||
ctx = UserContext(user_id="u-1", role="teacher")
|
||||
|
||||
@require_permission(PERMISSION_AI_CHAT, guard)
|
||||
async def handler(*, user_context: UserContext) -> str:
|
||||
return "ok"
|
||||
|
||||
result = await handler(user_context=ctx)
|
||||
assert result == "ok"
|
||||
|
||||
|
||||
async def test_require_permission_forbidden() -> None:
|
||||
guard = PermissionGuard(dev_mode=False)
|
||||
ctx = UserContext(user_id="u-1", role="student")
|
||||
|
||||
@require_permission(PERMISSION_AI_LESSON_GENERATE, guard)
|
||||
async def handler(*, user_context: UserContext) -> str:
|
||||
return "ok"
|
||||
|
||||
with pytest.raises(AIError) as exc_info:
|
||||
await handler(user_context=ctx)
|
||||
assert exc_info.value.code == ErrorCode.AI_FORBIDDEN
|
||||
|
||||
|
||||
async def test_require_permission_no_context_defaults_unauthenticated() -> None:
|
||||
guard = PermissionGuard(dev_mode=False)
|
||||
|
||||
@require_permission(PERMISSION_AI_CHAT, guard)
|
||||
async def handler() -> str:
|
||||
return "ok"
|
||||
|
||||
with pytest.raises(AIError) as exc_info:
|
||||
await handler()
|
||||
assert exc_info.value.code == ErrorCode.AI_UNAUTHORIZED
|
||||
|
||||
|
||||
async def test_require_permission_finds_ctx_in_args() -> None:
|
||||
guard = PermissionGuard(dev_mode=False)
|
||||
ctx = UserContext(user_id="u-1", role="teacher")
|
||||
|
||||
@require_permission(PERMISSION_AI_CHAT, guard)
|
||||
async def handler(user_context: UserContext) -> str:
|
||||
return "ok"
|
||||
|
||||
result = await handler(ctx)
|
||||
assert result == "ok"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WorkflowStateStore — Redis paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_workflow_create_with_redis() -> None:
|
||||
mock_redis = AsyncMock()
|
||||
mock_redis.setex = AsyncMock()
|
||||
store = WorkflowStateStore(redis=mock_redis)
|
||||
state = WorkflowState(user_id="u-1", topic="test")
|
||||
created = await store.create(state)
|
||||
mock_redis.setex.assert_called_once()
|
||||
assert created.workflow_id == state.workflow_id
|
||||
|
||||
|
||||
async def test_workflow_get_with_redis() -> None:
|
||||
mock_redis = AsyncMock()
|
||||
state = WorkflowState(user_id="u-1", topic="test")
|
||||
mock_redis.get = AsyncMock(return_value=json.dumps(state.to_dict()))
|
||||
store = WorkflowStateStore(redis=mock_redis)
|
||||
fetched = await store.get(state.workflow_id)
|
||||
assert fetched.user_id == "u-1"
|
||||
assert fetched.topic == "test"
|
||||
|
||||
|
||||
async def test_workflow_update_with_redis() -> None:
|
||||
mock_redis = AsyncMock()
|
||||
state = WorkflowState(user_id="u-1")
|
||||
mock_redis.get = AsyncMock(return_value=json.dumps(state.to_dict()))
|
||||
mock_redis.setex = AsyncMock()
|
||||
store = WorkflowStateStore(redis=mock_redis)
|
||||
updated = await store.update(state.workflow_id, status="analyzing")
|
||||
assert updated.status == "analyzing"
|
||||
mock_redis.setex.assert_called_once()
|
||||
|
||||
|
||||
async def test_workflow_delete_with_redis() -> None:
|
||||
mock_redis = AsyncMock()
|
||||
mock_redis.delete = AsyncMock()
|
||||
store = WorkflowStateStore(redis=mock_redis)
|
||||
await store.delete("wf-1")
|
||||
mock_redis.delete.assert_called_once()
|
||||
|
||||
|
||||
async def test_workflow_create_redis_error_degraded() -> None:
|
||||
mock_redis = AsyncMock()
|
||||
mock_redis.setex = AsyncMock(side_effect=RedisError("conn"))
|
||||
store = WorkflowStateStore(redis=mock_redis)
|
||||
state = WorkflowState(user_id="u-1")
|
||||
created = await store.create(state)
|
||||
assert created.workflow_id == state.workflow_id
|
||||
|
||||
|
||||
async def test_workflow_get_redis_error_degraded() -> None:
|
||||
mock_redis = AsyncMock()
|
||||
mock_redis.setex = AsyncMock(side_effect=RedisError("conn"))
|
||||
mock_redis.get = AsyncMock(side_effect=RedisError("conn"))
|
||||
store = WorkflowStateStore(redis=mock_redis)
|
||||
state = WorkflowState(user_id="u-1")
|
||||
await store.create(state)
|
||||
fetched = await store.get(state.workflow_id)
|
||||
assert fetched.user_id == "u-1"
|
||||
|
||||
|
||||
async def test_workflow_delete_redis_error_no_crash() -> None:
|
||||
mock_redis = AsyncMock()
|
||||
mock_redis.delete = AsyncMock(side_effect=RedisError("conn"))
|
||||
store = WorkflowStateStore(redis=mock_redis)
|
||||
await store.delete("wf-1")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# grpc_error_mapper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_grpc_error_mapper_ai_error() -> None:
|
||||
exc = AIError(ErrorCode.AI_UNAUTHORIZED, "not authenticated")
|
||||
code, msg, status = grpc_error_mapper(exc)
|
||||
assert code == "AI_UNAUTHORIZED"
|
||||
assert msg == "not authenticated"
|
||||
assert status == 8
|
||||
|
||||
|
||||
def test_grpc_error_mapper_unknown_exception() -> None:
|
||||
exc = RuntimeError("unexpected")
|
||||
code, msg, status = grpc_error_mapper(exc)
|
||||
assert code == "AI_INTERNAL_ERROR"
|
||||
assert msg == "Internal server error"
|
||||
assert status == 13
|
||||
|
||||
|
||||
def test_grpc_error_mapper_various_codes() -> None:
|
||||
assert grpc_error_mapper(AIError(ErrorCode.AI_FORBIDDEN, "denied"))[2] == 7
|
||||
assert grpc_error_mapper(AIRateLimitedError("user", 10))[2] == 9
|
||||
assert grpc_error_mapper(AIError(ErrorCode.AI_WORKFLOW_NOT_FOUND, "nf"))[2] == 5
|
||||
assert grpc_error_mapper(AIError(ErrorCode.AI_LLM_UNAVAILABLE, "down"))[2] == 14
|
||||
assert grpc_error_mapper(AIError(ErrorCode.AI_INVALID_MODEL, "bad"))[2] == 3
|
||||
assert grpc_error_mapper(AIError(ErrorCode.AI_WORKFLOW_STATE_INVALID, "bad"))[2] == 10
|
||||
assert grpc_error_mapper(AIError(ErrorCode.AI_INTERNAL_ERROR, "err"))[2] == 13
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Config — Settings properties
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_settings_defaults(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
for key in ("SERVICE_NAME", "HTTP_PORT", "GRPC_PORT", "DEV_MODE"):
|
||||
monkeypatch.delenv(key, raising=False)
|
||||
s = Settings(_env_file=None)
|
||||
assert s.service_name == "ai"
|
||||
assert s.http_port == 3008
|
||||
assert s.grpc_port == 50058
|
||||
assert s.dev_mode is False
|
||||
|
||||
|
||||
def test_settings_dev_mode() -> None:
|
||||
s = Settings(_env_file=None, dev_mode=True)
|
||||
assert s.is_dev is True
|
||||
|
||||
|
||||
def test_settings_providers_status() -> None:
|
||||
s = Settings(_env_file=None, openai_api_key="sk-test")
|
||||
status = s.providers_status
|
||||
assert set(status.keys()) == {"openai", "anthropic", "baichuan", "local_ollama"}
|
||||
assert status["openai"] is True
|
||||
|
||||
|
||||
def test_settings_llm_available() -> None:
|
||||
s = Settings(_env_file=None, openai_api_key="sk-test")
|
||||
assert s.llm_available is True
|
||||
|
||||
|
||||
def test_settings_provider_priority_list() -> None:
|
||||
s = Settings(_env_file=None)
|
||||
lst = s.provider_priority_list
|
||||
assert isinstance(lst, list)
|
||||
assert "openai" in lst
|
||||
assert "anthropic" in lst
|
||||
Reference in New Issue
Block a user