diff --git a/.gitignore b/.gitignore index ef6c1b0..1dfa6a0 100644 --- a/.gitignore +++ b/.gitignore @@ -9,3 +9,4 @@ .env /.cursor /openspec +TradEnv/ \ No newline at end of file diff --git a/.streamlit/config.toml b/.streamlit/config.toml new file mode 100644 index 0000000..f1f6de3 --- /dev/null +++ b/.streamlit/config.toml @@ -0,0 +1,6 @@ +[theme] +base = "light" + +[server] +port = 8503 +address = "127.0.0.1" \ No newline at end of file diff --git a/ai_agents.py b/ai_agents.py index 6ec1baa..75ea06e 100644 --- a/ai_agents.py +++ b/ai_agents.py @@ -55,10 +55,10 @@ class StockAnalysisAgents: print("💰 资金面分析师正在分析中...") # 如果有资金流向数据,显示数据来源 - if fund_flow_data and fund_flow_data.get('query_success'): - print(" ✓ 已获取问财资金流向数据") + if fund_flow_data and fund_flow_data.get('data_success'): + print(" ✓ 已获取资金流向数据(akshare数据源)") else: - print(" ⚠ 未获取到问财资金流向数据,将基于技术指标分析") + print(" ⚠ 未获取到资金流向数据,将基于技术指标分析") time.sleep(1) diff --git a/app.py b/app.py index ecc86b9..921c3c3 100644 --- a/app.py +++ b/app.py @@ -302,7 +302,7 @@ def main(): st.markdown("---") # 🎯 选股板块 - with st.expander("🎯 选股板块", expanded=False): + with st.expander("🎯 选股板块", expanded=True): st.markdown("**根据不同策略筛选优质股票**") if st.button("💰 主力选股", width='stretch', key="nav_main_force", help="基于主力资金流向的选股策略"): @@ -313,7 +313,7 @@ def main(): del st.session_state[key] # 📊 策略分析 - with st.expander("📊 策略分析", expanded=False): + with st.expander("📊 策略分析", expanded=True): st.markdown("**AI驱动的板块和龙虎榜策略**") if st.button("🎯 智策板块", width='stretch', key="nav_sector_strategy", help="AI板块策略分析"): @@ -331,7 +331,7 @@ def main(): del st.session_state[key] # 💼 投资管理 - with st.expander("💼 投资管理", expanded=False): + with st.expander("💼 投资管理", expanded=True): st.markdown("**持仓跟踪与实时监测**") if st.button("📊 持仓分析", width='stretch', key="nav_portfolio", help="投资组合分析与定时跟踪"): @@ -571,9 +571,9 @@ def main(): help="负责风险识别、风险评估、风险控制策略制定") with col3: - enable_sentiment = st.checkbox("📈 市场情绪分析师", value=False, + enable_sentiment = st.checkbox("📈 市场情绪分析师", value=True, help="负责市场情绪研究、ARBR指标分析(仅A股)") - enable_news = st.checkbox("📰 新闻分析师", value=False, + enable_news = st.checkbox("📰 新闻分析师", value=True, help="负责新闻事件分析、舆情研究(仅A股,qstock数据源)") # 显示已选择的分析师 @@ -1427,7 +1427,7 @@ def display_stock_chart(stock_data, stock_info): # 生成唯一的key chart_key = f"main_stock_chart_{stock_info.get('symbol', 'unknown')}_{int(time.time())}" - st.plotly_chart(fig, width='stretch', key=chart_key) + st.plotly_chart(fig, use_container_width=True, config={'responsive': True}, key=chart_key) # 成交量图 if 'Volume' in stock_data.columns: @@ -1448,7 +1448,7 @@ def display_stock_chart(stock_data, stock_info): # 生成唯一的key volume_key = f"volume_chart_{stock_info.get('symbol', 'unknown')}_{int(time.time())}" - st.plotly_chart(fig_volume, width='stretch', key=volume_key) + st.plotly_chart(fig_volume, use_container_width=True, config={'responsive': True}, key=volume_key) def display_agents_analysis(agents_results): """显示各分析师报告""" diff --git a/docs/toshare说明.md b/docs/toshare说明.md new file mode 100644 index 0000000..a46cb13 --- /dev/null +++ b/docs/toshare说明.md @@ -0,0 +1,100 @@ +Python SDK +下载SDK +下载并安装最新版tushare SDK 【安装和升级方法】 + +导入tushare + +import tushare as ts +这里注意, tushare版本需大于1.2.10 + +设置token + +ts.set_token('your token here') +以上方法只需要在第一次或者token失效后调用,完成调取tushare数据凭证的设置,正常情况下不需要重复设置。也可以忽略此步骤,直接用pro_api('your token')完成初始化 + +初始化pro接口 + +pro = ts.pro_api() +如果上一步骤ts.set_token('your token')无效或不想保存token到本地,也可以在初始化接口里直接设置token: + +pro = ts.pro_api('your token') +数据调取 + +以获取交易日历信息为例: + +df = pro.trade_cal(exchange='', start_date='20180901', end_date='20181001', fields='exchange,cal_date,is_open,pretrade_date', is_open='0') +或者 + +df = pro.query('trade_cal', exchange='', start_date='20180901', end_date='20181001', fields='exchange,cal_date,is_open,pretrade_date', is_open='0') +调取结果: + + exchange cal_date is_open pretrade_date +0 SSE 20180901 0 20180831 +1 SSE 20180902 0 20180831 +2 SSE 20180908 0 20180907 +3 SSE 20180909 0 20180907 +4 SSE 20180915 0 20180914 +5 SSE 20180916 0 20180914 +6 SSE 20180922 0 20180921 +7 SSE 20180923 0 20180921 +8 SSE 20180924 0 20180921 +9 SSE 20180929 0 20180928 +10 SSE 20180930 0 20180928 +11 SSE 20181001 0 20180928 + + +沪深港通资金流向 +接口:moneyflow_hsgt,可以通过数据工具调试和查看数据。 +描述:获取沪股通、深股通、港股通每日资金流向数据,每次最多返回300条记录,总量不限制。每天18~20点之间完成当日更新 +积分要求:2000积分起,5000积分每分钟可提取500次 + +输入参数 + +名称 类型 必选 描述 +trade_date str N 交易日期 (二选一) +start_date str N 开始日期 (二选一) +end_date str N 结束日期 +输出参数 + +名称 类型 描述 +trade_date str 交易日期 +ggt_ss float 港股通(上海) +ggt_sz float 港股通(深圳) +hgt float 沪股通(百万元) +sgt float 深股通(百万元) +north_money float 北向资金(百万元) +south_money float 南向资金(百万元) +接口用法 + + +pro = ts.pro_api() + +pro.moneyflow_hsgt(start_date='20180125', end_date='20180808') +或者 + + +pro.query('moneyflow_hsgt', trade_date='20180725') +数据样例 + + trade_date ggt_ss ggt_sz hgt sgt north_money south_money +0 20180808 -476.0 -188.0 962.68 799.94 1762.62 -664.0 +1 20180807 -261.0 177.0 2140.85 1079.82 3220.67 -84.0 +2 20180803 667.0 -32.0 -436.99 1088.07 651.08 635.0 +3 20180802 -1651.0 -366.0 874.97 -216.65 658.32 -2017.0 +4 20180801 -1443.0 -443.0 544.36 542.79 1087.15 -1886.0 +5 20180731 -299.0 -21.0 1923.72 1345.48 3269.20 -320.0 +6 20180730 -588.0 611.0 2536.54 146.24 2682.78 23.0 +7 20180727 -13.0 363.0 2182.84 533.06 2715.90 350.0 +8 20180726 -566.0 -339.0 1113.28 -567.47 545.81 -905.0 +9 20180725 319.0 370.0 1470.29 311.27 1781.56 689.0 +10 20180724 924.0 2312.0 1748.88 1053.52 2802.40 3236.0 +11 20180723 1628.0 1172.0 -279.96 334.82 54.86 2800.0 +12 20180720 2233.0 1773.0 606.33 1711.77 2318.10 4006.0 +13 20180719 456.0 206.0 1831.41 874.40 2705.81 662.0 +14 20180718 -181.0 261.0 126.80 -111.83 14.97 80.0 +15 20180717 -390.0 187.0 -90.32 -404.24 -494.56 -203.0 +16 20180716 -539.0 52.0 -457.00 487.60 30.60 -487.0 +17 20180713 -297.0 751.0 599.38 658.07 1257.45 454.0 +18 20180712 2635.0 1699.0 1695.62 269.56 1965.18 4334.0 +19 20180711 19.0 646.0 261.96 -339.20 -77.24 665.0 +20 20180710 668.0 889.0 514.05 262.33 776.38 1557.0 \ No newline at end of file diff --git a/fund_flow_akshare.py b/fund_flow_akshare.py index eb60052..ebe992c 100644 --- a/fund_flow_akshare.py +++ b/fund_flow_akshare.py @@ -36,7 +36,7 @@ class FundFlowAkshareDataFetcher: """资金流向数据获取类(使用akshare数据源)""" def __init__(self): - self.days = 20 # 获取最近20个交易日 + self.days = 30 # 获取最近30个交易日 self.available = True print("[OK] 资金流向数据获取器初始化成功(akshare数据源)") @@ -259,7 +259,7 @@ class FundFlowAkshareDataFetcher: # 添加统计汇总 text_parts.append(""" ═══════════════════════════════════════ -[统计汇总 - 最近20个交易日] +[统计汇总 - 最近30个交易日] ═══════════════════════════════════════ """) diff --git a/longhubang.db b/longhubang.db index f34c53d..8703486 100644 Binary files a/longhubang.db and b/longhubang.db differ diff --git a/longhubang_db.py b/longhubang_db.py index 25040e8..e9db828 100644 --- a/longhubang_db.py +++ b/longhubang_db.py @@ -7,6 +7,7 @@ import sqlite3 from datetime import datetime import json import pandas as pd +import logging class LonghubangDatabase: @@ -20,6 +21,10 @@ class LonghubangDatabase: db_path: 数据库文件路径 """ self.db_path = db_path + # 初始化日志 + self.logger = logging.getLogger(__name__) + if not self.logger.handlers: + logging.basicConfig(level=logging.INFO, format='[%(asctime)s] %(levelname)s %(name)s: %(message)s') self.init_database() def get_connection(self): @@ -100,7 +105,7 @@ class LonghubangDatabase: conn.commit() conn.close() - print("[智瞰龙虎] 数据库初始化完成") + self.logger.info("[智瞰龙虎] 数据库初始化完成") def save_longhubang_data(self, data_list): """ @@ -141,13 +146,13 @@ class LonghubangDatabase: )) saved_count += 1 except Exception as e: - print(f"保存记录失败: {e}") + self.logger.exception(f"保存记录失败: {e}", exc_info=True) continue conn.commit() conn.close() - print(f"[智瞰龙虎] 成功保存 {saved_count} 条龙虎榜记录") + self.logger.info(f"[智瞰龙虎] 成功保存 {saved_count} 条龙虎榜记录") return saved_count def get_longhubang_data(self, start_date=None, end_date=None, stock_code=None): @@ -321,7 +326,7 @@ class LonghubangDatabase: conn.commit() conn.close() - print(f"[智瞰龙虎] 分析报告已保存 (ID: {report_id})") + self.logger.info(f"[智瞰龙虎] 分析报告已保存 (ID: {report_id})") return report_id def get_analysis_reports(self, limit=10): @@ -365,31 +370,73 @@ class LonghubangDatabase: ''', (report_id,)) row = cursor.fetchone() + # 在关闭连接之前获取列名,避免关闭后访问游标属性报错 + columns = [desc[0] for desc in cursor.description] if cursor.description else [] conn.close() if row: - columns = [desc[0] for desc in cursor.description] report = dict(zip(columns, row)) # 解析JSON字段 if report.get('recommended_stocks'): try: report['recommended_stocks'] = json.loads(report['recommended_stocks']) - except: - pass + except Exception as e: + self.logger.warning(f"推荐股票JSON解析失败: {e}") # 解析analysis_content字段 if report.get('analysis_content'): try: report['analysis_content_parsed'] = json.loads(report['analysis_content']) - except: + except json.JSONDecodeError as e: # 如果不是JSON格式,保持原样 report['analysis_content_parsed'] = None + self.logger.debug(f"analysis_content字段不是有效JSON格式,将保持原始文本格式: {str(e)[:100]}") + except Exception as e: + report['analysis_content_parsed'] = None + self.logger.warning(f"analysis_content字段解析时发生未知错误: {str(e)[:100]}") return report return None + def delete_analysis_report(self, report_id): + """ + 删除分析报告 + + Args: + report_id: 报告ID + + Returns: + bool: 删除是否成功 + """ + conn = self.get_connection() + cursor = conn.cursor() + + try: + # 先删除相关的股票追踪记录 + cursor.execute('DELETE FROM stock_tracking WHERE analysis_id = ?', (report_id,)) + + # 删除分析报告 + cursor.execute('DELETE FROM longhubang_analysis WHERE id = ?', (report_id,)) + + deleted_count = cursor.rowcount + conn.commit() + + if deleted_count > 0: + self.logger.info(f"[智瞰龙虎] 成功删除分析报告 (ID: {report_id})") + return True + else: + self.logger.warning(f"[智瞰龙虎] 未找到要删除的分析报告 (ID: {report_id})") + return False + + except Exception as e: + self.logger.error(f"[智瞰龙虎] 删除分析报告失败: {e}") + conn.rollback() + return False + finally: + conn.close() + def update_stock_tracking(self, analysis_id, stock_code, current_price, status, notes=None): """ 更新股票追踪信息 diff --git a/longhubang_engine.py b/longhubang_engine.py index 649cf96..69ecd06 100644 --- a/longhubang_engine.py +++ b/longhubang_engine.py @@ -10,6 +10,7 @@ from longhubang_scoring import LonghubangScoring from typing import Dict, Any, List from datetime import datetime, timedelta import time +import logging class LonghubangEngine: @@ -27,7 +28,11 @@ class LonghubangEngine: self.database = LonghubangDatabase(db_path) self.agents = LonghubangAgents(model=model) self.scoring = LonghubangScoring() - print(f"[智瞰龙虎] 分析引擎初始化完成") + # 初始化日志 + self.logger = logging.getLogger(__name__) + if not self.logger.handlers: + logging.basicConfig(level=logging.INFO, format='[%(asctime)s] %(levelname)s %(name)s: %(message)s') + self.logger.info("[智瞰龙虎] 分析引擎初始化完成") def run_comprehensive_analysis(self, date=None, days=1) -> Dict[str, Any]: """ @@ -40,9 +45,9 @@ class LonghubangEngine: Returns: 完整的分析结果 """ - print("\n" + "=" * 60) - print("🚀 智瞰龙虎综合分析系统启动") - print("=" * 60) + self.logger.info("=" * 60) + self.logger.info("🚀 智瞰龙虎综合分析系统启动") + self.logger.info("=" * 60) results = { "success": False, @@ -55,8 +60,8 @@ class LonghubangEngine: try: # 阶段1: 获取龙虎榜数据 - print("\n[阶段1] 获取龙虎榜数据...") - print("-" * 60) + self.logger.info("[阶段1] 获取龙虎榜数据...") + self.logger.info("-" * 60) if date: data_list = [self.data_fetcher.get_longhubang_data(date)] @@ -65,21 +70,21 @@ class LonghubangEngine: data_list = self.data_fetcher.get_recent_days_data(days) if not data_list: - print("✗ 未获取到龙虎榜数据") + self.logger.error("未获取到龙虎榜数据") results["error"] = "未获取到龙虎榜数据" return results - - print(f"✓ 成功获取 {len(data_list)} 条龙虎榜记录") + + self.logger.info(f"成功获取 {len(data_list)} 条龙虎榜记录") # 阶段2: 保存数据到数据库 - print("\n[阶段2] 保存数据到数据库...") - print("-" * 60) + self.logger.info("[阶段2] 保存数据到数据库...") + self.logger.info("-" * 60) saved_count = self.database.save_longhubang_data(data_list) - print(f"✓ 保存 {saved_count} 条记录") + self.logger.info(f"保存 {saved_count} 条记录") # 阶段3: 数据分析和统计 - print("\n[阶段3] 数据分析和统计...") - print("-" * 60) + self.logger.info("[阶段3] 数据分析和统计...") + self.logger.info("-" * 60) summary = self.data_fetcher.analyze_data_summary(data_list) formatted_data = self.data_fetcher.format_data_for_ai(data_list, summary) @@ -89,83 +94,86 @@ class LonghubangEngine: "total_youzi": summary.get('total_youzi', 0), "summary": summary } - print(f"✓ 数据统计完成") + self.logger.info("数据统计完成") # 阶段3.5: AI智能评分排名 - print("\n[阶段3.5] AI智能评分排名...") - print("-" * 60) + self.logger.info("[阶段3.5] AI智能评分排名...") + self.logger.info("-" * 60) scoring_df = self.scoring.score_all_stocks(data_list) - results["scoring_ranking"] = scoring_df - print(f"✓ 完成 {len(scoring_df)} 只股票的智能评分排名") + # 转换为可序列化格式以避免UI/存储类型问题 + scoring_ranking_data: List[Dict[str, Any]] = [] + try: + if scoring_df is not None and hasattr(scoring_df, 'to_dict'): + scoring_ranking_data = scoring_df.to_dict('records') + self.logger.info(f"完成 {len(scoring_ranking_data)} 只股票的智能评分排名") + else: + self.logger.warning("评分结果为空或格式不支持转换") + except Exception as e: + self.logger.exception(f"评分排名数据转换失败: {e}", exc_info=True) + scoring_ranking_data = [] + results["scoring_ranking"] = scoring_ranking_data # 阶段4: AI分析师团队分析 - print("\n[阶段4] AI分析师团队工作中...") - print("-" * 60) + self.logger.info("[阶段4] AI分析师团队工作中...") + self.logger.info("-" * 60) agents_results = {} # 1. 游资行为分析师 - print("1/5 游资行为分析师...") + self.logger.info("1/5 游资行为分析师...") youzi_result = self.agents.youzi_behavior_analyst(formatted_data, summary) agents_results["youzi"] = youzi_result # 2. 个股潜力分析师 - print("2/5 个股潜力分析师...") + self.logger.info("2/5 个股潜力分析师...") stock_result = self.agents.stock_potential_analyst(formatted_data, summary) agents_results["stock"] = stock_result # 3. 题材追踪分析师 - print("3/5 题材追踪分析师...") + self.logger.info("3/5 题材追踪分析师...") theme_result = self.agents.theme_tracker_analyst(formatted_data, summary) agents_results["theme"] = theme_result # 4. 风险控制专家 - print("4/5 风险控制专家...") + self.logger.info("4/5 风险控制专家...") risk_result = self.agents.risk_control_specialist(formatted_data, summary) agents_results["risk"] = risk_result # 5. 首席策略师综合 - print("5/5 首席策略师综合分析...") + self.logger.info("5/5 首席策略师综合分析...") all_analyses = [youzi_result, stock_result, theme_result, risk_result] chief_result = self.agents.chief_strategist(all_analyses) agents_results["chief"] = chief_result results["agents_analysis"] = agents_results - print("\n✓ 所有AI分析师分析完成") + self.logger.info("所有AI分析师分析完成") # 阶段5: 提取推荐股票 - print("\n[阶段5] 提取推荐股票...") - print("-" * 60) + self.logger.info("[阶段5] 提取推荐股票...") + self.logger.info("-" * 60) recommended_stocks = self._extract_recommended_stocks( chief_result.get('analysis', ''), stock_result.get('analysis', ''), summary ) results["recommended_stocks"] = recommended_stocks - print(f"✓ 提取 {len(recommended_stocks)} 只推荐股票") + self.logger.info(f"提取 {len(recommended_stocks)} 只推荐股票") # 阶段6: 生成最终报告 - print("\n[阶段6] 生成最终报告...") - print("-" * 60) + self.logger.info("[阶段6] 生成最终报告...") + self.logger.info("-" * 60) final_report = self._generate_final_report(agents_results, summary, recommended_stocks) results["final_report"] = final_report - print("✓ 最终报告生成完成") + self.logger.info("最终报告生成完成") # 阶段7: 保存完整分析报告到数据库 - print("\n[阶段7] 保存完整分析报告...") - print("-" * 60) + self.logger.info("[阶段7] 保存完整分析报告...") + self.logger.info("-" * 60) data_date_range = self._get_date_range(data_list) # 转换评分排名数据为可序列化格式 - scoring_ranking_data = [] - if scoring_df is not None and hasattr(scoring_df, 'to_dict'): - try: - # 转换DataFrame为字典列表,确保所有数据都被序列化 - scoring_ranking_data = scoring_df.to_dict('records') - print(f"✓ 评分排名数据已转换: {len(scoring_ranking_data)} 条记录") - except Exception as e: - print(f"⚠ 评分排名数据转换失败: {e}") - scoring_ranking_data = [] + # 复用前面转换的评分数据 + # 若前面转换失败,此处不再重复转换,避免错误 # 构建完整的分析内容(结构化) full_analysis_content = { @@ -184,20 +192,18 @@ class LonghubangEngine: full_result=results # 传入完整结果 ) results["report_id"] = report_id - print(f"✓ 完整报告已保存 (ID: {report_id})") + self.logger.info(f"完整报告已保存 (ID: {report_id})") results["success"] = True - print("\n" + "=" * 60) - print("✓ 智瞰龙虎综合分析完成!") - print("=" * 60) + self.logger.info("=" * 60) + self.logger.info("✓ 智瞰龙虎综合分析完成!") + self.logger.info("=" * 60) except Exception as e: - print(f"\n✗ 分析过程出错: {e}") - import traceback - traceback.print_exc() + self.logger.exception(f"分析过程出错: {e}", exc_info=True) results["error"] = str(e) - + return results def _extract_recommended_stocks(self, chief_analysis: str, stock_analysis: str, summary: Dict) -> List[Dict]: diff --git a/longhubang_scoring.py b/longhubang_scoring.py index 57bf961..d939327 100644 --- a/longhubang_scoring.py +++ b/longhubang_scoring.py @@ -445,6 +445,7 @@ class LonghubangScoring: results.append({ '排名': 0, # 稍后填充 + '排名_display': '', # 用于显示奖牌 '股票名称': stock_info['name'], '股票代码': code, '综合评分': round(total_score, 1), @@ -467,13 +468,14 @@ class LonghubangScoring: df = df.sort_values('综合评分', ascending=False).reset_index(drop=True) df['排名'] = range(1, len(df) + 1) - # 添加奖牌 + # 添加奖牌显示 + df['排名_display'] = df['排名'].astype(str) if len(df) >= 1: - df.loc[0, '排名'] = '🥇 1' + df.loc[0, '排名_display'] = '🥇 1' if len(df) >= 2: - df.loc[1, '排名'] = '🥈 2' + df.loc[1, '排名_display'] = '🥈 2' if len(df) >= 3: - df.loc[2, '排名'] = '🥉 3' + df.loc[2, '排名_display'] = '🥉 3' return df diff --git a/longhubang_ui.py b/longhubang_ui.py index 1a80b88..db8ab46 100644 --- a/longhubang_ui.py +++ b/longhubang_ui.py @@ -146,10 +146,10 @@ def display_analysis_tab(): col1, col2, col3 = st.columns([2, 2, 2]) with col1: - analyze_button = st.button("🚀 开始分析", type="primary", use_container_width=True) + analyze_button = st.button("🚀 开始分析", type="primary", width='stretch') with col2: - if st.button("🔄 清除结果", use_container_width=True): + if st.button("🔄 清除结果", width='stretch'): if 'longhubang_result' in st.session_state: del st.session_state.longhubang_result st.success("已清除分析结果") @@ -335,13 +335,29 @@ def display_scoring_ranking(result): # 显示TOP10评分表格 st.markdown("### 🥇 TOP10 综合评分排名") + # 兼容历史数据与类型统一,避免 Arrow 序列化错误 + if isinstance(scoring_df, list): + scoring_df = pd.DataFrame(scoring_df) + + numeric_cols = ['排名','综合评分','资金含金量','净买入额','卖出压力','机构共振','加分项','顶级游资','买方数','净流入'] + for col in numeric_cols: + if col in scoring_df.columns: + scoring_df[col] = pd.to_numeric(scoring_df[col], errors='coerce') + + text_cols = ['股票名称','股票代码','机构参与'] + for col in text_cols: + if col in scoring_df.columns: + scoring_df[col] = scoring_df[col].astype(str) + top10_df = scoring_df.head(10).copy() + if '排名' in top10_df.columns: + top10_df['排名'] = pd.to_numeric(top10_df['排名'], errors='coerce').fillna(0).astype(int) # 格式化显示 st.dataframe( top10_df, column_config={ - "排名": st.column_config.TextColumn("排名", width="small"), + "排名": st.column_config.NumberColumn("排名", format="%d", width="small"), "股票名称": st.column_config.TextColumn("股票名称", width="medium"), "股票代码": st.column_config.TextColumn("代码", width="small"), "综合评分": st.column_config.NumberColumn( @@ -385,7 +401,7 @@ def display_scoring_ranking(result): "净流入": st.column_config.NumberColumn("净流入(元)", format="%.2f") }, hide_index=True, - use_container_width=True + width='stretch' ) # 一键批量分析功能 @@ -401,12 +417,15 @@ def display_scoring_ranking(result): "分析数量", options=[3, 5, 10], index=0, - help="选择分析前N只股票" + help="选择分析前N只股票", + key="batch_count_selector" ) + # 同步更新session_state中的batch_count + st.session_state.batch_count = batch_count with col_batch3: st.write("") # 占位 - if st.button("🚀 开始批量分析", type="primary", use_container_width=True): + if st.button("🚀 开始批量分析", type="primary", width='stretch'): # 提取股票代码 stock_codes = top10_df.head(batch_count)['股票代码'].tolist() @@ -439,25 +458,35 @@ def display_scoring_ranking(result): showlegend=False, height=400 ) - st.plotly_chart(fig1, use_container_width=True) + st.plotly_chart(fig1, config={'displayModeBar': False}, use_container_width=True) with col2: - # 五维评分雷达图(显示第一名) + # 五维评分雷达图(显示批量分析数量的股票) if len(top10_df) > 0: - first_place = top10_df.iloc[0] + display_count = min(5, len(top10_df)) - fig2 = go.Figure(data=go.Scatterpolar( - r=[ - first_place['资金含金量'] / 30 * 100, - first_place['净买入额'] / 25 * 100, - first_place['卖出压力'] / 20 * 100, - first_place['机构共振'] / 15 * 100, - first_place['加分项'] / 10 * 100 - ], - theta=['资金含金量', '净买入额', '卖出压力', '机构共振', '加分项'], - fill='toself', - name=first_place['股票名称'] - )) + fig2 = go.Figure() + + # 为每只股票添加雷达图 + colors = ['#FF6B6B', '#4ECDC4', '#45B7D1', '#96CEB4', '#FFEAA7'] + for i in range(display_count): + stock = top10_df.iloc[i] + + fig2.add_trace(go.Scatterpolar( + r=[ + stock['资金含金量'] / 30 * 100, + stock['净买入额'] / 25 * 100, + stock['卖出压力'] / 20 * 100, + stock['机构共振'] / 15 * 100, + stock['加分项'] / 10 * 100 + ], + theta=['资金含金量', '净买入额', '卖出压力', '机构共振', '加分项'], + fill='toself', + name=f"{stock['股票名称']}", + line_color=colors[i % len(colors)], + fillcolor=colors[i % len(colors)], + opacity=0.6 + )) fig2.update_layout( polar=dict( @@ -467,10 +496,17 @@ def display_scoring_ranking(result): ) ), showlegend=True, - title=f"🥇 {first_place['股票名称']} 五维评分", - height=400 + title=f"🏆 TOP{display_count} 五维评分对比", + height=400, + legend=dict( + orientation="h", + yanchor="auto", + y=-0.2, + xanchor="center", + x=0.5 + ) ) - st.plotly_chart(fig2, use_container_width=True) + st.plotly_chart(fig2, config={'displayModeBar': False}, use_container_width=True) st.markdown("---") @@ -480,7 +516,7 @@ def display_scoring_ranking(result): st.dataframe( scoring_df, column_config={ - "排名": st.column_config.TextColumn("排名", width="small"), + "排名": st.column_config.NumberColumn("排名", format="%d", width="small"), "股票名称": st.column_config.TextColumn("股票名称"), "股票代码": st.column_config.TextColumn("代码"), "综合评分": st.column_config.NumberColumn("综合评分", format="%.1f"), @@ -490,7 +526,7 @@ def display_scoring_ranking(result): "净流入": st.column_config.NumberColumn("净流入(元)", format="%.2f") }, hide_index=True, - use_container_width=True + width='stretch' ) @@ -523,7 +559,7 @@ def display_recommended_stocks(result): "reason": st.column_config.TextColumn("推荐理由") }, hide_index=True, - use_container_width=True + width='stretch' ) # 详细推荐理由 @@ -599,7 +635,7 @@ def display_data_details(result): "净流入金额": st.column_config.NumberColumn("净流入金额(元)", format="%.2f") }, hide_index=True, - use_container_width=True + width='stretch' ) # TOP股票 @@ -616,7 +652,7 @@ def display_data_details(result): "net_inflow": st.column_config.NumberColumn("净流入金额(元)", format="%.2f") }, hide_index=True, - use_container_width=True + width='stretch' ) # 热门概念 @@ -637,7 +673,7 @@ def display_data_details(result): "出现次数": st.column_config.NumberColumn("出现次数", format="%d") }, hide_index=True, - use_container_width=True + width='stretch' ) @@ -664,7 +700,7 @@ def display_visualizations(result): labels={'name': '股票名称', 'net_inflow': '净流入金额(元)'} ) fig.update_layout(xaxis_tickangle=-45) - st.plotly_chart(fig, use_container_width=True) + st.plotly_chart(fig, config={'displayModeBar': False}, use_container_width=True) # 热门概念图表 if summary.get('hot_concepts'): @@ -679,7 +715,7 @@ def display_visualizations(result): names='概念', title='热门概念出现次数分布' ) - st.plotly_chart(fig, use_container_width=True) + st.plotly_chart(fig, config={'displayModeBar': False}, use_container_width=True) def display_pdf_export_section(result): @@ -693,7 +729,7 @@ def display_pdf_export_section(result): st.info("💡 点击按钮生成并下载专业的PDF分析报告") with col2: - if st.button("📥 生成PDF", type="primary", use_container_width=True): + if st.button("📥 生成PDF", type="primary", width='stretch'): with st.spinner("正在生成PDF报告..."): try: generator = LonghubangPDFGenerator() @@ -709,7 +745,7 @@ def display_pdf_export_section(result): data=pdf_bytes, file_name=f"智瞰龙虎报告_{datetime.now().strftime('%Y%m%d_%H%M%S')}.pdf", mime="application/pdf", - use_container_width=True + width='stretch' ) st.success("✅ PDF报告生成成功!") @@ -780,7 +816,7 @@ def display_history_tab(): "hold_period": st.column_config.TextColumn("持有周期") }, hide_index=True, - use_container_width=True + width='stretch' ) st.markdown("---") @@ -818,6 +854,17 @@ def display_history_tab(): st.markdown("#### 🏆 AI智能评分排名 (TOP10)") df_scoring = pd.DataFrame(scoring_ranking[:10]) + # 类型统一,避免Arrow序列化错误 + numeric_cols = ['排名','综合评分','资金含金量','净买入额','卖出压力','机构共振','加分项','顶级游资','买方数','净流入'] + for col in numeric_cols: + if col in df_scoring.columns: + df_scoring[col] = pd.to_numeric(df_scoring[col], errors='coerce') + text_cols = ['股票名称','股票代码','机构参与'] + for col in text_cols: + if col in df_scoring.columns: + df_scoring[col] = df_scoring[col].astype(str) + if '排名' in df_scoring.columns: + df_scoring['排名'] = pd.to_numeric(df_scoring['排名'], errors='coerce').fillna(0).astype(int) # 显示完整的评分表格 st.dataframe( @@ -867,7 +914,7 @@ def display_history_tab(): "净流入": st.column_config.NumberColumn("净流入(元)", format="%.2f") }, hide_index=True, - use_container_width=True + width='stretch' ) # 显示评分说明 @@ -903,34 +950,88 @@ def display_history_tab(): st.markdown("#### 📄 原始分析内容") analysis_content = report_detail.get('analysis_content', '') if analysis_content: - st.text_area("", value=analysis_content[:2000], height=200, disabled=True) + st.text_area("原始分析内容", value=analysis_content[:2000], height=200, disabled=True) if len(analysis_content) > 2000: st.caption("(内容过长,仅显示前2000字符)") - # 导出按钮 + # 操作按钮 st.markdown("---") - col_export1, col_export2 = st.columns(2) + col_export1, col_export2, col_export3 = st.columns(3) with col_export1: if st.button(f"📥 导出为PDF", key=f"export_pdf_{report_id}"): st.info("PDF导出功能开发中...") with col_export2: - if st.button(f"📋 加载到分析页", key=f"load_report_{report_id}"): + # 使用session_state来管理按钮状态,避免需要点击两次的问题 + load_key = f"load_report_{report_id}" + if st.button(f"📋 加载到分析页", key=load_key): # 将历史报告加载到当前分析结果中 if analysis_content_parsed: # 重建完整的result结构 + scoring_data = analysis_content_parsed.get('scoring_ranking', []) + if scoring_data: + df_scoring = pd.DataFrame(scoring_data) + # 类型统一,避免Arrow序列化错误 + numeric_cols = ['排名','综合评分','资金含金量','净买入额','卖出压力','机构共振','加分项','顶级游资','买方数','净流入'] + for col in numeric_cols: + if col in df_scoring.columns: + df_scoring[col] = pd.to_numeric(df_scoring[col], errors='coerce') + text_cols = ['股票名称','股票代码','机构参与'] + for col in text_cols: + if col in df_scoring.columns: + df_scoring[col] = df_scoring[col].astype(str) + if '排名' in df_scoring.columns: + df_scoring['排名'] = pd.to_numeric(df_scoring['排名'], errors='coerce').fillna(0).astype(int) + else: + df_scoring = None + loaded_result = { "success": True, "timestamp": report_detail.get('analysis_date', ''), "data_info": analysis_content_parsed.get('data_info', {}), "agents_analysis": analysis_content_parsed.get('agents_analysis', {}), - "scoring_ranking": pd.DataFrame(analysis_content_parsed.get('scoring_ranking', [])) if analysis_content_parsed.get('scoring_ranking') else None, + "scoring_ranking": df_scoring, "final_report": analysis_content_parsed.get('final_report', {}), "recommended_stocks": report_detail.get('recommended_stocks', []) } st.session_state.longhubang_result = loaded_result + # 使用rerun来立即刷新页面状态 st.success('✅ 报告已加载到分析页面,请切换到"龙虎榜分析"标签查看') + st.rerun() + + with col_export3: + # 删除按钮 + delete_key = f"delete_report_{report_id}" + if st.button(f"🗑️ 删除报告", key=delete_key, type="secondary"): + # 使用session_state来管理删除确认状态 + st.session_state[f"confirm_delete_{report_id}"] = True + st.rerun() + + # 删除确认对话框 + if st.session_state.get(f"confirm_delete_{report_id}", False): + st.warning(f"⚠️ 确认删除报告 #{report_id}?此操作不可撤销!") + col_confirm1, col_confirm2 = st.columns(2) + + with col_confirm1: + if st.button(f"✅ 确认删除", key=f"confirm_delete_yes_{report_id}", type="primary"): + try: + # 调用数据库删除方法 - 修复属性名 + engine.database.delete_analysis_report(report_id) + st.success(f"✅ 报告 #{report_id} 已成功删除") + # 清除确认状态并刷新页面 + if f"confirm_delete_{report_id}" in st.session_state: + del st.session_state[f"confirm_delete_{report_id}"] + st.rerun() + except Exception as e: + st.error(f"❌ 删除失败: {str(e)}") + + with col_confirm2: + if st.button(f"❌ 取消", key=f"confirm_delete_no_{report_id}"): + # 清除确认状态 + if f"confirm_delete_{report_id}" in st.session_state: + del st.session_state[f"confirm_delete_{report_id}"] + st.rerun() except Exception as e: st.error(f"❌ 加载历史报告失败: {str(e)}") @@ -986,7 +1087,7 @@ def display_statistics_tab(): "total_net_inflow": st.column_config.NumberColumn("总净流入(元)", format="%.2f") }, hide_index=True, - use_container_width=True + width='stretch' ) st.markdown("---") @@ -1006,7 +1107,7 @@ def display_statistics_tab(): "total_net_inflow": st.column_config.NumberColumn("总净流入(元)", format="%.2f") }, hide_index=True, - use_container_width=True + width='stretch' ) except Exception as e: @@ -1026,7 +1127,7 @@ def run_longhubang_batch_analysis(): # 返回按钮 col_back, col_clear = st.columns(2) with col_back: - if st.button("🔙 返回龙虎榜分析", use_container_width=True): + if st.button("🔙 返回龙虎榜分析", width='stretch'): # 清除所有批量分析相关状态 if 'longhubang_batch_trigger' in st.session_state: del st.session_state.longhubang_batch_trigger @@ -1037,7 +1138,7 @@ def run_longhubang_batch_analysis(): st.rerun() with col_clear: - if st.button("🔄 重新分析", use_container_width=True): + if st.button("🔄 重新分析", width='stretch'): # 清除结果,保留触发标志和代码 if 'longhubang_batch_results' in st.session_state: del st.session_state.longhubang_batch_results @@ -1098,11 +1199,11 @@ def run_longhubang_batch_analysis(): start_analysis = False with col_confirm: - if st.button("🚀 确认开始分析", type="primary", use_container_width=True): + if st.button("🚀 确认开始分析", type="primary", width='stretch'): start_analysis = True with col_cancel: - if st.button("❌ 取消", type="secondary", use_container_width=True): + if st.button("❌ 取消", type="secondary", width='stretch'): # 清除所有批量分析相关状态 if 'longhubang_batch_trigger' in st.session_state: del st.session_state.longhubang_batch_trigger diff --git a/main_force_analysis.py b/main_force_analysis.py index 158f7a1..f7de82a 100644 --- a/main_force_analysis.py +++ b/main_force_analysis.py @@ -26,15 +26,19 @@ class MainForceAnalyzer: self.raw_stocks = None self.final_recommendations = [] - def run_full_analysis(self, start_date: str = None, days_ago: int = 90, - final_n: int = 5) -> Dict: + def run_full_analysis(self, start_date: str = None, days_ago: int = None, + final_n: int = None, max_range_change: float = None, + min_market_cap: float = None, max_market_cap: float = None) -> Dict: """ 运行完整的主力选股分析流程 - 整体批量分析 Args: start_date: 开始日期,格式如"2025年10月1日" - days_ago: 距今多少天,默认90天 - final_n: 最终精选N只,默认5只 + days_ago: 距今多少天 + final_n: 最终精选N只 + max_range_change: 最大涨跌幅限制 + min_market_cap: 最小市值限制 + max_market_cap: 最大市值限制 Returns: 分析结果字典 @@ -44,7 +48,15 @@ class MainForceAnalyzer: 'total_stocks': 0, 'filtered_stocks': 0, 'final_recommendations': [], - 'error': None + 'error': None, + 'params': { + 'start_date': start_date, + 'days_ago': days_ago, + 'final_n': final_n, + 'max_range_change': max_range_change, + 'min_market_cap': min_market_cap, + 'max_market_cap': max_market_cap + } } try: @@ -55,7 +67,9 @@ class MainForceAnalyzer: # 步骤1: 获取主力资金净流入前100名股票 success, raw_data, message = self.selector.get_main_force_stocks( start_date=start_date, - days_ago=days_ago + days_ago=days_ago, + min_market_cap=min_market_cap, + max_market_cap=max_market_cap ) if not success: @@ -67,9 +81,9 @@ class MainForceAnalyzer: # 步骤2: 智能筛选(涨幅、市值等) filtered_data = self.selector.filter_stocks( raw_data, - max_range_change=30.0, - min_market_cap=50.0, - max_market_cap=1300.0 + max_range_change=max_range_change, + min_market_cap=min_market_cap, + max_market_cap=max_market_cap ) result['filtered_stocks'] = len(filtered_data) diff --git a/main_force_batch.db b/main_force_batch.db index 7877f69..8e5b6a2 100644 Binary files a/main_force_batch.db and b/main_force_batch.db differ diff --git a/main_force_history_ui.py b/main_force_history_ui.py index c10fdbd..a4e4015 100644 --- a/main_force_history_ui.py +++ b/main_force_history_ui.py @@ -108,14 +108,26 @@ def display_batch_history(): '代码': r.get('symbol', 'N/A'), '名称': stock_info.get('name', stock_info.get('股票名称', 'N/A')), '评级': final_decision.get('rating', final_decision.get('investment_rating', 'N/A')), - '信心度': f"{final_decision.get('confidence_level', 0)}%", + '信心度': final_decision.get('confidence_level', 'N/A'), '进场区间': final_decision.get('entry_range', 'N/A'), '止盈位': final_decision.get('take_profit', 'N/A'), '止损位': final_decision.get('stop_loss', 'N/A') }) df = pd.DataFrame(table_data) - st.dataframe(df, use_container_width=True) + + # 类型统一,避免Arrow序列化错误 + numeric_cols = ['信心度', '止盈位', '止损位'] + for col in numeric_cols: + if col in df.columns: + df[col] = pd.to_numeric(df[col], errors='coerce') + + text_cols = ['代码', '名称', '评级', '进场区间'] + for col in text_cols: + if col in df.columns: + df[col] = df[col].astype(str) + + st.dataframe(df, width='content') # 显示详细分析(可展开) with st.expander("📊 查看详细分析报告"): @@ -147,7 +159,7 @@ def display_batch_history(): }) df_fail = pd.DataFrame(fail_data) - st.dataframe(df_fail, use_container_width=True) + st.dataframe(df_fail, width='content') # 操作按钮 col_del, col_reload = st.columns([1, 1]) diff --git a/main_force_pdf_generator.py b/main_force_pdf_generator.py index f9814cd..97c2375 100644 --- a/main_force_pdf_generator.py +++ b/main_force_pdf_generator.py @@ -451,6 +451,6 @@ def display_report_download_section(analyzer, result): data=csv, file_name=csv_filename, mime="text/csv", - use_container_width=True + width='content' ) diff --git a/main_force_selector.py b/main_force_selector.py index 3ba5306..3543009 100644 --- a/main_force_selector.py +++ b/main_force_selector.py @@ -5,6 +5,7 @@ 使用pywencai获取主力资金净流入前100名股票,并进行智能筛选 """ +from numpy.ma import minimum_fill_value import pandas as pd import pywencai from datetime import datetime, timedelta @@ -18,13 +19,16 @@ class MainForceStockSelector: self.raw_data = None self.filtered_stocks = None - def get_main_force_stocks(self, start_date: str = None, days_ago: int = 90) -> Tuple[bool, pd.DataFrame, str]: + def get_main_force_stocks(self, start_date: str = None, days_ago: int = None, + min_market_cap: float = None, max_market_cap: float = None) -> Tuple[bool, pd.DataFrame, str]: """ 获取主力资金净流入前100名股票 Args: start_date: 开始日期,格式如"2025年10月1日",如果不提供则使用days_ago - days_ago: 距今多少天,默认90天(约3个月) + days_ago: 距今多少天 + min_market_cap: 最小市值限制 + max_market_cap: 最大市值限制 Returns: (success, dataframe, message) @@ -44,21 +48,21 @@ class MainForceStockSelector: # 构建查询语句 - 使用多个备选方案,所有方案都要求计算区间涨跌幅 queries = [ # 方案1: 完整查询(最优) - f"{start_date}以来主力资金净流入排名,并计算区间涨跌幅,市值50-5000亿之间,非科创非st," + f"{start_date}以来主力资金净流入排名,并计算区间涨跌幅,市值{min_market_cap}-{max_market_cap}亿之间,非科创非st," f"所属同花顺行业,总市值,净利润,营收,市盈率,市净率," f"盈利能力评分,成长能力评分,营运能力评分,偿债能力评分," f"现金流评分,资产质量评分,流动性评分,资本充足性评分", # 方案2: 简化查询 - f"{start_date}以来主力资金净流入,并计算区间涨跌幅,市值50-5000亿,非科创非st," + f"{start_date}以来主力资金净流入,并计算区间涨跌幅,市值{min_market_cap}-{max_market_cap}亿,非科创非st," f"所属同花顺行业,总市值,净利润,营收,市盈率,市净率", # 方案3: 基础查询 - f"{start_date}以来主力资金净流入排名,并计算区间涨跌幅,市值50-5000亿,非科创非st," + f"{start_date}以来主力资金净流入排名,并计算区间涨跌幅,市值{min_market_cap}-{max_market_cap}亿,非科创非st," f"所属行业,总市值", # 方案4: 最简查询 - f"{start_date}以来主力资金净流入前100名,并计算区间涨跌幅,市值50-5000亿,非st非科创板,所属行业,总市值", + f"{start_date}以来主力资金净流入前100名,并计算区间涨跌幅,市值{min_market_cap}-{max_market_cap}亿,非st非科创板,所属行业,总市值", ] # 尝试不同的查询方案 @@ -132,17 +136,17 @@ class MainForceStockSelector: return None def filter_stocks(self, df: pd.DataFrame, - max_range_change: float = 30.0, - min_market_cap: float = 50.0, - max_market_cap: float = 1300.0) -> pd.DataFrame: + max_range_change: float = None, + min_market_cap: float = None, + max_market_cap: float = None) -> pd.DataFrame: """ - 智能筛选股票 + 智能筛选股票 - 基于涨跌幅和市值 Args: - df: 原始数据 - max_range_change: 区间涨跌幅上限(%),默认30% - min_market_cap: 最小市值(亿),默认50亿 - max_market_cap: 最大市值(亿),默认1300亿 + df: 原始股票数据DataFrame + max_range_change: 最大涨跌幅限制 + min_market_cap: 最小市值限制 + max_market_cap: 最大市值限制 Returns: 筛选后的DataFrame @@ -233,13 +237,13 @@ class MainForceStockSelector: self.filtered_stocks = filtered_df return filtered_df - def get_top_stocks(self, df: pd.DataFrame, top_n: int = 20) -> pd.DataFrame: + def get_top_stocks(self, df: pd.DataFrame, top_n: int = None) -> pd.DataFrame: """ - 获取主力资金净流入最多的前N只股票 + 获取主力资金净流入前N名股票 Args: - df: 筛选后的数据 - top_n: 取前N名,默认20 + df: 筛选后的股票数据 + top_n: 返回前N名 Returns: 前N名股票DataFrame diff --git a/main_force_ui.py b/main_force_ui.py index 05fdded..3df96fa 100644 --- a/main_force_ui.py +++ b/main_force_ui.py @@ -30,7 +30,7 @@ def display_main_force_selector(): st.markdown("## 🎯 主力选股 - 智能筛选优质标的") with col_history: st.write("") # 占位 - if st.button("📚 批量分析历史", use_container_width=True): + if st.button("📚 批量分析历史", width='content'): st.session_state.main_force_view_history = True st.rerun() @@ -102,8 +102,8 @@ def display_main_force_selector(): with col1: max_change = st.number_input( "最大涨跌幅(%)", - min_value=10.0, - max_value=100.0, + min_value=5.0, + max_value=200.0, value=30.0, step=5.0, help="过滤掉涨幅过高的股票,避免追高" @@ -121,7 +121,7 @@ def display_main_force_selector(): with col3: max_cap = st.number_input( "最大市值(亿)", - min_value=100.0, + min_value=50.0, max_value=50000.0, value=5000.0, step=100.0 @@ -137,7 +137,7 @@ def display_main_force_selector(): st.markdown("---") # 开始分析按钮 - if st.button("🚀 开始主力选股", type="primary", use_container_width=True): + if st.button("🚀 开始主力选股", type="primary", width='content'): with st.spinner("正在获取数据并分析,这可能需要几分钟..."): @@ -148,7 +148,10 @@ def display_main_force_selector(): result = analyzer.run_full_analysis( start_date=start_date, days_ago=days_ago, - final_n=final_n + final_n=final_n, + max_range_change=max_change, + min_market_cap=min_cap, + max_market_cap=max_cap ) # 保存结果到session_state @@ -277,7 +280,7 @@ def display_analysis_results(result: dict, analyzer): # 显示DataFrame display_df = analyzer.raw_stocks[final_cols].copy() - st.dataframe(display_df, use_container_width=True, height=400) + st.dataframe(display_df, width='content', height=400) # 显示统计 st.caption(f"共 {len(display_df)} 只候选股票,显示 {len(final_cols)} 个字段") @@ -309,7 +312,7 @@ def display_analysis_results(result: dict, analyzer): with col_batch3: st.write("") # 占位 - if st.button("🚀 开始批量分析", type="primary", use_container_width=True): + if st.button("🚀 开始批量分析", type="primary", width='content'): # 准备数据:按主力资金净流入排序 df_sorted = analyzer.raw_stocks.copy() @@ -497,7 +500,7 @@ def run_main_force_batch_analysis(): # 返回按钮 col_back, col_clear = st.columns(2) with col_back: - if st.button("🔙 返回主力选股", use_container_width=True): + if st.button("🔙 返回主力选股", width='content'): # 清除所有批量分析相关状态 if 'main_force_batch_trigger' in st.session_state: del st.session_state.main_force_batch_trigger @@ -508,7 +511,7 @@ def run_main_force_batch_analysis(): st.rerun() with col_clear: - if st.button("🔄 重新分析", use_container_width=True): + if st.button("🔄 重新分析", width='content'): # 清除结果,保留触发标志和代码 if 'main_force_batch_results' in st.session_state: del st.session_state.main_force_batch_results @@ -569,11 +572,11 @@ def run_main_force_batch_analysis(): start_analysis = False with col_confirm: - if st.button("🚀 确认开始分析", type="primary", use_container_width=True): + if st.button("🚀 确认开始分析", type="primary", width='content'): start_analysis = True with col_cancel: - if st.button("❌ 取消", type="secondary", use_container_width=True): + if st.button("❌ 取消", type="secondary", width='content'): # 清除所有批量分析相关状态 if 'main_force_batch_trigger' in st.session_state: del st.session_state.main_force_batch_trigger @@ -858,7 +861,19 @@ def display_main_force_batch_results(batch_results): }) df_display = pd.DataFrame(display_data) - st.dataframe(df_display, use_container_width=True, height=400) + + # 类型统一,避免Arrow序列化错误 + numeric_cols = ['信心度', '止盈位', '止损位', '目标价'] + for col in numeric_cols: + if col in df_display.columns: + df_display[col] = pd.to_numeric(df_display[col], errors='coerce') + + text_cols = ['股票代码', '股票名称', '评级', '进场区间'] + for col in text_cols: + if col in df_display.columns: + df_display[col] = df_display[col].astype(str) + + st.dataframe(df_display, width='content', height=400) # 详细分析结果(可展开) st.markdown("---") @@ -977,5 +992,5 @@ def display_main_force_batch_results(batch_results): }) df_failed = pd.DataFrame(failed_data) - st.dataframe(df_failed, use_container_width=True) + st.dataframe(df_failed, width='content') diff --git a/monitor_manager.py b/monitor_manager.py index e758f46..324feb4 100644 --- a/monitor_manager.py +++ b/monitor_manager.py @@ -145,7 +145,7 @@ def display_add_stock_section(): auto_take_profit = st.checkbox("自动止盈", value=True) # 添加按钮 - if st.button("✅ 添加监测", type="primary", use_container_width=True): + if st.button("✅ 添加监测", type="primary", width='stretch'): if symbol and entry_min > 0 and entry_max > 0 and entry_max > entry_min: try: # 准备数据 @@ -408,10 +408,10 @@ def display_edit_dialog(stock_id: int): col1, col2, col3 = st.columns(3) with col1: - submit = st.form_submit_button("✅ 保存修改", type="primary", use_container_width=True) + submit = st.form_submit_button("✅ 保存修改", type="primary", width='stretch') with col2: - cancel = st.form_submit_button("❌ 取消", use_container_width=True) + cancel = st.form_submit_button("❌ 取消", width='stretch') if submit: if entry_min > 0 and entry_max > 0 and entry_max > entry_min: @@ -482,7 +482,7 @@ def display_delete_confirm_dialog(stock_id: int): col1, col2, col3 = st.columns([1, 1, 1]) with col1: - if st.button("🗑️ 确认删除", type="primary", use_container_width=True, key=f"confirm_delete_{stock_id}"): + if st.button("🗑️ 确认删除", type="primary", width='stretch', key=f"confirm_delete_{stock_id}"): try: result = monitor_db.remove_monitored_stock(stock_id) if result: @@ -508,7 +508,7 @@ def display_delete_confirm_dialog(stock_id: int): st.rerun() with col2: - if st.button("❌ 取消", use_container_width=True, key=f"cancel_delete_{stock_id}"): + if st.button("❌ 取消", width='stretch', key=f"cancel_delete_{stock_id}"): del st.session_state.deleting_stock_id st.rerun() @@ -568,7 +568,7 @@ def display_notification_management(): # 测试邮件按钮 if email_config['configured']: - if st.button("📧 发送测试邮件", type="primary", use_container_width=True): + if st.button("📧 发送测试邮件", type="primary", width='stretch'): with st.spinner("正在发送测试邮件..."): success, message = notification_service.send_test_email() if success: @@ -577,7 +577,7 @@ def display_notification_management(): else: st.error(f"❌ {message}") else: - st.button("📧 发送测试邮件", type="primary", use_container_width=True, disabled=True) + st.button("📧 发送测试邮件", type="primary", width='stretch', disabled=True) st.caption("请先在.env文件中配置邮件参数") with col2: @@ -688,7 +688,7 @@ def display_miniqmt_status(): # 连接按钮 if qmt_status['enabled'] and not qmt_status['connected']: - if st.button("🔗 连接MiniQMT", type="primary", use_container_width=True): + if st.button("🔗 连接MiniQMT", type="primary", width='stretch'): success, msg = miniqmt.connect() if success: st.success(f"✅ {msg}") @@ -696,7 +696,7 @@ def display_miniqmt_status(): st.error(f"❌ {msg}") st.rerun() elif qmt_status['connected']: - if st.button("🔌 断开连接", use_container_width=True): + if st.button("🔌 断开连接", width='stretch'): if miniqmt.disconnect(): st.info("⏸️ 已断开MiniQMT连接") st.rerun() @@ -820,7 +820,7 @@ def display_scheduler_section(): col1, col2, col3 = st.columns([1, 1, 1]) with col1: - if st.button("💾 保存设置", type="primary", use_container_width=True): + if st.button("💾 保存设置", type="primary", width='stretch'): try: # 更新配置 scheduler.update_config( @@ -841,24 +841,24 @@ def display_scheduler_section(): with col2: if status['scheduler_running']: - if st.button("⏹️ 停止调度器", use_container_width=True): + if st.button("⏹️ 停止调度器", width='stretch'): scheduler.stop_scheduler() st.info("⏸️ 调度器已停止") time.sleep(0.5) st.rerun() else: if enabled: - if st.button("▶️ 启动调度器", type="secondary", use_container_width=True): + if st.button("▶️ 启动调度器", type="secondary", width='stretch'): scheduler.start_scheduler() st.success("✅ 调度器已启动") time.sleep(0.5) st.rerun() else: - st.button("▶️ 启动调度器", use_container_width=True, disabled=True) + st.button("▶️ 启动调度器", width='stretch', disabled=True) st.caption("请先启用定时调度") with col3: - if st.button("🔄 刷新状态", use_container_width=True): + if st.button("🔄 刷新状态", width='stretch'): st.rerun() def get_monitor_summary(): diff --git a/pdf_generator.py b/pdf_generator.py index d9a1292..69b9c47 100644 --- a/pdf_generator.py +++ b/pdf_generator.py @@ -270,7 +270,7 @@ def display_pdf_export_section(stock_info, agents_results, discussion_result, fi with col2: # 生成PDF报告按钮(使用股票代码作为key的一部分,确保唯一性) button_key = f"pdf_btn_{stock_info.get('symbol', 'unknown')}" - if st.button("📄 生成并下载PDF报告", type="primary", use_container_width=True, key=button_key): + if st.button("📄 生成并下载PDF报告", type="primary", width='content', key=button_key): with st.spinner("正在生成PDF报告..."): try: # 生成PDF内容 diff --git a/pdf_generator_fixed.py b/pdf_generator_fixed.py index ff49c28..9ea3ae0 100644 --- a/pdf_generator_fixed.py +++ b/pdf_generator_fixed.py @@ -291,7 +291,7 @@ def display_pdf_export_section(stock_info, agents_results, discussion_result, fi import uuid import time button_key = f"generate_report_btn_{int(time.time())}_{uuid.uuid4().hex[:8]}" - if st.button("📊 生成并下载报告", type="primary", use_container_width=True, key=button_key): + if st.button("📊 生成并下载报告", type="primary", width='content', key=button_key): with st.spinner("正在生成报告..."): try: # 生成Markdown内容 diff --git a/pdf_generator_pandoc.py b/pdf_generator_pandoc.py index 66ab10d..fde9066 100644 --- a/pdf_generator_pandoc.py +++ b/pdf_generator_pandoc.py @@ -267,7 +267,7 @@ def display_pdf_export_section(stock_info, agents_results, discussion_result, fi col1, col2, col3 = st.columns([1, 2, 1]) with col2: - if st.button("📊 生成并下载报告", type="primary", use_container_width=True, key="generate_report_btn"): + if st.button("📊 生成并下载报告", type="primary", width='content', key="generate_report_btn"): st.session_state.show_download_links = True with st.spinner("正在生成报告..."): success = generate_pdf_report(stock_info, agents_results, discussion_result, final_decision) diff --git a/portfolio_manager.py b/portfolio_manager.py index 879ab84..e23e85f 100644 --- a/portfolio_manager.py +++ b/portfolio_manager.py @@ -418,7 +418,18 @@ class PortfolioManager: # 使用正确的字段名 rating = final_decision.get("rating", "持有") - confidence = final_decision.get("confidence_level", 5.0) + # 确保信心度为float类型,避免Arrow序列化错误 + confidence_raw = final_decision.get("confidence_level", 5.0) + try: + confidence = float(confidence_raw) + # 确保信心度在合理范围内 + if confidence < 0: + confidence = 0.0 + elif confidence > 10: + confidence = 10.0 + except (ValueError, TypeError): + # 如果转换失败,使用默认值 + confidence = 5.0 current_price = stock_info.get("current_price", 0.0) target_price_str = final_decision.get("target_price", "") entry_range = final_decision.get("entry_range", "") diff --git a/portfolio_ui.py b/portfolio_ui.py index 2b08267..4d8fab8 100644 --- a/portfolio_ui.py +++ b/portfolio_ui.py @@ -286,7 +286,7 @@ def display_batch_analysis(): ) # 立即分析按钮 - if st.button("🚀 立即开始分析", type="primary", use_container_width=True): + if st.button("🚀 立即开始分析", type="primary", width='content'): with st.spinner("正在批量分析持仓股票..."): # 显示进度 progress_bar = st.progress(0) @@ -551,7 +551,7 @@ def display_scheduler_management(): with col_add: st.write("") # 占位,对齐按钮 st.write("") - if st.button("➕ 添加", type="primary", use_container_width=True): + if st.button("➕ 添加", type="primary", width='content'): time_str = new_time.strftime("%H:%M") if portfolio_scheduler.add_schedule_time(time_str): st.success(f"已添加 {time_str}") @@ -631,20 +631,20 @@ def display_scheduler_management(): with col_btn1: if is_running: - if st.button("⏹️ 停止调度器", type="secondary", use_container_width=True): + if st.button("⏹️ 停止调度器", type="secondary", width='content'): portfolio_scheduler.stop_scheduler() st.success("调度器已停止") time.sleep(0.5) st.rerun() else: - if st.button("▶️ 启动调度器", type="primary", use_container_width=True): + if st.button("▶️ 启动调度器", type="primary", width='content'): portfolio_scheduler.start_scheduler() st.success("调度器已启动") time.sleep(0.5) st.rerun() with col_btn2: - if st.button("🚀 立即执行一次", type="primary", use_container_width=True): + if st.button("🚀 立即执行一次", type="primary", width='content'): with st.spinner("正在执行持仓分析..."): try: portfolio_scheduler.run_analysis_now() @@ -653,7 +653,7 @@ def display_scheduler_management(): st.error(f"执行失败: {str(e)}") with col_btn3: - if st.button("🔄 刷新状态", use_container_width=True): + if st.button("🔄 刷新状态", width='content'): st.rerun() diff --git a/quarterly_report_data.py b/quarterly_report_data.py index 31ce36b..9874737 100644 --- a/quarterly_report_data.py +++ b/quarterly_report_data.py @@ -228,35 +228,54 @@ class QuarterlyReportDataFetcher: def _get_financial_indicators(self, symbol): """获取财务指标数据""" try: - # stock_financial_analysis_indicator - 财务指标 - df = ak.stock_financial_analysis_indicator(symbol=symbol) + # 使用stock_financial_abstract替代已失效的stock_financial_analysis_indicator + df = ak.stock_financial_abstract(symbol=symbol) if df is None or df.empty: print(f" 未找到财务指标数据") return None # 获取最近8期 - df = df.head(self.periods) + df = df.head(self.periods * 2) # 取更多数据以确保有足够的季度数据 - # 转换为字典列表 + # 提取关键财务指标 + key_indicators = [ + '净资产收益率(ROE)', '总资产报酬率(ROA)', '销售净利率', '销售毛利率', + '资产负债率', '流动比率', '速动比率', '应收账款周转率', '存货周转率', + '总资产周转率', '基本每股收益', '每股净资产', '每股现金流' + ] + + # 筛选出包含关键指标的行 + indicator_rows = df[df['指标'].isin(key_indicators)] + + if indicator_rows.empty: + print(f" 未找到关键财务指标数据") + return None + + # 获取日期列(排除'选项'和'指标'列) + date_columns = [col for col in df.columns if col not in ['选项', '指标']] + + # 转换为字典列表,每个字典代表一个时期的财务指标 data_list = [] - for idx, row in df.iterrows(): - item = {} - for col in df.columns: - value = row.get(col) - if value is None or (isinstance(value, float) and pd.isna(value)): - continue - try: - item[col] = str(value) - except: - item[col] = "N/A" - if item: - data_list.append(item) + for date_col in date_columns[:self.periods]: # 只取最近的periods期 + item = {'报告期': date_col} + for _, row in indicator_rows.iterrows(): + indicator_name = row['指标'] + value = row.get(date_col) + if value is not None and not (isinstance(value, float) and pd.isna(value)): + try: + # 尝试转换为字符串 + item[indicator_name] = str(value) + except: + item[indicator_name] = "N/A" + else: + item[indicator_name] = "N/A" + data_list.append(item) return { "data": data_list, "periods": len(data_list), - "columns": df.columns.tolist(), + "columns": ['报告期'] + key_indicators, "query_time": datetime.now().strftime('%Y-%m-%d %H:%M:%S') } diff --git a/run.py b/run.py index eaa2b3b..9beee4e 100644 --- a/run.py +++ b/run.py @@ -52,15 +52,15 @@ def main(): # 启动Streamlit应用 print("🌐 正在启动Web界面...") - print("📝 访问地址: http://localhost:8501") + print("📝 访问地址: http://localhost:8503") print("⏹️ 按 Ctrl+C 停止服务") print("=" * 50) try: subprocess.run([ sys.executable, "-m", "streamlit", "run", "app.py", - "--server.port", "8501", - "--server.address", "0.0.0.0" + "--server.port", "8503", + "--server.address", "127.0.0.1" ]) except KeyboardInterrupt: print("\n👋 感谢使用AI股票分析系统!") diff --git a/sector_strategy.db b/sector_strategy.db new file mode 100644 index 0000000..c092157 Binary files /dev/null and b/sector_strategy.db differ diff --git a/sector_strategy_data.py b/sector_strategy_data.py index 2e0027b..0df19c1 100644 --- a/sector_strategy_data.py +++ b/sector_strategy_data.py @@ -8,6 +8,13 @@ import pandas as pd from datetime import datetime, timedelta import warnings import time +import logging +import os +from dotenv import load_dotenv +from sector_strategy_db import SectorStrategyDatabase + +# 加载环境变量 +load_dotenv() warnings.filterwarnings('ignore') @@ -20,6 +27,18 @@ class SectorStrategyDataFetcher: self.max_retries = 3 # 最大重试次数 self.retry_delay = 2 # 重试延迟(秒) self.request_delay = 1 # 请求间隔(秒) + + # 初始化数据库和日志 + self.database = SectorStrategyDatabase() + self.logger = logging.getLogger(__name__) + + # 配置日志 + if not self.logger.handlers: + handler = logging.StreamHandler() + formatter = logging.Formatter('%(asctime)s - %(name)s - %(levelname)s - %(message)s') + handler.setFormatter(formatter) + self.logger.addHandler(handler) + self.logger.setLevel(logging.INFO) def _safe_request(self, func, *args, **kwargs): """安全的请求函数,包含重试机制""" @@ -102,6 +121,9 @@ class SectorStrategyDataFetcher: data["success"] = True print("[智策] ✓ 板块数据获取完成!") + # 保存原始数据到数据库 + self._save_raw_data_to_db(data) + except Exception as e: print(f"[智策] ✗ 数据获取出错: {e}") data["error"] = str(e) @@ -276,39 +298,111 @@ class SectorStrategyDataFetcher: return {} def _get_north_money_flow(self): - """获取北向资金流向""" + """获取北向资金流向(优先使用Tushare,失败时使用Akshare)""" + # 优先使用Tushare获取沪深港通资金流向 + self.ts_pro = None + tushare_token = os.getenv('TUSHARE_TOKEN', '') try: - # 获取沪深港通资金流向(使用重试机制) + # 初始化Tushare(如果尚未初始化) + if not hasattr(self, '_tushare_api'): + TUSHARE_TOKEN = os.getenv('TUSHARE_TOKEN', '') + if TUSHARE_TOKEN: + try: + import tushare as ts + ts.set_token(tushare_token) + self.ts_pro = ts.pro_api() + print(" [Tushare] ✅ 初始化成功") + except Exception as e: + print(f" [Tushare] 初始化失败: {e}") + self._tushare_api = None + else: + print(" [Tushare] 未配置Token") + self._tushare_api = None + + + # 如果Tushare可用,获取数据 + if hasattr(self, '_tushare_api') and self._tushare_api: + print(" [Tushare] 正在获取沪深港通资金流向...") + + # 获取最近30天的数据 + end_date = datetime.now() + start_date = end_date - timedelta(days=20) + + df = self._tushare_api.moneyflow_hsgt( + start_date=start_date.strftime('%Y%m%d'), + end_date=end_date.strftime('%Y%m%d') + ) + + if df is not None and not df.empty: + print(" [Tushare] ✅ 成功获取数据") + + # 按日期降序排列,获取最新数据 + df = df.sort_values('trade_date', ascending=False) + latest = df.iloc[0] + + # 转换数据格式以匹配原有结构 + north_flow = { + "date": str(latest['trade_date']), + "north_net_inflow": float(latest['north_money']), + "hgt_net_inflow": float(latest['hgt']), + "sgt_net_inflow": float(latest['sgt']), + "north_total_amount": float(latest['north_money']) # Tushare没有总成交金额,使用净流入作为近似值 + } + + # 获取历史趋势(最近20天) + history = [] + for idx, row in df.head(20).iterrows(): + history.append({ + "date": str(row['trade_date']), + "net_inflow": float(row['north_money']) + }) + north_flow["history"] = history + + return north_flow + else: + print(" [Tushare] ❌ 未获取到数据") + else: + print(" [Tushare] 不可用") + except Exception as e: + print(f" [Tushare] 获取北向资金失败: {e}") + + # Tushare失败,尝试使用Akshare + try: + print(" [Akshare] 正在获取沪深港通资金流向(备用数据源)...") df = self._safe_request(ak.stock_hsgt_fund_flow_summary_em) - if df is None or df.empty: - return {} - - # 获取最新数据 - latest = df.iloc[0] - - north_flow = { - "date": str(latest.get('日期', '')), - "north_net_inflow": latest.get('北向资金-成交净买额', 0), - "hgt_net_inflow": latest.get('沪股通-成交净买额', 0), - "sgt_net_inflow": latest.get('深股通-成交净买额', 0), - "north_total_amount": latest.get('北向资金-成交金额', 0) - } - - # 获取历史趋势(最近10天) - history = [] - for idx, row in df.head(10).iterrows(): - history.append({ - "date": str(row.get('日期', '')), - "net_inflow": row.get('北向资金-成交净买额', 0) - }) - north_flow["history"] = history - - return north_flow - + if df is not None and not df.empty: + print(" [Akshare] ✅ 成功获取数据") + + # 获取最新数据 + latest = df.iloc[0] + + north_flow = { + "date": str(latest.get('日期', '')), + "north_net_inflow": latest.get('北向资金-成交净买额', 0), + "hgt_net_inflow": latest.get('沪股通-成交净买额', 0), + "sgt_net_inflow": latest.get('深股通-成交净买额', 0), + "north_total_amount": latest.get('北向资金-成交金额', 0) + } + + # 获取历史趋势(最近20天) + history = [] + for idx, row in df.head(20).iterrows(): + history.append({ + "date": str(row.get('日期', '')), + "net_inflow": row.get('北向资金-成交净买额', 0) + }) + north_flow["history"] = history + + return north_flow + else: + print(" [Akshare] ❌ 未获取到数据") except Exception as e: - print(f" 获取北向资金失败: {e}") - return {} + print(f" [Akshare] 获取北向资金失败: {e}") + + # 所有数据源都失败 + print(" ❌ 所有数据源均获取失败") + return {} def _get_financial_news(self): """获取财经新闻""" @@ -438,6 +532,210 @@ class SectorStrategyDataFetcher: text_parts.append(f" {news['content'][:100]}...") return "\n".join(text_parts) + + def _save_raw_data_to_db(self, data): + """保存原始数据到数据库""" + try: + if not data.get("success"): + self.logger.warning("[智策数据] 数据获取失败,跳过保存") + return + + # 保存板块数据 + if data.get("sectors"): + # 将字典转换为DataFrame并映射必要列 + sectors_df = pd.DataFrame([ + { + '板块名称': v.get('name', k), + '涨跌幅': v.get('change_pct', 0), + '成交额': 0, + '总市值': v.get('total_market_cap', 0), + '市盈率': v.get('pe_ratio', 0), + '市净率': v.get('pb_ratio', 0), + '最新价': 0, + '成交量': 0, + 'turnover': v.get('turnover', 0) # 兼容保存方法中的fallback + } + for k, v in data["sectors"].items() + ]) + self.database.save_sector_raw_data( + data_date=datetime.now().strftime('%Y-%m-%d'), + data_type="industry", + data_df=sectors_df + ) + self.logger.info(f"[智策数据] 保存行业板块数据: {len(data['sectors'])} 个板块") + + # 保存概念板块数据 + if data.get("concepts"): + concepts_df = pd.DataFrame([ + { + '板块名称': v.get('name', k), + '涨跌幅': v.get('change_pct', 0), + '成交额': 0, + '总市值': v.get('total_market_cap', 0), + '市盈率': v.get('pe_ratio', 0), + '市净率': v.get('pb_ratio', 0), + '最新价': 0, + '成交量': 0, + 'turnover': v.get('turnover', 0) + } + for k, v in data["concepts"].items() + ]) + self.database.save_sector_raw_data( + data_date=datetime.now().strftime('%Y-%m-%d'), + data_type="concept", + data_df=concepts_df + ) + self.logger.info(f"[智策数据] 保存概念板块数据: {len(data['concepts'])} 个概念") + + # 保存资金流向数据 + if data.get("sector_fund_flow"): + flow_today = data["sector_fund_flow"].get("today", []) + fund_df = pd.DataFrame([ + { + '行业': item.get('sector', ''), + '主力净流入-净额': item.get('main_net_inflow', 0), + '主力净流入-净占比': item.get('main_net_inflow_pct', 0), + '超大单净流入-净额': item.get('super_large_net_inflow', 0), + '超大单净流入-净占比': item.get('super_large_net_inflow_pct', 0), + '大单净流入-净额': item.get('large_net_inflow', 0), + '大单净流入-净占比': item.get('large_net_inflow_pct', 0) + } + for item in flow_today + ]) + if not fund_df.empty: + self.database.save_sector_raw_data( + data_date=datetime.now().strftime('%Y-%m-%d'), + data_type="fund_flow", + data_df=fund_df + ) + self.logger.info("[智策数据] 保存资金流向数据") + + # 保存市场概况数据 + if data.get("market_overview"): + market = data["market_overview"] + mo_df = pd.DataFrame([ + {'名称': '上证指数', '最新价': market.get('sh_index', {}).get('close', 0), '涨跌幅': market.get('sh_index', {}).get('change_pct', 0), '成交量': market.get('sh_index', {}).get('volume', 0), '成交额': market.get('sh_index', {}).get('turnover', 0)}, + {'名称': '深证成指', '最新价': market.get('sz_index', {}).get('close', 0), '涨跌幅': market.get('sz_index', {}).get('change_pct', 0), '成交量': market.get('sz_index', {}).get('volume', 0), '成交额': market.get('sz_index', {}).get('turnover', 0)}, + {'名称': '创业板指', '最新价': market.get('cyb_index', {}).get('close', 0), '涨跌幅': market.get('cyb_index', {}).get('change_pct', 0), '成交量': market.get('cyb_index', {}).get('volume', 0), '成交额': market.get('cyb_index', {}).get('turnover', 0)} + ]) + self.database.save_sector_raw_data( + data_date=datetime.now().strftime('%Y-%m-%d'), + data_type="market_overview", + data_df=mo_df + ) + self.logger.info("[智策数据] 保存市场概况数据") + + # 保存北向资金数据 + # 注:north_flow结构与原始表不一致,此处暂不保存以避免歧义 + + # 保存新闻数据 + if data.get("news"): + self.database.save_news_data( + news_list=data["news"], + news_date=datetime.now().strftime('%Y-%m-%d'), + source="akshare" + ) + self.logger.info(f"[智策数据] 保存财经新闻: {len(data['news'])} 条") + + except Exception as e: + self.logger.error(f"[智策数据] 保存原始数据失败: {e}") + + def get_cached_data_with_fallback(self): + """获取缓存数据,支持回退机制""" + try: + # 首先尝试获取最新数据 + print("[智策] 尝试获取最新数据...") + fresh_data = self.get_all_sector_data() + + if fresh_data.get("success"): + return fresh_data + + # 如果获取失败,回退到缓存数据 + print("[智策] 获取最新数据失败,尝试加载缓存数据...") + cached_data = self._load_cached_data() + + if cached_data: + print("[智策] ✓ 成功加载缓存数据") + cached_data["from_cache"] = True + cached_data["cache_warning"] = "当前显示为缓存数据(24小时内),可能不是最新信息" + return cached_data + else: + print("[智策] ✗ 无可用缓存数据") + return { + "success": False, + "error": "无法获取数据且无可用缓存", + "timestamp": datetime.now().strftime('%Y-%m-%d %H:%M:%S') + } + + except Exception as e: + self.logger.error(f"[智策数据] 获取数据失败: {e}") + return { + "success": False, + "error": str(e), + "timestamp": datetime.now().strftime('%Y-%m-%d %H:%M:%S') + } + + def _load_cached_data(self): + """加载缓存数据""" + try: + # 获取最近的各类数据 + cached_data = { + "success": True, + "timestamp": datetime.now().strftime('%Y-%m-%d %H:%M:%S'), + "sectors": {}, + "concepts": {}, + "sector_fund_flow": {}, + "market_overview": {}, + "north_flow": {}, + "news": [] + } + + # 加载板块数据 + sectors_data = self.database.get_latest_raw_data("sectors") + if sectors_data: + cached_data["sectors"] = sectors_data.get("data_content", {}) + + # 加载概念数据 + concepts_data = self.database.get_latest_raw_data("concepts") + if concepts_data: + cached_data["concepts"] = concepts_data.get("data_content", {}) + + # 加载资金流向数据 + fund_flow_data = self.database.get_latest_raw_data("fund_flow") + if fund_flow_data: + cached_data["sector_fund_flow"] = fund_flow_data.get("data_content", {}) + + # 加载市场概况数据 + market_data = self.database.get_latest_raw_data("market_overview") + if market_data: + cached_data["market_overview"] = market_data.get("data_content", {}) + + # 加载北向资金数据 + north_data = self.database.get_latest_raw_data("north_flow") + if north_data: + cached_data["north_flow"] = north_data.get("data_content", {}) + + # 加载新闻数据 + news_data = self.database.get_latest_news_data() + if news_data: + # 仅传递内容列表给下游分析,避免结构不一致 + cached_data["news"] = news_data.get("data_content", []) + + # 检查是否有有效数据 + has_data = any([ + cached_data["sectors"], + cached_data["concepts"], + cached_data["sector_fund_flow"], + cached_data["market_overview"], + cached_data["north_flow"], + cached_data["news"] + ]) + + return cached_data if has_data else None + + except Exception as e: + self.logger.error(f"[智策数据] 加载缓存数据失败: {e}") + return None # 测试函数 diff --git a/sector_strategy_db.py b/sector_strategy_db.py new file mode 100644 index 0000000..38b8253 --- /dev/null +++ b/sector_strategy_db.py @@ -0,0 +1,934 @@ +""" +智策板块数据库模块 +用于存储板块策略历史数据和分析报告 +""" + +import sqlite3 +from datetime import datetime +import json +import pandas as pd +import logging + + +class SectorStrategyDatabase: + """智策板块数据库管理类""" + + def __init__(self, db_path='sector_strategy.db'): + """ + 初始化数据库 + + Args: + db_path: 数据库文件路径 + """ + self.db_path = db_path + # 初始化日志 + self.logger = logging.getLogger(__name__) + if not self.logger.handlers: + logging.basicConfig(level=logging.INFO, format='[%(asctime)s] %(levelname)s %(name)s: %(message)s') + self.init_database() + + def get_connection(self): + """获取数据库连接""" + return sqlite3.connect(self.db_path) + + def init_database(self): + """初始化数据库表""" + conn = self.get_connection() + cursor = conn.cursor() + + # 板块原始数据表 + cursor.execute(''' + CREATE TABLE IF NOT EXISTS sector_raw_data ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + data_date TEXT NOT NULL, + sector_code TEXT NOT NULL, + sector_name TEXT, + price REAL, + change_pct REAL, + volume REAL, + turnover REAL, + market_cap REAL, + pe_ratio REAL, + pb_ratio REAL, + data_type TEXT, + data_version INTEGER DEFAULT 1, + created_at TEXT DEFAULT CURRENT_TIMESTAMP, + UNIQUE(data_date, sector_code, data_type) + ) + ''') + + # 创建索引 + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_sector_data_date ON sector_raw_data(data_date) + ''') + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_sector_code ON sector_raw_data(sector_code) + ''') + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_data_type ON sector_raw_data(data_type) + ''') + cursor.execute(''' + CREATE INDEX IF NOT EXISTS idx_data_version ON sector_raw_data(data_version) + ''') + + # 板块新闻数据表 + cursor.execute(''' + CREATE TABLE IF NOT EXISTS sector_news_data ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + news_date TEXT NOT NULL, + title TEXT, + content TEXT, + source TEXT, + url TEXT, + related_sectors TEXT, + sentiment_score REAL, + importance_score REAL, + data_version INTEGER DEFAULT 1, + created_at TEXT DEFAULT CURRENT_TIMESTAMP + ) + ''') + + # AI分析报告表 + cursor.execute(''' + CREATE TABLE IF NOT EXISTS sector_analysis_reports ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + analysis_date TEXT NOT NULL, + data_date_range TEXT, + analysis_content TEXT, + recommended_sectors TEXT, + summary TEXT, + confidence_score REAL, + risk_level TEXT, + investment_horizon TEXT, + market_outlook TEXT, + created_at TEXT DEFAULT CURRENT_TIMESTAMP + ) + ''') + + # 板块追踪表(记录推荐板块的后续表现) + cursor.execute(''' + CREATE TABLE IF NOT EXISTS sector_tracking ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + analysis_id INTEGER, + sector_code TEXT NOT NULL, + sector_name TEXT, + recommended_date TEXT, + recommended_price REAL, + target_price REAL, + stop_loss_price REAL, + current_price REAL, + profit_loss_pct REAL, + status TEXT, + notes TEXT, + updated_at TEXT DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (analysis_id) REFERENCES sector_analysis_reports (id) + ) + ''') + + # 数据版本管理表 + cursor.execute(''' + CREATE TABLE IF NOT EXISTS data_versions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + data_type TEXT NOT NULL, + data_date TEXT NOT NULL, + version INTEGER NOT NULL, + status TEXT DEFAULT 'active', + fetch_success BOOLEAN DEFAULT 1, + error_message TEXT, + record_count INTEGER DEFAULT 0, + created_at TEXT DEFAULT CURRENT_TIMESTAMP, + UNIQUE(data_type, data_date, version) + ) + ''') + + conn.commit() + conn.close() + + self.logger.info("[智策板块] 数据库初始化完成") + + def save_raw_data(self, data_date, data_type, data_df, version=None): + """ + 保存原始数据 + + Args: + data_date: 数据日期 + data_type: 数据类型 (sector_data, news_data等) + data_df: 数据DataFrame + version: 数据版本号,如果为None则自动生成 + + Returns: + int: 数据版本号 + """ + conn = self.get_connection() + cursor = conn.cursor() + + try: + # 获取或生成版本号 + if version is None: + cursor.execute(''' + SELECT COALESCE(MAX(version), 0) + 1 + FROM data_versions + WHERE data_type = ? AND data_date = ? + ''', (data_type, data_date)) + version = cursor.fetchone()[0] + + # 保存数据 + if data_type == 'sector_data': + self._save_sector_data(cursor, data_date, data_df, version) + elif data_type == 'news_data': + self._save_news_data(cursor, data_date, data_df, version) + + # 记录版本信息 + cursor.execute(''' + INSERT OR REPLACE INTO data_versions + (data_type, data_date, version, status, fetch_success, record_count) + VALUES (?, ?, ?, 'active', 1, ?) + ''', (data_type, data_date, version, len(data_df))) + + conn.commit() + self.logger.info(f"[智策板块] 保存{data_type}数据成功 (日期: {data_date}, 版本: {version}, 记录数: {len(data_df)})") + return version + + except Exception as e: + conn.rollback() + # 记录失败版本 + cursor.execute(''' + INSERT OR REPLACE INTO data_versions + (data_type, data_date, version, status, fetch_success, error_message, record_count) + VALUES (?, ?, ?, 'failed', 0, ?, 0) + ''', (data_type, data_date, version or 1, str(e))) + conn.commit() + self.logger.error(f"[智策板块] 保存{data_type}数据失败: {e}") + raise + finally: + conn.close() + + def _save_sector_data(self, cursor, data_date, data_df, version): + """保存板块数据""" + for _, row in data_df.iterrows(): + cursor.execute(''' + INSERT OR REPLACE INTO sector_raw_data + (data_date, sector_code, sector_name, price, change_pct, volume, + turnover, market_cap, pe_ratio, pb_ratio, data_type, data_version) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'sector_data', ?) + ''', ( + data_date, + row.get('sector_code', ''), + row.get('sector_name', ''), + row.get('price', 0), + row.get('change_pct', 0), + row.get('volume', 0), + row.get('turnover', 0), + row.get('market_cap', 0), + row.get('pe_ratio', 0), + row.get('pb_ratio', 0), + version + )) + + def _save_news_data(self, cursor, data_date, data_df, version): + """保存新闻数据""" + for _, row in data_df.iterrows(): + cursor.execute(''' + INSERT OR REPLACE INTO sector_news_data + (news_date, title, content, source, url, related_sectors, + sentiment_score, importance_score, data_version) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) + ''', ( + data_date, + row.get('title', ''), + row.get('content', ''), + row.get('source', ''), + row.get('url', ''), + json.dumps(row.get('related_sectors', []), ensure_ascii=False), + row.get('sentiment_score', 0), + row.get('importance_score', 0), + version + )) + + def get_latest_data(self, data_type, data_date=None): + """ + 获取最新的成功数据 + + Args: + data_type: 数据类型 + data_date: 指定日期,如果为None则获取最新日期的数据 + + Returns: + pd.DataFrame: 数据DataFrame + """ + conn = self.get_connection() + + try: + # 获取最新成功的数据版本 + if data_date: + query = ''' + SELECT version FROM data_versions + WHERE data_type = ? AND data_date = ? AND fetch_success = 1 + ORDER BY version DESC LIMIT 1 + ''' + params = [data_type, data_date] + else: + query = ''' + SELECT data_date, version FROM data_versions + WHERE data_type = ? AND fetch_success = 1 + ORDER BY data_date DESC, version DESC LIMIT 1 + ''' + params = [data_type] + + version_df = pd.read_sql_query(query, conn, params=params) + + if version_df.empty: + self.logger.warning(f"[智策板块] 未找到{data_type}的成功数据") + return pd.DataFrame() + + if data_date is None: + data_date = version_df.iloc[0]['data_date'] + version = version_df.iloc[0]['version'] + + # 获取具体数据 + if data_type == 'sector_data': + data_query = ''' + SELECT * FROM sector_raw_data + WHERE data_date = ? AND data_version = ? + ORDER BY sector_code + ''' + elif data_type == 'news_data': + data_query = ''' + SELECT * FROM sector_news_data + WHERE news_date = ? AND data_version = ? + ORDER BY importance_score DESC + ''' + else: + return pd.DataFrame() + + data_df = pd.read_sql_query(data_query, conn, params=[data_date, version]) + self.logger.info(f"[智策板块] 获取{data_type}数据成功 (日期: {data_date}, 版本: {version}, 记录数: {len(data_df)})") + return data_df + + except Exception as e: + self.logger.error(f"[智策板块] 获取{data_type}数据失败: {e}") + return pd.DataFrame() + finally: + conn.close() + + def save_analysis_report(self, data_date_range, analysis_content, + recommended_sectors, summary, confidence_score=None, + risk_level=None, investment_horizon=None, market_outlook=None): + """ + 保存AI分析报告 + + Args: + data_date_range: 数据日期范围 + analysis_content: 分析内容(JSON字符串或字典) + recommended_sectors: 推荐板块列表 + summary: 摘要 + confidence_score: 置信度分数 + risk_level: 风险等级 + investment_horizon: 投资周期 + market_outlook: 市场展望 + + Returns: + int: 报告ID + """ + conn = self.get_connection() + cursor = conn.cursor() + + # 如果传入的是字典,转换为JSON字符串 + if isinstance(analysis_content, dict): + analysis_content = json.dumps(analysis_content, ensure_ascii=False, indent=2) + + cursor.execute(''' + INSERT INTO sector_analysis_reports + (analysis_date, data_date_range, analysis_content, recommended_sectors, + summary, confidence_score, risk_level, investment_horizon, market_outlook) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) + ''', ( + datetime.now().strftime('%Y-%m-%d %H:%M:%S'), + data_date_range, + analysis_content, + json.dumps(recommended_sectors, ensure_ascii=False), + summary, + confidence_score, + risk_level, + investment_horizon, + market_outlook + )) + + report_id = cursor.lastrowid + + conn.commit() + conn.close() + + self.logger.info(f"[智策板块] 分析报告已保存 (ID: {report_id})") + return report_id + + def get_analysis_reports(self, limit=10): + """ + 获取历史分析报告 + + Args: + limit: 返回数量 + + Returns: + pd.DataFrame: 报告列表 + """ + conn = self.get_connection() + + query = ''' + SELECT * FROM sector_analysis_reports + ORDER BY created_at DESC + LIMIT ? + ''' + + df = pd.read_sql_query(query, conn, params=[limit]) + conn.close() + + return df + + def get_analysis_report(self, report_id): + """ + 获取单个分析报告详情 + + Args: + report_id: 报告ID + + Returns: + dict: 报告详情 + """ + conn = self.get_connection() + cursor = conn.cursor() + + cursor.execute(''' + SELECT * FROM sector_analysis_reports WHERE id = ? + ''', (report_id,)) + + row = cursor.fetchone() + columns = [desc[0] for desc in cursor.description] if cursor.description else [] + conn.close() + + if row: + report = dict(zip(columns, row)) + + # 解析JSON字段 + try: + if report.get('analysis_content'): + report['analysis_content_parsed'] = json.loads(report['analysis_content']) + if report.get('recommended_sectors'): + report['recommended_sectors_parsed'] = json.loads(report['recommended_sectors']) + except json.JSONDecodeError as e: + self.logger.warning(f"[智策板块] JSON解析失败: {e}") + + return report + + return None + + def delete_analysis_report(self, report_id): + """ + 删除分析报告 + + Args: + report_id: 报告ID + + Returns: + bool: 删除是否成功 + """ + conn = self.get_connection() + cursor = conn.cursor() + + try: + # 删除相关的追踪记录 + cursor.execute('DELETE FROM sector_tracking WHERE analysis_id = ?', (report_id,)) + + # 删除报告 + cursor.execute('DELETE FROM sector_analysis_reports WHERE id = ?', (report_id,)) + + deleted_count = cursor.rowcount + conn.commit() + + if deleted_count > 0: + self.logger.info(f"[智策板块] 报告删除成功 (ID: {report_id})") + return True + else: + self.logger.warning(f"[智策板块] 未找到要删除的报告 (ID: {report_id})") + return False + + except Exception as e: + conn.rollback() + self.logger.error(f"[智策板块] 删除报告失败: {e}") + return False + finally: + conn.close() + + def get_data_versions(self, data_type, limit=10): + """ + 获取数据版本历史 + + Args: + data_type: 数据类型 + limit: 返回数量 + + Returns: + pd.DataFrame: 版本历史 + """ + conn = self.get_connection() + + query = ''' + SELECT * FROM data_versions + WHERE data_type = ? + ORDER BY data_date DESC, version DESC + LIMIT ? + ''' + + df = pd.read_sql_query(query, conn, params=[data_type, limit]) + conn.close() + + return df + + def save_sector_raw_data(self, data_date, data_type, data_df): + """ + 保存板块原始数据 + + Args: + data_date: 数据日期 + data_type: 数据类型 ('industry', 'concept', 'fund_flow', 'market_overview', 'north_fund', 'news') + data_df: 数据DataFrame + """ + # 兼容不同数据结构的空值判断 + is_empty = False + if data_df is None: + is_empty = True + elif hasattr(data_df, 'empty'): + is_empty = data_df.empty + elif isinstance(data_df, (list, tuple, set, dict)): + is_empty = len(data_df) == 0 + if is_empty: + self.logger.warning(f"[智策板块] {data_type}数据为空,跳过保存") + return + + conn = self.get_connection() + cursor = conn.cursor() + + try: + # 获取下一个版本号 + version = self._get_next_version(data_date, data_type) + + # 根据数据类型保存数据 + if data_type in ['industry', 'concept']: + self._save_sector_data_raw(cursor, data_date, data_df, data_type, version) + elif data_type == 'fund_flow': + self._save_fund_flow_data(cursor, data_date, data_df, version) + elif data_type == 'market_overview': + self._save_market_overview_data(cursor, data_date, data_df, version) + elif data_type == 'north_fund': + self._save_north_fund_data(cursor, data_date, data_df, version) + elif data_type == 'news': + self._save_news_data_raw(cursor, data_date, data_df, version) + + # 记录版本信息 + cursor.execute(''' + INSERT OR REPLACE INTO data_versions + (data_date, data_type, version, fetch_success, record_count) + VALUES (?, ?, ?, 1, ?) + ''', (data_date, data_type, version, len(data_df))) + + conn.commit() + self.logger.info(f"[智策板块] {data_type}数据保存成功 (日期: {data_date}, 版本: {version}, 记录数: {len(data_df)})") + + except Exception as e: + conn.rollback() + self.logger.error(f"[智策板块] 保存{data_type}数据失败: {e}") + raise + finally: + conn.close() + + def _save_sector_data_raw(self, cursor, data_date, data_df, data_type, version): + """保存板块原始数据""" + for _, row in data_df.iterrows(): + cursor.execute(''' + INSERT OR REPLACE INTO sector_raw_data + (data_date, sector_code, sector_name, price, change_pct, volume, + turnover, market_cap, pe_ratio, pb_ratio, data_type, data_version) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ''', ( + data_date, + str(row.get('板块代码', row.get('sector_code', ''))), + str(row.get('板块名称', row.get('sector_name', ''))), + float(row.get('最新价', row.get('price', 0))) if pd.notna(row.get('最新价', row.get('price', 0))) else 0, + float(row.get('涨跌幅', row.get('change_pct', 0))) if pd.notna(row.get('涨跌幅', row.get('change_pct', 0))) else 0, + float(row.get('成交量', row.get('volume', 0))) if pd.notna(row.get('成交量', row.get('volume', 0))) else 0, + float(row.get('成交额', row.get('turnover', 0))) if pd.notna(row.get('成交额', row.get('turnover', 0))) else 0, + float(row.get('总市值', row.get('market_cap', 0))) if pd.notna(row.get('总市值', row.get('market_cap', 0))) else 0, + float(row.get('市盈率', row.get('pe_ratio', 0))) if pd.notna(row.get('市盈率', row.get('pe_ratio', 0))) else 0, + float(row.get('市净率', row.get('pb_ratio', 0))) if pd.notna(row.get('市净率', row.get('pb_ratio', 0))) else 0, + data_type, + version + )) + + def _save_fund_flow_data(self, cursor, data_date, data_df, version): + """保存资金流向数据""" + for _, row in data_df.iterrows(): + cursor.execute(''' + INSERT OR REPLACE INTO sector_raw_data + (data_date, sector_code, sector_name, price, change_pct, volume, + turnover, market_cap, pe_ratio, pb_ratio, data_type, data_version) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'fund_flow', ?) + ''', ( + data_date, + str(row.get('行业', '')), + str(row.get('行业', '')), + float(row.get('主力净流入-净额', 0)) if pd.notna(row.get('主力净流入-净额', 0)) else 0, + float(row.get('主力净流入-净占比', 0)) if pd.notna(row.get('主力净流入-净占比', 0)) else 0, + float(row.get('超大单净流入-净额', 0)) if pd.notna(row.get('超大单净流入-净额', 0)) else 0, + float(row.get('超大单净流入-净占比', 0)) if pd.notna(row.get('超大单净流入-净占比', 0)) else 0, + float(row.get('大单净流入-净额', 0)) if pd.notna(row.get('大单净流入-净额', 0)) else 0, + float(row.get('大单净流入-净占比', 0)) if pd.notna(row.get('大单净流入-净占比', 0)) else 0, + 0, + version + )) + + def _save_market_overview_data(self, cursor, data_date, data_df, version): + """保存市场概况数据""" + for _, row in data_df.iterrows(): + cursor.execute(''' + INSERT OR REPLACE INTO sector_raw_data + (data_date, sector_code, sector_name, price, change_pct, volume, + turnover, market_cap, pe_ratio, pb_ratio, data_type, data_version) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'market_overview', ?) + ''', ( + data_date, + str(row.get('名称', '')), + str(row.get('名称', '')), + float(row.get('最新价', 0)) if pd.notna(row.get('最新价', 0)) else 0, + float(row.get('涨跌幅', 0)) if pd.notna(row.get('涨跌幅', 0)) else 0, + float(row.get('成交量', 0)) if pd.notna(row.get('成交量', 0)) else 0, + float(row.get('成交额', 0)) if pd.notna(row.get('成交额', 0)) else 0, + 0, 0, 0, + version + )) + + def _save_north_fund_data(self, cursor, data_date, data_df, version): + """保存北向资金数据""" + for _, row in data_df.iterrows(): + cursor.execute(''' + INSERT OR REPLACE INTO sector_raw_data + (data_date, sector_code, sector_name, price, change_pct, volume, + turnover, market_cap, pe_ratio, pb_ratio, data_type, data_version) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'north_fund', ?) + ''', ( + data_date, + str(row.get('代码', '')), + str(row.get('名称', '')), + float(row.get('收盘价', 0)) if pd.notna(row.get('收盘价', 0)) else 0, + float(row.get('涨跌幅', 0)) if pd.notna(row.get('涨跌幅', 0)) else 0, + float(row.get('持股数量', 0)) if pd.notna(row.get('持股数量', 0)) else 0, + float(row.get('持股市值', 0)) if pd.notna(row.get('持股市值', 0)) else 0, + float(row.get('持股变化', 0)) if pd.notna(row.get('持股变化', 0)) else 0, + 0, 0, + version + )) + + def _save_news_data_raw(self, cursor, data_date, data_df, version): + """保存新闻数据""" + for _, row in data_df.iterrows(): + cursor.execute(''' + INSERT OR REPLACE INTO sector_news_data + (news_date, title, content, source, url, related_sectors, + sentiment_score, importance_score, data_version) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) + ''', ( + data_date, + str(row.get('新闻标题', row.get('title', ''))), + str(row.get('新闻内容', row.get('content', ''))), + str(row.get('新闻来源', row.get('source', ''))), + str(row.get('新闻链接', row.get('url', ''))), + json.dumps([], ensure_ascii=False), # 暂时为空 + 0, # 暂时为0 + 0, # 暂时为0 + version + )) + + def cleanup_old_data(self, data_type, keep_days=30): + """ + 清理旧数据,保留指定天数的数据 + + Args: + data_type: 数据类型 + keep_days: 保留天数 + + Returns: + int: 删除的记录数 + """ + conn = self.get_connection() + cursor = conn.cursor() + + try: + cutoff_date = (datetime.now() - pd.Timedelta(days=keep_days)).strftime('%Y-%m-%d') + + if data_type == 'sector_data': + cursor.execute(''' + DELETE FROM sector_raw_data + WHERE data_date < ? + ''', (cutoff_date,)) + elif data_type == 'news_data': + cursor.execute(''' + DELETE FROM sector_news_data + WHERE news_date < ? + ''', (cutoff_date,)) + + deleted_count = cursor.rowcount + + # 同时清理版本记录 + cursor.execute(''' + DELETE FROM data_versions + WHERE data_type = ? AND data_date < ? + ''', (data_type, cutoff_date)) + + conn.commit() + self.logger.info(f"[智策板块] 清理{data_type}旧数据完成,删除{deleted_count}条记录") + return deleted_count + + except Exception as e: + conn.rollback() + self.logger.error(f"[智策板块] 清理{data_type}旧数据失败: {e}") + return 0 + finally: + conn.close() + + # ===================== + # 缓存与最近数据读取接口 + # ===================== + def save_news_data(self, news_list, news_date, source="akshare"): + """ + 保存新闻列表(字典列表)到数据库,用于非DataFrame场景 + Args: + news_list: [{title, content, url, related_sectors, sentiment_score, importance_score}] + news_date: 新闻日期字符串 + source: 新闻来源 + """ + if not news_list: + self.logger.warning("[智策板块] 新闻列表为空,跳过保存") + return 0 + + conn = self.get_connection() + cursor = conn.cursor() + try: + # 版本号按日期累加 + version = self._get_next_version(news_date, 'news') + inserted = 0 + for item in news_list: + cursor.execute(''' + INSERT OR REPLACE INTO sector_news_data + (news_date, title, content, source, url, related_sectors, + sentiment_score, importance_score, data_version) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) + ''', ( + str(news_date), + str(item.get('title', '')), + str(item.get('content', '')), + str(item.get('source', source)), + str(item.get('url', '')), + json.dumps(item.get('related_sectors', []), ensure_ascii=False), + float(item.get('sentiment_score', 0) or 0), + float(item.get('importance_score', 0) or 0), + version + )) + inserted += 1 + + # 记录版本信息 + cursor.execute(''' + INSERT OR REPLACE INTO data_versions + (data_date, data_type, version, fetch_success, record_count) + VALUES (?, ?, ?, 1, ?) + ''', (str(news_date), 'news', version, inserted)) + + conn.commit() + self.logger.info(f"[智策板块] 保存新闻数据成功 (日期: {news_date}, 版本: {version}, 记录数: {inserted})") + return inserted + except Exception as e: + conn.rollback() + self.logger.error(f"[智策板块] 保存新闻数据失败: {e}") + return 0 + finally: + conn.close() + + def _get_next_version(self, data_date: str, data_type: str) -> int: + """获取指定日期与类型的下一个版本号""" + conn = self.get_connection() + cursor = conn.cursor() + try: + cursor.execute(''' + SELECT COALESCE(MAX(version), 0) + 1 FROM data_versions + WHERE data_type = ? AND data_date = ? + ''', (data_type, data_date)) + next_version = cursor.fetchone()[0] or 1 + return int(next_version) + finally: + conn.close() + + def get_latest_raw_data(self, key: str, within_hours: int = 24): + """ + 获取最近within_hours小时内的原始数据并组装为分析所需结构 + Args: + key: 'sectors' | 'concepts' | 'fund_flow' | 'market_overview' | 'north_flow' + within_hours: 有效缓存时长(小时) + Returns: + dict 或 None + """ + # 将key映射到内部data_type + key_map = { + 'sectors': 'industry', + 'concepts': 'concept', + 'fund_flow': 'fund_flow', + 'market_overview': 'market_overview', + 'north_flow': 'north_fund' + } + data_type = key_map.get(key) + if not data_type: + return None + + conn = self.get_connection() + try: + cutoff = (pd.Timestamp.now() - pd.Timedelta(hours=within_hours)).strftime('%Y-%m-%d %H:%M:%S') + # 选取最近版本的数据(同一天可能有多版本) + # 先查最近有效版本记录 + version_df = pd.read_sql_query(''' + SELECT data_date, version FROM data_versions + WHERE data_type = ? AND fetch_success = 1 + AND datetime(created_at) >= datetime(?) + ORDER BY data_date DESC, version DESC LIMIT 1 + ''', conn, params=[data_type, cutoff]) + + if version_df.empty: + return None + + data_date = version_df.iloc[0]['data_date'] + version = int(version_df.iloc[0]['version']) + + # 读取具体行 + raw_df = pd.read_sql_query(''' + SELECT * FROM sector_raw_data + WHERE data_type = ? AND data_date = ? AND data_version = ? + ''', conn, params=[data_type, data_date, version]) + + if raw_df.empty: + return None + + # 组装成预期结构 + if key in ['sectors', 'concepts']: + result = {} + for _, row in raw_df.iterrows(): + name = str(row.get('sector_name', '')) + result[name] = { + 'name': name, + 'change_pct': float(row.get('change_pct', 0) or 0), + 'price': float(row.get('price', 0) or 0), + 'volume': float(row.get('volume', 0) or 0), + 'turnover': float(row.get('turnover', 0) or 0), + 'market_cap': float(row.get('market_cap', 0) or 0), + 'pe_ratio': float(row.get('pe_ratio', 0) or 0), + 'pb_ratio': float(row.get('pb_ratio', 0) or 0), + } + return { + 'data_date': data_date, + 'data_content': result + } + + if key == 'fund_flow': + today = [] + for _, row in raw_df.iterrows(): + name = str(row.get('sector_name', '')) + today.append({ + 'sector': name, + 'main_net_inflow': float(row.get('price', 0) or 0), # 映射自主力净额 + 'main_net_inflow_pct': float(row.get('change_pct', 0) or 0), + 'super_large_net_inflow': float(row.get('volume', 0) or 0), + 'super_large_net_inflow_pct': float(row.get('turnover', 0) or 0), + 'large_net_inflow': float(row.get('market_cap', 0) or 0), + 'large_net_inflow_pct': float(row.get('pe_ratio', 0) or 0), + 'medium_net_inflow': 0, + 'small_net_inflow': 0 + }) + return { + 'data_date': data_date, + 'data_content': { + 'today': today + } + } + + if key == 'market_overview': + overview = {} + for _, row in raw_df.iterrows(): + name = str(row.get('sector_name', '')) + entry = { + 'price': float(row.get('price', 0) or 0), + 'change_pct': float(row.get('change_pct', 0) or 0), + 'turnover': float(row.get('turnover', 0) or 0), + 'volume': float(row.get('volume', 0) or 0) + } + # 简单映射:名称包含上证/深证/创业板 + if '上证' in name or '沪指' in name or 'SH' in name: + overview['sh_index'] = entry + elif '深证' in name or 'SZ' in name: + overview['sz_index'] = entry + elif '创业' in name or 'CYB' in name: + overview['cyb_index'] = entry + return { + 'data_date': data_date, + 'data_content': overview + } + + if key == 'north_flow': + # 北向资金结构差异较大,返回最简结构用于提示 + total_value = float(raw_df['turnover'].sum()) if not raw_df.empty else 0 + return { + 'data_date': data_date, + 'data_content': { + 'north_total_amount': total_value, + 'history': [] + } + } + + return None + except Exception as e: + self.logger.error(f"[智策板块] 获取最近原始数据失败: {e}") + return None + finally: + conn.close() + + def get_latest_news_data(self, within_hours: int = 24): + """获取最近within_hours小时的新闻列表""" + conn = self.get_connection() + try: + cutoff = (pd.Timestamp.now() - pd.Timedelta(hours=within_hours)).strftime('%Y-%m-%d %H:%M:%S') + df = pd.read_sql_query(''' + SELECT * FROM sector_news_data + WHERE datetime(created_at) >= datetime(?) + ORDER BY importance_score DESC, created_at DESC + ''', conn, params=[cutoff]) + if df.empty: + return None + news = [] + for _, row in df.iterrows(): + try: + related = json.loads(row.get('related_sectors', '[]')) + except Exception: + related = [] + news.append({ + 'title': row.get('title', ''), + 'content': row.get('content', ''), + 'source': row.get('source', ''), + 'url': row.get('url', ''), + 'related_sectors': related, + 'sentiment_score': float(row.get('sentiment_score', 0) or 0), + 'importance_score': float(row.get('importance_score', 0) or 0), + 'news_date': row.get('news_date', '') + }) + return { + 'data_date': df.iloc[0]['news_date'] if not df.empty else None, + 'data_content': news + } + except Exception as e: + self.logger.error(f"[智策板块] 获取最近新闻数据失败: {e}") + return None + finally: + conn.close() \ No newline at end of file diff --git a/sector_strategy_engine.py b/sector_strategy_engine.py index 7fdfb9c..5e39364 100644 --- a/sector_strategy_engine.py +++ b/sector_strategy_engine.py @@ -4,10 +4,13 @@ """ from sector_strategy_agents import SectorStrategyAgents +from sector_strategy_db import SectorStrategyDatabase from deepseek_client import DeepSeekClient from typing import Dict, Any import time import json +import pandas as pd +import logging class SectorStrategyEngine: @@ -17,8 +20,81 @@ class SectorStrategyEngine: self.model = model self.agents = SectorStrategyAgents(model=model) self.deepseek_client = DeepSeekClient(model=model) + self.database = SectorStrategyDatabase() + self.logger = logging.getLogger(__name__) + if not self.logger.handlers: + logging.basicConfig(level=logging.INFO, format='[%(asctime)s] %(levelname)s %(name)s: %(message)s') print(f"[智策引擎] 初始化完成 (模型: {model})") + def save_raw_data_with_fallback(self, data_type, data_df, data_date=None): + """ + 保存原始数据,支持失败回退机制 + + Args: + data_type: 数据类型 + data_df: 数据DataFrame + data_date: 数据日期,默认为今天 + + Returns: + tuple: (success, version, message) + """ + if data_date is None: + data_date = time.strftime("%Y-%m-%d") + + try: + is_empty = False + if data_df is None: + is_empty = True + elif hasattr(data_df, 'empty'): + is_empty = data_df.empty + elif isinstance(data_df, (list, tuple, set, dict)): + is_empty = len(data_df) == 0 + if is_empty: + self.logger.warning(f"[智策引擎] {data_type}数据为空,跳过保存") + return False, None, "数据为空" + + version = self.database.save_raw_data(data_date, data_type, data_df) + return True, version, f"保存成功,版本: {version}" + + except Exception as e: + self.logger.error(f"[智策引擎] 保存{data_type}数据失败: {e}") + return False, None, str(e) + + def get_data_with_fallback(self, data_type, data_date=None): + """ + 获取数据,支持失败时回退到历史数据 + + Args: + data_type: 数据类型 + data_date: 数据日期,默认为今天 + + Returns: + tuple: (data_df, is_fallback, message) + """ + if data_date is None: + data_date = time.strftime("%Y-%m-%d") + + try: + # 尝试获取指定日期的数据 + data_df = self.database.get_latest_data(data_type, data_date) + + if not data_df.empty: + return data_df, False, f"获取{data_date}数据成功" + + # 如果指定日期没有数据,获取最新的历史数据 + self.logger.warning(f"[智策引擎] {data_date}的{data_type}数据不存在,尝试获取历史数据") + data_df = self.database.get_latest_data(data_type) + + if not data_df.empty: + fallback_date = data_df.iloc[0].get('data_date', '未知日期') + return data_df, True, f"回退到{fallback_date}的历史数据" + else: + return pd.DataFrame(), True, "无可用的历史数据" + + except Exception as e: + self.logger.error(f"[智策引擎] 获取{data_type}数据失败: {e}") + return pd.DataFrame(), True, str(e) + def run_comprehensive_analysis(self, data: Dict) -> Dict[str, Any]: """ 运行综合分析流程 @@ -102,6 +178,24 @@ class SectorStrategyEngine: results["success"] = True + # 4. 保存分析报告 + print("\n[阶段4] 保存分析报告...") + print("-" * 60) + try: + report_id = self.save_analysis_report(results, data) + results["report_id"] = report_id + print(f"✓ 分析报告已保存 (ID: {report_id})") + # 保存后读取报告详情并回传到结果,用于主页面动态渲染 + try: + saved_report = self.database.get_analysis_report(report_id) + if saved_report: + results["saved_report"] = saved_report + except Exception as fetch_e: + self.logger.warning(f"[智策引擎] 获取保存报告详情失败: {fetch_e}") + except Exception as e: + print(f"⚠ 保存分析报告失败: {e}") + self.logger.error(f"[智策引擎] 保存分析报告失败: {e}") + print("\n" + "=" * 60) print("✓ 智策综合分析完成!") print("=" * 60) @@ -320,6 +414,155 @@ class SectorStrategyEngine: except Exception as e: print(f" ⚠ JSON解析失败: {e},返回文本格式") return {"prediction_text": response} + + def save_analysis_report(self, results: Dict, original_data: Dict) -> int: + """ + 保存分析报告到数据库 + + Args: + results: 分析结果 + original_data: 原始数据 + + Returns: + int: 报告ID + """ + try: + # 提取数据日期范围 + data_date_range = f"{time.strftime('%Y-%m-%d')} 数据分析" + + # 提取推荐板块 + recommended_sectors = [] + predictions = results.get("final_predictions", {}) + + if isinstance(predictions, dict): + # 从预测结果中提取推荐板块 + hot_sectors = predictions.get("hot_sectors", []) + rotation_sectors = predictions.get("rotation_opportunities", []) + + for sector in hot_sectors[:5]: # 取前5个热门板块 + if isinstance(sector, dict): + recommended_sectors.append({ + "sector_name": sector.get("name", ""), + "reason": sector.get("reason", ""), + "confidence": sector.get("confidence", ""), + "type": "热门板块" + }) + + for sector in rotation_sectors[:3]: # 取前3个轮动机会 + if isinstance(sector, dict): + recommended_sectors.append({ + "sector_name": sector.get("name", ""), + "reason": sector.get("reason", ""), + "confidence": sector.get("confidence", ""), + "type": "轮动机会" + }) + + # 生成摘要 + summary = self._generate_report_summary(results) + + # 提取其他信息 + confidence_score = self._extract_confidence_score(results) + risk_level = self._extract_risk_level(results) + investment_horizon = self._extract_investment_horizon(results) + market_outlook = self._extract_market_outlook(results) + + # 保存到数据库 + report_id = self.database.save_analysis_report( + data_date_range=data_date_range, + analysis_content=results, + recommended_sectors=recommended_sectors, + summary=summary, + confidence_score=confidence_score, + risk_level=risk_level, + investment_horizon=investment_horizon, + market_outlook=market_outlook + ) + + return report_id + + except Exception as e: + self.logger.error(f"[智策引擎] 保存分析报告失败: {e}") + raise + + def _generate_report_summary(self, results: Dict) -> str: + """生成报告摘要""" + try: + predictions = results.get("final_predictions", {}) + if isinstance(predictions, dict): + # 从summary中提取市场趋势信息 + summary_info = predictions.get("summary", {}) + market_trend = summary_info.get("market_view", "") if isinstance(summary_info, dict) else "" + + # 从long_short.bullish中计算热门板块数量 + long_short_info = predictions.get("long_short", {}) + bullish_sectors = long_short_info.get("bullish", []) if isinstance(long_short_info, dict) else [] + hot_sectors_count = len(bullish_sectors) + + # 如果有看多板块信息,则添加到摘要中 + if bullish_sectors and isinstance(bullish_sectors, list): + # 提取前3个看多板块名称 + bullish_names = [sector.get("sector", "") for sector in bullish_sectors[:3] if isinstance(sector, dict)] + if bullish_names: + bullish_text = ",".join(bullish_names) + return f"市场趋势: {market_trend},识别{hot_sectors_count}个热门板块机会,看多板块: {bullish_text}" + + return f"市场趋势: {market_trend},识别{hot_sectors_count}个热门板块机会" + else: + return "智策板块分析报告" + except: + return "智策板块分析报告" + + def _extract_confidence_score(self, results: Dict) -> float: + """提取置信度分数""" + try: + predictions = results.get("final_predictions", {}) + if isinstance(predictions, dict): + return predictions.get("confidence_score", 0.75) + return 0.75 + except: + return 0.75 + + def _extract_risk_level(self, results: Dict) -> str: + """提取风险等级""" + try: + predictions = results.get("final_predictions", {}) + if isinstance(predictions, dict): + return predictions.get("risk_level", "中等") + return "中等" + except: + return "中等" + + def _extract_investment_horizon(self, results: Dict) -> str: + """提取投资周期""" + try: + predictions = results.get("final_predictions", {}) + if isinstance(predictions, dict): + return predictions.get("investment_horizon", "短期") + return "短期" + except: + return "短期" + + def _extract_market_outlook(self, results: Dict) -> str: + """提取市场展望""" + try: + predictions = results.get("final_predictions", {}) + if isinstance(predictions, dict): + return predictions.get("market_outlook", "谨慎乐观") + return "谨慎乐观" + except: + return "谨慎乐观" + + def get_historical_reports(self, limit=10): + """获取历史报告""" + return self.database.get_analysis_reports(limit) + + def get_report_detail(self, report_id): + """获取报告详情""" + return self.database.get_analysis_report(report_id) + + def delete_report(self, report_id): + """删除报告""" + return self.database.delete_analysis_report(report_id) # 测试函数 diff --git a/sector_strategy_ui.py b/sector_strategy_ui.py index 6855db9..8b03514 100644 --- a/sector_strategy_ui.py +++ b/sector_strategy_ui.py @@ -4,19 +4,39 @@ """ import streamlit as st +import time import plotly.graph_objects as go import plotly.express as px import pandas as pd from datetime import datetime, time as dt_time import time import base64 +import json from sector_strategy_data import SectorStrategyDataFetcher from sector_strategy_engine import SectorStrategyEngine from sector_strategy_pdf import SectorStrategyPDFGenerator +from sector_strategy_db import SectorStrategyDatabase from sector_strategy_scheduler import sector_strategy_scheduler +def _parse_json_field(value, default): + """将可能的JSON字符串安全转换为Python对象""" + try: + if isinstance(value, (dict, list)): + return value + if value is None: + return default + if isinstance(value, str): + v = value.strip() + if not v: + return default + return json.loads(v) + return default + except Exception: + return default + + def display_sector_strategy(): """显示智策板块分析主界面""" @@ -29,6 +49,19 @@ def display_sector_strategy(): st.markdown("---") + # 创建标签页 + tab1, tab2 = st.tabs(["📊 智策分析", "📋 历史报告"]) + + with tab1: + display_analysis_tab() + + with tab2: + display_history_tab() + + +def display_analysis_tab(): + """显示分析标签页""" + # 定时任务设置区域 display_scheduler_settings() @@ -92,12 +125,12 @@ def display_sector_strategy(): with col2: st.write("") st.write("") - analyze_button = st.button("🚀 开始智策分析", type="primary", use_container_width=True) + analyze_button = st.button("🚀 开始智策分析", type="primary", width='content') with col3: st.write("") st.write("") - if st.button("🔄 清除结果", use_container_width=True): + if st.button("🔄 清除结果", width='content'): if 'sector_strategy_result' in st.session_state: del st.session_state.sector_strategy_result st.success("已清除分析结果") @@ -123,6 +156,108 @@ def display_sector_strategy(): st.error(f"❌ 分析失败: {result.get('error', '未知错误')}") +def display_history_tab(): + """显示历史报告标签页""" + + st.markdown("### 📋 智策历史报告") + st.markdown("查看和管理历史分析报告") + + try: + # 初始化引擎以获取历史报告 + engine = SectorStrategyEngine() + + # 获取历史报告 + reports = engine.get_historical_reports(limit=20) + + if reports.empty: + st.info("📝 暂无历史报告") + st.markdown(""" + **提示**: + - 运行智策分析后,报告将自动保存到历史记录中 + - 您可以在此查看和管理所有历史分析报告 + """) + return + + st.success(f"📊 共找到 {len(reports)} 份历史报告") + + # 报告列表(精简摘要展示) + for i, report in reports.iterrows(): + report_id = report['id'] if 'id' in report else None + created_at = report['created_at'] if 'created_at' in report else '' + data_date_range = report['data_date_range'] if 'data_date_range' in report else '' + summary = report['summary'] if 'summary' in report else '智策板块分析报告' + confidence_score = report['confidence_score'] if 'confidence_score' in report else 0 + risk_level = report['risk_level'] if 'risk_level' in report else '中等' + market_outlook = report['market_outlook'] if 'market_outlook' in report else '谨慎乐观' + + with st.container(): + st.markdown(f"**📊 报告 #{report_id}**") + st.caption(f"生成时间: {created_at} | 数据区间: {data_date_range}") + + col1, col2, col3 = st.columns([1, 1, 1]) + with col1: + st.metric("置信度", f"{confidence_score:.1%}") + with col2: + st.metric("风险等级", risk_level) + with col3: + st.metric("市场展望", market_outlook) + + # 操作区:加载到分析视图 / 删除 + op1, op2 = st.columns([1, 1]) + with op1: + if st.button("📥 加载到分析视图", key=f"load_{report_id}"): + # 获取报告详情并写入session以展示到分析视图 + detail = engine.get_report_detail(report_id) + if detail and isinstance(detail.get('analysis_content_parsed'), dict): + st.session_state.sector_strategy_result = detail['analysis_content_parsed'] + st.session_state.sector_strategy_result_source = 'from_history' + st.session_state.loaded_report_id = report_id + st.success("✅ 已加载到分析视图,请切换到‘智策分析’标签查看") + time.sleep(0.5) + st.rerun() + else: + st.error("❌ 加载失败:报告内容缺失") + with op2: + if st.button(f"🗑️ 删除", key=f"delete_{report_id}"): + if engine.delete_report(report_id): + st.success("报告已删除") + st.rerun() + else: + st.error("删除失败") + + # 改进的摘要展示逻辑,突出看多板块信息 + st.markdown("**📝 报告摘要**") + summary_text = summary or "智策板块分析报告" + + # 解析摘要中的看多板块信息 + if "看多板块:" in summary_text: + parts = summary_text.split(",看多板块:") + main_summary = parts[0] + bullish_info = parts[1] if len(parts) > 1 else "" + + # 显示主要摘要信息 + st.markdown(f"🔹 {main_summary}") + + # 特别突出显示看多板块 + if bullish_info: + st.markdown(f"📈 **看多板块**: :green[{bullish_info}]") + else: + # 原有的简单展示方式 + short = summary_text if len(summary_text) <= 120 else (summary_text[:120] + "...") + with st.expander(f"{short}", expanded=False): + st.write(summary_text) + + st.markdown("-") + + except Exception as e: + st.error(f"❌ 加载历史报告失败: {e}") + + +def display_report_detail(report_id): + """详细报告页面已移除:保留占位以避免旧调用报错""" + st.info("当前版本仅提供报告摘要,详细页面已移除。") + + def run_sector_strategy_analysis(model="deepseek-chat"): """运行智策分析""" @@ -136,7 +271,8 @@ def run_sector_strategy_analysis(model="deepseek-chat"): progress_bar.progress(10) fetcher = SectorStrategyDataFetcher() - data = fetcher.get_all_sector_data() + # 使用带缓存回退的获取逻辑 + data = fetcher.get_cached_data_with_fallback() if not data.get("success"): st.error("❌ 数据获取失败") @@ -145,7 +281,7 @@ def run_sector_strategy_analysis(model="deepseek-chat"): progress_bar.progress(30) status_text.text("✓ 数据获取完成") - # 显示数据摘要 + # 显示数据摘要(含缓存提示) display_data_summary(data) # 2. 运行AI分析 @@ -154,6 +290,13 @@ def run_sector_strategy_analysis(model="deepseek-chat"): engine = SectorStrategyEngine(model=model) result = engine.run_comprehensive_analysis(data) + # 传递缓存元信息到结果以便页面提示 + if data.get("from_cache") or data.get("cache_warning"): + result["cache_meta"] = { + "from_cache": bool(data.get("from_cache")), + "cache_warning": data.get("cache_warning", ""), + "data_timestamp": data.get("timestamp") + } progress_bar.progress(90) @@ -185,6 +328,9 @@ def run_sector_strategy_analysis(model="deepseek-chat"): def display_data_summary(data): """显示数据摘要""" st.subheader("📊 市场数据概览") + # 缓存提示横幅 + if data.get("from_cache") or data.get("cache_warning"): + st.warning(data.get("cache_warning", "当前数据来自缓存,可能不是最新信息")) col1, col2, col3, col4 = st.columns(4) @@ -216,11 +362,70 @@ def display_data_summary(data): st.metric("概念板块", concepts_count) +def display_saved_report_summary(saved_report: dict): + """在主页面显示保存的报告摘要(标题、时间、关键指标)""" + st.subheader("📝 报告摘要") + summary = saved_report.get('summary', '智策板块分析报告') + created_at = saved_report.get('created_at', '') + data_date_range = saved_report.get('data_date_range', '') + confidence_score = saved_report.get('confidence_score', 0) + risk_level = saved_report.get('risk_level', '中等') + market_outlook = saved_report.get('market_outlook', '谨慎乐观') + st.caption(f"生成时间: {created_at} | 数据区间: {data_date_range}") + + # 使用改进的摘要展示逻辑,突出看多板块信息 + summary_text = summary or "智策板块分析报告" + + # 解析摘要中的看多板块信息 + if "看多板块:" in summary_text: + parts = summary_text.split(",看多板块:") + main_summary = parts[0] + bullish_info = parts[1] if len(parts) > 1 else "" + + # 显示主要摘要信息 + st.markdown(f"🔹 {main_summary}") + + # 特别突出显示看多板块 + if bullish_info: + st.markdown(f"📈 **看多板块**: :green[{bullish_info}]") + else: + # 原有的简单展示方式 + st.info(summary_text) + + col1, col2, col3 = st.columns(3) + with col1: + st.metric("置信度", f"{confidence_score:.1%}") + with col2: + st.metric("风险等级", risk_level) + with col3: + st.metric("市场展望", market_outlook) + + def display_analysis_results(result): """显示分析结果""" st.success("✅ 智策分析完成!") st.info(f"📅 分析时间: {result.get('timestamp', 'N/A')}") + # 显示缓存提示(如果本次分析使用了缓存数据) + cache_meta = result.get("cache_meta") + if cache_meta and (cache_meta.get("from_cache") or cache_meta.get("cache_warning")): + st.warning(cache_meta.get("cache_warning", "当前分析基于缓存数据,可能不是最新信息")) + + # 如果内容源自历史报告,给出返回入口 + if st.session_state.get('sector_strategy_result_source') == 'from_history': + loaded_id = st.session_state.get('loaded_report_id') + st.info(f"🗂️ 当前展示为历史报告内容(ID: {loaded_id})") + if st.button("↩️ 返回历史报告列表"): + # 清除已加载的历史报告并返回 + for key in ['sector_strategy_result', 'sector_strategy_result_source', 'loaded_report_id']: + if key in st.session_state: + del st.session_state[key] + st.rerun() + + # 显示引擎回传的保存报告摘要(用于主页面动态更新) + saved_report = result.get("saved_report") + if saved_report: + display_saved_report_summary(saved_report) # PDF导出功能 display_pdf_export_section(result) @@ -525,7 +730,7 @@ def display_visualizations(predictions): title='板块多空信心度对比') fig.update_layout(height=400) - st.plotly_chart(fig, use_container_width=True, key="sector_confidence") + st.plotly_chart(fig, use_container_width=True, config={'responsive': True}, key="sector_confidence") st.markdown("---") @@ -562,7 +767,7 @@ def display_visualizations(predictions): title='板块热度分布图') fig.update_layout(height=400) - st.plotly_chart(fig, use_container_width=True, key="sector_heat") + st.plotly_chart(fig, use_container_width=True, config={'responsive': True}, key="sector_heat") def display_pdf_export_section(result): @@ -575,7 +780,7 @@ def display_pdf_export_section(result): st.write("将分析报告导出为PDF文件,方便保存和分享") with col2: - if st.button("📥 生成PDF报告", type="primary", use_container_width=True): + if st.button("📥 生成PDF报告", type="primary", width='content'): with st.spinner("正在生成PDF报告..."): try: # 生成PDF @@ -600,12 +805,12 @@ def display_pdf_export_section(result): # 如果已经生成了PDF,显示下载按钮 if 'sector_pdf_data' in st.session_state: st.download_button( - label="💾 下载PDF", - data=st.session_state.sector_pdf_data, - file_name=st.session_state.sector_pdf_filename, - mime="application/pdf", - use_container_width=True - ) + label="💾 下载PDF", + data=st.session_state.sector_pdf_data, + file_name=st.session_state.sector_pdf_filename, + mime="application/pdf", + width='content' + ) def display_scheduler_settings(): @@ -653,7 +858,7 @@ def display_scheduler_settings(): with col_a: if not status['running']: - if st.button("▶️ 启动", use_container_width=True, type="primary"): + if st.button("▶️ 启动", width='content', type="primary"): if sector_strategy_scheduler.start(schedule_time_str): st.success(f"✅ 定时任务已启动!每天 {schedule_time_str} 运行") time.sleep(1) @@ -661,7 +866,7 @@ def display_scheduler_settings(): else: st.error("❌ 启动失败") else: - if st.button("⏹️ 停止", use_container_width=True): + if st.button("⏹️ 停止", width='content'): if sector_strategy_scheduler.stop(): st.success("✅ 定时任务已停止") time.sleep(1) @@ -670,13 +875,13 @@ def display_scheduler_settings(): st.error("❌ 停止失败") with col_b: - if st.button("🔄 立即运行", use_container_width=True): + if st.button("🔄 立即运行", width='content'): with st.spinner("正在运行分析..."): sector_strategy_scheduler.manual_run() st.success("✅ 手动分析完成!") with col_c: - if st.button("📧 测试邮件", use_container_width=True): + if st.button("📧 测试邮件", width='content'): test_email_notification() # 邮件配置检查 diff --git a/sector_strategy_ui_clean.py b/sector_strategy_ui_clean.py deleted file mode 100644 index e69de29..0000000 diff --git a/smart_monitor_ui.py b/smart_monitor_ui.py index cae708d..919748d 100644 --- a/smart_monitor_ui.py +++ b/smart_monitor_ui.py @@ -354,7 +354,7 @@ def render_monitor_tasks(): notify_email = st.text_input("通知邮箱(可选)") # 添加任务按钮(表单提交按钮) - submitted = st.form_submit_button("➕ 添加任务", type="primary", use_container_width=True) + submitted = st.form_submit_button("➕ 添加任务", type="primary", width='stretch') if submitted: # 验证必填项(form中直接使用局部变量) @@ -564,7 +564,7 @@ def render_position_management(): "profit_loss_pct": "盈亏%" }, hide_index=True, - use_container_width=True + width='stretch' ) # 单只股票操作 @@ -652,7 +652,7 @@ def render_history(): "profit_loss": "盈亏" }, hide_index=True, - use_container_width=True + width='stretch' ) # 通知记录 @@ -801,7 +801,7 @@ def _render_task_kline_and_decisions(task: Dict, db: SmartMonitorDB, engine): ) # 显示图表 - st.plotly_chart(fig, use_container_width=True) + st.plotly_chart(fig, use_container_width=True, config={'responsive': True}) st.caption(f"📅 数据时间范围:{kline_data['日期'].min()} ~ {kline_data['日期'].max()}") else: diff --git a/stock_analysis.db b/stock_analysis.db index e51ce36..b38972d 100644 Binary files a/stock_analysis.db and b/stock_analysis.db differ diff --git a/stock_data.py b/stock_data.py index 5c628ec..c5b3268 100644 --- a/stock_data.py +++ b/stock_data.py @@ -125,7 +125,7 @@ class StockDataFetcher: print(f"[Akshare] 获取个股详细信息失败: {e}") # 如果akshare失败,尝试从tushare获取 if self.data_source_manager.tushare_available and info['name'] == '未知': - print(f"[Tushare] 尝试获取基本信息(备用数据源)...") + print(f"[Tushare] 尝试获取基本信息(tushare)...") try: ts_code = self.data_source_manager._convert_to_ts_code(symbol) df = self.data_source_manager.tushare_api.daily_basic( @@ -141,65 +141,65 @@ class StockDataFetcher: except Exception as te: print(f"[Tushare] ❌ 获取失败: {te}") - # 方法2: 尝试获取实时价格和涨跌幅(如果网络允许) - try: - # 使用更简单的接口获取实时价格 - real_time_data = ak.stock_zh_a_spot_em() - if real_time_data is not None and not real_time_data.empty: - stock_real_time = real_time_data[real_time_data['代码'] == symbol] - if not stock_real_time.empty: - row = stock_real_time.iloc[0] - info['current_price'] = row.get('最新价', 'N/A') - info['change_percent'] = row.get('涨跌幅', 'N/A') - if info['name'] == '未知': - info['name'] = row.get('名称', '未知') + # 方法2: 尝试获取历史价格和涨跌幅(如果网络允许) + # try: + # # 使用更简单的接口获取实时价格 + # real_time_data = ak.stock_zh_a_spot_em() + # if real_time_data is not None and not real_time_data.empty: + # stock_real_time = real_time_data[real_time_data['代码'] == symbol] + # if not stock_real_time.empty: + # row = stock_real_time.iloc[0] + # info['current_price'] = row.get('最新价', 'N/A') + # info['change_percent'] = row.get('涨跌幅', 'N/A') + # if info['name'] == '未知': + # info['name'] = row.get('名称', '未知') - # 如果实时数据中有市盈率和市净率,优先使用 - if '市盈率-动态' in row and info['pe_ratio'] == 'N/A': - try: - pe_val = row['市盈率-动态'] - if pe_val and pe_val != '-': - pe_val = float(pe_val) - if 0 < pe_val <= 1000: - info['pe_ratio'] = pe_val - except: - pass + # # 如果实时数据中有市盈率和市净率,优先使用 + # if '市盈率-动态' in row and info['pe_ratio'] == 'N/A': + # try: + # pe_val = row['市盈率-动态'] + # if pe_val and pe_val != '-': + # pe_val = float(pe_val) + # if 0 < pe_val <= 1000: + # info['pe_ratio'] = pe_val + # except: + # pass - if '市净率' in row and info['pb_ratio'] == 'N/A': - try: - pb_val = row['市净率'] - if pb_val and pb_val != '-': - pb_val = float(pb_val) - if 0 < pb_val <= 100: - info['pb_ratio'] = pb_val - except: - pass + # if '市净率' in row and info['pb_ratio'] == 'N/A': + # try: + # pb_val = row['市净率'] + # if pb_val and pb_val != '-': + # pb_val = float(pb_val) + # if 0 < pb_val <= 100: + # info['pb_ratio'] = pb_val + # except: + # pass - except Exception as e: - print(f"[Akshare] 获取实时数据失败: {e}") - # 如果实时数据获取失败,尝试使用数据源管理器获取历史数据(支持tushare备用) - try: - print(f"[数据源管理器] 尝试获取最近交易数据...") - hist_data = self.data_source_manager.get_stock_hist_data( - symbol=symbol, - start_date=(datetime.now() - timedelta(days=5)).strftime('%Y%m%d'), - end_date=datetime.now().strftime('%Y%m%d'), - adjust='qfq' - ) - - if hist_data is not None and not hist_data.empty: - # 标准化列名 - if 'close' in hist_data.columns: - latest = hist_data.iloc[-1] - info['current_price'] = latest['close'] - # 计算涨跌幅 - if len(hist_data) > 1: - prev_close = hist_data.iloc[-2]['close'] - change_pct = ((latest['close'] - prev_close) / prev_close) * 100 - info['change_percent'] = round(change_pct, 2) - print(f"[数据源管理器] ✅ 成功获取价格数据") - except Exception as e2: - print(f"获取历史数据也失败: {e2}") + # except Exception as e: + # print(f"[Akshare] 获取实时数据失败: {e}") + # # 如果实时数据获取失败,尝试使用数据源管理器获取历史数据(支持tushare备用) + try: + print(f"[数据源管理器] 尝试获取最近交易数据...") + hist_data = self.data_source_manager.get_stock_hist_data( + symbol=symbol, + start_date=(datetime.now() - timedelta(days=30)).strftime('%Y%m%d'), + end_date=datetime.now().strftime('%Y%m%d'), + adjust='qfq' + ) + + if hist_data is not None and not hist_data.empty: + # 标准化列名 + if 'close' in hist_data.columns: + latest = hist_data.iloc[-1] + info['current_price'] = latest['close'] + # 计算涨跌幅 + if len(hist_data) > 1: + prev_close = hist_data.iloc[-2]['close'] + change_pct = ((latest['close'] - prev_close) / prev_close) * 100 + info['change_percent'] = round(change_pct, 2) + print(f"[数据源管理器] ✅ 成功获取价格数据") + except Exception as e2: + print(f"获取历史数据也失败: {e2}") # 方法3: 使用百度估值数据获取市盈率和市净率 if info['pe_ratio'] == 'N/A': @@ -638,25 +638,40 @@ class StockDataFetcher: # 4. 获取主要财务指标 try: - financial_indicators = ak.stock_financial_analysis_indicator(symbol=symbol) - if financial_indicators is not None and not financial_indicators.empty: - latest_data = financial_indicators.iloc[0] + financial_abstract = ak.stock_financial_abstract(symbol=symbol) + if financial_abstract is not None and not financial_abstract.empty: + # 提取关键财务指标 + key_indicators = [ + '净资产收益率(ROE)', '总资产报酬率(ROA)', '销售毛利率', '销售净利率', + '资产负债率', '流动比率', '速动比率', '存货周转率', '应收账款周转率', + '总资产周转率', '营业收入同比增长', '净利润同比增长' + ] - financial_data["financial_ratios"] = { - "报告期": latest_data.get('报告期', 'N/A'), - "净资产收益率ROE": latest_data.get('净资产收益率', 'N/A'), - "总资产收益率ROA": latest_data.get('总资产收益率', 'N/A'), - "销售毛利率": latest_data.get('销售毛利率', 'N/A'), - "销售净利率": latest_data.get('销售净利率', 'N/A'), - "资产负债率": latest_data.get('资产负债率', 'N/A'), - "流动比率": latest_data.get('流动比率', 'N/A'), - "速动比率": latest_data.get('速动比率', 'N/A'), - "存货周转率": latest_data.get('存货周转率', 'N/A'), - "应收账款周转率": latest_data.get('应收账款周转率', 'N/A'), - "总资产周转率": latest_data.get('总资产周转率', 'N/A'), - "营业收入同比增长": latest_data.get('营业收入同比增长', 'N/A'), - "净利润同比增长": latest_data.get('净利润同比增长', 'N/A'), - } + # 筛选出包含关键指标的行 + indicator_rows = financial_abstract[financial_abstract['指标'].isin(key_indicators)] + + if not indicator_rows.empty: + # 获取最新的报告期数据(第一列日期) + date_columns = [col for col in financial_abstract.columns if col not in ['选项', '指标']] + if date_columns: + latest_date = date_columns[0] # 最新日期列 + + # 构建财务比率字典 + financial_ratios = {"报告期": latest_date} + + # 提取每个指标的最新值 + for _, row in indicator_rows.iterrows(): + indicator_name = row['指标'] + value = row.get(latest_date, 'N/A') + if value is not None and not (isinstance(value, float) and pd.isna(value)): + try: + financial_ratios[indicator_name] = str(value) + except: + financial_ratios[indicator_name] = "N/A" + else: + financial_ratios[indicator_name] = "N/A" + + financial_data["financial_ratios"] = financial_ratios except Exception as e: print(f"获取财务指标失败: {e}") diff --git a/stock_monitor.db b/stock_monitor.db index 8dc8c2e..4f44a9b 100644 Binary files a/stock_monitor.db and b/stock_monitor.db differ diff --git a/test_main_force_batch_db.py b/test_main_force_batch_db.py deleted file mode 100644 index 4f1a26e..0000000 --- a/test_main_force_batch_db.py +++ /dev/null @@ -1,137 +0,0 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- -""" -主力选股批量分析数据库功能测试 -""" - -from main_force_batch_db import batch_db -import json - -def test_database(): - """测试数据库功能""" - - print("=" * 60) - print("主力选股批量分析数据库功能测试") - print("=" * 60) - - # 测试1: 保存批量分析结果 - print("\n📝 测试1: 保存批量分析结果") - test_results = [ - { - "symbol": "000001", - "success": True, - "stock_info": {"股票名称": "平安银行"}, - "final_decision": { - "investment_rating": "买入", - "confidence_level": 85, - "entry_range": "10.0-10.5", - "take_profit": "12.0", - "stop_loss": "9.5" - } - }, - { - "symbol": "600036", - "success": True, - "stock_info": {"股票名称": "招商银行"}, - "final_decision": { - "investment_rating": "持有", - "confidence_level": 75, - "entry_range": "35.0-36.0", - "take_profit": "40.0", - "stop_loss": "33.0" - } - }, - { - "symbol": "600519", - "success": False, - "error": "数据获取失败" - } - ] - - try: - record_id = batch_db.save_batch_analysis( - batch_count=3, - analysis_mode="sequential", - success_count=2, - failed_count=1, - total_time=180.5, - results=test_results - ) - print(f"✅ 保存成功,记录ID: {record_id}") - except Exception as e: - print(f"❌ 保存失败: {str(e)}") - return - - # 测试2: 获取统计信息 - print("\n📊 测试2: 获取统计信息") - try: - stats = batch_db.get_statistics() - print(f"✅ 统计信息:") - print(f" 总记录数: {stats['total_records']}") - print(f" 分析股票总数: {stats['total_stocks_analyzed']}") - print(f" 成功数: {stats['total_success']}") - print(f" 失败数: {stats['total_failed']}") - print(f" 成功率: {stats['success_rate']}%") - print(f" 平均耗时: {stats['average_time']}秒") - except Exception as e: - print(f"❌ 获取统计信息失败: {str(e)}") - - # 测试3: 获取历史记录列表 - print("\n📚 测试3: 获取历史记录列表") - try: - history = batch_db.get_all_history(limit=5) - print(f"✅ 获取到 {len(history)} 条记录") - for idx, record in enumerate(history[:3], 1): - print(f"\n 记录{idx}:") - print(f" - ID: {record['id']}") - print(f" - 时间: {record['analysis_date']}") - print(f" - 数量: {record['batch_count']}") - print(f" - 成功: {record['success_count']}") - print(f" - 失败: {record['failed_count']}") - print(f" - 耗时: {record['total_time']}秒") - except Exception as e: - print(f"❌ 获取历史记录失败: {str(e)}") - - # 测试4: 获取单条记录 - print(f"\n🔍 测试4: 获取单条记录 (ID: {record_id})") - try: - record = batch_db.get_record_by_id(record_id) - if record: - print(f"✅ 获取成功") - print(f" 分析时间: {record['analysis_date']}") - print(f" 结果数量: {len(record['results'])}") - print(f" 成功股票: {[r['symbol'] for r in record['results'] if r.get('success')]}") - print(f" 失败股票: {[r['symbol'] for r in record['results'] if not r.get('success')]}") - else: - print(f"❌ 记录不存在") - except Exception as e: - print(f"❌ 获取记录失败: {str(e)}") - - # 测试5: 删除记录 - print(f"\n🗑️ 测试5: 删除记录 (ID: {record_id})") - confirm = input(" 是否删除测试记录? (y/n): ") - if confirm.lower() == 'y': - try: - success = batch_db.delete_record(record_id) - if success: - print(f"✅ 删除成功") - else: - print(f"❌ 删除失败") - except Exception as e: - print(f"❌ 删除失败: {str(e)}") - else: - print(" 跳过删除") - - print("\n" + "=" * 60) - print("测试完成!") - print("=" * 60) - print("\n💡 提示:") - print(" - 数据库文件: main_force_batch.db") - print(" - 可使用 SQLite 工具查看数据库内容") - print(" - 在Streamlit应用中点击'📚 批量分析历史'查看UI界面") - print("=" * 60) - - -if __name__ == "__main__": - test_database() - diff --git a/test_smart_monitor_data.py b/test_smart_monitor_data.py deleted file mode 100644 index 315c850..0000000 --- a/test_smart_monitor_data.py +++ /dev/null @@ -1,103 +0,0 @@ -#!/usr/bin/env python -# -*- coding: utf-8 -*- -""" -智能盯盘 - 数据获取功能测试脚本 - -用于验证实时行情数据获取是否正常工作 -""" - -import logging -import sys - -# 配置日志 -logging.basicConfig( - level=logging.INFO, - format='%(asctime)s - %(levelname)s - %(message)s' -) - -def test_data_fetcher(): - """测试数据获取功能""" - print("\n" + "="*70) - print("智能盯盘 - 数据获取功能测试") - print("="*70 + "\n") - - try: - from smart_monitor_data import SmartMonitorDataFetcher - - # 创建数据获取器 - fetcher = SmartMonitorDataFetcher() - print("✅ 数据获取器初始化成功\n") - - # 测试股票列表(可以自行修改) - test_stocks = [ - ('600519', '贵州茅台'), - ('000001', '平安银行'), - ('002167', '东方锆业') - ] - - for stock_code, stock_name in test_stocks: - print(f"\n{'─'*70}") - print(f"测试股票: {stock_code} ({stock_name})") - print(f"{'─'*70}") - - # 获取实时行情 - quote = fetcher.get_realtime_quote(stock_code) - - if quote: - print(f"\n✅ 数据获取成功:") - print(f" 📌 股票代码: {quote['code']}") - print(f" 📌 股票名称: {quote['name']}") - print(f" 💰 当前价格: ¥{quote['current_price']:.2f}") - print(f" 📊 涨跌幅: {quote['change_pct']:+.2f}%") - print(f" 💵 涨跌额: ¥{quote['change_amount']:+.2f}") - print(f" 📦 成交量: {quote['volume']:.0f}手") - print(f" 💸 成交额: ¥{quote['amount']/10000:.2f}万") - print(f" 📈 最高: ¥{quote['high']:.2f}") - print(f" 📉 最低: ¥{quote['low']:.2f}") - print(f" 🔓 今开: ¥{quote['open']:.2f}") - print(f" 🔒 昨收: ¥{quote['pre_close']:.2f}") - print(f" 🔄 换手率: {quote['turnover_rate']:.2f}%") - print(f" ⏰ 更新时间: {quote['update_time']}") - print(f" 🌐 数据源: {quote['data_source']}") - - # 验证数据是否有效(不全为0) - if quote['current_price'] > 0: - print(f"\n ✅ 数据有效性检查: 通过") - else: - print(f"\n ⚠️ 数据有效性检查: 价格为0,可能是非交易时间") - else: - print(f"\n❌ 获取 {stock_code} 的数据失败") - return False - - print("\n" + "="*70) - print("✅ 所有测试通过!数据获取功能正常工作") - print("="*70 + "\n") - return True - - except ImportError as e: - print(f"❌ 导入模块失败: {e}") - print("请确保 smart_monitor_data.py 文件存在") - return False - except Exception as e: - print(f"❌ 测试过程中出现错误: {e}") - import traceback - traceback.print_exc() - return False - - -if __name__ == '__main__': - print("\n💡 提示:") - print(" - 此脚本用于测试智能盯盘的数据获取功能") - print(" - 需要网络连接以访问AKShare API") - print(" - 如果所有数据都正常显示,说明修复成功") - print(" - 如果仍然显示0,请检查网络或查看日志\n") - - success = test_data_fetcher() - - if success: - print("🎉 恭喜!数据获取功能测试通过,可以正常使用智能盯盘了!\n") - sys.exit(0) - else: - print("⚠️ 测试失败,请查看上方错误信息并联系技术支持\n") - sys.exit(1) - diff --git a/test_tushare_monitoring.py b/test_tushare_monitoring.py deleted file mode 100644 index 33753c6..0000000 --- a/test_tushare_monitoring.py +++ /dev/null @@ -1,247 +0,0 @@ -""" -测试Tushare数据源是否能满足AI盯盘监控要求 -测试内容: -1. 实时行情数据获取 -2. 技术指标计算 -3. K线图数据获取 -""" - -import logging -import os -from dotenv import load_dotenv - -# 设置日志 -logging.basicConfig( - level=logging.INFO, - format='%(asctime)s - %(name)s - %(levelname)s - %(message)s' -) - -# 加载环境变量 -load_dotenv() - -def test_tushare_token(): - """测试Tushare Token是否配置""" - token = os.getenv('TUSHARE_TOKEN', '') - if token: - print(f"✅ Tushare Token已配置: {token[:10]}...") - return True - else: - print("❌ Tushare Token未配置") - return False - - -def test_realtime_quote(stock_code='000063'): - """测试实时行情获取(Tushare降级)""" - print(f"\n{'='*60}") - print(f"测试1: 实时行情数据 - {stock_code}") - print(f"{'='*60}") - - from smart_monitor_data import SmartMonitorDataFetcher - - fetcher = SmartMonitorDataFetcher() - - # 强制使用Tushare(模拟AKShare失败) - print("正在通过Tushare获取实时行情...") - quote = fetcher._get_realtime_quote_from_tushare(stock_code) - - if quote: - print("✅ 实时行情获取成功!") - print(f" 股票名称: {quote.get('stock_name', 'N/A')}") - print(f" 当前价格: ¥{quote.get('current_price', 0):.2f}") - print(f" 涨跌幅: {quote.get('change_pct', 0):+.2f}%") - print(f" 成交量: {quote.get('volume', 0):,}手") - print(f" 换手率: {quote.get('turnover_rate', 0):.2f}%") - print(f" 数据来源: {quote.get('data_source', 'N/A')}") - return True - else: - print("❌ 实时行情获取失败") - return False - - -def test_technical_indicators(stock_code='000063'): - """测试技术指标计算(Tushare降级)""" - print(f"\n{'='*60}") - print(f"测试2: 技术指标计算 - {stock_code}") - print(f"{'='*60}") - - from smart_monitor_data import SmartMonitorDataFetcher - - fetcher = SmartMonitorDataFetcher() - - # 强制使用Tushare - print("正在通过Tushare获取历史数据并计算技术指标...") - indicators = fetcher._get_technical_indicators_from_tushare(stock_code) - - if indicators: - print("✅ 技术指标计算成功!") - print(f"\n均线系统:") - print(f" MA5: {indicators.get('ma5', 0):.2f}") - print(f" MA20: {indicators.get('ma20', 0):.2f}") - print(f" MA60: {indicators.get('ma60', 0):.2f}") - print(f" 趋势: {indicators.get('trend', 'N/A')}") - - print(f"\nMACD指标:") - print(f" DIF: {indicators.get('macd_dif', 0):.4f}") - print(f" DEA: {indicators.get('macd_dea', 0):.4f}") - print(f" MACD: {indicators.get('macd', 0):.4f}") - - print(f"\nRSI指标:") - print(f" RSI6: {indicators.get('rsi6', 0):.2f}") - print(f" RSI12: {indicators.get('rsi12', 0):.2f}") - print(f" RSI24: {indicators.get('rsi24', 0):.2f}") - - print(f"\nKDJ指标:") - print(f" K: {indicators.get('kdj_k', 0):.2f}") - print(f" D: {indicators.get('kdj_d', 0):.2f}") - print(f" J: {indicators.get('kdj_j', 0):.2f}") - - print(f"\n布林带:") - print(f" 上轨: {indicators.get('boll_upper', 0):.2f}") - print(f" 中轨: {indicators.get('boll_mid', 0):.2f}") - print(f" 下轨: {indicators.get('boll_lower', 0):.2f}") - print(f" 位置: {indicators.get('boll_position', 'N/A')}") - - return True - else: - print("❌ 技术指标计算失败") - return False - - -def test_kline_data(stock_code='000063'): - """测试K线图数据获取(Tushare降级)""" - print(f"\n{'='*60}") - print(f"测试3: K线图数据 - {stock_code}") - print(f"{'='*60}") - - from smart_monitor_kline import SmartMonitorKline - from smart_monitor_data import SmartMonitorDataFetcher - - kline = SmartMonitorKline() - fetcher = SmartMonitorDataFetcher() - - # 使用Tushare获取K线数据 - print("正在通过Tushare获取K线数据(60天)...") - df = kline._get_kline_from_tushare(stock_code, days=60, ts_pro=fetcher.ts_pro) - - if df is not None and not df.empty: - print(f"✅ K线数据获取成功!") - print(f" 数据条数: {len(df)}条") - print(f" 日期范围: {df['日期'].min()} ~ {df['日期'].max()}") - print(f"\n数据列:") - for col in df.columns: - print(f" - {col}") - - print(f"\n最近5条数据预览:") - print(df.tail(5)[['日期', '开盘', '最高', '最低', '收盘', '成交量']].to_string()) - - return True - else: - print("❌ K线数据获取失败") - return False - - -def test_full_monitoring_flow(stock_code='000063'): - """测试完整的监控流程(使用Tushare)""" - print(f"\n{'='*60}") - print(f"测试4: 完整监控流程 - {stock_code}") - print(f"{'='*60}") - - from smart_monitor_data import SmartMonitorDataFetcher - - fetcher = SmartMonitorDataFetcher() - - # 1. 获取实时行情 - print("\n步骤1: 获取实时行情...") - quote = fetcher.get_realtime_quote(stock_code, retry=1) - if not quote: - print(" ❌ 实时行情获取失败") - return False - print(f" ✅ 当前价: ¥{quote.get('current_price', 0):.2f}") - - # 2. 计算技术指标 - print("\n步骤2: 计算技术指标...") - indicators = fetcher.get_technical_indicators(stock_code, retry=1) - if not indicators: - print(" ❌ 技术指标计算失败") - return False - print(f" ✅ MA5: {indicators.get('ma5', 0):.2f}, 趋势: {indicators.get('trend', 'N/A')}") - - # 3. 综合数据 - print("\n步骤3: 获取综合数据...") - comprehensive_data = fetcher.get_comprehensive_data(stock_code) - if not comprehensive_data: - print(" ❌ 综合数据获取失败") - return False - - print(" ✅ 综合数据包含:") - print(f" - 实时行情: {comprehensive_data.get('realtime_quote') is not None}") - print(f" - 技术指标: {comprehensive_data.get('technical_indicators') is not None}") - - print("\n✅ 完整监控流程测试通过!") - print(" Tushare可以满足AI盯盘的监控要求") - return True - - -def main(): - """主测试函数""" - print("="*60) - print("Tushare数据源监控能力测试") - print("="*60) - - # 检查Token - if not test_tushare_token(): - print("\n❌ 请在.env文件中配置TUSHARE_TOKEN") - return - - # 测试股票代码 - test_stock = '000063' # 中兴通讯 - - results = [] - - # 测试1: 实时行情 - results.append(("实时行情获取", test_realtime_quote(test_stock))) - - # 测试2: 技术指标 - results.append(("技术指标计算", test_technical_indicators(test_stock))) - - # 测试3: K线数据 - results.append(("K线图数据", test_kline_data(test_stock))) - - # 测试4: 完整流程 - results.append(("完整监控流程", test_full_monitoring_flow(test_stock))) - - # 汇总结果 - print(f"\n{'='*60}") - print("测试结果汇总") - print(f"{'='*60}") - - for test_name, result in results: - status = "✅ 通过" if result else "❌ 失败" - print(f"{test_name:20} {status}") - - all_passed = all(result for _, result in results) - - if all_passed: - print(f"\n{'='*60}") - print("🎉 所有测试通过!") - print(f"{'='*60}") - print("✅ Tushare完全可以满足AI盯盘的监控要求") - print("✅ 数据源降级策略工作正常") - print("✅ 可以在AKShare IP被封时使用Tushare作为备用") - print("\n建议:") - print("1. 保持Tushare Token配置在.env文件中") - print("2. AKShare重试次数已设置为1次,减少IP封禁风险") - print("3. Tushare 10000积分可以支持日常监控需求") - else: - print(f"\n{'='*60}") - print("⚠️ 部分测试失败") - print(f"{'='*60}") - print("请检查:") - print("1. Tushare Token是否有效") - print("2. Tushare积分是否足够") - print("3. 网络连接是否正常") - - -if __name__ == '__main__': - main() - diff --git a/启动系统.bat b/启动系统.bat new file mode 100644 index 0000000..810824a --- /dev/null +++ b/启动系统.bat @@ -0,0 +1,7 @@ +@echo off +set VENV_PATH=.\TradEnv +set PYTHON_EXE="%VENV_PATH%\python.exe" +set STREAMLIT_MODULE="streamlit.cli" +cd /d ..\AIagentsStock +%PYTHON_EXE% -m streamlit run app.py +pause \ No newline at end of file