from __future__ import annotations import math import re from datetime import datetime, time as dt_time from typing import Any from backend.data.numbers import finite_number as _number from backend.data.providers.tushare_industries import _match_sector_row from backend.data.providers.tushare_transport import TushareError class SectorMixin: def sector_snapshot( self, identifier: str, requested_date: str, realtime_expected: bool | None = None, ) -> dict[str, Any]: trade_date, _ = self.resolve_trade_context(requested_date) raw_identifier = identifier.strip() if not raw_identifier: raise TushareError("Sector identifier is empty") errors = [] now = datetime.now().astimezone() if realtime_expected is None: realtime_expected = ( trade_date == now.strftime("%Y%m%d") and dt_time(9, 15) <= now.time().replace(tzinfo=None) <= dt_time(15, 5) ) try: dc_params = {"trade_date": trade_date} if re.fullmatch(r"[A-Z0-9.]+", raw_identifier.upper()) and "." in raw_identifier: dc_params["ts_code"] = raw_identifier.upper() else: dc_params["name"] = raw_identifier dc_rows = self.query( "dc_index", dc_params, "ts_code,trade_date,name,leading,leading_code,pct_change,leading_pct," "total_mv,turnover_rate,up_num,down_num", ) if not dc_rows and "name" in dc_params: dc_rows = self.query( "dc_index", {"trade_date": trade_date}, "ts_code,trade_date,name,leading,leading_code,pct_change,leading_pct," "total_mv,turnover_rate,up_num,down_num", ) dc_row = _match_sector_row(dc_rows, raw_identifier) if dc_row and not realtime_expected: change = _number(dc_row.get("pct_change")) actual_trade_date = str(dc_row.get("trade_date") or "") return { "code": dc_row.get("ts_code") or "", "name": dc_row.get("name") or raw_identifier, "leader": dc_row.get("leading") or "--", "leader_code": dc_row.get("leading_code") or "", "leading_pct": _number(dc_row.get("leading_pct")), "change": change, "turnover_rate": _number(dc_row.get("turnover_rate")), "up_count": int(_number(dc_row.get("up_num"))), "down_count": int(_number(dc_row.get("down_num"))), "total_mv": _number(dc_row.get("total_mv")), "strength": round(max(0, min(100, 50 + change * 5)), 1), "amount_billion": 0, "count": 0, "max_streak": 0, "source": "tushare_dc", "trade_date": actual_trade_date, "realtime": False, "precise": actual_trade_date == trade_date, } except TushareError as exc: errors.append(f"DC: {exc}") ts_code = raw_identifier.upper() if re.fullmatch(r"\d{6}", ts_code): ts_code = f"{ts_code}.TI" try: if re.fullmatch(r"\d{6}\.TI", ts_code): index_rows = self.query( "ths_index", {"ts_code": ts_code}, "ts_code,name,count,exchange,list_date,type", ) else: index_rows = self.query( "ths_index", {}, "ts_code,name,count,exchange,list_date,type", ) basic = _match_sector_row(index_rows, raw_identifier) if not basic: raise TushareError(f"No THS sector returned for {raw_identifier}") except TushareError as exc: errors.append(f"THS: {exc}") raise TushareError("; ".join(errors)) from exc actual_code = str(basic.get("ts_code") or ts_code) if realtime_expected: try: realtime_sector = self._realtime_sector_snapshot(actual_code, basic, trade_date) if realtime_sector: return realtime_sector except TushareError as exc: errors.append(f"THS realtime members: {exc}") daily_rows = self.query( "ths_daily", {"ts_code": actual_code, "trade_date": trade_date}, "ts_code,trade_date,close,pct_change,vol,turnover_rate,total_mv,float_mv", ) daily = daily_rows[0] if daily_rows else {} actual_trade_date = str(daily.get("trade_date") or "") change = _number(daily.get("pct_change")) return { "code": actual_code, "name": basic.get("name") or raw_identifier, "leader": "--", "change": change, "leading_pct": change, "turnover_rate": _number(daily.get("turnover_rate")), "up_count": 0, "down_count": 0, "strength": round(max(0, min(100, 50 + change * 5)), 1), "amount_billion": 0, "count": 0, "max_streak": 0, "source": "tushare_ths", "trade_date": actual_trade_date, "realtime": False, "precise": actual_trade_date == trade_date, } def _realtime_sector_snapshot( self, sector_code: str, basic: dict[str, Any], trade_date: str, ) -> dict[str, Any] | None: members = self.query( "ths_member", {"ts_code": sector_code, "is_new": "Y"}, "ts_code,con_code,con_name,is_new", ) codes = [str(row.get("con_code") or "") for row in members if row.get("con_code")] if not codes: return None quotes = self.query("rt_k", {"ts_code": ",".join(codes)}, "") valid = [] 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, } ) minimum = max(1, math.ceil(len(codes) * 0.9)) if len(valid) < minimum: raise TushareError( f"Realtime sector coverage is insufficient ({len(valid)}/{len(codes)})" ) up_count = sum(item["change"] > 0 for item in valid) down_count = sum(item["change"] < 0 for item in valid) flat_count = len(valid) - up_count - down_count leader = max(valid, key=lambda item: item["change"]) change = sum(item["change"] for item in valid) / len(valid) amount_billion = sum(_number(item.get("amount")) for item in valid) / 100000000 self._ensure_realtime_market_cache(trade_date) with self._realtime_reference_lock: references = list(self._realtime_reference_cache.values()) market_rows = list((self._latest_realtime_market.get(trade_date) or {}).get("rows") or []) capital_map: dict[str, dict[str, Any]] = {} for reference in reversed(references): capital_map = { str(item.get("ts_code") or ""): item for item in reference.get("capital_rows") or [] } if capital_map: break 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 return { "code": sector_code, "name": basic.get("name") or sector_code, "leader": str(leader.get("name") or "--").strip(), "leader_code": leader.get("ts_code") or "", "leading_pct": round(leader["change"], 3), "change": round(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": flat_count, "member_count": len(codes), "quote_count": len(valid), "coverage": round(len(valid) / len(codes) * 100, 1), "strength": round(max(0, min(100, 50 + 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_ths_members", "trade_date": trade_date, "realtime": True, "precise": True, "methodology": "同花顺行业最新成分股的 rt_k 等权涨跌、宽度与成交额聚合", }