fix: validate Shenwan constituent coverage

This commit is contained in:
leefer
2026-07-24 08:45:23 +08:00
parent 3c79a4976b
commit fde2728a86
4 changed files with 340 additions and 21 deletions
+39 -7
View File
@@ -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
+74 -1
View File
@@ -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 = {
+7 -1
View File
@@ -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]
+220 -12
View File
@@ -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 "")