ForcePilot/backend/test/unit/channel/domain/service/test_pipeline.py
Kris 6c95dc006a test: 新增渠道模块全链路单元测试用例与目录结构
完成渠道模块的单元测试目录搭建,新增多个领域模型、端口、中间件、事件、服务以及基础设施层的单元测试文件,同时补充了conftest.py的环境变量配置,完善测试基础环境。
2026-05-30 21:55:35 +08:00

182 lines
6.1 KiB
Python

from __future__ import annotations
import pytest
from yuxi.channel.domain.service.pipeline import Pipeline
from yuxi.channel.domain.service.message_context import MessageContext, Span
from yuxi.channel.domain.middleware.middleware import Middleware, CallNext
from yuxi.channel.domain.model.message.unified_message import UnifiedMessage
from yuxi.channel.domain.model.message.peer import Peer
from yuxi.channel.domain.model.shared.channel_type import ChannelType
class _FakeMiddleware:
def __init__(self, name: str) -> None:
self._name = name
@property
def name(self) -> str:
return self._name
async def process(self, ctx: MessageContext, call_next: CallNext) -> MessageContext:
return await call_next(ctx)
class _AbortMiddleware:
def __init__(self, name: str) -> None:
self._name = name
@property
def name(self) -> str:
return self._name
async def process(self, ctx: MessageContext, call_next: CallNext) -> MessageContext:
ctx.abort("aborted", "TEST")
return ctx
class _SkipMiddleware:
def __init__(self, name: str) -> None:
self._name = name
@property
def name(self) -> str:
return self._name
async def process(self, ctx: MessageContext, call_next: CallNext) -> MessageContext:
ctx.skip()
return ctx
class _ErrorMiddleware:
def __init__(self, name: str) -> None:
self._name = name
@property
def name(self) -> str:
return self._name
async def process(self, ctx: MessageContext, call_next: CallNext) -> MessageContext:
raise RuntimeError("boom")
class _FakeMetrics:
def __init__(self) -> None:
self.middleware_durations: list[tuple[str, str, float]] = []
self.pipeline_duration: tuple[str, float] | None = None
self.pipeline_total: str | None = None
self.pipeline_aborted: tuple[str, str] | None = None
async def record_middleware_duration(self, middleware: str, channel_type: str, duration_s: float) -> None:
self.middleware_durations.append((middleware, channel_type, duration_s))
async def record_pipeline_duration(self, channel_type: str, duration_s: float) -> None:
self.pipeline_duration = (channel_type, duration_s)
async def record_pipeline_total(self, channel_type: str) -> None:
self.pipeline_total = channel_type
async def record_pipeline_aborted(self, channel_type: str, reason: str) -> None:
self.pipeline_aborted = (channel_type, reason)
@pytest.fixture
def sample_message() -> UnifiedMessage:
return UnifiedMessage(
message_id="msg-1",
channel_type=ChannelType.FEISHU,
sender=Peer(id="user-1", name="Alice"),
content="hello",
)
@pytest.mark.asyncio
async def test_pipeline_empty(sample_message: UnifiedMessage) -> None:
pipeline = Pipeline([])
ctx = MessageContext(message=sample_message)
result = await pipeline.execute(ctx)
assert result.message == sample_message
assert result.spans == []
@pytest.mark.asyncio
async def test_pipeline_single_middleware(sample_message: UnifiedMessage) -> None:
pipeline = Pipeline([_FakeMiddleware("test")])
ctx = MessageContext(message=sample_message)
result = await pipeline.execute(ctx)
assert result.message == sample_message
assert len(result.spans) == 1
assert result.spans[0].middleware_name == "test"
assert result.spans[0].aborted is False
assert result.spans[0].skipped is False
@pytest.mark.asyncio
async def test_pipeline_multiple_middlewares(sample_message: UnifiedMessage) -> None:
pipeline = Pipeline([_FakeMiddleware("m1"), _FakeMiddleware("m2")])
ctx = MessageContext(message=sample_message)
result = await pipeline.execute(ctx)
assert len(result.spans) == 2
assert result.spans[0].middleware_name == "m1"
assert result.spans[1].middleware_name == "m2"
@pytest.mark.asyncio
async def test_pipeline_abort_stops_chain(sample_message: UnifiedMessage) -> None:
pipeline = Pipeline([_AbortMiddleware("abort"), _FakeMiddleware("after")])
ctx = MessageContext(message=sample_message)
result = await pipeline.execute(ctx)
assert result.is_aborted is True
assert len(result.spans) == 1
assert result.spans[0].middleware_name == "abort"
assert result.spans[0].aborted is True
@pytest.mark.asyncio
async def test_pipeline_skip_stops_chain(sample_message: UnifiedMessage) -> None:
pipeline = Pipeline([_SkipMiddleware("skip"), _FakeMiddleware("after")])
ctx = MessageContext(message=sample_message)
result = await pipeline.execute(ctx)
assert result.is_skipped is True
assert len(result.spans) == 1
assert result.spans[0].middleware_name == "skip"
assert result.spans[0].skipped is True
@pytest.mark.asyncio
async def test_pipeline_exception_raises(sample_message: UnifiedMessage) -> None:
pipeline = Pipeline([_ErrorMiddleware("error")])
ctx = MessageContext(message=sample_message)
with pytest.raises(RuntimeError, match="boom"):
await pipeline.execute(ctx)
@pytest.mark.asyncio
async def test_pipeline_with_metrics(sample_message: UnifiedMessage) -> None:
metrics = _FakeMetrics()
pipeline = Pipeline([_FakeMiddleware("m1")], metrics=metrics)
ctx = MessageContext(message=sample_message, channel_type="feishu")
await pipeline.execute(ctx)
assert len(metrics.middleware_durations) == 1
assert metrics.middleware_durations[0][0] == "m1"
assert metrics.pipeline_total == "feishu"
assert metrics.pipeline_duration is not None
assert metrics.pipeline_duration[0] == "feishu"
@pytest.mark.asyncio
async def test_pipeline_aborted_metrics(sample_message: UnifiedMessage) -> None:
metrics = _FakeMetrics()
pipeline = Pipeline([_AbortMiddleware("abort")], metrics=metrics)
ctx = MessageContext(message=sample_message, channel_type="web")
await pipeline.execute(ctx)
assert metrics.pipeline_aborted is not None
assert metrics.pipeline_aborted[0] == "web"
assert metrics.pipeline_aborted[1] == "TEST"
def test_pipeline_middlewares_property() -> None:
m1 = _FakeMiddleware("m1")
pipeline = Pipeline([m1])
assert pipeline.middlewares == [m1]