"""默认入站中间件单元测试。""" from __future__ import annotations from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock import pytest from yuxi.channel.constants import InboundRejectionReason from yuxi.channel.message.dedupe import MessageDeduper from yuxi.channel.middlewares.inbound import ( CreateRunMiddleware, DedupeMiddleware, RouteMiddleware, SecurityMiddleware, SessionMiddleware, TransactionMiddleware, ) from yuxi.channel.middlewares.protocols import InboundContext, InboundResult from yuxi.channel.plugins.protocol import BindingRoute, InboundMessage async def _noop_next() -> InboundResult: return InboundResult(accepted=True) @pytest.fixture def deduper_mock(): d = MagicMock(spec=MessageDeduper) d.is_processed = AsyncMock(return_value=False) d.clear_processed = AsyncMock() return d @pytest.fixture def plugin_mock(): return MagicMock() @pytest.fixture def base_ctx(plugin_mock): return InboundContext( channel_type="feishu", account_id="acc-1", inbound=InboundMessage( channel_type="feishu", account_id="acc-1", channel_message_id="m-1", sender_id="u-1", ), config={}, config_mw={}, plugin=plugin_mock, ) @pytest.fixture def db_session_mock(): return MagicMock() @pytest.fixture def pg_manager_mock(monkeypatch, db_session_mock): @asynccontextmanager async def ctx(): yield db_session_mock mock = MagicMock() mock.get_async_session_context = ctx monkeypatch.setattr("yuxi.channel.middlewares.inbound.pg_manager", mock) return mock async def test_dedupe_middleware_duplicate(deduper_mock, base_ctx): deduper_mock.is_processed.return_value = True mw = DedupeMiddleware(deduper_mock) result = await mw.process(base_ctx, _noop_next) assert result.accepted is False assert result.reason == InboundRejectionReason.DUPLICATE deduper_mock.is_processed.assert_awaited_once_with("m-1") async def test_dedupe_middleware_not_duplicate(deduper_mock, base_ctx): mw = DedupeMiddleware(deduper_mock) result = await mw.process(base_ctx, _noop_next) assert result.accepted is True deduper_mock.is_processed.assert_awaited_once_with("m-1") async def test_security_middleware_allowed(base_ctx): from yuxi.channel.security.policy import SecurityCheckResult policy = MagicMock() policy.check = AsyncMock(return_value=SecurityCheckResult(allowed=True)) mw = SecurityMiddleware(policy, MagicMock()) result = await mw.process(base_ctx, _noop_next) assert result.accepted is True policy.check.assert_awaited_once_with({}, base_ctx.plugin, base_ctx.inbound) async def test_security_middleware_rate_limited_clears_dedupe(base_ctx, deduper_mock): from yuxi.channel.security.policy import SecurityCheckResult policy = MagicMock() policy.check = AsyncMock( return_value=SecurityCheckResult(allowed=False, reason=InboundRejectionReason.RATE_LIMITED) ) mw = SecurityMiddleware(policy, deduper_mock) result = await mw.process(base_ctx, _noop_next) assert result.accepted is False assert result.reason == InboundRejectionReason.RATE_LIMITED deduper_mock.clear_processed.assert_awaited_once_with("m-1") async def test_security_middleware_dm_pairing(base_ctx): from yuxi.channel.security.policy import SecurityCheckResult policy = MagicMock() policy.check = AsyncMock( return_value=SecurityCheckResult( allowed=False, reason=InboundRejectionReason.DM_PAIRING_REQUIRED, pairing_code="code-1", qr_content="qr-1", ) ) mw = SecurityMiddleware(policy, MagicMock()) result = await mw.process(base_ctx, _noop_next) assert result.accepted is False assert result.reason == InboundRejectionReason.DM_PAIRING_REQUIRED assert result.pairing_code == "code-1" assert result.qr_content == "qr-1" async def test_security_middleware_other_rejection(base_ctx): from yuxi.channel.security.policy import SecurityCheckResult policy = MagicMock() policy.check = AsyncMock( return_value=SecurityCheckResult(allowed=False, reason=InboundRejectionReason.GROUP_NOT_ALLOWED) ) mw = SecurityMiddleware(policy, MagicMock()) result = await mw.process(base_ctx, _noop_next) assert result.accepted is False assert result.reason == InboundRejectionReason.GROUP_NOT_ALLOWED async def test_transaction_middleware_injects_db_and_commits(base_ctx, pg_manager_mock, db_session_mock): mw = TransactionMiddleware(MagicMock()) called = False async def next_mw(): nonlocal called called = True assert base_ctx.db is db_session_mock return InboundResult(accepted=True) result = await mw.process(base_ctx, next_mw) assert called is True assert result.accepted is True async def test_transaction_middleware_rollback_and_clears_dedupe_on_exception( base_ctx, pg_manager_mock, deduper_mock ): mw = TransactionMiddleware(deduper_mock) async def next_mw(): raise RuntimeError("boom") with pytest.raises(RuntimeError, match="boom"): await mw.process(base_ctx, next_mw) deduper_mock.clear_processed.assert_awaited_once_with("m-1") async def test_session_middleware_resolves_and_sets_ctx(base_ctx): session_manager = MagicMock() session = MagicMock() route = BindingRoute(agent_id="agent-1", session_key="sk-1", matched_by="default") session_manager.resolve = AsyncMock(return_value=(session, route)) mw = SessionMiddleware(session_manager) result = await mw.process(base_ctx, _noop_next) assert result.accepted is True assert base_ctx.session is session assert base_ctx.route is route session_manager.resolve.assert_awaited_once_with( base_ctx.config, base_ctx.plugin, base_ctx.inbound, base_ctx.db ) async def test_route_middleware_resolves_when_route_missing(base_ctx): binding_router = MagicMock() base_ctx.session = MagicMock() base_ctx.route = None route = BindingRoute(agent_id="agent-1", session_key="sk-1", matched_by="default") binding_router.resolve = AsyncMock(return_value=route) mw = RouteMiddleware(binding_router) result = await mw.process(base_ctx, _noop_next) assert result.accepted is True assert base_ctx.route is route binding_router.resolve.assert_awaited_once_with( base_ctx.session, base_ctx.config, base_ctx.plugin, base_ctx.inbound, base_ctx.db ) async def test_route_middleware_skips_when_route_present(base_ctx): binding_router = MagicMock() base_ctx.session = MagicMock() route = BindingRoute(agent_id="agent-1", session_key="sk-1", matched_by="default") base_ctx.route = route mw = RouteMiddleware(binding_router) result = await mw.process(base_ctx, _noop_next) assert result.accepted is True assert base_ctx.route is route binding_router.resolve.assert_not_called() async def test_create_run_middleware_creates_run_and_records_message(base_ctx): base_ctx.session = MagicMock() base_ctx.route = BindingRoute(agent_id="agent-1", session_key="sk-1", matched_by="default") create_run = AsyncMock(return_value={"run_id": "run-1"}) record_message = AsyncMock() mw = CreateRunMiddleware(create_run, record_message) result = await mw.process(base_ctx, _noop_next) assert result.accepted is True assert result.run_id == "run-1" assert base_ctx.run_result == {"run_id": "run-1"} create_run.assert_awaited_once_with( base_ctx.session, base_ctx.route, base_ctx.inbound, base_ctx.db ) record_message.assert_awaited_once_with( base_ctx.session, base_ctx.inbound, base_ctx.db )