diff --git a/server.py b/server.py index 1b1f281..e593cdc 100644 --- a/server.py +++ b/server.py @@ -68,7 +68,7 @@ from sentiment_engine import ( ) from strategy_tracking import StrategyTrackingService from trade_journal import TradeJournalService -from tushare_client import TushareClient, TushareError +from tushare_client import TushareClient, TushareError, _sector_coverage_issue LEGACY_SECRET_KEYS = { @@ -1729,6 +1729,22 @@ class DashboardService: sector_date = str(sector.get("trade_date") or "").replace("-", "") sector_coverage = float(sector.get("coverage") or 0) + sector_explained_count = int( + sector.get("explained_count") + if sector.get("explained_count") is not None + else sector.get("quote_count") or 0 + ) + sector_explained_coverage = float( + sector.get("explained_coverage") + if sector.get("explained_coverage") is not None + else sector_coverage + ) + sector_coverage_issue = _sector_coverage_issue( + int(sector.get("member_count") or 0), + int(sector.get("quote_count") or 0), + sector_explained_coverage, + sector_explained_count, + ) sector_common = [] if not sector: sector_common.append("未取得申万二级行业归属") @@ -1748,8 +1764,8 @@ class DashboardService: sector_inner.append(str(sector.get("inner_error") or sector.get("error") or "行业内核数据未通过校验")) if not sector.get("outer_precise", sector.get("precise")): sector_outer.append(str(sector.get("outer_error") or sector.get("error") or "行业外显数据未通过校验")) - if sector and sector_coverage < 90: - sector_inner.append(f"行业成分行情覆盖率仅 {sector_coverage:.1f}%,低于 90%") + if sector and sector_coverage_issue and sector_coverage_issue not in sector_inner: + sector_inner.append(sector_coverage_issue) if sector.get("realtime") and not sector.get("relative_turnover"): sector_inner.append("缺少行业相对全市场换手活跃度") @@ -1850,7 +1866,7 @@ class DashboardService: invalid_fields[3].update(required[3]) invalid_fields[4].update(required[4]) else: - if not sector.get("inner_precise", sector.get("precise")) or sector_coverage < 90: + if not sector.get("inner_precise", sector.get("precise")) or sector_coverage_issue: invalid_fields[3].update(key for key in required[3] if key != "sector_name") if sector.get("realtime") and not sector.get("relative_turnover"): invalid_fields[3].add("sector_relative_turnover") @@ -2292,6 +2308,22 @@ class DashboardService: sector = sector or {} sector_date = str(sector.get("trade_date") or "").replace("-", "") sector_coverage = float(sector.get("coverage") or 0) + sector_explained_count = int( + sector.get("explained_count") + if sector.get("explained_count") is not None + else sector.get("quote_count") or 0 + ) + sector_explained_coverage = float( + sector.get("explained_coverage") + if sector.get("explained_coverage") is not None + else sector_coverage + ) + sector_coverage_issue = _sector_coverage_issue( + int(sector.get("member_count") or 0), + int(sector.get("quote_count") or 0), + sector_explained_coverage, + sector_explained_count, + ) if not sector: issues.append("行业层缺少申万二级行业归属") elif sector.get("taxonomy") != "sw_l2": @@ -2308,8 +2340,8 @@ class DashboardService: issues.append("行业内核缺少可核验的成分行情") if not sector.get("outer_precise", sector.get("precise")): issues.append("行业外显缺少申万官方行情") - if sector and sector_coverage < 90: - issues.append("行业成分行情覆盖率不足90%") + if sector and sector_coverage_issue: + issues.append(sector_coverage_issue) if sector.get("realtime") and not sector.get("relative_turnover"): issues.append("行业内核缺少相对全市场换手活跃度") @@ -2742,7 +2774,7 @@ class DashboardService: and cached.get("inner_precise", cached.get("precise")) and cached.get("outer_precise", cached.get("precise")) and not cached.get("realtime") - and int(cached.get("schema_version") or 0) >= 4 + and int(cached.get("schema_version") or 0) >= 6 ) if market_mode != "intraday" and cached_valid: return cached diff --git a/tests/test_heaven_realtime.py b/tests/test_heaven_realtime.py index 8c0ea1d..6df74a4 100644 --- a/tests/test_heaven_realtime.py +++ b/tests/test_heaven_realtime.py @@ -8,7 +8,11 @@ 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 +from tushare_client import ( + TushareClient, + _filter_members_by_listing, + _sector_coverage_issue, +) class HeavenMarketLineTests(unittest.TestCase): @@ -196,6 +200,75 @@ class HeavenMarketLineTests(unittest.TestCase): 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 = { diff --git a/tests/test_market_mode.py b/tests/test_market_mode.py index 941e348..b45269e 100644 --- a/tests/test_market_mode.py +++ b/tests/test_market_mode.py @@ -5,6 +5,8 @@ import unittest from datetime import datetime, timedelta, timezone from pathlib import Path +from tushare_client import _sector_coverage_issue + def load_method(name: str): source = Path("server.py").read_text(encoding="utf-8") @@ -19,7 +21,11 @@ def load_method(name: str): and node.name == name ) module = ast.Module(body=[method], type_ignores=[]) - namespace = {"datetime": datetime, "Any": object} + namespace = { + "datetime": datetime, + "Any": object, + "_sector_coverage_issue": _sector_coverage_issue, + } exec(compile(ast.fix_missing_locations(module), "server.py", "exec"), namespace) return namespace[name] diff --git a/tushare_client.py b/tushare_client.py index 0e88453..ad99240 100644 --- a/tushare_client.py +++ b/tushare_client.py @@ -30,6 +30,10 @@ class TushareClient: _capital_cache: ClassVar[dict[str, dict[str, Any]]] = {} _latest_realtime_market: ClassVar[dict[str, dict[str, Any]]] = {} _stock_activity_cache: ClassVar[dict[str, dict[str, Any]]] = {} + _stock_listing_cache: ClassVar[dict[str, Any]] = {} + _stock_listing_lock: ClassVar[Lock] = Lock() + _suspension_cache: ClassVar[dict[str, dict[str, str] | None]] = {} + _suspension_lock: ClassVar[Lock] = Lock() def query( self, @@ -713,33 +717,62 @@ class TushareClient: members = self._sw_sector_members(sector_code, trade_date) if not members: raise TushareError(f"No Shenwan members returned for {sector_code}") + raw_member_count = len(members) + members, excluded_members = _filter_members_by_listing( + members, + self._stock_listing_reference(), + trade_date, + ) + if not members: + raise TushareError(f"No listed Shenwan members returned for {sector_code}") if realtime_expected: - return self._sw_realtime_sector_snapshot( + snapshot = self._sw_realtime_sector_snapshot( industry, members, trade_date, previous_trade_date, finalized=False, ) + snapshot.update({ + "raw_member_count": raw_member_count, + "excluded_member_count": len(excluded_members), + "excluded_members": excluded_members, + }) + return snapshot member_set = {str(item.get("ts_code") or "") for item in members} + member_names = { + str(item.get("ts_code") or ""): str(item.get("name") or "") + for item in members + } member_rows = [ row for row in self._load_daily(trade_date) if str(row.get("ts_code") or "") in member_set ] + quoted_codes = {str(row.get("ts_code") or "") for row in member_rows} + suspended_members = self._confirmed_suspended_members( + members, quoted_codes, trade_date + ) up_count = sum(_number(row.get("pct_chg")) > 0 for row in member_rows) down_count = sum(_number(row.get("pct_chg")) < 0 for row in member_rows) leader = max(member_rows, key=lambda row: _number(row.get("pct_chg")), default={}) + leader_code = str(leader.get("ts_code") or "") equal_change = ( sum(_number(row.get("pct_chg")) for row in member_rows) / len(member_rows) if member_rows else 0 ) coverage = len(member_rows) / max(len(members), 1) * 100 - inner_precise = coverage >= 90 - inner_error = "" if inner_precise else ( - f"Shenwan member daily coverage is insufficient ({len(member_rows)}/{len(members)})" + explained_count = len(member_rows) + len(suspended_members) + explained_coverage = explained_count / max(len(members), 1) * 100 + coverage_issue = _sector_coverage_issue( + len(members), + len(member_rows), + explained_coverage, + explained_count, ) + inner_precise = not coverage_issue + inner_error = coverage_issue amount_billion = sum(_number(row.get("amount")) for row in member_rows) / 100000 rows = self.query( "sw_daily", @@ -768,8 +801,8 @@ class TushareClient: return { "code": sector_code, "name": industry.get("l2_name") or daily.get("name") or sector_code, - "leader": str(leader.get("name") or "--"), - "leader_code": str(leader.get("ts_code") or ""), + "leader": str(leader.get("name") or member_names.get(leader_code) or "--"), + "leader_code": leader_code, "leading_pct": round(_number(leader.get("pct_chg")), 3), "change": round(official_change, 3) if official_change is not None else None, "member_equal_change": round(equal_change, 3), @@ -778,8 +811,15 @@ class TushareClient: "down_count": down_count, "flat_count": len(member_rows) - up_count - down_count, "member_count": len(members), + "raw_member_count": raw_member_count, + "excluded_member_count": len(excluded_members), + "excluded_members": excluded_members, "quote_count": len(member_rows), "coverage": round(coverage, 1), + "explained_count": explained_count, + "explained_coverage": round(explained_coverage, 1), + "suspended_count": len(suspended_members), + "suspended_members": suspended_members, "strength": round(max(0, min(100, 50 + (official_change if official_change is not None else equal_change) * 5)), 1), "amount_billion": round(amount_billion, 2), "count": 0, @@ -799,7 +839,7 @@ class TushareClient: "precise": inner_precise and outer_precise, "inner_error": inner_error, "outer_error": outer_error, - "schema_version": 4, + "schema_version": 6, "methodology": "外显使用申万二级行业官方日线;内核独立使用当日成分日线宽度与等权涨跌聚合", } @@ -826,6 +866,89 @@ class TushareClient: deduped[code] = row return list(deduped.values()) + def _stock_listing_reference(self) -> dict[str, dict[str, Any]]: + now = datetime.now().astimezone() + with self._stock_listing_lock: + loaded_at = self._stock_listing_cache.get("loaded_at") + cached = self._stock_listing_cache.get("rows") + if ( + isinstance(loaded_at, datetime) + and isinstance(cached, dict) + and now - loaded_at < timedelta(hours=6) + ): + return cached + + rows: list[dict[str, Any]] = [] + try: + for status in ("L", "D", "P"): + rows.extend(self.query( + "stock_basic", + {"list_status": status}, + "ts_code,name,list_status,list_date,delist_date", + )) + except TushareError: + # Unknown status must remain in the denominator so a reference-data + # failure cannot silently improve coverage. + return {} + reference = { + str(row.get("ts_code") or ""): dict(row) + for row in rows + if row.get("ts_code") + } + with self._stock_listing_lock: + type(self)._stock_listing_cache = {"loaded_at": now, "rows": reference} + return reference + + def _confirmed_suspended_members( + self, + members: list[dict[str, Any]], + quoted_codes: set[str], + trade_date: str, + ) -> list[dict[str, str]]: + suspended: list[dict[str, str]] = [] + for member in members: + code = str(member.get("ts_code") or "") + if not code or code in quoted_codes: + continue + cache_key = f"{trade_date}:{code}" + with self._suspension_lock: + cached = self._suspension_cache.get(cache_key, "missing") + if cached == "missing": + try: + rows = self.query( + "suspend_d", + {"ts_code": code}, + "ts_code,suspend_date,resume_date,ann_date,suspend_reason,reason_type", + ) + except TushareError: + rows = [] + active = [ + row for row in rows + if str(row.get("suspend_date") or "") + and str(row.get("suspend_date") or "") <= trade_date + and ( + not str(row.get("resume_date") or "") + or trade_date < str(row.get("resume_date") or "") + ) + ] + row = max( + active, + key=lambda item: str(item.get("suspend_date") or ""), + default=None, + ) + cached = ({ + "ts_code": code, + "name": str(member.get("name") or code), + "suspend_date": str(row.get("suspend_date") or ""), + "resume_date": str(row.get("resume_date") or ""), + "reason": str(row.get("suspend_reason") or row.get("reason_type") or "已确认停牌"), + } if row else None) + with self._suspension_lock: + type(self)._suspension_cache[cache_key] = cached + if isinstance(cached, dict): + suspended.append(cached) + return suspended + def _sw_realtime_sector_snapshot( self, industry: dict[str, Any], @@ -892,12 +1015,26 @@ class TushareClient: inner_error = str(exc) coverage = len(valid) / max(len(codes), 1) * 100 - inner_precise = bool(codes) and coverage >= 90 + valid_codes = {str(item.get("ts_code") or "") for item in valid} + suspended_members = self._confirmed_suspended_members( + members, valid_codes, trade_date + ) + explained_count = len(valid) + len(suspended_members) + explained_coverage = explained_count / max(len(codes), 1) * 100 + coverage_issue = _sector_coverage_issue( + len(codes), len(valid), explained_coverage, explained_count + ) + inner_precise = bool(codes) and not coverage_issue if not inner_precise and not inner_error: - inner_error = f"Shenwan realtime coverage is insufficient ({len(valid)}/{len(codes)})" + inner_error = coverage_issue or "申万实时有效成分为空" up_count = sum(item["change"] > 0 for item in valid) down_count = sum(item["change"] < 0 for item in valid) leader = max(valid, key=lambda item: item["change"], default={}) + leader_code = str(leader.get("ts_code") or "") + member_names = { + str(item.get("ts_code") or ""): str(item.get("name") or "") + for item in members + } equal_change = sum(item["change"] for item in valid) / len(valid) if valid else 0 amount_billion = sum(_number(item.get("amount")) for item in valid) / 100000000 try: @@ -935,8 +1072,8 @@ class TushareClient: return { "code": sector_code, "name": str(industry.get("l2_name") or sw_row.get("name") or ""), - "leader": str(leader.get("name") or "--").strip(), - "leader_code": str(leader.get("ts_code") or ""), + "leader": str(leader.get("name") or member_names.get(leader_code) or "--").strip(), + "leader_code": leader_code, "leading_pct": round(_number(leader.get("change")), 3), "change": round(official_change, 3) if official_change is not None else None, "member_equal_change": round(equal_change, 3), @@ -949,6 +1086,10 @@ class TushareClient: "member_count": len(codes), "quote_count": len(valid), "coverage": round(coverage, 1), + "explained_count": explained_count, + "explained_coverage": round(explained_coverage, 1), + "suspended_count": len(suspended_members), + "suspended_members": suspended_members, "strength": round(max(0, min(100, 50 + (official_change if official_change is not None else equal_change) * 5)), 1), "amount_billion": round(amount_billion, 2), "count": sum(item["change"] >= 9.5 for item in valid), @@ -969,7 +1110,7 @@ class TushareClient: "precise": inner_precise and outer_precise, "inner_error": inner_error, "outer_error": outer_error, - "schema_version": 4, + "schema_version": 6, "methodology": "外显使用申万官方 rt_sw_k;内核独立使用申万成分 rt_k 宽度与相对换手聚合", } @@ -1624,6 +1765,73 @@ def _text(value: Any) -> str: return str(value or "").strip() +def _filter_members_by_listing( + members: list[dict[str, Any]], + listing_reference: dict[str, dict[str, Any]], + trade_date: str, +) -> tuple[list[dict[str, Any]], list[dict[str, str]]]: + eligible: list[dict[str, Any]] = [] + excluded: list[dict[str, str]] = [] + for member in members: + code = str(member.get("ts_code") or "") + listing = listing_reference.get(code) + if not listing: + eligible.append(member) + continue + list_date = str(listing.get("list_date") or "") + delist_date = str(listing.get("delist_date") or "") + reason = "" + effective_date = "" + if delist_date and delist_date <= trade_date: + reason = "目标日期前已退市" + effective_date = delist_date + elif list_date and list_date > trade_date: + reason = "目标日期尚未上市" + effective_date = list_date + if not reason: + eligible.append(member) + continue + excluded.append({ + "ts_code": code, + "name": str(member.get("name") or listing.get("name") or code), + "reason": reason, + "effective_date": effective_date, + }) + return eligible, excluded + + +def _sector_coverage_issue( + member_count: int, + quote_count: int, + coverage: float | None = None, + explained_count: int | None = None, +) -> str: + members = max(0, int(member_count or 0)) + quotes = max(0, min(int(quote_count or 0), members)) + if members <= 0: + if coverage is not None and float(coverage) >= 90: + return "" + if coverage is not None: + return "行业成分行情覆盖率低于90%" + return "申万有效成分为空" + explained = quotes if explained_count is None else max( + quotes, min(int(explained_count or 0), members) + ) + actual_coverage = ( + float(coverage) + if coverage is not None + else explained / members * 100 + ) + missing = members - explained + if members <= 7 and missing: + return f"小型行业有效成分状态仅确认 {explained}/{members},要求全部可解释" + if members <= 20 and (actual_coverage < 90 or missing > 1): + return f"中型行业有效成分状态仅确认 {explained}/{members},要求覆盖率至少90%且最多缺1只" + if members > 20 and actual_coverage < 90: + return f"行业有效成分状态仅确认 {explained}/{members},覆盖率低于90%" + return "" + + def _membership_active_on(row: dict[str, Any], trade_date: str) -> bool: start = str(row.get("in_date") or "") end = str(row.get("out_date") or "")