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>
This commit is contained in:
@@ -0,0 +1,163 @@
|
||||
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)
|
||||
Reference in New Issue
Block a user