tradingSystem/app/repo/industry_repo.py

138 lines
5.8 KiB
Python

# -*- coding: utf-8 -*-
"""
行业板块表 gp_hybk (199 库, 只读)
==================================
数据源由另一项目的持仓行业分析给出 (2026-07-31 参考其实现确认):
库 DB_MYSQL_URL (192.168.18.199) —— PMS 里原本只用来读大盘指数 zs_day_data,
`session.fetch_all(..., source="index")` 直接可用, 不新增连接。
表 gp_hybk
列 gp_code (股票代码) / bk_code (板块代码, 数值型) / bk_name (板块名)
口径 bk_code 前缀 881 = 二级行业, 884 = 三级行业
特点 **一只票会对应多条同级行业**
PMS 用**三级 (884)** 做行业集中度硬拦截, 且**每票只认一个主行业** ——
`sizer.check_caps` 的 `ctx["sector"]` 是单值, 累计口径 (`sector_names_map` /
`sector_mv_map`) 也按单值建。主行业的选法必须**稳定**: 同一只票每次都要得到同一个行业,
否则今天算「工程机械」明天算「专用设备」, 集中度累计会自己跳。故定为
**bk_code 升序取第一个** —— 纯粹按数值排, 不依赖查询返回次序, 也不依赖表里有没有主次标记。
代码写法未确认 (点式 600000.SH / 前缀式 SH600000 / 纯数字 600000 都可能), 所以
`probe()` 会三种各试一遍, 认出哪种能命中就缓存下来, 之后批量查一律用那种。
"""
from __future__ import annotations
import logging
import threading
import time
from app.db.session import fetch_all
from app.repo.downstream_repo import to_dot, to_prefix
logger = logging.getLogger("pms.hybk")
TABLE = "gp_hybk"
SOURCE = "index" # DB_MYSQL_URL
LEVEL_PREFIX = {"l2": "881", "l3": "884"}
BATCH = 400 # 一次 IN 多少个代码 (持仓+候选也就几十, 留足余量)
# 代码写法探测样本: 沪/深各一, 一定在任何行业表里
PROBE_CODES = ("600000.SH", "000001.SZ", "600519.SH")
_fmt = {"at": 0.0, "form": None, "error": None, "columns": None}
_lock = threading.Lock()
FMT_TTL = 600.0
def _forms(ts_code: str) -> dict:
dot = to_dot(ts_code)
return {"dot": dot, "prefix": to_prefix(dot), "num": dot.split(".")[0]}
def _as(ts_code: str, form: str) -> str:
return _forms(ts_code).get(form) or to_dot(ts_code)
def _rows_for(codes: list) -> list:
"""一次 IN 查询。不用 CAST/LIKE —— 前缀判断放 Python 侧, 免得 SQL 方言与单表守卫扯皮。"""
if not codes:
return []
keys = [f"c{i}" for i in range(len(codes))]
sql = (f"SELECT gp_code, bk_code, bk_name FROM {TABLE} "
f"WHERE gp_code IN ({', '.join(':' + k for k in keys)})")
return fetch_all(sql, dict(zip(keys, codes)), source=SOURCE)
def probe(force: bool = False) -> dict:
"""认代码写法。返回 {form, error, columns, tried}; form=None 表示这张表用不了。
三种写法各拿 PROBE_CODES 试一次, 谁先命中用谁。命中不了但也没报错, 说明表能查、
只是这些票不在里面 (或代码写法是第四种) —— 这跟"表根本查不了"要分开报, 不能都吞成
"没有行业"。行业约束是硬拦截, 拿不到就该明着停, 见 UPSTREAM_PLAN_API.md §7.4。
"""
now = time.time()
if not force and _fmt["form"] is not None and now - _fmt["at"] < FMT_TTL:
return dict(_fmt, tried=list(PROBE_CODES))
err = None
for form in ("dot", "prefix", "num"):
codes = [_as(c, form) for c in PROBE_CODES]
try:
rows = _rows_for(codes)
except Exception as e:
err = f"{form}({codes[0]}): {type(e).__name__}: {str(e)[:180]}"
continue
if rows:
cols = sorted(rows[0].keys())
with _lock:
_fmt.update({"at": now, "form": form, "error": None, "columns": cols})
logger.info("[行业] gp_hybk 代码写法 = %s (样本 %s), 列 %s", form, codes[0], cols)
return dict(_fmt, tried=codes)
with _lock:
_fmt.update({"at": now, "form": None, "columns": None,
"error": err or f"三种代码写法都没命中 (试了 {list(PROBE_CODES)}); "
f"表能查但里面没有这些票, 或 gp_code 是第四种写法"})
return dict(_fmt, tried=list(PROBE_CODES))
def invalidate():
with _lock:
_fmt.update({"at": 0.0, "form": None, "error": None, "columns": None})
def fetch_industries(codes, level: str = "l3") -> dict:
"""{ts_code(点式): [(bk_code, bk_name), ...]}, 已按 bk_code 升序。
只回该 level 的板块; 查不到的票不出现在结果里 (调用方据此判 None)。
"""
prefix = LEVEL_PREFIX.get(level, LEVEL_PREFIX["l3"])
p = probe()
form = p.get("form")
if not form:
raise RuntimeError(p.get("error") or "gp_hybk 不可用")
want = {}
for c in codes or []:
d = to_dot(c)
if d:
want.setdefault(_as(d, form), d)
out = {}
keys = list(want)
for i in range(0, len(keys), BATCH):
for r in _rows_for(keys[i:i + BATCH]):
raw = str(r.get("gp_code") or "").strip().upper()
dot = want.get(raw) or want.get(raw.upper())
bk = str(r.get("bk_code") or "").strip()
name = str(r.get("bk_name") or "").strip()
if not dot or not name or not bk.startswith(prefix):
continue
out.setdefault(dot, [])
if (bk, name) not in out[dot]:
out[dot].append((bk, name))
for v in out.values():
v.sort(key=lambda x: x[0]) # bk_code 升序 —— 主行业选取要可复现
return out
def primary_industry_map(codes, level: str = "l3") -> dict:
"""{ts_code: 主行业名}。主行业 = 该级板块里 bk_code 最小的那个 (见模块头部)。"""
return {c: v[0][1] for c, v in fetch_industries(codes, level=level).items() if v}