Qionk/a-share-quant
0
1"""2价格预测 - 数据输入3支持 AKShare API 获取 和 Excel 手动上传两种方式4akshare 失败时自动切换到腾讯财经备用源,并尝试通过网易财经补充换手率5"""6 7import io8import time9import json10import subprocess11import pandas as pd12import numpy as np13import akshare as ak14 15REQUIRED_COLUMNS = ["日期", "开盘", "最高", "最低", "收盘", "成交量", "成交额", "涨跌幅"]16 17COLUMN_MAP = {18 "日期": "date", "开盘": "open", "最高": "high", "最低": "low",19 "收盘": "close", "成交量": "volume", "成交额": "amount", "涨跌幅": "pct_change",20}21 22MAX_RETRIES = 223RETRY_DELAY = 124 25 26def _market_prefix(stock_code: str) -> str:27 """根据股票代码判断市场前缀"""28 if stock_code.startswith(("6", "9")):29 return "sh"30 return "sz"31 32 33def _fetch_via_tencent(stock_code: str, start_date: str, end_date: str) -> pd.DataFrame:34 """35 备用数据源:通过腾讯财经 newfqkline 接口获取前复权日K数据(含换手率)36 """37 import requests as _req38 39 prefix = _market_prefix(stock_code)40 symbol = f"{prefix}{stock_code}"41 sd = f"{start_date[:4]}-{start_date[4:6]}-{start_date[6:8]}"42 ed = f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}"43 44 all_klines = []45 current_end = ed46 for _ in range(20):47 # newfqkline 接口:返回含换手率的完整日K48 url = (f"https://proxy.finance.qq.com/ifzqgtimg/appstock/app/newfqkline/get"49 f"?param={symbol},day,{sd},{current_end},800,qfq")50 try:51 resp = _req.get(url, timeout=10)52 if resp.status_code != 200:53 break54 data = resp.json()55 except Exception:56 # fallback 旧接口57 try:58 url_old = (f"https://web.ifzq.gtimg.cn/appstock/app/fqkline/get"59 f"?param={symbol},day,{sd},{current_end},800,qfq")60 resp = _req.get(url_old, timeout=10)61 data = resp.json()62 except Exception:63 break64 65 stock_data = data.get("data", {}).get(symbol, {})66 klines = stock_data.get("qfqday") or stock_data.get("day", [])67 if not klines:68 break69 70 all_klines = klines + all_klines71 earliest = klines[0][0]72 if earliest <= sd:73 break74 from datetime import datetime as dt, timedelta75 prev = (dt.strptime(earliest, "%Y-%m-%d") - timedelta(days=1)).strftime("%Y-%m-%d")76 current_end = prev77 time.sleep(0.3)78 79 if not all_klines:80 return pd.DataFrame()81 82 # 补齐尾部83 last_fetched = all_klines[-1][0] if all_klines else None84 if last_fetched and last_fetched < ed:85 url = (f"https://proxy.finance.qq.com/ifzqgtimg/appstock/app/newfqkline/get"86 f"?param={symbol},day,{last_fetched},{ed},30,qfq")87 try:88 resp = _req.get(url, timeout=10)89 if resp.status_code == 200:90 data2 = resp.json()91 klines2 = data2.get("data", {}).get(symbol, {})92 klines2 = klines2.get("qfqday") or klines2.get("day", [])93 if klines2:94 extra = [k for k in klines2 if k[0] > last_fetched]95 all_klines.extend(extra)96 except Exception:97 pass98 99 # 去重(按日期)100 seen = set()101 unique = []102 for k in all_klines:103 if k[0] not in seen:104 seen.add(k[0])105 unique.append(k)106 unique.sort(key=lambda x: x[0])107 108 # newfqkline 字段: [date, open, close, high, low, volume, {ma}, turnover%, amount, ?]109 # 旧接口字段: [date, open, close, high, low, volume]110 rows = []111 for row in unique:112 entry = {113 "date": row[0],114 "open": row[1],115 "close": row[2],116 "high": row[3],117 "low": row[4],118 "volume": row[5],119 }120 # newfqkline 有更多字段(第7个是dict/ma,第8个是换手率)121 if len(row) >= 8 and not isinstance(row[7], dict):122 try:123 turnover_val = float(row[7])124 if turnover_val > 0:125 entry["turnover"] = turnover_val126 except (ValueError, TypeError):127 pass128 if len(row) >= 9:129 try:130 amount_val = float(row[8]) if not isinstance(row[8], dict) else None131 if amount_val and amount_val > 0:132 entry["amount"] = amount_val * 10000 # 腾讯单位是万元133 except (ValueError, TypeError):134 pass135 rows.append(entry)136 137 df = pd.DataFrame(rows)138 df["date"] = pd.to_datetime(df["date"])139 df = df.sort_values("date").set_index("date")140 141 for col in ["open", "close", "high", "low", "volume"]:142 if col in df.columns:143 df[col] = pd.to_numeric(df[col], errors="coerce")144 145 # 计算缺失的列146 df["pct_change"] = df["close"].pct_change() * 100147 if "amount" not in df.columns:148 df["amount"] = df["close"] * df["volume"] * 100149 150 return df151 152 153def _fetch_turnover_netease(stock_code: str, start_date: str, end_date: str) -> pd.Series:154 """155 通过网易财经接口获取换手率数据,返回 Series(index=DatetimeIndex, values=turnover%)。156 失败返回 None。157 """158 try:159 prefix = "0" + stock_code if stock_code.startswith(("6", "9")) else "1" + stock_code160 sd = f"{start_date[:4]}{start_date[4:6]}{start_date[6:8]}"161 ed = f"{end_date[:4]}{end_date[4:6]}{end_date[6:8]}"162 url = (f"https://quotes.money.163.com/service/chddata.html"163 f"?code={prefix}&start={sd}&end={ed}&fields=TURNOVER")164 165 content = None166 # 优先用 requests(Streamlit Cloud 无 curl)167 try:168 import requests169 resp = requests.get(url, timeout=15)170 if resp.status_code == 200 and resp.content:171 content = resp.content.decode("gb2312", errors="ignore")172 except Exception:173 pass174 # fallback: curl175 if not content:176 result = subprocess.run(177 ["curl", "-s", "-m", "15", "-L", url],178 capture_output=True, text=True, timeout=20,179 )180 if result.returncode == 0 and result.stdout.strip():181 content = result.stdout182 183 if not content:184 return None185 from io import StringIO186 ndf = pd.read_csv(StringIO(content), engine="python")187 if ndf.empty:188 return None189 # 网易列名:日期, 股票代码, 名称, 换手率190 date_col = ndf.columns[0]191 turnover_col = [c for c in ndf.columns if "换手" in c]192 if not turnover_col:193 return None194 ndf[date_col] = pd.to_datetime(ndf[date_col])195 ndf = ndf.sort_values(date_col).set_index(date_col)196 s = pd.to_numeric(ndf[turnover_col[0]], errors="coerce")197 s.index.name = None198 return s199 except Exception:200 return None201 202 203def load_from_akshare(stock_code: str, start_date: str, end_date: str) -> pd.DataFrame:204 """205 获取个股日线数据206 优先用腾讯财经 newfqkline(稳定、含换手率),失败后切换 AKShare 备用207 """208 last_err = None209 210 # 主源:腾讯财经 newfqkline(前复权,含换手率)211 try:212 df = _fetch_via_tencent(stock_code, start_date, end_date)213 if df is not None and not df.empty:214 if "pct_change" not in df.columns or df["pct_change"].isna().all():215 df["pct_change"] = df["close"].pct_change() * 100216 return df217 except Exception as e:218 last_err = e219 220 # 备用源:AKShare(东方财富,可能被限流)221 for attempt in range(MAX_RETRIES):222 try:223 df = ak.stock_zh_a_hist(224 symbol=stock_code, period="daily",225 start_date=start_date, end_date=end_date,226 adjust="qfq",227 )228 if df is not None and not df.empty:229 df = df.rename(columns={230 "日期": "date", "开盘": "open", "收盘": "close",231 "最高": "high", "最低": "low", "成交量": "volume",232 "成交额": "amount", "涨跌幅": "pct_change", "换手率": "turnover",233 })234 df["date"] = pd.to_datetime(df["date"])235 df = df.sort_values("date").set_index("date")236 needed = ["open", "high", "low", "close", "volume", "amount", "pct_change"]237 for col in needed:238 if col not in df.columns:239 if col == "pct_change":240 df["pct_change"] = df["close"].pct_change() * 100241 return df242 except Exception as e:243 last_err = e244 if attempt < MAX_RETRIES - 1:245 time.sleep(RETRY_DELAY * (attempt + 1))246 247 msg = f"未获取到 {stock_code} 的数据(主备数据源均失败)"248 if last_err:249 msg += f"\n原始错误: {type(last_err).__name__}: {last_err}"250 raise ValueError(msg)251 252 253STOCK_NAME_CACHE = {254 "601869": "长飞光纤",255 "603601": "再升科技",256 "601138": "工业富联",257}258 259 260def get_stock_name(stock_code: str) -> str:261 """查询股票名称(优先用缓存,失败时回退到代码)"""262 if stock_code in STOCK_NAME_CACHE:263 return STOCK_NAME_CACHE[stock_code]264 try:265 info = ak.stock_info_a_code_name()266 info.columns = ["code", "name"]267 match = info[info["code"] == stock_code]268 if not match.empty:269 name = match.iloc[0]["name"]270 STOCK_NAME_CACHE[stock_code] = name271 return name272 except Exception:273 pass274 return stock_code275 276 277def validate_dataframe(df: pd.DataFrame) -> tuple:278 """279 验证上传的 DataFrame 格式280 返回: (是否通过, 错误信息列表)281 """282 errors = []283 284 missing = [c for c in REQUIRED_COLUMNS if c not in df.columns]285 if missing:286 errors.append(f"缺少必需列: {', '.join(missing)}")287 return False, errors288 289 try:290 pd.to_datetime(df["日期"])291 except Exception:292 errors.append("'日期' 列格式无法解析,请使用 YYYY-MM-DD 或 YYYYMMDD 格式")293 294 numeric_cols = [c for c in REQUIRED_COLUMNS if c != "日期"]295 for col in numeric_cols:296 if not pd.api.types.is_numeric_dtype(df[col]):297 try:298 pd.to_numeric(df[col])299 except Exception:300 errors.append(f"'{col}' 列包含非数值数据")301 302 if len(df) < 30:303 errors.append(f"数据行数不足: {len(df)} 行(至少需要 30 行)")304 305 na_counts = df[REQUIRED_COLUMNS].isna().sum()306 cols_with_na = na_counts[na_counts > 0]307 if not cols_with_na.empty:308 for col, cnt in cols_with_na.items():309 errors.append(f"'{col}' 列有 {cnt} 个缺失值")310 311 return len(errors) == 0, errors312 313 314def load_from_excel(uploaded_file) -> pd.DataFrame:315 """316 读取上传的 Excel 文件并标准化317 返回标准化 DataFrame(date 为 DatetimeIndex)318 """319 for encoding in ["utf-8", "gbk", "gb2312"]:320 try:321 df = pd.read_excel(uploaded_file, engine=None)322 break323 except Exception:324 uploaded_file.seek(0)325 continue326 else:327 raise ValueError("无法读取 Excel 文件,请确认文件格式正确")328 329 ok, errors = validate_dataframe(df)330 if not ok:331 raise ValueError("Excel 数据验证失败:\n" + "\n".join(f" - {e}" for e in errors))332 333 result = pd.DataFrame()334 for cn_col, en_col in COLUMN_MAP.items():335 if cn_col in df.columns:336 result[en_col] = df[cn_col]337 338 result["date"] = pd.to_datetime(result["date"])339 result = result.sort_values("date").set_index("date")340 341 for col in ["open", "high", "low", "close", "volume", "amount", "pct_change"]:342 if col in result.columns:343 result[col] = pd.to_numeric(result[col], errors="coerce")344 345 result = result.dropna(subset=["close"])346 return result347 348 349def generate_template() -> bytes:350 """生成可下载的 Excel 模板(含示例数据和说明)"""351 dates = pd.bdate_range("2024-01-02", periods=5)352 sample = pd.DataFrame({353 "日期": dates.strftime("%Y-%m-%d"),354 "开盘": [10.50, 10.80, 10.60, 10.90, 11.00],355 "最高": [10.90, 10.95, 10.85, 11.10, 11.20],356 "最低": [10.40, 10.55, 10.50, 10.75, 10.90],357 "收盘": [10.80, 10.60, 10.80, 11.05, 11.10],358 "成交量": [500000, 450000, 520000, 600000, 550000],359 "成交额": [5300000, 4800000, 5500000, 6500000, 6100000],360 "涨跌幅": [1.50, -1.85, 1.89, 2.31, 0.45],361 })362 363 buf = io.BytesIO()364 with pd.ExcelWriter(buf, engine="xlsxwriter") as writer:365 sample.to_excel(writer, sheet_name="日线数据", index=False)366 367 instructions = pd.DataFrame({368 "说明": [369 "请按照'日线数据'工作表的格式填入数据",370 "日期格式: YYYY-MM-DD",371 "涨跌幅单位: %(如 1.50 表示涨 1.50%)",372 "成交量单位: 股",373 "成交额单位: 元",374 "所有列均为必填项",375 "至少需要 30 个交易日的数据",376 ]377 })378 instructions.to_excel(writer, sheet_name="填写说明", index=False)379 380 buf.seek(0)381 return buf.getvalue()382 