189 lines
6.0 KiB
Python
189 lines
6.0 KiB
Python
"""M8 MCP 服务模块单元测试。"""
|
||
|
||
from unittest.mock import MagicMock, patch
|
||
|
||
from vectorstore.models import SearchFilter
|
||
|
||
# --------------------------------------------------------------------------- #
|
||
# _fmt_results
|
||
# --------------------------------------------------------------------------- #
|
||
|
||
|
||
class TestFmtResults:
|
||
"""_fmt_results 格式化测试。"""
|
||
|
||
def test_empty_hits(self):
|
||
from mcp_server.server import _fmt_results
|
||
result = _fmt_results([], "test query")
|
||
assert "未找到" in result
|
||
|
||
def test_with_results(self):
|
||
from mcp_server.server import _fmt_results
|
||
hits = [{
|
||
"title": "Fed Holds Rates",
|
||
"title_zh": "美联储维持利率",
|
||
"source": "reuters",
|
||
"score": 0.95,
|
||
"url": "https://example.com/1",
|
||
"events": [{
|
||
"event_type": "央行决议",
|
||
"sentiment": "neutral",
|
||
"importance": 5,
|
||
"stock_codes": [],
|
||
"summary_zh": "美联储维持利率不变",
|
||
}],
|
||
}]
|
||
result = _fmt_results(hits, "Fed")
|
||
assert "美联储维持利率" in result
|
||
assert "reuters" in result
|
||
assert "0.95" in result
|
||
|
||
def test_with_stock_codes(self):
|
||
from mcp_server.server import _fmt_results
|
||
hits = [{
|
||
"title": "Apple Earnings",
|
||
"title_zh": "苹果财报",
|
||
"source": "reuters",
|
||
"score": 0.88,
|
||
"url": "https://example.com/1",
|
||
"events": [{
|
||
"event_type": "财报披露",
|
||
"sentiment": "positive",
|
||
"importance": 4,
|
||
"stock_codes": ["AAPL"],
|
||
"summary_zh": "苹果财报超预期",
|
||
}],
|
||
}]
|
||
result = _fmt_results(hits, "AAPL")
|
||
assert "AAPL" in result
|
||
assert "🟢利好" in result
|
||
|
||
|
||
# --------------------------------------------------------------------------- #
|
||
# _load_today_events
|
||
# --------------------------------------------------------------------------- #
|
||
|
||
|
||
class TestLoadTodayEvents:
|
||
"""_load_today_events 测试。"""
|
||
|
||
def test_no_data_dir(self):
|
||
from mcp_server.server import _load_today_events
|
||
events = _load_today_events("20990101")
|
||
assert events == []
|
||
|
||
@patch("mcp_server.server.Path.is_dir")
|
||
@patch("mcp_server.server.Path.glob")
|
||
def test_loads_high_importance_only(self, mock_glob, mock_is_dir):
|
||
import json
|
||
|
||
from mcp_server.server import _load_today_events
|
||
|
||
mock_is_dir.return_value = True
|
||
# 创建 mock 文件
|
||
mock_file = MagicMock()
|
||
mock_file.name = "test.json"
|
||
mock_file.read_text.return_value = json.dumps({
|
||
"title": "Test",
|
||
"title_zh": "测试",
|
||
"url": "https://x.com/1",
|
||
"source_id": "reuters",
|
||
"events": [
|
||
{"importance": 5, "event_type": "央行决议", "sentiment": "neutral",
|
||
"stock_codes": [], "summary_zh": "t1"},
|
||
{"importance": 2, "event_type": "其他", "sentiment": "neutral",
|
||
"stock_codes": [], "summary_zh": "t2"},
|
||
],
|
||
})
|
||
mock_glob.return_value = [mock_file]
|
||
|
||
events = _load_today_events("20260621")
|
||
# 只有 importance ≥ 4 的
|
||
assert len(events) == 1
|
||
assert events[0]["importance"] == 5
|
||
|
||
|
||
# --------------------------------------------------------------------------- #
|
||
# SearchFilter
|
||
# --------------------------------------------------------------------------- #
|
||
|
||
|
||
class TestSearchFilterMCP:
|
||
"""SearchFilter 用于 MCP 的测试。"""
|
||
|
||
def test_stock_filter(self):
|
||
f = SearchFilter(stock_codes=["AAPL"])
|
||
assert "AAPL" in f.stock_codes
|
||
|
||
def test_sentiment_filter(self):
|
||
f = SearchFilter(sentiment="positive")
|
||
assert f.sentiment == "positive"
|
||
|
||
|
||
# --------------------------------------------------------------------------- #
|
||
# _search(mock)
|
||
# --------------------------------------------------------------------------- #
|
||
|
||
|
||
class TestSearch:
|
||
"""_search 函数测试(mock backend)。"""
|
||
|
||
@patch("mcp_server.server._get_backend")
|
||
def test_search_returns_formatted(self, mock_backend):
|
||
from mcp_server.server import _search
|
||
|
||
# Mock backend
|
||
be = MagicMock()
|
||
mock_backend.return_value = be
|
||
|
||
# Mock embed_batch
|
||
with patch("mcp_server.server.embed_batch") as mock_embed:
|
||
mock_embed.return_value = [[0.1] * 1024]
|
||
|
||
# Mock vector_store.query
|
||
mock_result = MagicMock()
|
||
mock_result.title = "Test"
|
||
mock_result.title_zh = "测试"
|
||
mock_result.url = "https://x.com/1"
|
||
mock_result.source_id = "reuters"
|
||
mock_result.score = 0.9
|
||
mock_result.publish_time = "2026-06-21"
|
||
mock_result.events = []
|
||
mock_result.content_zh_preview = ""
|
||
be.vector_store.query.return_value = [mock_result]
|
||
|
||
hits = _search("test")
|
||
assert len(hits) == 1
|
||
assert hits[0]["title"] == "Test"
|
||
|
||
|
||
# --------------------------------------------------------------------------- #
|
||
# MCP 模块导入
|
||
# --------------------------------------------------------------------------- #
|
||
|
||
|
||
class TestMCPImport:
|
||
"""MCP 服务模块导入测试。"""
|
||
|
||
def test_mcp_object_imports(self):
|
||
"""确保 mcp FastMCP 对象可导入。"""
|
||
from mcp_server.server import mcp
|
||
assert mcp is not None
|
||
assert mcp.name is not None
|
||
|
||
def test_tools_registered(self):
|
||
"""确保 5 个工具函数存在且已注册。"""
|
||
from mcp_server.server import (
|
||
get_stats,
|
||
get_today_events,
|
||
search_by_sentiment,
|
||
search_by_stock,
|
||
search_news,
|
||
)
|
||
# 验证函数存在且可调用
|
||
assert callable(search_news)
|
||
assert callable(search_by_stock)
|
||
assert callable(search_by_sentiment)
|
||
assert callable(get_today_events)
|
||
assert callable(get_stats)
|