Files
xiaobaifupan/tests/test_heaven_realtime.py
T

310 lines
12 KiB
Python

from __future__ import annotations
import http.client
import json
import unittest
from unittest.mock import MagicMock, patch
from heaven_engine import _market_line_scores, build_manual_market_hexagram
from realtime_aggregator import WebRealtimeAggregator
from server import DashboardService
from tushare_client import TushareClient
class HeavenMarketLineTests(unittest.TestCase):
def test_index_external_uses_only_current_index_change(self):
dashboard = {
"overview": {
"up_count": 1740,
"down_count": 3710,
"amount_billion": 27180.8,
"sentiment_score": 27,
"seal_rate": 57.3,
"limit_up_count": 55,
"limit_down_count": 267,
},
"sectors": [],
"sector_rotation": [],
}
index_context = {
"aggregate": {
"average_pct_chg": 0.187,
"average_return_5d": -5.606,
}
}
scores = _market_line_scores(
dashboard,
[{"amount_billion": 25000}],
index_context,
{},
{},
[],
)
self.assertAlmostEqual(scores[5]["score"], 0.187 / 3)
self.assertGreater(scores[5]["score"], 0)
self.assertIn("不参与外显阴阳", scores[5]["evidence"][1])
def test_manual_calibration_preserves_six_lines_and_marks_user_evidence(self):
chart = build_manual_market_hexagram(
[8, 7, 6, 9, 8, 7],
"20260722",
{"name": "元件", "taxonomy": "sw_l2"},
{"code": "002141", "name": "贤丰控股"},
{},
"人工核对",
)
self.assertTrue(chart["manual_calibration"])
self.assertEqual([line["value"] for line in chart["hexagram"]["lines"]], [8, 7, 6, 9, 8, 7])
self.assertEqual(chart["hexagram"]["moving_lines"], [3, 4])
self.assertIn("用户手动校准", chart["hexagram"]["lines"][0]["evidence"][0])
def test_quantitative_supplement_repairs_only_failed_lines_and_recalculates(self):
trade_date = "20260722"
dashboard = {
"overview": {
"sentiment_score": 32,
"seal_rate": 48,
"amount_billion": 16500,
"up_count": 1800,
"down_count": 3500,
"limit_up_count": 35,
"limit_down_count": 192,
},
"limits": [{"amount_billion": 12}, {"amount_billion": 25}],
"sectors": [],
"sector_rotation": [],
}
history = [{"amount_billion": 16000}, {"amount_billion": 15800}]
index_context = {
"trade_date": trade_date,
"source": "tushare",
"realtime": False,
"precise": False,
"indices": [
{"ts_code": code, "trade_date": "20260721", "pct_chg": 0.1}
for code in ("000001.SH", "399001.SZ", "399006.SZ")
],
}
sector = {"taxonomy": "sw_l2", "precise": False, "error": "行业日线尚未返回"}
stock = {
"code": "002141", "name": "贤丰控股", "trade_date": trade_date,
"data_source": "tushare", "realtime": False, "precise": True,
"amount_billion": 20, "turnover_rate": 8.5, "seal_amount_million": 0,
"open_times": 0, "change": 2.4, "streak": 0, "status": "普通",
}
automatic = DashboardService._heaven_line_checks(
trade_date, dashboard, history, index_context, sector, stock, "closed", {}
)
self.assertEqual(
[item["line"] for item in automatic if not item["passed"]], [3, 4, 6]
)
manual = {
"sector_name": "元件", "sector_up_count": 18, "sector_down_count": 42,
"sector_coverage": 96, "sector_member_equal_change": -2.2,
"sector_change": -2.6, "sector_leading_pct": 3.1,
"index_sh_change": -0.9, "index_sz_change": -1.4, "index_cy_change": -1.8,
}
merged = DashboardService._apply_heaven_manual_data(
dashboard, index_context, sector, stock, manual, "closed", trade_date, "002141"
)
repaired = DashboardService._heaven_line_checks(
trade_date, merged[0], history, merged[1], merged[2], merged[3], "closed", manual
)
self.assertTrue(all(item["passed"] for item in repaired))
self.assertEqual(
[item["line"] for item in repaired if item["status"] == "manual"], [3, 4, 6]
)
self.assertEqual(repaired[5]["line_value"], 8)
self.assertAlmostEqual(repaired[5]["score"], (-0.9 - 1.4 - 1.8) / 3 / 3, places=3)
def test_sector_inner_and_outer_have_independent_quality_gates(self):
trade_date = "20260722"
dashboard = {
"overview": {
"sentiment_score": 30, "seal_rate": 50, "amount_billion": 15000,
"up_count": 2000, "down_count": 3000,
"limit_up_count": 40, "limit_down_count": 80,
},
"limits": [], "sectors": [], "sector_rotation": [],
}
history = [{"amount_billion": 14800}, {"amount_billion": 14900}]
indices = {
"trade_date": trade_date, "source": "tushare", "realtime": False,
"precise": True,
"indices": [
{"ts_code": code, "trade_date": trade_date, "pct_chg": -1}
for code in ("000001.SH", "399001.SZ", "399006.SZ")
],
"aggregate": {"average_pct_chg": -1},
}
sector = {
"name": "元件", "code": "801083.SI", "taxonomy": "sw_l2",
"trade_date": trade_date, "realtime": False, "finalized": True,
"inner_precise": True, "outer_precise": False, "precise": False,
"coverage": 98, "up_count": 19, "down_count": 46,
"member_equal_change": -2.15, "leading_pct": 5.2,
"outer_error": "申万日线尚未发布", "source": "tushare_member_daily",
}
stock = {
"code": "002141", "trade_date": trade_date, "precise": True,
"realtime": False, "data_source": "tushare", "amount_billion": 10,
"turnover_rate": 5, "seal_amount_million": 0, "open_times": 0,
"change": 2, "streak": 0, "status": "普通",
}
checks = DashboardService._heaven_line_checks(
trade_date, dashboard, history, indices, sector, stock, "closed", {}
)
self.assertTrue(checks[2]["passed"])
self.assertFalse(checks[3]["passed"])
self.assertIn("申万日线尚未发布", checks[3]["reasons"])
manual = {"sector_change": -3.85}
merged = DashboardService._apply_heaven_manual_data(
dashboard, indices, sector, stock, manual, "closed", trade_date, "002141"
)
repaired = DashboardService._heaven_line_checks(
trade_date, merged[0], history, merged[1], merged[2], merged[3], "closed", manual
)
self.assertTrue(repaired[3]["passed"])
self.assertEqual(repaired[3]["status"], "manual")
self.assertEqual(
[field["key"] for field in repaired[3]["fields"] if field["manual"]],
["sector_change"],
)
missing_leader_sector = {**merged[2], "leading_pct": None}
still_blocked = DashboardService._heaven_line_checks(
trade_date,
merged[0],
history,
merged[1],
missing_leader_sector,
merged[3],
"closed",
manual,
)
self.assertFalse(still_blocked[3]["passed"])
self.assertIn("需补充:行业领涨股涨跌幅", still_blocked[3]["reasons"])
class ShenwanMembershipTests(unittest.TestCase):
@patch.object(TushareClient, "query")
def test_latest_effective_membership_wins_over_stale_is_new_row(self, query: MagicMock):
stale_y = {
"l1_code": "801010.SI", "l1_name": "农林牧渔",
"l2_code": "801018.SI", "l2_name": "动物保健Ⅱ",
"l3_code": "850181.SI", "l3_name": "动物保健Ⅲ",
"ts_code": "002141.SZ", "in_date": "20240730", "out_date": None, "is_new": "Y",
}
current_y = {
"l1_code": "801080.SI", "l1_name": "电子",
"l2_code": "801083.SI", "l2_name": "元件",
"l3_code": "850822.SI", "l3_name": "印制电路板",
"ts_code": "002141.SZ", "in_date": "20260701", "out_date": None, "is_new": "Y",
}
closed_n = {**stale_y, "out_date": "20260630", "is_new": "N"}
query.side_effect = lambda _api, params, _fields: (
[stale_y, current_y] if params["is_new"] == "Y" else [closed_n]
)
industry = TushareClient("token").sw_stock_industry("002141.SZ", "20260722")
self.assertEqual(industry["l2_code"], "801083.SI")
self.assertEqual(industry["l2_name"], "元件")
class RealtimeAggregatorTests(unittest.TestCase):
def setUp(self):
WebRealtimeAggregator._response_cache.clear()
@staticmethod
def _response(payload: dict) -> MagicMock:
response = MagicMock()
response.headers.get.return_value = "application/json"
response.read.return_value = json.dumps(payload).encode("utf-8")
context = MagicMock()
context.__enter__.return_value = response
return context
@patch("realtime_aggregator.urllib.request.urlopen")
def test_transport_failure_is_retried(self, urlopen: MagicMock):
urlopen.side_effect = [
http.client.RemoteDisconnected("temporary disconnect"),
self._response({"rc": 0, "data": {"diff": []}}),
]
aggregator = WebRealtimeAggregator(retry_delay_seconds=0)
payload = aggregator._get_json("https://example.test", {}, "https://example.test")
self.assertEqual(payload["rc"], 0)
self.assertEqual(urlopen.call_count, 2)
@patch("realtime_aggregator.urllib.request.urlopen")
def test_recent_success_is_used_after_retries_fail(self, urlopen: MagicMock):
aggregator = WebRealtimeAggregator(retry_delay_seconds=0)
urlopen.return_value = self._response({"rc": 0, "data": {"diff": []}})
aggregator._get_json("https://example.test", {}, "https://example.test")
urlopen.side_effect = http.client.RemoteDisconnected("temporary disconnect")
payload = aggregator._get_json("https://example.test", {}, "https://example.test")
self.assertIn("_aggregate_cache", payload)
self.assertEqual(urlopen.call_count, 4)
@patch.object(WebRealtimeAggregator, "_get_text")
def test_tencent_indices_include_verifiable_quote_times(self, get_text: MagicMock):
def quote_line(
symbol: str,
name: str,
code: str,
price: str,
previous_close: str,
quote_time: str,
change_amount: str,
change: str,
amount: str,
) -> str:
fields = [""] * 38
fields[1] = name
fields[2] = code
fields[3] = price
fields[4] = previous_close
fields[5] = price
fields[30] = quote_time
fields[31] = change_amount
fields[32] = change
fields[33] = price
fields[34] = price
fields[37] = amount
return f'v_{symbol}="{"~".join(fields)}";'
get_text.return_value = (
'\n'.join(
[
quote_line("sh000001", "上证指数", "000001", "3796.28", "3764.15", "20260720155402", "32.13", "0.85", "129465190"),
quote_line("sz399001", "深证成指", "399001", "13610.23", "13706.88", "20260720155330", "-96.65", "-0.71", "140747525"),
quote_line("sz399006", "创业板指", "399006", "3443.10", "3428.63", "20260720155345", "14.47", "0.42", "67120487"),
]
),
0,
)
rows = WebRealtimeAggregator().tencent_indices()
self.assertEqual(len(rows), 3)
self.assertEqual(rows[0]["source"], "tencent_qt")
self.assertEqual(rows[0]["quote_time"][:10], "2026-07-20")
self.assertAlmostEqual(rows[0]["amount_billion"], 12946.52)
if __name__ == "__main__":
unittest.main()