from __future__ import annotations from unittest.mock import AsyncMock, MagicMock import pytest from yuxi.channel.application.pipeline.middlewares.auth_middleware import AuthMiddleware from yuxi.channel.application.pipeline.middlewares.validation_middleware import ValidationMiddleware from yuxi.channel.application.pipeline.middlewares.keyword_filter_middleware import KeywordFilterMiddleware from yuxi.channel.application.pipeline.middlewares.dedup_middleware import DedupMiddleware from yuxi.channel.application.pipeline.middlewares.rate_limit_middleware import RateLimitMiddleware from yuxi.channel.application.pipeline.middlewares.access_policy_middleware import AccessPolicyMiddleware from yuxi.channel.application.pipeline.middlewares.mention_gate_middleware import MentionGateMiddleware from yuxi.channel.application.pipeline.middlewares.enqueue_mq_middleware import EnqueueMQMiddleware from yuxi.channel.domain.service.message_context import MessageContext class TestAuthMiddleware: @pytest.fixture def mock_auth_service(self): service = AsyncMock() service.authenticate.return_value = (True, "") return service @pytest.fixture def auth_middleware(self, mock_auth_service): return AuthMiddleware(auth_service=mock_auth_service) @pytest.mark.asyncio async def test_auth_pass(self, auth_middleware, mock_auth_service): ctx = MagicMock(spec=MessageContext) ctx.headers = {"Authorization": "Bearer token"} result = await auth_middleware.process(ctx) assert result.is_aborted is False @pytest.mark.asyncio async def test_auth_fail(self, auth_middleware, mock_auth_service): mock_auth_service.authenticate.return_value = (False, "auth failed") ctx = MagicMock(spec=MessageContext) ctx.headers = {"Authorization": "Bearer wrong"} result = await auth_middleware.process(ctx) assert result.is_aborted is True assert result.abort_code == "AUTH_FAILED" @pytest.mark.asyncio async def test_no_auth_header(self, auth_middleware): ctx = MagicMock(spec=MessageContext) ctx.headers = {} result = await auth_middleware.process(ctx) assert result.is_aborted is False class TestValidationMiddleware: @pytest.fixture def validation_middleware(self): return ValidationMiddleware() @pytest.mark.asyncio async def test_valid_message(self, validation_middleware): ctx = MagicMock(spec=MessageContext) ctx.message.message_id = "msg123" ctx.message.channel_type = "web" ctx.message.content = "Hello" result = await validation_middleware.process(ctx) assert result.is_aborted is False @pytest.mark.asyncio async def test_missing_message_id(self, validation_middleware): ctx = MagicMock(spec=MessageContext) ctx.message.message_id = "" ctx.message.channel_type = "web" ctx.message.content = "Hello" result = await validation_middleware.process(ctx) assert result.is_aborted is True assert result.abort_code == "VALIDATION_ERROR" @pytest.mark.asyncio async def test_missing_channel_type(self, validation_middleware): ctx = MagicMock(spec=MessageContext) ctx.message.message_id = "msg123" ctx.message.channel_type = "" ctx.message.content = "Hello" result = await validation_middleware.process(ctx) assert result.is_aborted is True @pytest.mark.asyncio async def test_missing_content(self, validation_middleware): ctx = MagicMock(spec=MessageContext) ctx.message.message_id = "msg123" ctx.message.channel_type = "web" ctx.message.content = "" result = await validation_middleware.process(ctx) assert result.is_aborted is True class TestKeywordFilterMiddleware: @pytest.fixture def mock_matcher(self): matcher = MagicMock() matcher.match.return_value = False return matcher @pytest.fixture def keyword_middleware(self, mock_matcher): return KeywordFilterMiddleware(keyword_matcher=mock_matcher) @pytest.mark.asyncio async def test_no_match(self, keyword_middleware, mock_matcher): ctx = MagicMock(spec=MessageContext) ctx.message.content = "Hello world" result = await keyword_middleware.process(ctx) assert result.is_aborted is False @pytest.mark.asyncio async def test_match_found(self, keyword_middleware, mock_matcher): mock_matcher.match.return_value = True ctx = MagicMock(spec=MessageContext) ctx.message.content = "bad word" result = await keyword_middleware.process(ctx) assert result.is_aborted is True assert result.abort_code == "KEYWORD_BLOCKED" @pytest.mark.asyncio async def test_no_matcher(self): middleware = KeywordFilterMiddleware() ctx = MagicMock(spec=MessageContext) ctx.message.content = "Hello" result = await middleware.process(ctx) assert result.is_aborted is False class TestDedupMiddleware: @pytest.fixture def mock_cache(self): cache = AsyncMock() cache.exists.return_value = False return cache @pytest.fixture def dedup_middleware(self, mock_cache): return DedupMiddleware(cache_port=mock_cache) @pytest.mark.asyncio async def test_new_message(self, dedup_middleware, mock_cache): ctx = MagicMock(spec=MessageContext) ctx.message.message_id = "msg123" result = await dedup_middleware.process(ctx) assert result.is_aborted is False mock_cache.set.assert_awaited_once() @pytest.mark.asyncio async def test_duplicate_message(self, dedup_middleware, mock_cache): mock_cache.exists.return_value = True ctx = MagicMock(spec=MessageContext) ctx.message.message_id = "msg123" result = await dedup_middleware.process(ctx) assert result.is_aborted is True assert result.abort_code == "DEDUP" @pytest.mark.asyncio async def test_no_cache(self): middleware = DedupMiddleware() ctx = MagicMock(spec=MessageContext) ctx.message.message_id = "msg123" result = await middleware.process(ctx) assert result.is_aborted is False class TestRateLimitMiddleware: @pytest.fixture def mock_rate_limiter(self): limiter = AsyncMock() limiter.is_locked.return_value = (False, 0) limiter.check_and_incr.return_value = True return limiter @pytest.fixture def rate_limit_middleware(self, mock_rate_limiter): return RateLimitMiddleware(rate_limit_port=mock_rate_limiter) @pytest.mark.asyncio async def test_within_limit(self, rate_limit_middleware, mock_rate_limiter): ctx = MagicMock(spec=MessageContext) ctx.message.sender.id = "user1" result = await rate_limit_middleware.process(ctx) assert result.is_aborted is False @pytest.mark.asyncio async def test_locked(self, rate_limit_middleware, mock_rate_limiter): mock_rate_limiter.is_locked.return_value = (True, 120) ctx = MagicMock(spec=MessageContext) ctx.message.sender.id = "user1" result = await rate_limit_middleware.process(ctx) assert result.is_aborted is True assert result.abort_code == "RATE_LIMITED" @pytest.mark.asyncio async def test_exceeded(self, rate_limit_middleware, mock_rate_limiter): mock_rate_limiter.check_and_incr.return_value = False ctx = MagicMock(spec=MessageContext) ctx.message.sender.id = "user1" result = await rate_limit_middleware.process(ctx) assert result.is_aborted is True @pytest.mark.asyncio async def test_no_limiter(self): middleware = RateLimitMiddleware() ctx = MagicMock(spec=MessageContext) ctx.message.sender.id = "user1" result = await middleware.process(ctx) assert result.is_aborted is False class TestAccessPolicyMiddleware: @pytest.fixture def access_policy_middleware(self): return AccessPolicyMiddleware( access_policies={"web": {"dm_policy": "open"}}, allow_from={"web": ["user1"]}, ) @pytest.mark.asyncio async def test_allow_open_policy(self, access_policy_middleware): ctx = MagicMock(spec=MessageContext) ctx.message.channel_type = "web" ctx.message.sender.id = "user1" ctx.message.metadata.is_group = False result = await access_policy_middleware.process(ctx) assert result.is_aborted is False @pytest.mark.asyncio async def test_allowlist_allowed(self): middleware = AccessPolicyMiddleware( access_policies={"web": {"dm_policy": "allowlist"}}, allow_from={"web": ["user1"]}, ) ctx = MagicMock(spec=MessageContext) ctx.message.channel_type = "web" ctx.message.sender.id = "user1" ctx.message.metadata.is_group = False result = await middleware.process(ctx) assert result.is_aborted is False @pytest.mark.asyncio async def test_allowlist_denied(self): middleware = AccessPolicyMiddleware( access_policies={"web": {"dm_policy": "allowlist"}}, allow_from={"web": ["user1"]}, ) ctx = MagicMock(spec=MessageContext) ctx.message.channel_type = "web" ctx.message.sender.id = "user2" ctx.message.metadata.is_group = False result = await middleware.process(ctx) assert result.is_aborted is True assert result.abort_code == "ACCESS_DENIED" @pytest.mark.asyncio async def test_no_policies(self): middleware = AccessPolicyMiddleware() ctx = MagicMock(spec=MessageContext) ctx.message.channel_type = "web" ctx.message.sender.id = "user1" ctx.message.metadata.is_group = False result = await middleware.process(ctx) assert result.is_aborted is False @pytest.mark.asyncio async def test_group_message_allowed(self, access_policy_middleware): ctx = MagicMock(spec=MessageContext) ctx.message.channel_type = "web" ctx.message.sender.id = "user1" ctx.message.metadata.is_group = True result = await access_policy_middleware.process(ctx) assert result.is_aborted is False class TestMentionGateMiddleware: @pytest.fixture def mention_gate_middleware(self): return MentionGateMiddleware( default_require_mention=False, mention_patterns=["@bot"], ) @pytest.mark.asyncio async def test_no_mention_required(self, mention_gate_middleware): ctx = MagicMock(spec=MessageContext) ctx.message.content = "Hello" ctx.message.metadata.is_group = False result = await mention_gate_middleware.process(ctx) assert result.is_aborted is False @pytest.mark.asyncio async def test_mention_required_and_found(self): middleware = MentionGateMiddleware( default_require_mention=True, mention_patterns=["@bot"], ) ctx = MagicMock(spec=MessageContext) ctx.message.content = "Hello @bot" ctx.message.metadata.is_group = True result = await middleware.process(ctx) assert result.is_aborted is False @pytest.mark.asyncio async def test_mention_required_not_found(self): middleware = MentionGateMiddleware( default_require_mention=True, mention_patterns=["@bot"], ) ctx = MagicMock(spec=MessageContext) ctx.message.content = "Hello" ctx.message.metadata.is_group = True result = await middleware.process(ctx) assert result.is_aborted is True assert result.abort_code == "MENTION_REQUIRED" @pytest.mark.asyncio async def test_dm_not_required(self): middleware = MentionGateMiddleware( default_require_mention=True, mention_patterns=["@bot"], ) ctx = MagicMock(spec=MessageContext) ctx.message.content = "Hello" ctx.message.metadata.is_group = False result = await middleware.process(ctx) assert result.is_aborted is False class TestEnqueueMQMiddleware: @pytest.fixture def mock_queue(self): return AsyncMock() @pytest.fixture def enqueue_middleware(self, mock_queue): return EnqueueMQMiddleware(queue_port=mock_queue) @pytest.mark.asyncio async def test_enqueue_success(self, enqueue_middleware, mock_queue): ctx = MagicMock(spec=MessageContext) ctx.message.to_dict.return_value = {"message_id": "msg123"} ctx.trace_id = "trace123" result = await enqueue_middleware.process(ctx) assert result.is_aborted is False mock_queue.enqueue.assert_awaited_once() @pytest.mark.asyncio async def test_enqueue_no_queue(self): middleware = EnqueueMQMiddleware() ctx = MagicMock(spec=MessageContext) ctx.message.to_dict.return_value = {"message_id": "msg123"} ctx.trace_id = "trace123" result = await middleware.process(ctx) assert result.is_aborted is False @pytest.mark.asyncio async def test_enqueue_exception(self, enqueue_middleware, mock_queue): mock_queue.enqueue.side_effect = Exception("queue error") ctx = MagicMock(spec=MessageContext) ctx.message.to_dict.return_value = {"message_id": "msg123"} ctx.trace_id = "trace123" result = await enqueue_middleware.process(ctx) assert result.is_aborted is True assert result.abort_code == "ENQUEUE_ERROR"