from __future__ import annotations from typing import Any from backend.bootstrap.config import normalize_date, validate_text from backend.data.providers.tushare_client import TushareError from backend.features.sentiment.engine import ( build_sentiment_history, latest_contiguous_history, ) class RotationServiceMixin: def rotation_history(self, trade_date: str, limit: int = 9) -> dict[str, Any]: normalized_date = normalize_date(trade_date) # 板块轮动固定展示最近 9 个交易日,按由近到远排列。 limit = 9 snapshots = self.database.list_snapshot_payloads(normalized_date, 240) by_trade_date: dict[str, dict[str, Any]] = {} for snapshot in snapshots: meta = snapshot.get("meta") or {} actual_date = str(meta.get("trade_date") or snapshot.get("_snapshot_date") or "") compact_date = actual_date.replace("-", "") if len(compact_date) == 8: by_trade_date[compact_date] = snapshot sentiment_dates = { str(row.get("trade_date") or "").replace("-", "") for row in latest_contiguous_history(build_sentiment_history(snapshots)) } ordered_dates = sorted( date_key for date_key in by_trade_date if not sentiment_dates or date_key in sentiment_dates )[-limit:][::-1] rows = [] for date_key in ordered_dates: snapshot = by_trade_date[date_key] sector_context = { str(item.get("name") or ""): item for item in snapshot.get("sectors") or [] } sectors = [] for item in (snapshot.get("sector_rotation") or [])[:12]: name = str(item.get("name") or "").strip() context = sector_context.get(name, {}) sectors.append( { "name": name, "rank": int(item.get("rank") or len(sectors) + 1), "trend": item.get("trend") or "持平", "count": int(item.get("count") or 0), "strength": float(item.get("strength") or context.get("strength") or 0), "change": float(context.get("change") or 0), "leader": item.get("leader") or context.get("leader") or "--", } ) rows.append( { "trade_date": f"{date_key[:4]}-{date_key[4:6]}-{date_key[6:]}", "sectors": sectors, } ) return { "trade_date": rows[0]["trade_date"] if rows else normalized_date, "available_days": len(ordered_dates), "requested_days": limit, "rows": rows, } def rotation_sector_members(self, trade_date: str, sector_name: str) -> dict[str, Any]: normalized_date = normalize_date(trade_date) sector_name = validate_text(sector_name, "板块名称", 60, required=True) dashboard = self.get_dashboard(normalized_date) actual_date = normalize_date( str((dashboard.get("meta") or {}).get("trade_date") or normalized_date) ) cache_key = f"{actual_date}:{sector_name}" cached = self.database.get_data_snapshot("rotation_sector_members_v1", cache_key) if cached: cached["meta"] = {**(cached.get("meta") or {}), "cached": True} return cached if not self.configured: raise ValueError("板块成分数据暂不可用。") representative = next( ( item for item in dashboard.get("limits") or [] if str(item.get("sector") or "").strip() == sector_name ), None, ) if not representative: raise ValueError("未找到该板块的代表股票,暂时无法核验成分股。") raw_code = str(representative.get("ts_code") or representative.get("code") or "") if "." in raw_code: ts_code = raw_code elif raw_code.startswith(("4", "8", "92")): ts_code = f"{raw_code}.BJ" elif raw_code.startswith(("6", "68", "90")): ts_code = f"{raw_code}.SH" else: ts_code = f"{raw_code}.SZ" client = self._tushare_client() try: industry = client.sw_stock_industry(ts_code, actual_date) sector_code = str(industry.get("l2_code") or "") members = client.sw_sector_members(sector_code, actual_date) except TushareError as exc: raise ValueError(f"该板块成分股暂不可用:{exc}") from exc daily_rows = self.database.daily_bars_for_date(actual_date) if len(daily_rows) < 1000: try: daily_rows = client.query( "daily", {"trade_date": actual_date}, "ts_code,trade_date,open,high,low,close,pct_chg,vol,amount", ) if daily_rows: self.database.upsert_daily_bars(daily_rows) except TushareError: daily_rows = self.database.daily_bars_for_date(actual_date) daily_map = {str(item.get("ts_code") or ""): item for item in daily_rows} rows = [] for member in members: member_code = str(member.get("ts_code") or "") quote = daily_map.get(member_code) or {} rows.append( { "code": member_code.split(".")[0], "ts_code": member_code, "name": str(member.get("name") or "--"), "change": quote.get("pct_chg"), "open": quote.get("open"), "close": quote.get("close"), "amount_billion": ( round(float(quote.get("amount") or 0) / 100000, 2) if quote else None ), "quoted": bool(quote), } ) rows.sort( key=lambda item: ( bool(item.get("quoted")), float(item.get("change") or -999), float(item.get("amount_billion") or 0), ), reverse=True, ) result = { "meta": { "trade_date": self._display_compact_date(actual_date), "sector_name": str(industry.get("l2_name") or sector_name), "sector_code": sector_code, "member_count": len(rows), "quoted_count": sum(bool(item.get("quoted")) for item in rows), "cached": False, }, "rows": rows, } self.database.save_data_snapshot( "rotation_sector_members_v1", cache_key, "tushare", result ) return result