from __future__ import annotations from datetime import datetime, timedelta from typing import Any from backend.data.numbers import finite_number as _number from backend.data.providers.tushare_transport import TushareError class ShenwanIndustryMixin: def sw_stock_industry(self, ts_code: str, trade_date: str) -> dict[str, Any]: """Return the Shenwan industry active for a stock on trade_date.""" rows = [] for is_new in ("Y", "N"): rows.extend( self.query( "index_member_all", {"ts_code": ts_code, "is_new": is_new}, "l1_code,l1_name,l2_code,l2_name,l3_code,l3_name," "ts_code,name,in_date,out_date,is_new", ) ) rows = _reconcile_membership_rows(rows) matched = [row for row in rows if _membership_active_on(row, trade_date)] if not matched: matched = [ row for row in rows if row.get("is_new") == "Y" and str(row.get("in_date") or "") <= trade_date ] if not matched: raise TushareError(f"No Shenwan industry returned for {ts_code}") row = max( matched, key=lambda item: ( str(item.get("in_date") or ""), 1 if item.get("is_new") == "Y" else 0, str(item.get("l3_code") or item.get("l2_code") or ""), ), ) return { "l1_code": str(row.get("l1_code") or ""), "l1_name": str(row.get("l1_name") or ""), "l2_code": str(row.get("l2_code") or ""), "l2_name": str(row.get("l2_name") or ""), "l3_code": str(row.get("l3_code") or ""), "l3_name": str(row.get("l3_name") or ""), "in_date": str(row.get("in_date") or ""), "out_date": str(row.get("out_date") or ""), "is_new": str(row.get("is_new") or ""), } def sw_sector_snapshot( self, ts_code: str, requested_date: str, realtime_expected: bool = False, allow_realtime_close: bool = False, ) -> dict[str, Any]: """Build the single Shenwan L2 sector context used by heaven trend.""" trade_date, previous_trade_date = self.resolve_trade_context(requested_date) industry = self.sw_stock_industry(ts_code, trade_date) sector_code = str(industry.get("l2_code") or "") if not sector_code: raise TushareError(f"Shenwan L2 code is unavailable for {ts_code}") 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: 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 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", {"ts_code": sector_code, "trade_date": trade_date}, "ts_code,trade_date,name,close,pct_change,vol,amount,pe,pb,float_mv,total_mv", ) daily = rows[0] if rows else {} actual_trade_date = str(daily.get("trade_date") or "") outer_precise = actual_trade_date == trade_date outer_error = "" if outer_precise else ( f"No Shenwan daily returned for {sector_code} on {trade_date}" ) if not outer_precise and allow_realtime_close: try: return self._sw_realtime_sector_snapshot( industry, members, trade_date, previous_trade_date, finalized=True, ) except TushareError as exc: outer_error = f"{outer_error}; realtime close fallback failed: {exc}" official_change = _number(daily.get("pct_change")) if outer_precise else None return { "code": sector_code, "name": industry.get("l2_name") or daily.get("name") or sector_code, "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), "turnover_rate": 0, "up_count": up_count, "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, "max_streak": 0, "source": "tushare_sw_daily+member_daily" if outer_precise else "tushare_member_daily", "inner_source": "tushare_member_daily", "outer_source": "tushare_sw_daily" if outer_precise else "unavailable", "taxonomy": "sw_l2", "industry": industry, "trade_date": trade_date, "inner_trade_date": trade_date if member_rows else "", "outer_trade_date": actual_trade_date, "realtime": False, "finalized": True, "inner_precise": inner_precise, "outer_precise": outer_precise, "precise": inner_precise and outer_precise, "inner_error": inner_error, "outer_error": outer_error, "schema_version": 6, "methodology": "外显使用申万二级行业官方日线;内核独立使用当日成分日线宽度与等权涨跌聚合", } def _sw_sector_members( self, sector_code: str, trade_date: str, ) -> list[dict[str, Any]]: rows = [] for is_new in ("Y", "N"): rows.extend( self.query( "index_member_all", {"l2_code": sector_code, "is_new": is_new}, "l2_code,l2_name,ts_code,name,in_date,out_date,is_new", ) ) deduped: dict[str, dict[str, Any]] = {} for row in _reconcile_membership_rows(rows): code = str(row.get("ts_code") or "") if code and _membership_active_on(row, trade_date): current = deduped.get(code) if current is None or str(row.get("in_date") or "") > str(current.get("in_date") or ""): deduped[code] = row return list(deduped.values()) def sw_sector_members(self, sector_code: str, trade_date: str) -> list[dict[str, Any]]: """Return constituents active in a Shenwan L2 industry on the target date.""" return self._sw_sector_members(sector_code, trade_date) 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], members: list[dict[str, Any]], trade_date: str, previous_trade_date: str, finalized: bool = False, ) -> dict[str, Any]: sector_code = str(industry.get("l2_code") or "") sw_rows = self.query( "rt_sw_k", {"ts_code": sector_code}, "ts_code,name,trade_time,close,pre_close,high,open,low,vol,amount,pct_change", ) sw_row = sw_rows[0] if sw_rows else {} trade_time = str(sw_row.get("trade_time") or "") quote_date = trade_time[:10].replace("-", "") quote_clock = trade_time[11:19] if len(trade_time) >= 19 else "" outer_precise = bool(sw_row and quote_date == trade_date) if finalized and (not quote_clock or quote_clock < "15:00:00"): outer_precise = False official_change = _number(sw_row.get("pct_change")) if not official_change: close = _number(sw_row.get("close")) pre_close = _number(sw_row.get("pre_close")) official_change = (close / pre_close - 1) * 100 if close and pre_close else 0 if not outer_precise: official_change = None outer_error = "" if not sw_row: outer_error = f"No Shenwan realtime index returned for {sector_code}" elif quote_date != trade_date: outer_error = f"Shenwan realtime index date is {quote_date or 'unknown'}, expected {trade_date}" elif finalized and (not quote_clock or quote_clock < "15:00:00"): outer_error = f"Shenwan realtime index is not a close snapshot ({trade_time})" valid: list[dict[str, Any]] = [] codes: list[str] = [] reference: dict[str, Any] = {} inner_error = "" try: reference = self._load_realtime_reference(trade_date, previous_trade_date) active_codes = { str(row.get("ts_code") or "") for row in reference.get("basic_rows") or [] if row.get("ts_code") } codes = [ str(row.get("ts_code") or "") for row in members if str(row.get("ts_code") or "") in active_codes ] if codes: quotes = self.query("rt_k", {"ts_code": ",".join(codes)}, "") for row in quotes: close = _number(row.get("close")) previous_close = _number(row.get("pre_close")) if close <= 0 or previous_close <= 0: continue valid.append({**row, "change": (close / previous_close - 1) * 100}) else: inner_error = f"No active Shenwan members returned for {sector_code}" except TushareError as exc: inner_error = str(exc) coverage = len(valid) / max(len(codes), 1) * 100 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 = 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: self._ensure_realtime_market_cache(trade_date) with self._realtime_reference_lock: market_rows = list( (self._latest_realtime_market.get(trade_date) or {}).get("rows") or [] ) except TushareError as exc: market_rows = [] inner_precise = False inner_error = inner_error or str(exc) capital_map = { str(item.get("ts_code") or ""): item for item in reference.get("capital_rows") or [] } sector_turnovers = [] for item in valid: capital = capital_map.get(str(item.get("ts_code") or ""), {}) float_share = _number(capital.get("float_share")) if float_share: sector_turnovers.append(_number(item.get("vol")) / float_share / 100) market_turnovers = [] for item in market_rows: capital = capital_map.get(str(item.get("ts_code") or ""), {}) float_share = _number(capital.get("float_share")) if float_share: market_turnovers.append(_number(item.get("vol")) / float_share / 100) average_turnover = sum(sector_turnovers) / len(sector_turnovers) if sector_turnovers else 0 market_turnover = sum(market_turnovers) / len(market_turnovers) if market_turnovers else 0 relative_turnover = average_turnover / market_turnover if market_turnover else 0 if not relative_turnover: inner_precise = False inner_error = inner_error or "Shenwan member relative turnover is unavailable" return { "code": sector_code, "name": str(industry.get("l2_name") or sw_row.get("name") 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), "turnover_rate": round(average_turnover, 4), "market_turnover_rate": round(market_turnover, 4), "relative_turnover": round(relative_turnover, 4), "up_count": up_count, "down_count": down_count, "flat_count": len(valid) - up_count - down_count, "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), "max_streak": 0, "source": "tushare_rt_sw_k+sw_members_rt_k", "inner_source": "tushare_sw_members+rt_k", "outer_source": "tushare_rt_sw_k", "taxonomy": "sw_l2", "industry": industry, "trade_date": trade_date, "inner_trade_date": trade_date if valid else "", "outer_trade_date": quote_date, "trade_time": trade_time, "realtime": True, "finalized": finalized, "inner_precise": inner_precise, "outer_precise": outer_precise, "precise": inner_precise and outer_precise, "inner_error": inner_error, "outer_error": outer_error, "schema_version": 6, "methodology": "外显使用申万官方 rt_sw_k;内核独立使用申万成分 rt_k 宽度与相对换手聚合", } 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 "") return (not start or start <= trade_date) and (not end or end > trade_date) def _reconcile_membership_rows(rows: list[dict[str, Any]]) -> list[dict[str, Any]]: """Merge duplicate Y/N membership rows before evaluating their date interval.""" reconciled: dict[tuple[str, str, str, str, str], dict[str, Any]] = {} for raw in rows: row = dict(raw) key = ( str(row.get("ts_code") or ""), str(row.get("l1_code") or ""), str(row.get("l2_code") or ""), str(row.get("l3_code") or ""), str(row.get("in_date") or ""), ) current = reconciled.get(key) if current is None: reconciled[key] = row continue current_end = str(current.get("out_date") or "") candidate_end = str(row.get("out_date") or "") if candidate_end and not current_end: current["out_date"] = candidate_end current["is_new"] = row.get("is_new") or current.get("is_new") for field, value in row.items(): if not current.get(field) and value not in (None, ""): current[field] = value return list(reconciled.values()) def _match_sector_row(rows: list[dict[str, Any]], identifier: str) -> dict[str, Any] | None: if not rows: return None target = identifier.strip().upper() code_match = next( (row for row in rows if str(row.get("ts_code") or "").strip().upper() == target), None, ) if code_match: return code_match def normalized(value: Any) -> str: text = str(value or "").strip().replace(" ", "") for suffix in ("板块", "概念", "行业"): text = text.removesuffix(suffix) aliases = { "元器件": "元件", "电子元器件": "元件", } return aliases.get(text, text) target_name = normalized(identifier) exact = [row for row in rows if normalized(row.get("name")) == target_name] if exact: return min(exact, key=_sector_match_priority) fuzzy = [ row for row in rows if target_name and ( target_name in normalized(row.get("name")) or normalized(row.get("name")) in target_name ) ] return min( fuzzy, key=lambda row: (len(normalized(row.get("name"))), *_sector_match_priority(row)), ) if fuzzy else None def _sector_match_priority(row: dict[str, Any]) -> tuple[int, int, int]: code = str(row.get("ts_code") or "") exchange = str(row.get("exchange") or "").upper() return ( 0 if exchange == "A" else 1, 0 if code.startswith("881") else 1, 0 if _number(row.get("count")) > 0 else 1, )