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,363 @@
|
||||
import pandas as pd
|
||||
from django.http import HttpResponse
|
||||
from rest_framework.response import Response
|
||||
try:
|
||||
from .config import TS_TOKEN, START_DATE, END_DATE
|
||||
except (ImportError, SystemError):
|
||||
from config import TS_TOKEN, START_DATE, END_DATE
|
||||
|
||||
from .data_source import get_tushare_pro
|
||||
|
||||
# 统一入口(全局单例,向后兼容旧代码直接访问 pro)
|
||||
pro = get_tushare_pro()
|
||||
|
||||
def dataCorrect(df, columns=None):
|
||||
"""
|
||||
检查并修正DataFrame中的NaN或None值。
|
||||
1. 如果前后有数字,该字段取前后值的均值填入
|
||||
2. 如果仅前或后有数字,该字段取copy前或后数字
|
||||
3. 若前后都为NaN,则设为0
|
||||
4. 若出发点修改,打印修改情况到屏幕
|
||||
|
||||
Parameters:
|
||||
df (pd.DataFrame): 要处理的数据框
|
||||
columns (list): 要检查的列名列表,如果为None则检查所有列
|
||||
"""
|
||||
# 如果没有指定列,则检查所有列
|
||||
if columns is None:
|
||||
columns = df.columns
|
||||
# 遍历指定列
|
||||
for col in columns:
|
||||
# 遍历每一行
|
||||
for i in range(len(df)):
|
||||
# 检查当前值是否为NaN或None
|
||||
if pd.isna(df.at[i, col]):
|
||||
# 获取前后值
|
||||
prev_val = df.at[i-1, col] if i > 0 else None
|
||||
next_val = df.at[i+1, col] if i < len(df) - 1 else None
|
||||
|
||||
# 检查前后值是否为有效数字
|
||||
prev_valid = prev_val is not None and not pd.isna(prev_val)
|
||||
next_valid = next_val is not None and not pd.isna(next_val)
|
||||
# 如果前后都有值,取均值
|
||||
if prev_valid and next_valid:
|
||||
new_val = (prev_val + next_val) / 2
|
||||
df.at[i, col] = new_val
|
||||
#print(f"修正 {col} 列第 {i} 行: 前后值均值填充为 {new_val}")
|
||||
|
||||
# 如果只有前值,取前值
|
||||
elif prev_valid:
|
||||
df.at[i, col] = prev_val
|
||||
#print(f"修正 {col} 列第 {i} 行: 前值填充为 {prev_val}")
|
||||
|
||||
# 如果只有后值,取后值
|
||||
elif next_valid:
|
||||
df.at[i, col] = next_val
|
||||
#print(f"修正 {col} 列第 {i} 行: 后值填充为 {next_val}")
|
||||
|
||||
# 如果前后都没有有效值,设为0
|
||||
else:
|
||||
df.at[i, col] = 0
|
||||
#print(f"修正 {col} 列第 {i} 行: 前后均无有效值,设为0")
|
||||
|
||||
return df
|
||||
|
||||
def dataMerge(*dfs):
|
||||
"""
|
||||
合并多个股票数据DataFrame,确保合并后的数据ts_code和trade_date匹配。
|
||||
|
||||
Parameters:
|
||||
*dfs: 可变数量的DataFrame参数,每个DataFrame应包含ts_code和trade_date列
|
||||
|
||||
Returns:
|
||||
pd.DataFrame: 合并后的DataFrame,包含所有输入DataFrame的列。如果某行的ts_code或trade_date不匹配,则用NaN填充。
|
||||
"""
|
||||
if len(dfs) < 2:
|
||||
raise ValueError("至少需要提供2个DataFrame进行合并")
|
||||
|
||||
# 检查所有DataFrame是否都有ts_code和trade_date列
|
||||
for df in dfs:
|
||||
if 'ts_code' not in df.columns or 'trade_date' not in df.columns:
|
||||
raise ValueError("所有DataFrame必须包含ts_code和trade_date列")
|
||||
|
||||
# 初始化合并结果为第一个DataFrame
|
||||
merged_df = dfs[0].copy()
|
||||
|
||||
# 逐个合并剩余的DataFrame
|
||||
for df in dfs[1:]:
|
||||
merged_df = pd.merge(merged_df, df, on=['ts_code', 'trade_date'], how='outer')
|
||||
|
||||
return merged_df
|
||||
|
||||
def tscodeCheck(tscode):
|
||||
"""
|
||||
检查并修正股票代码格式。
|
||||
|
||||
Parameters:
|
||||
tscode (str): 股票代码
|
||||
|
||||
Returns:
|
||||
str: 格式正确的股票代码 (如 '000001.SZ')
|
||||
|
||||
Raises:
|
||||
ValueError: 如果输入不符合要求
|
||||
"""
|
||||
if not isinstance(tscode, str):
|
||||
raise ValueError("输入必须是字符串")
|
||||
|
||||
tscode = tscode.upper() # 转为大写
|
||||
|
||||
# 检查长度
|
||||
if len(tscode) < 6:
|
||||
raise ValueError("股票代码长度不能小于6位")
|
||||
elif len(tscode) == 6:
|
||||
if not tscode.isdigit():
|
||||
raise ValueError("6位股票代码必须全为数字")
|
||||
# 根据开头添加后缀
|
||||
if tscode.startswith('00') or tscode.startswith('3'):
|
||||
return f"{tscode}.SZ"
|
||||
elif tscode.startswith(('60','68')): # 68为科创板
|
||||
return f"{tscode}.SH"
|
||||
elif tscode.startswith(('8', '9')):
|
||||
return f"{tscode}.BJ"
|
||||
else:
|
||||
raise ValueError("未知的6位股票代码开头")
|
||||
elif len(tscode) == 9:
|
||||
prefix = tscode[:6]
|
||||
suffix = tscode[-3:]
|
||||
if not prefix.isdigit():
|
||||
raise ValueError("9位股票代码前6位必须为数字")
|
||||
if suffix not in ('.SZ', '.SH', '.BJ'):
|
||||
raise ValueError("9位股票代码后缀必须是.SZ/.SH/.BJ")
|
||||
return tscode
|
||||
else:
|
||||
raise ValueError("股票代码长度必须为6位或9位")
|
||||
|
||||
|
||||
|
||||
def get_trading_dates(trade_date):
|
||||
"""
|
||||
获取上一个季度财报日期到给定日期的所有交易日期
|
||||
|
||||
Parameters:
|
||||
trade_date (str): 给定日期 (格式yyyymmdd)
|
||||
Returns:
|
||||
list: 交易日期列表 (格式yyyymmdd)
|
||||
"""
|
||||
# 将输入日期转为datetime
|
||||
try:
|
||||
current_date = pd.to_datetime(trade_date, format='%Y%m%d')
|
||||
except:
|
||||
raise ValueError("trade_date格式应为yyyymmdd")
|
||||
|
||||
# 计算上一个季度财报日期 (3/31, 6/30, 9/30, 12/31)
|
||||
year = current_date.year
|
||||
month = current_date.month
|
||||
if month < 4:
|
||||
prev_report_date = pd.Timestamp(year-1, 12, 31)
|
||||
elif month < 7:
|
||||
prev_report_date = pd.Timestamp(year, 3, 31)
|
||||
elif month < 10:
|
||||
prev_report_date = pd.Timestamp(year, 6, 30)
|
||||
else:
|
||||
prev_report_date = pd.Timestamp(year, 9, 30)
|
||||
|
||||
# 获取交易日历
|
||||
calendar_df = pro.trade_cal(exchange='', start_date=prev_report_date.strftime('%Y%m%d'), end_date=trade_date)
|
||||
|
||||
# 过滤交易日历
|
||||
calendar_df['cal_date_dt'] = pd.to_datetime(calendar_df['cal_date'], format='%Y%m%d')
|
||||
filtered_dates = calendar_df[
|
||||
(calendar_df['cal_date_dt'] > prev_report_date) &
|
||||
(calendar_df['cal_date_dt'] <= current_date) &
|
||||
(calendar_df['is_open'] == 1)
|
||||
]
|
||||
|
||||
# 返回日期列表 (格式yyyymmdd)
|
||||
return filtered_dates['cal_date'].tolist()
|
||||
|
||||
def get_next_report_date(trade_date):
|
||||
"""
|
||||
获取给定日期的下一个财报日期
|
||||
|
||||
Parameters:
|
||||
trade_date (str): 给定日期 (格式yyyymmdd)
|
||||
|
||||
Returns:
|
||||
str: 下一个财报日期 (格式yyyymmdd)
|
||||
"""
|
||||
try:
|
||||
current_date = pd.to_datetime(trade_date, format='%Y%m%d')
|
||||
except:
|
||||
raise ValueError("trade_date格式应为yyyymmdd")
|
||||
|
||||
year = current_date.year
|
||||
month = current_date.month
|
||||
day = current_date.day
|
||||
|
||||
if month < 3 or (month == 3 and day < 31):
|
||||
return f"{year}0331"
|
||||
elif month < 6 or (month == 6 and day < 30):
|
||||
return f"{year}0630"
|
||||
elif month < 9 or (month == 9 and day < 30):
|
||||
return f"{year}0930"
|
||||
elif month < 12 or (month == 12 and day < 31):
|
||||
return f"{year}1231"
|
||||
else:
|
||||
return f"{year+1}0331"
|
||||
|
||||
def viewFunc_singleParam(request, data_func, param_name='tscode', default_value=None):
|
||||
"""
|
||||
通用单参数视图包装器。
|
||||
|
||||
:param request: Django request 对象
|
||||
:param data_func: 接收单个参数的数据处理函数
|
||||
:param param_name: URL 查询参数名
|
||||
:param default_value: 参数默认值
|
||||
:return: Response
|
||||
"""
|
||||
param = request.GET.get(param_name, default_value)
|
||||
if not param:
|
||||
return Response({'error': f'缺少 {param_name} 参数'}, status=400)
|
||||
|
||||
try:
|
||||
data = data_func(param)
|
||||
dict_data = data.to_dict(orient='records') if isinstance(data, pd.DataFrame) else data
|
||||
return Response(dict_data)
|
||||
except ImportError:
|
||||
return Response({'error': f'模块不存在'}, status=500)
|
||||
except Exception as e:
|
||||
return Response({'error': str(e)}, status=500)
|
||||
|
||||
|
||||
def viewFunc_tsCodeAndDate(request,data_func):
|
||||
"""
|
||||
通用数据处理函数
|
||||
:param request: Django request对象
|
||||
:param data_func: 数据处理函数(需返回Response)
|
||||
:return: Response
|
||||
"""
|
||||
#通过url 传入 tscode 变量
|
||||
tscode = request.GET.get('tscode','000001.SZ') # 从URL获取industry参数
|
||||
tscode = tscodeCheck(tscode)
|
||||
if not tscode:
|
||||
return Response({'error': '缺少 tscode 参数'}, status=400)
|
||||
|
||||
start_date = request.GET.get('start_date', START_DATE)
|
||||
end_date = request.GET.get('end_date', END_DATE)
|
||||
|
||||
try:
|
||||
data = data_func(tscode, start_date, end_date) # 调用函数
|
||||
dict_data = data.to_dict(orient='records') if isinstance(data, pd.DataFrame) else data
|
||||
return Response(dict_data)
|
||||
except ImportError:
|
||||
return Response({'error': '模块 stock_basic 不存在'}, status=500)
|
||||
except Exception as e:
|
||||
return Response({'error': str(e)}, status=500)
|
||||
|
||||
def date_format_correction(date_str):
|
||||
"""
|
||||
日期格式矫正函数:
|
||||
1. 如果输入格式为yyyy-mm-dd则返回yyyymmdd格式
|
||||
2. 如果输入格式已经是yyyymmdd则直接返回
|
||||
3. 其他情况返回None
|
||||
|
||||
Parameters:
|
||||
date_str (str): 日期字符串
|
||||
|
||||
Returns:
|
||||
str: yyyymmdd格式的日期字符串,如果输入格式不正确则返回None
|
||||
"""
|
||||
try:
|
||||
# 先尝试解析yyyymmdd格式
|
||||
pd.to_datetime(date_str, format='%Y%m%d')
|
||||
return date_str
|
||||
except ValueError:
|
||||
try:
|
||||
# 再尝试解析yyyy-mm-dd格式
|
||||
date_obj = pd.to_datetime(date_str, format='%Y-%m-%d')
|
||||
return date_obj.strftime('%Y%m%d')
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
def get_released_report_dates(start_date: str, end_date: str) -> list:
|
||||
"""
|
||||
根据起止日期生成所有财报季末日,并去除尚未可能披露的日期。
|
||||
|
||||
:param start_date: 开始日期,格式为 "YYYYMMDD"
|
||||
:param end_date: 结束日期,格式为 "YYYYMMDD"
|
||||
:param today: 当前日期,格式为 "YYYYMMDD",默认为系统当前日期
|
||||
:return: 已披露的财报季末日期列表,格式为 ["20240331", "20240630", ...]
|
||||
"""
|
||||
|
||||
today = datetime.today()
|
||||
|
||||
# 财报发布日期最晚规则
|
||||
report_deadlines = {
|
||||
"0331": "0430",
|
||||
"0630": "0831",
|
||||
"0930": "1031",
|
||||
"1231": "0430" # 次年4月30日
|
||||
}
|
||||
|
||||
# 生成所有财报季度末日期
|
||||
years = range(int(start_date[:4]), int(end_date[:4]) + 1)
|
||||
report_dates = [
|
||||
f"{year}{quarter}"
|
||||
for year in years
|
||||
for quarter in ["0331", "0630", "0930", "1231"]
|
||||
]
|
||||
|
||||
# 筛选日期范围内的
|
||||
report_dates = [d for d in report_dates if start_date <= d <= end_date]
|
||||
|
||||
def is_report_released(report_date_str):
|
||||
report_dt = datetime.strptime(report_date_str, "%Y%m%d")
|
||||
year = report_dt.year
|
||||
q_end = report_date_str[4:]
|
||||
|
||||
if q_end == "1231":
|
||||
deadline = datetime.strptime(f"{year+1}0430", "%Y%m%d")
|
||||
else:
|
||||
deadline = datetime.strptime(f"{year}{report_deadlines[q_end]}", "%Y%m%d")
|
||||
|
||||
return today >= deadline
|
||||
|
||||
# 返回已发布的日期
|
||||
return [d for d in report_dates if is_report_released(d)]
|
||||
|
||||
'''
|
||||
写一个函数,调用stock_basic API(说明文档:https://tushare.pro/document/2?doc_id=25) 获取个股清单
|
||||
给定交易所代码,仅查询当前仍然上市的股票
|
||||
返回:
|
||||
- ts_code: TS代码
|
||||
- symbol: 股票代码
|
||||
- name: 股票名称
|
||||
- fullname: 股票全称
|
||||
- exchange: 交易所代码 SSE上交所 SZSE深交所 BSE北交所
|
||||
- list_status: 上市状态 L上市 D退市 P暂停上市
|
||||
'''
|
||||
def get_stock_basic(exchange='SSE'):
|
||||
"""
|
||||
获取指定交易所的上市股票基本信息
|
||||
|
||||
Parameters:
|
||||
exchange (str): 交易所代码 SSE上交所 SZSE深交所 BSE北交所
|
||||
|
||||
Returns:
|
||||
pd.DataFrame: 包含股票基本信息的DataFrame
|
||||
"""
|
||||
# 调用stock_basic接口
|
||||
df = pro.stock_basic(exchange=exchange, list_status='L',
|
||||
fields='ts_code,symbol,name,fullname,exchange,list_status')
|
||||
|
||||
return df
|
||||
|
||||
if __name__ == "__main__":
|
||||
# 测试当前日期之前的财报
|
||||
print("\n测试2: 当前日期之前的财报")
|
||||
result = get_released_report_dates("20230101", "20251230")
|
||||
print(f"结果: {result}")
|
||||
Reference in New Issue
Block a user