Files
wecom_it_smart_desk/backend/tests/test_p2_p3.py
T

732 lines
27 KiB
Python
Raw Normal View History

# =============================================================================
# 企微IT智能服务台 — 复杂场景重构第二阶段 P2/P3 单元测试
# =============================================================================
# 测试范围:TokenCounter、ContextCompressor、SnapshotService、CorrectionService
# 创建日期: 2026-07-11
# =============================================================================
import os
import sys
# 设置测试环境变量(在导入 app 之前)
os.environ.setdefault("DEV_MODE", "true")
os.environ.setdefault("WECOM_SSO_CALLBACK_BASE", "https://test.example.com")
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from datetime import datetime, timezone
# =============================================================================
# TokenCounter 测试
# =============================================================================
class TestTokenCounter:
"""TokenCounter 工具类测试。"""
def test_count_tokens_empty(self):
"""空文本/None 返回 0。"""
from app.utils.token_counter import TokenCounter
assert TokenCounter.count_tokens("") == 0
assert TokenCounter.count_tokens(None) == 0
def test_count_tokens_non_empty(self):
"""非空文本返回正数。"""
from app.utils.token_counter import TokenCounter
assert TokenCounter.count_tokens("Hello world") > 0
assert TokenCounter.count_tokens("你好世界") > 0
def test_count_tokens_long_text(self):
"""长文本 token 数应大于短文本。"""
from app.utils.token_counter import TokenCounter
short = TokenCounter.count_tokens("hi")
long = TokenCounter.count_tokens("This is a much longer piece of text " * 100)
assert long > short
def test_count_messages_tokens(self):
"""消息列表 token 计数包含每条4 token overhead。"""
from app.utils.token_counter import TokenCounter
messages = [
{"role": "user", "content": "你好"},
{"role": "assistant", "content": "您好,有什么可以帮您?"},
]
total = TokenCounter.count_messages_tokens(messages)
assert total > 0
# 应该比单独文本token之和大(因为每条有4 token overhead
text_only = sum(
TokenCounter.count_tokens(m["content"]) + TokenCounter.count_tokens(m["role"])
for m in messages
)
assert total == text_only + 8 # 2条消息 x 4 overhead
def test_count_messages_tokens_empty(self):
"""空消息列表返回 0。"""
from app.utils.token_counter import TokenCounter
assert TokenCounter.count_messages_tokens([]) == 0
def test_is_precise_returns_bool(self):
"""is_precise 返回布尔值。"""
from app.utils.token_counter import TokenCounter
assert isinstance(TokenCounter.is_precise(), bool)
def test_count_tokens_consistency(self):
"""相同输入应返回相同结果(幂等性)。"""
from app.utils.token_counter import TokenCounter
text = "这是一段测试文本 for consistency check"
result1 = TokenCounter.count_tokens(text)
result2 = TokenCounter.count_tokens(text)
assert result1 == result2
# =============================================================================
# SnapshotService 测试
# =============================================================================
class TestSnapshotService:
"""SnapshotService 快照服务测试。"""
def _make_mock_item(self, name="工号", value="12345", version=1, item_id="item-1"):
"""创建 mock 信息项。"""
mock_item = MagicMock()
mock_item.name = name
mock_item.value = value
mock_item.version = version
mock_item.id = item_id
return mock_item
@pytest.mark.asyncio
async def test_create_snapshot(self):
"""测试创建快照。"""
from app.services.automation.snapshot_service import SnapshotService
db = AsyncMock()
mock_item = self._make_mock_item()
mock_result = MagicMock()
mock_result.scalars.return_value.all.return_value = [mock_item]
db.execute.return_value = mock_result
svc = SnapshotService(db)
snapshot = await svc.create_snapshot("session-1", "工号", [])
assert snapshot is not None
assert snapshot.trigger_item_key == "工号"
assert snapshot.session_id == "session-1"
assert snapshot.is_undone is False
db.add.assert_called_once()
db.flush.assert_called_once()
@pytest.mark.asyncio
async def test_create_snapshot_multiple_items(self):
"""测试多信息项快照。"""
from app.services.automation.snapshot_service import SnapshotService
db = AsyncMock()
item1 = self._make_mock_item("工号", "10001", 1, "id-1")
item2 = self._make_mock_item("姓名", "张三", 2, "id-2")
mock_result = MagicMock()
mock_result.scalars.return_value.all.return_value = [item1, item2]
db.execute.return_value = mock_result
svc = SnapshotService(db)
snapshot = await svc.create_snapshot("session-1", "工号", [])
assert snapshot is not None
assert snapshot.trigger_item_key == "工号"
@pytest.mark.asyncio
async def test_undo_limit_exceeded(self):
"""测试撤销次数超限。"""
from app.services.automation.snapshot_service import (
SnapshotService,
MAX_UNDO_COUNT,
)
db = AsyncMock()
# Mock: 已撤销次数 >= MAX_UNDO_COUNT
mock_count_result = MagicMock()
mock_count_result.scalar.return_value = MAX_UNDO_COUNT
db.execute.return_value = mock_count_result
svc = SnapshotService(db)
with pytest.raises(ValueError, match="超限"):
await svc.undo_correction("session-1")
@pytest.mark.asyncio
async def test_undo_no_snapshot(self):
"""测试无可撤销快照。"""
from app.services.automation.snapshot_service import SnapshotService
db = AsyncMock()
# Mock: 已撤销次数 = 0 (未超限)
mock_count_result = MagicMock()
mock_count_result.scalar.return_value = 0
# Mock: 无未撤销快照
mock_snapshot_result = MagicMock()
mock_snapshot_result.scalar_one_or_none.return_value = None
# db.execute 需要返回不同结果(先 count,后 snapshot
db.execute.side_effect = [mock_count_result, mock_snapshot_result]
svc = SnapshotService(db)
with pytest.raises(ValueError, match="无可撤销"):
await svc.undo_correction("session-1")
@pytest.mark.asyncio
async def test_get_snapshot_history_empty(self):
"""测试空快照历史。"""
from app.services.automation.snapshot_service import SnapshotService
db = AsyncMock()
mock_result = MagicMock()
mock_result.scalars.return_value.all.return_value = []
db.execute.return_value = mock_result
svc = SnapshotService(db)
history = await svc.get_snapshot_history("session-1")
assert history == []
@pytest.mark.asyncio
async def test_get_version_diff_item_not_found(self):
"""测试版本对比时信息项不存在。"""
from app.services.automation.snapshot_service import SnapshotService
db = AsyncMock()
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = None
db.execute.return_value = mock_result
svc = SnapshotService(db)
with pytest.raises(ValueError, match="不存在"):
await svc.get_version_diff("session-1", "不存在项", 1, 2)
@pytest.mark.asyncio
async def test_get_latest_snapshot_none(self):
"""测试无最新快照。"""
from app.services.automation.snapshot_service import SnapshotService
db = AsyncMock()
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = None
db.execute.return_value = mock_result
svc = SnapshotService(db)
snapshot = await svc.get_latest_snapshot("session-1")
assert snapshot is None
# =============================================================================
# CorrectionService 测试
# =============================================================================
class TestCorrectionService:
"""CorrectionService 纠错服务测试。"""
@pytest.mark.asyncio
async def test_batch_correct_empty_corrections(self):
"""测试空更正列表抛异常。"""
from app.services.automation.correction_service import CorrectionService
db = AsyncMock()
svc = CorrectionService(db)
with pytest.raises(Exception):
await svc.batch_correct("session-1", [])
@pytest.mark.asyncio
async def test_check_dependencies(self):
"""测试依赖检查。"""
from app.services.automation.correction_service import CorrectionService
db = AsyncMock()
# Mock: 有一个依赖信息项
mock_item = MagicMock()
mock_item.name = "设备分配人"
mock_item.value = "张三"
mock_item.derived_from = ["工号"]
mock_item.session_id = "session-1"
mock_result = MagicMock()
mock_result.scalars.return_value.all.return_value = [mock_item]
db.execute.return_value = mock_result
svc = CorrectionService(db)
warnings = await svc.check_dependencies("session-1", "工号")
assert len(warnings) == 1
assert warnings[0]["item_key"] == "设备分配人"
assert warnings[0]["derived_from"] == "工号"
@pytest.mark.asyncio
async def test_check_dependencies_no_deps(self):
"""测试无依赖项时返回空列表。"""
from app.services.automation.correction_service import CorrectionService
db = AsyncMock()
mock_item = MagicMock()
mock_item.name = "工号"
mock_item.value = "12345"
mock_item.derived_from = None
mock_result = MagicMock()
mock_result.scalars.return_value.all.return_value = [mock_item]
db.execute.return_value = mock_result
svc = CorrectionService(db)
warnings = await svc.check_dependencies("session-1", "工号")
assert warnings == []
@pytest.mark.asyncio
async def test_get_version_chain_empty(self):
"""测试空版本链(信息项不存在)。"""
from app.services.automation.correction_service import CorrectionService
db = AsyncMock()
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = None
db.execute.return_value = mock_result
svc = CorrectionService(db)
chain = await svc.get_version_chain("session-1", "不存在项")
assert chain == []
@pytest.mark.asyncio
async def test_get_version_chain_with_history(self):
"""测试带历史记录的版本链。"""
from app.services.automation.correction_service import CorrectionService
db = AsyncMock()
mock_item = MagicMock()
mock_item.name = "工号"
mock_item.value = "99999"
mock_item.version = 3
mock_item.derived_from = None
mock_item.correction_reason = "用户更正"
mock_item.updated_at = datetime(2026, 7, 11, 10, 0, 0, tzinfo=timezone.utc)
mock_item.update_history = [
{
"version": 1,
"old_value": "11111",
"new_value": "22222",
"action": "correct",
"reason": "初次更正",
"timestamp": "2026-07-11T09:00:00+00:00",
},
{
"version": 2,
"old_value": "22222",
"new_value": "99999",
"action": "correct",
"reason": "再次更正",
"timestamp": "2026-07-11T09:30:00+00:00",
},
]
mock_result = MagicMock()
mock_result.scalar_one_or_none.return_value = mock_item
db.execute.return_value = mock_result
svc = CorrectionService(db)
chain = await svc.get_version_chain("session-1", "工号")
# 2 history entries + 1 current = 3
assert len(chain) == 3
# First entry
assert chain[0]["version"] == 1
assert chain[0]["value"] == "11111"
assert chain[0]["new_value"] == "22222"
assert chain[0]["action"] == "correct"
# Last entry (current)
assert chain[-1]["version"] == 3
assert chain[-1]["value"] == "99999"
assert chain[-1]["action"] == "current"
@pytest.mark.asyncio
async def test_get_correction_history_empty(self):
"""测试空更正历史。"""
from app.services.automation.correction_service import CorrectionService
db = AsyncMock()
mock_result = MagicMock()
mock_result.scalars.return_value.all.return_value = []
db.execute.return_value = mock_result
svc = CorrectionService(db)
history = await svc.get_correction_history("session-1")
assert history == []
# =============================================================================
# ContextCompressor 测试
# =============================================================================
class TestContextCompressor:
"""ContextCompressor 上下文压缩引擎测试。"""
def test_should_compress_under_threshold(self):
"""测试未超阈值不压缩。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
# 少量消息不应触发压缩(阈值默认6000)
messages = [{"role": "user", "content": "你好"}]
assert compressor.should_compress(messages) is False
def test_should_compress_over_threshold(self):
"""测试超阈值触发压缩。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
compressor.threshold = 10 # 设置极低阈值
messages = [{"role": "user", "content": "这是一段很长的对话内容" * 10}]
assert compressor.should_compress(messages) is True
def test_should_compress_empty_messages(self):
"""测试空消息列表不压缩。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
assert compressor.should_compress([]) is False
def test_count_tokens_method(self):
"""测试 count_tokens 方法。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
messages = [{"role": "user", "content": "Hello"}]
assert compressor.count_tokens(messages) > 0
def test_count_tokens_empty(self):
"""测试空消息 token 数为 0。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
assert compressor.count_tokens([]) == 0
def test_extract_key_info_empty(self):
"""测试空信息项提取。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
result = compressor._extract_key_info([], [], "")
assert "暂无" in result["info_items"]
assert "暂无" in result["actions"]
def test_extract_key_info_with_items(self):
"""测试有信息项时正确提取。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
mock_item = MagicMock()
mock_item.name = "工号"
mock_item.value = "12345"
mock_item.version = 2
mock_item.is_filled = True
mock_action = MagicMock()
mock_action.title = "终端扫描"
mock_action.action_type = "virus_scan"
mock_action.status = "success"
result = compressor._extract_key_info([mock_item], [mock_action], "扫描中")
assert "工号" in result["info_items"]
assert "12345" in result["info_items"]
assert "v2" in result["info_items"]
assert "终端扫描" in result["actions"]
assert "" in result["actions"]
def test_extract_key_info_unfilled_item(self):
"""测试未填写信息项被过滤。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
mock_item = MagicMock()
mock_item.name = "工号"
mock_item.value = "12345"
mock_item.version = 1
mock_item.is_filled = False # 未填写
result = compressor._extract_key_info([mock_item], [], "")
# 未填写项应被过滤,显示"暂无"
assert "暂无" in result["info_items"]
def test_get_recent_messages(self):
"""测试获取最近N轮对话。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
messages = [
{"role": "user", "content": "msg1"},
{"role": "assistant", "content": "reply1"},
{"role": "user", "content": "msg2"},
{"role": "assistant", "content": "reply2"},
{"role": "user", "content": "msg3"},
{"role": "assistant", "content": "reply3"},
]
# 取最近2轮 = 4条消息
recent = compressor._get_recent_messages(messages, 2)
assert len(recent) == 4
assert recent[-1]["content"] == "reply3"
def test_get_recent_messages_empty(self):
"""测试空消息列表。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
recent = compressor._get_recent_messages([], 2)
assert recent == []
def test_get_recent_messages_more_than_available(self):
"""测试请求轮数超过实际轮数。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
messages = [
{"role": "user", "content": "msg1"},
{"role": "assistant", "content": "reply1"},
]
# 请求5轮但只有1轮 = 2条消息
recent = compressor._get_recent_messages(messages, 5)
assert len(recent) == 2
def test_format_recent_messages_empty(self):
"""测试格式化空消息。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
result = compressor._format_recent_messages([])
assert "" in result
def test_format_recent_messages_non_empty(self):
"""测试格式化非空消息。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
messages = [
{"role": "user", "content": "你好"},
{"role": "assistant", "content": "您好"},
]
result = compressor._format_recent_messages(messages)
assert "user" in result
assert "assistant" in result
assert "你好" in result
def test_truncate_messages(self):
"""测试降级截断方法。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
messages = [
{"role": "user", "content": "msg1"},
{"role": "assistant", "content": "reply1"},
{"role": "user", "content": "msg2"},
{"role": "assistant", "content": "reply2"},
]
result = compressor._truncate_messages(messages, [], [], "")
assert len(result) >= 1
assert result[0]["role"] == "system"
assert "降级" in result[0]["content"]
@pytest.mark.asyncio
async def test_compress_with_no_llm(self):
"""测试无LLM客户端时的压缩流程。"""
from app.services.automation.context_compressor import ContextCompressor
db = AsyncMock()
compressor = ContextCompressor(db)
compressor.threshold = 5 # 极低阈值触发压缩
messages = []
for i in range(10):
messages.append({"role": "user", "content": f"消息内容{i}"})
messages.append({"role": "assistant", "content": f"回复内容{i}"})
result = await compressor.compress("session-1", messages, [], [])
assert "compressed_messages" in result
assert "tokens_before" in result
assert "tokens_after" in result
assert "compression_ratio" in result
assert "compression_level" in result
assert "summary" in result
assert "duration_ms" in result
# 压缩后消息应少于原消息
assert len(result["compressed_messages"]) <= len(messages)
# =============================================================================
# 模型层测试
# =============================================================================
class TestP2P3Models:
"""P2/P3 数据模型测试。"""
def test_context_compression_model_fields(self):
"""测试 ContextCompression 模型字段。"""
from app.models.automation import ContextCompression
# 验证表名
assert ContextCompression.__tablename__ == "auto_context_compressions"
# 验证列存在
columns = ContextCompression.__table__.columns
col_names = {c.name for c in columns}
expected = {
"id", "session_id", "tokens_before", "tokens_after",
"compression_ratio", "task_node", "duration_ms",
"compression_level", "summary", "created_at",
}
assert expected.issubset(col_names)
def test_information_snapshot_model_fields(self):
"""测试 InformationSnapshot 模型字段。"""
from app.models.automation import InformationSnapshot
assert InformationSnapshot.__tablename__ == "auto_information_snapshots"
columns = InformationSnapshot.__table__.columns
col_names = {c.name for c in columns}
expected = {
"id", "session_id", "trigger_item_key", "snapshot_data",
"correction_ids", "is_undone", "created_at",
}
assert expected.issubset(col_names)
def test_information_item_p2p3_fields(self):
"""测试 InformationItem 模型新增 P2/P3 字段。"""
from app.models.automation import InformationItem
columns = InformationItem.__table__.columns
col_names = {c.name for c in columns}
assert "derived_from" in col_names
assert "correction_reason" in col_names
# =============================================================================
# Schema 测试
# =============================================================================
class TestP2P3Schemas:
"""P2/P3 Schema 测试。"""
def test_batch_correct_request_schema(self):
"""测试批量更正请求 Schema。"""
from app.schemas.automation import BatchCorrectRequest
req = BatchCorrectRequest(
corrections=[{"field": "工号", "new_value": "99999"}],
reason="用户更正",
)
assert len(req.corrections) == 1
assert req.reason == "用户更正"
def test_batch_correct_response_schema(self):
"""测试批量更正响应 Schema。"""
from app.schemas.automation import BatchCorrectResponse
resp = BatchCorrectResponse(
corrected_items=[{"name": "工号", "value": "99999", "version": 2}],
snapshot_id=1,
dependency_warnings=[],
)
assert resp.snapshot_id == 1
assert len(resp.corrected_items) == 1
def test_undo_correction_response_schema(self):
"""测试撤销更正响应 Schema。"""
from app.schemas.automation import UndoCorrectionResponse
resp = UndoCorrectionResponse(
undone_items=["工号"],
restored_values={"工号": "12345"},
snapshot_id=1,
remaining_undo_count=4,
)
assert resp.remaining_undo_count == 4
assert "工号" in resp.undone_items
def test_version_diff_request_schema(self):
"""测试版本对比请求 Schema。"""
from app.schemas.automation import VersionDiffRequest
req = VersionDiffRequest(v1=1, v2=2)
assert req.v1 == 1
assert req.v2 == 2
def test_version_diff_response_schema(self):
"""测试版本对比响应 Schema。"""
from app.schemas.automation import VersionDiffResponse
resp = VersionDiffResponse(
item_key="工号",
v1=1,
v1_value="11111",
v2=2,
v2_value="22222",
changed=True,
)
assert resp.changed is True
assert resp.v1_value == "11111"
def test_compression_log_item_schema(self):
"""测试压缩日志项 Schema。"""
from app.schemas.automation import CompressionLogItem
log = CompressionLogItem(
id=1,
session_id="session-1",
tokens_before=8000,
tokens_after=3000,
compression_ratio=0.38,
task_node="扫描中",
duration_ms=150,
compression_level=2,
)
assert log.tokens_before == 8000
assert log.compression_level == 2
def test_compression_log_list_response_schema(self):
"""测试压缩日志列表响应 Schema。"""
from app.schemas.automation import CompressionLogListResponse
resp = CompressionLogListResponse(logs=[], total=0)
assert resp.total == 0
assert resp.logs == []
# =============================================================================
# 常量测试
# =============================================================================
class TestP2P3Constants:
"""P2/P3 常量测试。"""
def test_error_codes_exist(self):
"""测试新增错误码存在。"""
from app.constants import AutomationErrorCode
assert AutomationErrorCode.COMPRESSION_FAILED == 4017
assert AutomationErrorCode.UNDO_LIMIT_EXCEEDED == 4018
assert AutomationErrorCode.BATCH_CORRECT_FAILED == 4019
def test_error_messages_exist(self):
"""测试新增错误消息存在。"""
from app.constants import AUTOMATION_ERROR_MESSAGES, AutomationErrorCode
assert AutomationErrorCode.COMPRESSION_FAILED in AUTOMATION_ERROR_MESSAGES
assert AutomationErrorCode.UNDO_LIMIT_EXCEEDED in AUTOMATION_ERROR_MESSAGES
assert AutomationErrorCode.BATCH_CORRECT_FAILED in AUTOMATION_ERROR_MESSAGES
def test_error_message_content(self):
"""测试错误消息内容。"""
from app.constants import automation_error_message, AutomationErrorCode
msg = automation_error_message(AutomationErrorCode.UNDO_LIMIT_EXCEEDED)
assert "超限" in msg
assert "撤销" in msg