Files
simonandClaude Opus 4.7 271a9343a5 Initial commit: cc-cursor 全链路量化研究平台
7 Sprints 全部完成:
  Sprint 0: 基础设施 (DataManager + MariaDB)
  Sprint 1: 因子引擎 (34因子/12分类)
  Sprint 2: VectorBT 回测 (5策略+截面)
  Sprint 3: Optuna 优化 (+Walk-Forward)
  Sprint 4: ML 模型 (LightGBM+CatBoost)
  Sprint 5: Qwen 情绪因子 (三源新闻+日期对齐)
  Sprint 6: Agent 系统 (4Agent+日报.md/.html)

生产加固 (15项): Tushare双源fallback, SSH自动恢复, pool_pre_ping,
  save_daily先删后插, load_dotenv绝对路径, 日报5d/20d修复,
  RiskAgent改上证指数, 昨日对比+数据截止, mac_report utf8mb4,
  CLAUDE-*.md 9条已知Bug, demo全参数化, djapi数据源归一化,
  indexDatas API修正

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-06-07 15:59:05 +08:00

163 lines
5.9 KiB
Python

import pandas as pd
# 将项目根目录添加到 sys.path
#project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
#sys.path.append(project_root)
# 导入当站目录的config文件
try:
# 尝试相对导入(作为包的一部分)
from .config import TS_TOKEN, START_DATE, END_DATE
from .stock_utils import dataCorrect,get_trading_dates,tscodeCheck
except (ImportError, SystemError):
# 失败则使用绝对导入(直接运行脚本)
from config import TS_TOKEN, START_DATE, END_DATE
from stock_utils import dataCorrect,get_trading_dates,tscodeCheck
from .data_source import get_tushare_pro
pro = get_tushare_pro()
def getStockBasic(TS_CODE,START_DATE=START_DATE,END_DATE=END_DATE):
"""
从tushare daily接口获取单只股票的所有返回数据,并进行数据修正。
Parameters:
TS_CODE (str): 股票代码,格式为 '股票代码.SZ' 或 '股票代码.SH',例如 '000001.SZ'
Returns:
pd.DataFrame: 包含股票日线数据的DataFrame,字段说明如下:
ts_code: 股票代码
trade_date: 交易日期
open: 开盘价
high: 最高价
low: 最低价
close: 收盘价
pre_close: 前日收盘价
change: 涨跌额
pct_chg: 涨跌幅(百分比)
vol: 成交量(手)
amount: 成交额(千元)
Raises:
Exception: 如果从Tushare接口获取数据时发生错误。
"""
try:
# 获取股票日线数据
df = pro.daily(ts_code=TS_CODE, start_date=START_DATE, end_date=END_DATE)
# 需要修正的列
cols = ['open', 'high', 'low', 'close', 'pre_close', 'change', 'pct_chg', 'vol', 'amount']
# 调用dataCorrect函数进行数据修正
df = dataCorrect(df, cols)
return df
except Exception as e:
# 捕获并处理异常
print(f"错误: 获取股票 {TS_CODE} 的日线数据时发生错误: {e}")
return pd.DataFrame()
def getStockInfo(TS_CODE):
"""
从tushare的stock_basic接口获取单只股票的基本信息。
Parameters:
TS_CODE (str): 股票代码,格式为 '股票代码.SZ' 或 '股票代码.SH',例如 '000001.SZ'
Returns:
pd.DataFrame: 包含股票基本信息的DataFrame,字段说明如下:
ts_code: 股票代码
symbol: 股票代码(不带后缀)
name: 股票名称
area: 所在地域
industry: 所属行业
market: 市场类型(主板/创业板/科创板等)
list_date: 上市日期
fullname: 股票全称
enname: 英文全称
exchange: 交易所代码
curr_type: 交易货币
list_status: 上市状态
is_hs: 是否沪深港通标的
Raises:
Exception: 如果从Tushare接口获取数据时发生错误。
"""
try:
TS_CODE=tscodeCheck(TS_CODE)
# 调用stock_basic接口获取股票基本信息
df = pro.stock_basic(ts_code=TS_CODE)
# 如果没有获取到数据,返回空的DataFrame
if df.empty:
print(f"警告: 未找到股票 {TS_CODE} 的基本信息")
return pd.DataFrame()
return df
except Exception as e:
# 捕获并处理异常
print(f"错误: 获取股票 {TS_CODE} 的基本信息时发生错误: {e}")
return pd.DataFrame()
def getStockListByIndustry(industry):
"""
根据行业名称获取股票列表。
该函数首先通过 `index_classify` 接口获取行业分类的级别和行业代码,
然后使用 `index_member_all` 接口获取该行业下的所有股票列表。
Parameters:
industry (str): 行业名称,例如 "银行"、"医药" 等。
Returns:
pd.DataFrame: 包含股票代码 (ts_code) 和股票名称 (name) 的 DataFrame。
如果未找到匹配的行业或股票,返回空的 DataFrame。
Raises:
Exception: 如果从 Tushare 接口获取数据时发生错误。
"""
try:
# 1. 获取行业分类信息
# 使用 index_classify 接口获取所有行业分类信息
industry_df = pro.index_classify(level='', src='SW2021')
# 过滤出与给定行业名称匹配的行业
industry_info = industry_df[industry_df['industry_name'] == industry]
# 如果没有找到匹配的行业,返回空的 DataFrame
if industry_info.empty:
print(f"警告: 未找到行业 '{industry}' 的分类信息")
return pd.DataFrame(columns=['ts_code', 'name'])
# 获取行业代码和级别
industry_code = industry_info.iloc[0]['index_code']
industry_level = industry_info.iloc[0]['level']
print(industry_info)
# 2. 根据行业级别调用 index_member_all 接口
if industry_level == 'L1':
stock_list_df = pro.index_member_all(l1_code=industry_code)
elif industry_level == 'L2':
stock_list_df = pro.index_member_all(l2_code=industry_code)
elif industry_level == 'L3':
stock_list_df = pro.index_member_all(l3_code=industry_code)
else:
print(f"警告: 未知的行业级别 '{industry_level}'")
return pd.DataFrame(columns=['ts_code', 'name'])
# 如果没有找到股票,返回空的 DataFrame
if stock_list_df.empty:
print(f"警告: 行业 '{industry}' 下没有找到股票")
return pd.DataFrame(columns=['ts_code', 'name'])
# 3. 返回股票代码和名称
return stock_list_df[['ts_code', 'name']]
except Exception as e:
# 捕获并处理异常
print(f"错误: 获取行业 '{industry}' 的股票列表时发生错误: {e}")
return pd.DataFrame(columns=['ts_code', 'name'])
if __name__ == '__main__':
df = getStockListByIndustry('果蔬加工')
print(df)