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)