245 lines
7.7 KiB
Python
245 lines
7.7 KiB
Python
"""默认入站中间件单元测试。"""
|
|
|
|
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
|
|
)
|