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, _filter_members_by_listing, _sector_coverage_issue, ) 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): def test_confirmed_delisted_members_are_removed_for_the_target_date(self): members = [ {"ts_code": "601318.SH", "name": "中国平安"}, {"ts_code": "601319.SH", "name": "中国人保"}, {"ts_code": "000627.SZ", "name": "退市成员"}, {"ts_code": "999999.SZ", "name": "状态未知成员"}, ] reference = { "601318.SH": {"list_date": "20070228", "delist_date": ""}, "601319.SH": {"list_date": "20181030", "delist_date": ""}, "000627.SZ": {"list_date": "19961112", "delist_date": "20250930"}, } eligible, excluded = _filter_members_by_listing( members, reference, "20260723" ) self.assertEqual( [item["ts_code"] for item in eligible], ["601318.SH", "601319.SH", "999999.SZ"], ) self.assertEqual(excluded[0]["ts_code"], "000627.SZ") self.assertEqual(excluded[0]["reason"], "目标日期前已退市") def test_member_is_kept_for_dates_before_its_delisting(self): members = [{"ts_code": "000627.SZ", "name": "历史有效成员"}] reference = { "000627.SZ": {"list_date": "19961112", "delist_date": "20250930"} } eligible, excluded = _filter_members_by_listing( members, reference, "20250929" ) self.assertEqual(eligible, members) self.assertEqual(excluded, []) def test_sector_coverage_gate_adapts_to_member_count(self): self.assertEqual(_sector_coverage_issue(5, 5, 100), "") self.assertIn("全部可解释", _sector_coverage_issue(5, 4, 80)) self.assertEqual(_sector_coverage_issue(5, 4, 100, 5), "") self.assertEqual(_sector_coverage_issue(10, 9, 90), "") self.assertIn("至少90%", _sector_coverage_issue(9, 8, 88.9)) self.assertIn("最多缺1只", _sector_coverage_issue(20, 18, 90)) self.assertEqual(_sector_coverage_issue(50, 45, 90), "") self.assertIn("低于90%", _sector_coverage_issue(50, 44, 88)) @patch.object(TushareClient, "query") def test_confirmed_suspension_explains_a_missing_quote(self, query: MagicMock): TushareClient._suspension_cache.clear() query.return_value = [{ "ts_code": "601319.SH", "suspend_date": "20260720", "resume_date": "20260725", "suspend_reason": "重大事项", }] members = [ {"ts_code": "601318.SH", "name": "中国平安"}, {"ts_code": "601319.SH", "name": "中国人保"}, ] suspended = TushareClient("token")._confirmed_suspended_members( members, {"601318.SH"}, "20260723" ) self.assertEqual(len(suspended), 1) self.assertEqual(suspended[0]["ts_code"], "601319.SH") self.assertEqual(suspended[0]["reason"], "重大事项") @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()