@@ -2,9 +2,9 @@
|
||||
业务服务层
|
||||
"""
|
||||
from app.services.kline import KlineService
|
||||
from app.services.analysis import AnalysisService
|
||||
from app.services.backtest import BacktestService
|
||||
from app.services.strategy_compiler import StrategyCompiler
|
||||
from app.services.fast_analysis import FastAnalysisService
|
||||
|
||||
__all__ = ['KlineService', 'AnalysisService', 'BacktestService', 'StrategyCompiler']
|
||||
__all__ = ['KlineService', 'BacktestService', 'StrategyCompiler', 'FastAnalysisService']
|
||||
|
||||
|
||||
@@ -1,148 +0,0 @@
|
||||
# 多智能体分析系统
|
||||
|
||||
基于 TradingAgents 架构优化的多智能体股票分析系统。
|
||||
|
||||
## 架构特点
|
||||
|
||||
### 1. 多智能体协作
|
||||
- **分析师团队**:市场分析师、基本面分析师、新闻分析师、情绪分析师、风险分析师
|
||||
- **研究团队**:看涨研究员、看跌研究员
|
||||
- **交易团队**:交易员、风险分析师(激进/中性/保守)
|
||||
|
||||
### 2. 工作流程
|
||||
```
|
||||
分析师团队分析 → 研究辩论 → 交易员决策 → 风险辩论 → 最终决策
|
||||
```
|
||||
|
||||
### 3. 记忆系统
|
||||
- 使用 SQLite 存储历史决策
|
||||
- 基于文本相似度检索历史经验
|
||||
- 支持从交易结果中学习
|
||||
|
||||
### 4. 工具调用
|
||||
- 智能体可以主动获取数据
|
||||
- 支持多数据源(yfinance, finnhub, ccxt, 腾讯接口)
|
||||
|
||||
## 使用方法
|
||||
|
||||
### 基本使用
|
||||
|
||||
```python
|
||||
from app.services.analysis import AnalysisService
|
||||
|
||||
# 使用多智能体架构(默认)
|
||||
service = AnalysisService(use_multi_agent=True)
|
||||
result = service.analyze("USStock", "AAPL", "zh-CN")
|
||||
|
||||
# 使用传统架构(向后兼容)
|
||||
service = AnalysisService(use_multi_agent=False)
|
||||
result = service.analyze("USStock", "AAPL", "zh-CN")
|
||||
```
|
||||
|
||||
### 环境变量配置
|
||||
|
||||
在 `.env` 文件中设置:
|
||||
|
||||
```bash
|
||||
# 是否启用多智能体架构(默认:True)
|
||||
USE_MULTI_AGENT=True
|
||||
|
||||
# 最大辩论轮数(默认:2)
|
||||
MAX_DEBATE_ROUNDS=2
|
||||
```
|
||||
|
||||
### 反思学习
|
||||
|
||||
```python
|
||||
from app.services.analysis import reflect_analysis
|
||||
|
||||
# 从交易结果中学习
|
||||
reflect_analysis(
|
||||
market="USStock",
|
||||
symbol="AAPL",
|
||||
decision="BUY",
|
||||
returns=1000.0, # 收益
|
||||
result="交易成功,收益 10%"
|
||||
)
|
||||
```
|
||||
|
||||
## 智能体说明
|
||||
|
||||
### 分析师智能体
|
||||
|
||||
- **MarketAnalyst**: 技术分析,计算技术指标
|
||||
- **FundamentalAnalyst**: 基本面分析,财务数据
|
||||
- **NewsAnalyst**: 新闻事件分析
|
||||
- **SentimentAnalyst**: 市场情绪分析
|
||||
- **RiskAnalyst**: 风险评估
|
||||
|
||||
### 研究员智能体
|
||||
|
||||
- **BullResearcher**: 构建看涨论据
|
||||
- **BearResearcher**: 构建看跌论据
|
||||
|
||||
### 交易员智能体
|
||||
|
||||
- **TraderAgent**: 综合所有分析,做出交易决策
|
||||
|
||||
### 风险分析师智能体
|
||||
|
||||
- **RiskyAnalyst**: 激进风险分析
|
||||
- **NeutralAnalyst**: 中性风险分析
|
||||
- **SafeAnalyst**: 保守风险分析
|
||||
|
||||
## 返回结果格式
|
||||
|
||||
多智能体模式返回的完整结果包括:
|
||||
|
||||
```json
|
||||
{
|
||||
"overview": {
|
||||
"overallScore": 75,
|
||||
"recommendation": "BUY",
|
||||
"confidence": 82,
|
||||
"dimensionScores": {...},
|
||||
"report": "..."
|
||||
},
|
||||
"fundamental": {...},
|
||||
"technical": {...},
|
||||
"news": {...},
|
||||
"sentiment": {...},
|
||||
"risk": {...},
|
||||
"debate": {
|
||||
"bull": {...},
|
||||
"bear": {...},
|
||||
"research_decision": "..."
|
||||
},
|
||||
"trader_decision": {
|
||||
"decision": "BUY",
|
||||
"confidence": 85,
|
||||
"trading_plan": {...}
|
||||
},
|
||||
"risk_debate": {
|
||||
"risky": {...},
|
||||
"neutral": {...},
|
||||
"safe": {...}
|
||||
},
|
||||
"final_decision": {
|
||||
"decision": "BUY",
|
||||
"confidence": 85,
|
||||
"reasoning": "..."
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## 优势
|
||||
|
||||
1. **多角度分析**:多个智能体从不同角度分析,减少盲点
|
||||
2. **辩论机制**:看涨/看跌辩论,发现潜在问题
|
||||
3. **风险控制**:多维度风险分析,提高决策质量
|
||||
4. **持续学习**:记忆系统支持从历史经验中学习
|
||||
5. **向后兼容**:可以切换到传统模式
|
||||
|
||||
## 注意事项
|
||||
|
||||
1. 多智能体模式会产生更多的 API 调用,成本较高
|
||||
2. 分析时间会比传统模式长
|
||||
3. 记忆系统需要 SQLite 数据库支持
|
||||
4. 建议在生产环境中根据需求选择合适的模式
|
||||
@@ -1,32 +0,0 @@
|
||||
"""
|
||||
多智能体分析系统
|
||||
基于 TradingAgents 架构优化
|
||||
"""
|
||||
from .base_agent import BaseAgent
|
||||
from .analyst_agents import (
|
||||
MarketAnalyst,
|
||||
FundamentalAnalyst,
|
||||
NewsAnalyst,
|
||||
SentimentAnalyst,
|
||||
RiskAnalyst
|
||||
)
|
||||
from .researcher_agents import BullResearcher, BearResearcher
|
||||
from .trader_agent import TraderAgent
|
||||
from .risk_agents import RiskyAnalyst, NeutralAnalyst, SafeAnalyst
|
||||
from .memory import AgentMemory
|
||||
|
||||
__all__ = [
|
||||
'BaseAgent',
|
||||
'MarketAnalyst',
|
||||
'FundamentalAnalyst',
|
||||
'NewsAnalyst',
|
||||
'SentimentAnalyst',
|
||||
'RiskAnalyst',
|
||||
'BullResearcher',
|
||||
'BearResearcher',
|
||||
'TraderAgent',
|
||||
'RiskyAnalyst',
|
||||
'NeutralAnalyst',
|
||||
'SafeAnalyst',
|
||||
'AgentMemory',
|
||||
]
|
||||
@@ -1,440 +0,0 @@
|
||||
"""
|
||||
Analyst agents.
|
||||
Includes: market/technical, fundamental, news, sentiment, risk analysts.
|
||||
"""
|
||||
import json
|
||||
from typing import Dict, Any
|
||||
from .base_agent import BaseAgent
|
||||
from .tools import AgentTools
|
||||
from app.services.llm import LLMService
|
||||
|
||||
logger = __import__('app.utils.logger', fromlist=['get_logger']).get_logger(__name__)
|
||||
|
||||
|
||||
class MarketAnalyst(BaseAgent):
|
||||
"""Market / technical analyst."""
|
||||
|
||||
def __init__(self, memory=None):
|
||||
super().__init__("MarketAnalyst", memory)
|
||||
self.tools = AgentTools()
|
||||
self.llm_service = LLMService()
|
||||
|
||||
def analyze(self, context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Run technical analysis."""
|
||||
market = context.get('market')
|
||||
symbol = context.get('symbol')
|
||||
language = context.get('language', 'zh-CN')
|
||||
model = context.get('model')
|
||||
base_data = context.get('base_data', {})
|
||||
|
||||
# Kline + current price
|
||||
kline_data = base_data.get('kline_data') or self.tools.get_stock_data(market, symbol, days=30)
|
||||
current_price = base_data.get('current_price') or self.tools.get_current_price(market, symbol)
|
||||
|
||||
# Technical indicators
|
||||
indicators = {}
|
||||
if kline_data:
|
||||
indicators = self.tools.calculate_technical_indicators(kline_data)
|
||||
|
||||
# Prompts
|
||||
lang_instruction = self._get_language_instruction(language)
|
||||
system_prompt = f"""You are a professional technical analyst. Please analyze the technical trends of the stock or cryptocurrency, including:
|
||||
{lang_instruction}
|
||||
1. Technical Indicator Signals (MACD signals, RSI range, trend strength, important MA directions)
|
||||
2. Technical Score (0-100). Be objective. Do not default to 75.
|
||||
- 0-40: Bearish/Weak
|
||||
- 41-60: Neutral
|
||||
- 61-100: Bullish/Strong
|
||||
3. Technical Analysis Report (about 300 words)
|
||||
|
||||
Please strictly return the result in JSON format as follows:
|
||||
{{
|
||||
"score": 75,
|
||||
"indicators": {{
|
||||
"MACD": "Golden Cross/Death Cross or Flat",
|
||||
"RSI(14)": "75 (Overbought)",
|
||||
"MA20": "Upward/Downward/Flat",
|
||||
"Support/Resistance": "Support: 150.00, Resistance: 165.50"
|
||||
}},
|
||||
"report": "Technical analysis report content..."
|
||||
}}"""
|
||||
|
||||
user_prompt = f"""Please perform technical analysis for {symbol} in {market} market.
|
||||
|
||||
**Current Price:**
|
||||
{json.dumps(current_price, ensure_ascii=False, indent=2) if current_price else 'No Data'}
|
||||
|
||||
**Kline Data (Last 30 days):**
|
||||
{json.dumps(kline_data[-10:] if kline_data else [], ensure_ascii=False, indent=2) if kline_data else 'No Data'}
|
||||
|
||||
**Calculated Technical Indicators:**
|
||||
{json.dumps(indicators, ensure_ascii=False, indent=2) if indicators else 'No Data'}
|
||||
|
||||
Please analyze the short-term and medium-term trends based on the above Kline data and price movements."""
|
||||
|
||||
# LLM call
|
||||
result = self.llm_service.safe_call_llm(
|
||||
system_prompt,
|
||||
user_prompt,
|
||||
{"score": 50, "indicators": {}, "report": "Failed to parse technical analysis"},
|
||||
model=model
|
||||
)
|
||||
|
||||
return {
|
||||
"type": "technical",
|
||||
"data": result,
|
||||
"indicators": indicators
|
||||
}
|
||||
|
||||
def _get_language_instruction(self, language: str) -> str:
|
||||
"""Return an English language instruction string for the LLM prompt."""
|
||||
language_map = {
|
||||
'zh-CN': 'Answer in Simplified Chinese.',
|
||||
'zh-TW': 'Answer in Traditional Chinese.',
|
||||
'en-US': 'Answer in English.',
|
||||
'ja-JP': 'Answer in Japanese.',
|
||||
'ko-KR': 'Answer in Korean.',
|
||||
'vi-VN': 'Answer in Vietnamese.',
|
||||
'th-TH': 'Answer in Thai.',
|
||||
'ar-SA': 'Answer in Arabic.',
|
||||
'fr-FR': 'Answer in French.',
|
||||
'de-DE': 'Answer in German.'
|
||||
}
|
||||
return language_map.get(language, 'Answer in English.')
|
||||
|
||||
|
||||
class FundamentalAnalyst(BaseAgent):
|
||||
"""Fundamental analyst."""
|
||||
|
||||
def __init__(self, memory=None):
|
||||
super().__init__("FundamentalAnalyst", memory)
|
||||
self.tools = AgentTools()
|
||||
self.llm_service = LLMService()
|
||||
|
||||
def analyze(self, context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Run fundamental analysis."""
|
||||
market = context.get('market')
|
||||
symbol = context.get('symbol')
|
||||
language = context.get('language', 'zh-CN')
|
||||
model = context.get('model')
|
||||
base_data = context.get('base_data', {})
|
||||
|
||||
# Fundamental data
|
||||
fundamental_data = base_data.get('fundamental_data') or self.tools.get_fundamental_data(market, symbol)
|
||||
company_data = base_data.get('company_data') or self.tools.get_company_data(market, symbol, language=language)
|
||||
|
||||
# Memory
|
||||
situation = f"{market}:{symbol} fundamental analysis"
|
||||
memory_meta = {
|
||||
"market": market,
|
||||
"symbol": symbol,
|
||||
"timeframe": context.get("timeframe"),
|
||||
"features": context.get("memory_features") or {},
|
||||
}
|
||||
memories = self.get_memories(situation, n_matches=None, metadata=memory_meta)
|
||||
memory_prompt = self.format_memories_for_prompt(memories)
|
||||
|
||||
lang_instruction = self._get_language_instruction(language)
|
||||
system_prompt = f"""You are a fundamental analyst. Please analyze the financial status and industry position of the stock, including:
|
||||
{lang_instruction}
|
||||
1. Financial Indicators (Select key P/E, P/B, ROE, Revenue Growth)
|
||||
2. Fundamental Score (0-100). Be objective. Do not default to 80.
|
||||
- 0-40: Poor/Overvalued
|
||||
- 41-60: Fair/Neutral
|
||||
- 61-100: Good/Undervalued
|
||||
3. Fundamental Analysis Report (about 300 words)
|
||||
|
||||
{memory_prompt}
|
||||
|
||||
Please strictly return the result in JSON format as follows:
|
||||
{{
|
||||
"score": 80,
|
||||
"financials": {{
|
||||
"P/E": "25.3",
|
||||
"P/B": "4.2",
|
||||
"ROE": "18.5%",
|
||||
"Market Cap": "1200.5 B"
|
||||
}},
|
||||
"report": "Fundamental analysis report content..."
|
||||
}}"""
|
||||
|
||||
user_prompt = f"""Please perform fundamental analysis for {symbol} in {market} market.
|
||||
|
||||
**Basic Company Info:**
|
||||
{json.dumps(company_data, ensure_ascii=False, indent=2) if company_data else 'No Data'}
|
||||
|
||||
**Raw Fundamental Indicators:**
|
||||
{json.dumps(fundamental_data, ensure_ascii=False, indent=2) if fundamental_data else 'No Data'}
|
||||
|
||||
Please analyze based on the above data."""
|
||||
|
||||
result = self.llm_service.safe_call_llm(
|
||||
system_prompt,
|
||||
user_prompt,
|
||||
{"score": 50, "financials": {}, "report": "Failed to parse fundamental analysis"},
|
||||
model=model
|
||||
)
|
||||
|
||||
return {
|
||||
"type": "fundamental",
|
||||
"data": result
|
||||
}
|
||||
|
||||
def _get_language_instruction(self, language: str) -> str:
|
||||
language_map = {
|
||||
'zh-CN': 'Answer in Simplified Chinese.',
|
||||
'zh-TW': 'Answer in Traditional Chinese.',
|
||||
'en-US': 'Answer in English.',
|
||||
'ja-JP': 'Answer in Japanese.',
|
||||
'ko-KR': 'Answer in Korean.',
|
||||
'vi-VN': 'Answer in Vietnamese.',
|
||||
'th-TH': 'Answer in Thai.',
|
||||
'ar-SA': 'Answer in Arabic.',
|
||||
'fr-FR': 'Answer in French.',
|
||||
'de-DE': 'Answer in German.'
|
||||
}
|
||||
return language_map.get(language, 'Answer in English.')
|
||||
|
||||
|
||||
class NewsAnalyst(BaseAgent):
|
||||
"""News analyst."""
|
||||
|
||||
def __init__(self, memory=None):
|
||||
super().__init__("NewsAnalyst", memory)
|
||||
self.llm_service = LLMService()
|
||||
|
||||
def analyze(self, context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Run news analysis."""
|
||||
market = context.get('market')
|
||||
symbol = context.get('symbol')
|
||||
language = context.get('language', 'zh-CN')
|
||||
model = context.get('model')
|
||||
base_data = context.get('base_data', {})
|
||||
|
||||
company_data = base_data.get('company_data', {})
|
||||
news_data = base_data.get('news_data', [])
|
||||
|
||||
lang_instruction = self._get_language_instruction(language)
|
||||
system_prompt = f"""You are a professional news intelligence analyst. Your task is to extract key intelligence that has a major impact on stock prices from massive market information.
|
||||
|
||||
Analysis Requirements:
|
||||
{lang_instruction}
|
||||
1. **Filter Noise**: News data may contain internet results from search engines. Please carefully discriminate and ignore ads, duplicate content, or irrelevant noise. Prioritize authoritative financial media and official announcements.
|
||||
2. **Timeliness**: Focus on breaking news within the last 48 hours. For old news, lower its weight unless there are new developments.
|
||||
3. **Deep Interpretation**: Do not just repeat news titles. Analyze the logic behind the event and its specific impact on company fundamentals or market sentiment (e.g., Earnings Beat -> Profit Improvement -> Valuation Increase).
|
||||
4. **Scoring**: News Score (0-100).
|
||||
- 0-40: Major Negative (e.g., financial fraud, regulatory crackdown, core business damage)
|
||||
- 41-59: Neutral or minor impact
|
||||
- 60-100: Positive (e.g., strong earnings, major partnership, policy support)
|
||||
- Higher scores indicate more positive news.
|
||||
|
||||
Please strictly return the result in JSON format as follows:
|
||||
{{
|
||||
"score": 70,
|
||||
"events": [
|
||||
{{
|
||||
"title": "Event Title",
|
||||
"impact": "Positive/Negative/Neutral",
|
||||
"summary": "Event summary and deep impact analysis...",
|
||||
"date": "2023-10-27"
|
||||
}}
|
||||
],
|
||||
"report": "Comprehensive news analysis report, including overall judgment on recent market public opinion..."
|
||||
}}"""
|
||||
|
||||
user_prompt = f"""Please perform in-depth news intelligence analysis for {symbol} in {market} market.
|
||||
|
||||
**Company Background:**
|
||||
{json.dumps(company_data, ensure_ascii=False, indent=2) if company_data else 'No Data'}
|
||||
|
||||
**Latest Intelligence Sources (Contains professional financial news and web search results, please discriminate):**
|
||||
{json.dumps(news_data, ensure_ascii=False, indent=2) if news_data else 'No directly related news'}
|
||||
|
||||
Please analyze based on the above intelligence. If the provided "Latest Intelligence Sources" contain no substantive content or only irrelevant noise, please explicitly state "No valid news available" and deduce logically based on the industry trends and macro market environment of the stock."""
|
||||
|
||||
result = self.llm_service.safe_call_llm(
|
||||
system_prompt,
|
||||
user_prompt,
|
||||
{"score": 50, "events": [], "report": "Failed to parse news analysis"},
|
||||
model=model
|
||||
)
|
||||
|
||||
return {
|
||||
"type": "news",
|
||||
"data": result
|
||||
}
|
||||
|
||||
def _get_language_instruction(self, language: str) -> str:
|
||||
language_map = {
|
||||
'zh-CN': 'Answer in Simplified Chinese.',
|
||||
'zh-TW': 'Answer in Traditional Chinese.',
|
||||
'en-US': 'Answer in English.',
|
||||
'ja-JP': 'Answer in Japanese.',
|
||||
'ko-KR': 'Answer in Korean.',
|
||||
'vi-VN': 'Answer in Vietnamese.',
|
||||
'th-TH': 'Answer in Thai.',
|
||||
'ar-SA': 'Answer in Arabic.',
|
||||
'fr-FR': 'Answer in French.',
|
||||
'de-DE': 'Answer in German.'
|
||||
}
|
||||
return language_map.get(language, 'Answer in English.')
|
||||
|
||||
|
||||
class SentimentAnalyst(BaseAgent):
|
||||
"""Sentiment analyst."""
|
||||
|
||||
def __init__(self, memory=None):
|
||||
super().__init__("SentimentAnalyst", memory)
|
||||
self.llm_service = LLMService()
|
||||
|
||||
def analyze(self, context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Run sentiment analysis."""
|
||||
market = context.get('market')
|
||||
symbol = context.get('symbol')
|
||||
language = context.get('language', 'zh-CN')
|
||||
model = context.get('model')
|
||||
base_data = context.get('base_data', {})
|
||||
|
||||
current_price = base_data.get('current_price', {})
|
||||
kline_data = base_data.get('kline_data', [])
|
||||
|
||||
lang_instruction = self._get_language_instruction(language)
|
||||
system_prompt = f"""You are a market sentiment analyst. Please analyze the market sentiment and popularity of the stock, including:
|
||||
{lang_instruction}
|
||||
1. Sentiment Heat Indicators (e.g., analyst ratings, social media discussion volume, put/call ratio)
|
||||
2. Sentiment Score (0-100, higher means more optimistic). Be objective. Do not default to 65.
|
||||
- 0-40: Bearish/Fearful
|
||||
- 41-60: Neutral
|
||||
- 61-100: Bullish/Greedy
|
||||
3. Sentiment Analysis Report (about 300 words)
|
||||
|
||||
Please strictly return the result in JSON format as follows:
|
||||
{{
|
||||
"score": 65,
|
||||
"scores": {{
|
||||
"Analyst Rating": 90,
|
||||
"Social Media Heat": 85,
|
||||
"Market Sentiment Index": 70
|
||||
}},
|
||||
"report": "Sentiment analysis report content..."
|
||||
}}
|
||||
Note: All values in the 'scores' dictionary must be integers between 0-100 (pure numbers), representing optimism or heat, without any text or percentage signs."""
|
||||
|
||||
user_prompt = f"""Please perform sentiment analysis for {symbol} in {market} market.
|
||||
Based on the current Kline trends and price fluctuations, combined with your existing knowledge, evaluate whether the market sentiment for this stock is bullish, bearish, or neutral.
|
||||
|
||||
**Current Price:**
|
||||
{json.dumps(current_price, ensure_ascii=False, indent=2) if current_price else 'No Data'}
|
||||
|
||||
**Kline Data (Recent Trends):**
|
||||
{json.dumps(kline_data[-5:] if kline_data else [], ensure_ascii=False, indent=2) if kline_data else 'No Data'}
|
||||
|
||||
Please evaluate market sentiment."""
|
||||
|
||||
result = self.llm_service.safe_call_llm(
|
||||
system_prompt,
|
||||
user_prompt,
|
||||
{"score": 50, "scores": {"Analyst Rating": 50, "Social Media Heat": 50, "Market Sentiment Index": 50}, "report": "Failed to parse sentiment analysis"},
|
||||
model=model
|
||||
)
|
||||
|
||||
return {
|
||||
"type": "sentiment",
|
||||
"data": result
|
||||
}
|
||||
|
||||
def _get_language_instruction(self, language: str) -> str:
|
||||
language_map = {
|
||||
'zh-CN': 'Answer in Simplified Chinese.',
|
||||
'zh-TW': 'Answer in Traditional Chinese.',
|
||||
'en-US': 'Answer in English.',
|
||||
'ja-JP': 'Answer in Japanese.',
|
||||
'ko-KR': 'Answer in Korean.',
|
||||
'vi-VN': 'Answer in Vietnamese.',
|
||||
'th-TH': 'Answer in Thai.',
|
||||
'ar-SA': 'Answer in Arabic.',
|
||||
'fr-FR': 'Answer in French.',
|
||||
'de-DE': 'Answer in German.'
|
||||
}
|
||||
return language_map.get(language, 'Answer in English.')
|
||||
|
||||
|
||||
class RiskAnalyst(BaseAgent):
|
||||
"""Risk analyst."""
|
||||
|
||||
def __init__(self, memory=None):
|
||||
super().__init__("RiskAnalyst", memory)
|
||||
self.llm_service = LLMService()
|
||||
|
||||
def analyze(self, context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Run risk assessment."""
|
||||
market = context.get('market')
|
||||
symbol = context.get('symbol')
|
||||
language = context.get('language', 'zh-CN')
|
||||
model = context.get('model')
|
||||
base_data = context.get('base_data', {})
|
||||
|
||||
current_price = base_data.get('current_price', {})
|
||||
kline_data = base_data.get('kline_data', [])
|
||||
fundamental_data = base_data.get('fundamental_data', {})
|
||||
|
||||
lang_instruction = self._get_language_instruction(language)
|
||||
system_prompt = f"""You are a risk management analyst. Please evaluate the investment risks of the stock, including:
|
||||
{lang_instruction}
|
||||
1. Risk Indicators (Volatility, Liquidity, Concentration Risk, Systemic Risk Exposure)
|
||||
2. Risk Score (0-100, higher score means lower risk/safer). Be objective. Do not default to 60.
|
||||
- 0-40: High Risk (Dangerous)
|
||||
- 41-60: Moderate Risk
|
||||
- 61-100: Low Risk (Safe)
|
||||
3. Risk Assessment Report (about 300 words)
|
||||
|
||||
Please strictly return the result in JSON format as follows:
|
||||
{{
|
||||
"score": 60,
|
||||
"metrics": {{
|
||||
"Volatility (Beta)": "1.2 (Higher than market)",
|
||||
"Liquidity": "Good",
|
||||
"Concentration Risk": "Low (Diversified business)"
|
||||
}},
|
||||
"report": "Risk assessment report content..."
|
||||
}}"""
|
||||
|
||||
user_prompt = f"""Please perform risk assessment for {symbol} in {market} market.
|
||||
|
||||
**Current Price:**
|
||||
{json.dumps(current_price, ensure_ascii=False, indent=2) if current_price else 'No Data'}
|
||||
|
||||
**Kline Data (Volatility Analysis):**
|
||||
{json.dumps(kline_data, ensure_ascii=False, indent=2) if kline_data else 'No Data'}
|
||||
|
||||
**Fundamental Data (Debt Risk):**
|
||||
{json.dumps(fundamental_data, ensure_ascii=False, indent=2) if fundamental_data else 'No Data'}
|
||||
|
||||
Please evaluate investment risks based on price volatility, fundamental data, etc."""
|
||||
|
||||
result = self.llm_service.safe_call_llm(
|
||||
system_prompt,
|
||||
user_prompt,
|
||||
{"score": 50, "metrics": {}, "report": "Failed to parse risk assessment"},
|
||||
model=model
|
||||
)
|
||||
|
||||
return {
|
||||
"type": "risk",
|
||||
"data": result
|
||||
}
|
||||
|
||||
def _get_language_instruction(self, language: str) -> str:
|
||||
language_map = {
|
||||
'zh-CN': 'Answer in Simplified Chinese.',
|
||||
'zh-TW': 'Answer in Traditional Chinese.',
|
||||
'en-US': 'Answer in English.',
|
||||
'ja-JP': 'Answer in Japanese.',
|
||||
'ko-KR': 'Answer in Korean.',
|
||||
'vi-VN': 'Answer in Vietnamese.',
|
||||
'th-TH': 'Answer in Thai.',
|
||||
'ar-SA': 'Answer in Arabic.',
|
||||
'fr-FR': 'Answer in French.',
|
||||
'de-DE': 'Answer in German.'
|
||||
}
|
||||
return language_map.get(language, 'Answer in English.')
|
||||
@@ -1,96 +0,0 @@
|
||||
"""
|
||||
智能体基类
|
||||
"""
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Dict, Any, Optional, List
|
||||
import os
|
||||
from app.utils.logger import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class BaseAgent(ABC):
|
||||
"""智能体基类,所有分析智能体都继承此类"""
|
||||
|
||||
def __init__(self, name: str, memory: Optional[Any] = None):
|
||||
"""
|
||||
初始化智能体
|
||||
|
||||
Args:
|
||||
name: 智能体名称
|
||||
memory: 记忆系统实例(可选)
|
||||
"""
|
||||
self.name = name
|
||||
self.memory = memory
|
||||
self.logger = get_logger(f"{__name__}.{name}")
|
||||
|
||||
@abstractmethod
|
||||
def analyze(self, context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
执行分析任务
|
||||
|
||||
Args:
|
||||
context: 分析上下文,包含市场、代码、基础数据等
|
||||
|
||||
Returns:
|
||||
分析结果字典
|
||||
"""
|
||||
pass
|
||||
|
||||
def get_memories(self, situation: str, n_matches: Optional[int] = None, metadata: Optional[Dict[str, Any]] = None) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
从记忆中检索相似情况
|
||||
|
||||
Args:
|
||||
situation: 当前情况描述
|
||||
n_matches: 返回的匹配数量
|
||||
|
||||
Returns:
|
||||
匹配的历史记录列表
|
||||
"""
|
||||
if n_matches is None:
|
||||
try:
|
||||
n_matches = int(os.getenv("AGENT_MEMORY_TOP_K", "5") or 5)
|
||||
except Exception:
|
||||
n_matches = 5
|
||||
if self.memory:
|
||||
# New memory API supports metadata; older implementations will ignore extra args.
|
||||
try:
|
||||
return self.memory.get_memories(situation, n_matches=n_matches, metadata=metadata)
|
||||
except TypeError:
|
||||
return self.memory.get_memories(situation, n_matches=n_matches)
|
||||
return []
|
||||
|
||||
def format_memories_for_prompt(self, memories: List[Dict[str, Any]]) -> str:
|
||||
"""
|
||||
格式化记忆为提示词
|
||||
|
||||
Args:
|
||||
memories: 记忆列表
|
||||
|
||||
Returns:
|
||||
格式化的字符串
|
||||
"""
|
||||
if not memories:
|
||||
return "No prior experience available."
|
||||
|
||||
lines = ["Prior experience (most relevant first):"]
|
||||
for i, mem in enumerate(memories, 1):
|
||||
rec = mem.get("recommendation") or "N/A"
|
||||
res = mem.get("result") or ""
|
||||
ret = mem.get("returns")
|
||||
created_at = mem.get("created_at")
|
||||
# Keep created_at as-is (SQLite string), but include it for traceability.
|
||||
meta_bits = []
|
||||
if created_at:
|
||||
meta_bits.append(f"at {created_at}")
|
||||
if ret is not None and ret != "":
|
||||
meta_bits.append(f"returns={ret}%")
|
||||
meta_s = f" ({', '.join(meta_bits)})" if meta_bits else ""
|
||||
|
||||
if res:
|
||||
lines.append(f"{i}. {rec}{meta_s}\n outcome: {res}")
|
||||
else:
|
||||
lines.append(f"{i}. {rec}{meta_s}")
|
||||
|
||||
return "\n".join(lines)
|
||||
@@ -1,745 +0,0 @@
|
||||
"""
|
||||
Multi-agent coordinator.
|
||||
Orchestrates agents and the overall analysis workflow.
|
||||
"""
|
||||
from typing import Dict, Any, List, Optional
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from .analyst_agents import (
|
||||
MarketAnalyst, FundamentalAnalyst, NewsAnalyst,
|
||||
SentimentAnalyst, RiskAnalyst
|
||||
)
|
||||
from .researcher_agents import BullResearcher, BearResearcher
|
||||
from .trader_agent import TraderAgent
|
||||
from .risk_agents import RiskyAnalyst, NeutralAnalyst, SafeAnalyst
|
||||
from .memory import AgentMemory
|
||||
from .reflection import ReflectionService
|
||||
from app.utils.logger import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class AgentCoordinator:
|
||||
"""Multi-agent coordinator."""
|
||||
|
||||
@staticmethod
|
||||
def _is_zh(language: str) -> bool:
|
||||
return str(language or "").lower().startswith("zh")
|
||||
|
||||
@classmethod
|
||||
def _t(cls, language: str, en: str, zh: str) -> str:
|
||||
"""Pick a localized string for user-facing fields (not logs)."""
|
||||
return zh if cls._is_zh(language) else en
|
||||
|
||||
def __init__(self, enable_memory: bool = True, max_debate_rounds: int = 2):
|
||||
"""
|
||||
Args:
|
||||
enable_memory: Enable memory/reflection
|
||||
max_debate_rounds: Max debate rounds
|
||||
"""
|
||||
self.enable_memory = enable_memory
|
||||
self.max_debate_rounds = max_debate_rounds
|
||||
|
||||
# Reflection service
|
||||
self.reflection_service = ReflectionService() if enable_memory else None
|
||||
|
||||
# Memory stores
|
||||
if enable_memory:
|
||||
self.memories = {
|
||||
'market': AgentMemory('market_analyst'),
|
||||
'fundamental': AgentMemory('fundamental_analyst'),
|
||||
'news': AgentMemory('news_analyst'),
|
||||
'sentiment': AgentMemory('sentiment_analyst'),
|
||||
'risk': AgentMemory('risk_analyst'),
|
||||
'bull': AgentMemory('bull_researcher'),
|
||||
'bear': AgentMemory('bear_researcher'),
|
||||
'trader': AgentMemory('trader_agent'),
|
||||
}
|
||||
else:
|
||||
self.memories = {}
|
||||
|
||||
# Initialize agents
|
||||
self._init_agents()
|
||||
|
||||
def _init_agents(self):
|
||||
"""Initialize all agents."""
|
||||
# Analyst agents
|
||||
self.market_analyst = MarketAnalyst(
|
||||
memory=self.memories.get('market')
|
||||
)
|
||||
self.fundamental_analyst = FundamentalAnalyst(
|
||||
memory=self.memories.get('fundamental')
|
||||
)
|
||||
self.news_analyst = NewsAnalyst(
|
||||
memory=self.memories.get('news')
|
||||
)
|
||||
self.sentiment_analyst = SentimentAnalyst(
|
||||
memory=self.memories.get('sentiment')
|
||||
)
|
||||
self.risk_analyst = RiskAnalyst(
|
||||
memory=self.memories.get('risk')
|
||||
)
|
||||
|
||||
# Researcher agents
|
||||
self.bull_researcher = BullResearcher(
|
||||
memory=self.memories.get('bull')
|
||||
)
|
||||
self.bear_researcher = BearResearcher(
|
||||
memory=self.memories.get('bear')
|
||||
)
|
||||
|
||||
# Trader agent
|
||||
self.trader_agent = TraderAgent(
|
||||
memory=self.memories.get('trader')
|
||||
)
|
||||
|
||||
# Risk debate agents
|
||||
self.risky_analyst = RiskyAnalyst()
|
||||
self.neutral_analyst = NeutralAnalyst()
|
||||
self.safe_analyst = SafeAnalyst()
|
||||
|
||||
def run_analysis_stream(self, market: str, symbol: str, language: str = 'zh-CN', model: str = None, timeframe: str = "1D", on_progress=None):
|
||||
"""
|
||||
Run the full multi-agent analysis workflow with progress callbacks.
|
||||
|
||||
Args:
|
||||
on_progress: Callback function that receives (agent_name: str, status: str, result: Optional[dict])
|
||||
status can be: 'started', 'completed', 'error'
|
||||
|
||||
Yields:
|
||||
Progress events as dicts: { 'agent': str, 'status': str, 'result': dict or None }
|
||||
"""
|
||||
logger.info(f"Multi-agent stream analysis start: {market}:{symbol}, model={model}, language={language}")
|
||||
|
||||
def emit_progress(agent: str, status: str, result: dict = None):
|
||||
"""Emit progress event."""
|
||||
if on_progress:
|
||||
on_progress(agent, status, result)
|
||||
|
||||
# Build base context
|
||||
from .tools import AgentTools
|
||||
tools = AgentTools()
|
||||
|
||||
emit_progress('data_collection', 'started', None)
|
||||
|
||||
# 1) Base data
|
||||
current_price = tools.get_current_price(market, symbol)
|
||||
company_data = tools.get_company_data(market, symbol, language=language)
|
||||
|
||||
# Normalize timeframe
|
||||
tf = (timeframe or "1D").strip()
|
||||
|
||||
# 2) Kline + fundamentals
|
||||
kline_data = tools.get_stock_data(market, symbol, days=30, timeframe=tf)
|
||||
fundamental_data = tools.get_fundamental_data(market, symbol)
|
||||
indicators = tools.calculate_technical_indicators(kline_data or [])
|
||||
|
||||
# 3) News (Finnhub + web search)
|
||||
company_name = company_data.get('name', symbol) if company_data else symbol
|
||||
news_data = tools.get_news(market, symbol, days=7, company_name=company_name)
|
||||
|
||||
base_data = {
|
||||
"market": market,
|
||||
"symbol": symbol,
|
||||
"current_price": current_price,
|
||||
"kline_data": kline_data,
|
||||
"fundamental_data": fundamental_data,
|
||||
"company_data": company_data,
|
||||
"news_data": news_data,
|
||||
"indicators": indicators,
|
||||
}
|
||||
|
||||
context = {
|
||||
"market": market,
|
||||
"symbol": symbol,
|
||||
"language": language,
|
||||
"model": model,
|
||||
"timeframe": tf,
|
||||
"memory_features": {
|
||||
"timeframe": tf,
|
||||
"price": (current_price or {}).get("price"),
|
||||
"changePercent": (current_price or {}).get("changePercent"),
|
||||
"indicators": indicators,
|
||||
},
|
||||
"base_data": base_data
|
||||
}
|
||||
|
||||
emit_progress('data_collection', 'completed', None)
|
||||
|
||||
# Phase 1: Analysts (parallel but report individually)
|
||||
logger.info("Phase 1: Analyst team")
|
||||
|
||||
# Fundamental Analyst
|
||||
emit_progress('fundamental', 'started', None)
|
||||
fundamental_report = self.fundamental_analyst.analyze(context)
|
||||
emit_progress('fundamental', 'completed', fundamental_report.get('data', {}))
|
||||
|
||||
# Technical Analyst (market_analyst)
|
||||
emit_progress('technical', 'started', None)
|
||||
market_report = self.market_analyst.analyze(context)
|
||||
emit_progress('technical', 'completed', market_report.get('data', {}))
|
||||
|
||||
# News Analyst
|
||||
emit_progress('news', 'started', None)
|
||||
news_report = self.news_analyst.analyze(context)
|
||||
emit_progress('news', 'completed', news_report.get('data', {}))
|
||||
|
||||
# Sentiment Analyst
|
||||
emit_progress('sentiment', 'started', None)
|
||||
sentiment_report = self.sentiment_analyst.analyze(context)
|
||||
emit_progress('sentiment', 'completed', sentiment_report.get('data', {}))
|
||||
|
||||
# Risk Analyst
|
||||
emit_progress('risk', 'started', None)
|
||||
risk_report = self.risk_analyst.analyze(context)
|
||||
emit_progress('risk', 'completed', risk_report.get('data', {}))
|
||||
|
||||
# Update context with analyst outputs
|
||||
context.update({
|
||||
"market_report": market_report,
|
||||
"fundamental_report": fundamental_report,
|
||||
"news_report": news_report,
|
||||
"sentiment_report": sentiment_report,
|
||||
"risk_report": risk_report,
|
||||
})
|
||||
|
||||
# Phase 2: Bull/Bear debate
|
||||
logger.info("Phase 2: Research debate")
|
||||
|
||||
emit_progress('debate_bull', 'started', None)
|
||||
bull_argument = self.bull_researcher.analyze(context)
|
||||
emit_progress('debate_bull', 'completed', bull_argument.get('data', {}))
|
||||
|
||||
emit_progress('debate_bear', 'started', None)
|
||||
bear_argument = self.bear_researcher.analyze(context)
|
||||
emit_progress('debate_bear', 'completed', bear_argument.get('data', {}))
|
||||
|
||||
context["bull_argument"] = bull_argument
|
||||
context["bear_argument"] = bear_argument
|
||||
|
||||
# Research manager decision
|
||||
emit_progress('debate_research', 'started', None)
|
||||
research_decision = self._make_research_decision(bull_argument, bear_argument, context)
|
||||
context["research_decision"] = research_decision
|
||||
emit_progress('debate_research', 'completed', {'research_decision': research_decision})
|
||||
|
||||
# Phase 3: Trader decision
|
||||
logger.info("Phase 3: Trader decision")
|
||||
emit_progress('trader_decision', 'started', None)
|
||||
trader_result = self.trader_agent.analyze(context)
|
||||
trader_plan = trader_result.get('data', {}).get('trading_plan', {})
|
||||
context["trader_plan"] = trader_plan
|
||||
emit_progress('trader_decision', 'completed', trader_result.get('data', {}))
|
||||
|
||||
# Phase 4: Risk debate
|
||||
logger.info("Phase 4: Risk debate")
|
||||
|
||||
emit_progress('risk_debate_risky', 'started', None)
|
||||
risky_result = self.risky_analyst.analyze(context)
|
||||
emit_progress('risk_debate_risky', 'completed', risky_result.get('data', {}))
|
||||
|
||||
emit_progress('risk_debate_neutral', 'started', None)
|
||||
neutral_result = self.neutral_analyst.analyze(context)
|
||||
emit_progress('risk_debate_neutral', 'completed', neutral_result.get('data', {}))
|
||||
|
||||
emit_progress('risk_debate_safe', 'started', None)
|
||||
safe_result = self.safe_analyst.analyze(context)
|
||||
emit_progress('risk_debate_safe', 'completed', safe_result.get('data', {}))
|
||||
|
||||
# Final decision
|
||||
logger.info("Phase 5: Final decision")
|
||||
emit_progress('final_decision', 'started', None)
|
||||
final_decision = self._make_risk_decision(risky_result, neutral_result, safe_result, trader_result, context)
|
||||
emit_progress('final_decision', 'completed', final_decision)
|
||||
|
||||
# Generate overview
|
||||
emit_progress('overview', 'started', None)
|
||||
overview = self._generate_overview(context, final_decision)
|
||||
emit_progress('overview', 'completed', overview)
|
||||
|
||||
# Record for reflection
|
||||
if self.reflection_service and final_decision.get('decision') in ['BUY', 'SELL', 'HOLD']:
|
||||
try:
|
||||
self.reflection_service.record_analysis(
|
||||
market=market,
|
||||
symbol=symbol,
|
||||
price=base_data.get('current_price', {}).get('price'),
|
||||
decision=final_decision.get('decision'),
|
||||
confidence=final_decision.get('confidence', 50),
|
||||
reasoning=final_decision.get('reasoning', ''),
|
||||
check_days=7
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Record reflection failed: {e}")
|
||||
|
||||
# Build final result
|
||||
debate_data = {
|
||||
"bull": bull_argument.get('data', {}) if bull_argument.get('data') else {},
|
||||
"bear": bear_argument.get('data', {}) if bear_argument.get('data') else {},
|
||||
"research_decision": research_decision if research_decision else "Analyzing..."
|
||||
}
|
||||
|
||||
trader_decision_data = trader_result.get('data', {}) if trader_result.get('data') else {
|
||||
"decision": "HOLD",
|
||||
"confidence": 50,
|
||||
"reasoning": "Analyzing...",
|
||||
"trading_plan": {},
|
||||
"report": "Analyzing..."
|
||||
}
|
||||
|
||||
risk_debate_data = {
|
||||
"risky": risky_result.get('data', {}) if risky_result.get('data') else {},
|
||||
"neutral": neutral_result.get('data', {}) if neutral_result.get('data') else {},
|
||||
"safe": safe_result.get('data', {}) if safe_result.get('data') else {}
|
||||
}
|
||||
|
||||
if not final_decision or (isinstance(final_decision, dict) and len(final_decision) == 0):
|
||||
final_decision = {
|
||||
"decision": "HOLD",
|
||||
"confidence": 50,
|
||||
"reasoning": "Analyzing...",
|
||||
"risk_summary": {},
|
||||
"recommendation": "Analyzing..."
|
||||
}
|
||||
|
||||
result = {
|
||||
"overview": overview,
|
||||
"fundamental": fundamental_report.get('data', {}),
|
||||
"technical": market_report.get('data', {}),
|
||||
"news": news_report.get('data', {}),
|
||||
"sentiment": sentiment_report.get('data', {}),
|
||||
"risk": risk_report.get('data', {}),
|
||||
"debate": debate_data,
|
||||
"trader_decision": trader_decision_data,
|
||||
"risk_debate": risk_debate_data,
|
||||
"final_decision": final_decision,
|
||||
"error": None
|
||||
}
|
||||
|
||||
logger.info(f"Multi-agent stream analysis completed: {market}:{symbol}")
|
||||
return result
|
||||
|
||||
def run_analysis(self, market: str, symbol: str, language: str = 'zh-CN', model: str = None, timeframe: str = "1D") -> Dict[str, Any]:
|
||||
"""
|
||||
Run the full multi-agent analysis workflow.
|
||||
"""
|
||||
logger.info(f"Multi-agent analysis start: {market}:{symbol}, model={model}, language={language}")
|
||||
|
||||
# Build base context
|
||||
from .tools import AgentTools
|
||||
tools = AgentTools()
|
||||
|
||||
# 1) Base data
|
||||
current_price = tools.get_current_price(market, symbol)
|
||||
company_data = tools.get_company_data(market, symbol, language=language)
|
||||
|
||||
# Normalize timeframe
|
||||
tf = (timeframe or "1D").strip()
|
||||
|
||||
# 2) Kline + fundamentals
|
||||
kline_data = tools.get_stock_data(market, symbol, days=30, timeframe=tf)
|
||||
fundamental_data = tools.get_fundamental_data(market, symbol)
|
||||
indicators = tools.calculate_technical_indicators(kline_data or [])
|
||||
|
||||
# 3) News (Finnhub + web search)
|
||||
company_name = company_data.get('name', symbol) if company_data else symbol
|
||||
news_data = tools.get_news(market, symbol, days=7, company_name=company_name)
|
||||
|
||||
base_data = {
|
||||
"market": market,
|
||||
"symbol": symbol,
|
||||
"current_price": current_price,
|
||||
"kline_data": kline_data,
|
||||
"fundamental_data": fundamental_data,
|
||||
"company_data": company_data,
|
||||
"news_data": news_data,
|
||||
"indicators": indicators,
|
||||
}
|
||||
|
||||
context = {
|
||||
"market": market,
|
||||
"symbol": symbol,
|
||||
"language": language,
|
||||
"model": model,
|
||||
"timeframe": tf,
|
||||
# Compact structured features for memory retrieval.
|
||||
"memory_features": {
|
||||
"timeframe": tf,
|
||||
"price": (current_price or {}).get("price"),
|
||||
"changePercent": (current_price or {}).get("changePercent"),
|
||||
"indicators": indicators,
|
||||
},
|
||||
"base_data": base_data
|
||||
}
|
||||
|
||||
# Phase 1: Analysts (parallel)
|
||||
logger.info("Phase 1: Analyst team (parallel)")
|
||||
with ThreadPoolExecutor(max_workers=5) as executor:
|
||||
future_market = executor.submit(self.market_analyst.analyze, context)
|
||||
future_fundamental = executor.submit(self.fundamental_analyst.analyze, context)
|
||||
future_news = executor.submit(self.news_analyst.analyze, context)
|
||||
future_sentiment = executor.submit(self.sentiment_analyst.analyze, context)
|
||||
future_risk = executor.submit(self.risk_analyst.analyze, context)
|
||||
|
||||
market_report = future_market.result()
|
||||
fundamental_report = future_fundamental.result()
|
||||
news_report = future_news.result()
|
||||
sentiment_report = future_sentiment.result()
|
||||
risk_report = future_risk.result()
|
||||
|
||||
# Update context with analyst outputs
|
||||
context.update({
|
||||
"market_report": market_report,
|
||||
"fundamental_report": fundamental_report,
|
||||
"news_report": news_report,
|
||||
"sentiment_report": sentiment_report,
|
||||
"risk_report": risk_report,
|
||||
})
|
||||
|
||||
# Phase 2: Bull/Bear debate (parallel)
|
||||
logger.info("Phase 2: Research debate (parallel)")
|
||||
with ThreadPoolExecutor(max_workers=2) as executor:
|
||||
future_bull = executor.submit(self.bull_researcher.analyze, context)
|
||||
future_bear = executor.submit(self.bear_researcher.analyze, context)
|
||||
|
||||
bull_argument = future_bull.result()
|
||||
bear_argument = future_bear.result()
|
||||
|
||||
context["bull_argument"] = bull_argument
|
||||
context["bear_argument"] = bear_argument
|
||||
|
||||
# Research manager decision (lightweight, based on debate)
|
||||
research_decision = self._make_research_decision(
|
||||
bull_argument, bear_argument, context
|
||||
)
|
||||
context["research_decision"] = research_decision
|
||||
|
||||
# Phase 3: Trader decision
|
||||
logger.info("Phase 3: Trader decision")
|
||||
trader_result = self.trader_agent.analyze(context)
|
||||
trader_plan = trader_result.get('data', {}).get('trading_plan', {})
|
||||
context["trader_plan"] = trader_plan
|
||||
|
||||
# Phase 4: Risk debate (parallel)
|
||||
logger.info("Phase 4: Risk debate (parallel)")
|
||||
with ThreadPoolExecutor(max_workers=3) as executor:
|
||||
future_risky = executor.submit(self.risky_analyst.analyze, context)
|
||||
future_neutral = executor.submit(self.neutral_analyst.analyze, context)
|
||||
future_safe = executor.submit(self.safe_analyst.analyze, context)
|
||||
|
||||
risky_result = future_risky.result()
|
||||
neutral_result = future_neutral.result()
|
||||
safe_result = future_safe.result()
|
||||
|
||||
# Risk manager final decision
|
||||
final_decision = self._make_risk_decision(
|
||||
risky_result, neutral_result, safe_result,
|
||||
trader_result, context
|
||||
)
|
||||
|
||||
# Record analysis result for later reflection/validation
|
||||
if self.reflection_service and final_decision.get('decision') in ['BUY', 'SELL', 'HOLD']:
|
||||
try:
|
||||
self.reflection_service.record_analysis(
|
||||
market=market,
|
||||
symbol=symbol,
|
||||
price=base_data.get('current_price', {}).get('price'),
|
||||
decision=final_decision.get('decision'),
|
||||
confidence=final_decision.get('confidence', 50),
|
||||
reasoning=final_decision.get('reasoning', ''),
|
||||
check_days=7 # Validate after 7 days by default
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Record reflection failed: {e}")
|
||||
|
||||
# Build final result (defensive defaults to keep frontend stable)
|
||||
debate_data = {
|
||||
"bull": bull_argument.get('data', {}) if bull_argument.get('data') else {},
|
||||
"bear": bear_argument.get('data', {}) if bear_argument.get('data') else {},
|
||||
"research_decision": research_decision if research_decision else "Analyzing..."
|
||||
}
|
||||
|
||||
trader_decision_data = trader_result.get('data', {}) if trader_result.get('data') else {
|
||||
"decision": "HOLD",
|
||||
"confidence": 50,
|
||||
"reasoning": "Analyzing...",
|
||||
"trading_plan": {},
|
||||
"report": "Analyzing..."
|
||||
}
|
||||
|
||||
risk_debate_data = {
|
||||
"risky": risky_result.get('data', {}) if risky_result.get('data') else {},
|
||||
"neutral": neutral_result.get('data', {}) if neutral_result.get('data') else {},
|
||||
"safe": safe_result.get('data', {}) if safe_result.get('data') else {}
|
||||
}
|
||||
|
||||
# Ensure final_decision is present
|
||||
if not final_decision or (isinstance(final_decision, dict) and len(final_decision) == 0):
|
||||
final_decision = {
|
||||
"decision": "HOLD",
|
||||
"confidence": 50,
|
||||
"reasoning": "Analyzing...",
|
||||
"risk_summary": {},
|
||||
"recommendation": "Analyzing..."
|
||||
}
|
||||
|
||||
result = {
|
||||
"overview": self._generate_overview(context, final_decision),
|
||||
"fundamental": fundamental_report.get('data', {}),
|
||||
"technical": market_report.get('data', {}),
|
||||
"news": news_report.get('data', {}),
|
||||
"sentiment": sentiment_report.get('data', {}),
|
||||
"risk": risk_report.get('data', {}),
|
||||
"debate": debate_data,
|
||||
"trader_decision": trader_decision_data,
|
||||
"risk_debate": risk_debate_data,
|
||||
"final_decision": final_decision,
|
||||
"error": None
|
||||
}
|
||||
|
||||
logger.info(f"Multi-agent analysis completed: {market}:{symbol}")
|
||||
logger.info(
|
||||
"Result fields - debate=%s, trader_decision=%s, risk_debate=%s, final_decision=%s",
|
||||
bool(result.get('debate')),
|
||||
bool(result.get('trader_decision')),
|
||||
bool(result.get('risk_debate')),
|
||||
bool(result.get('final_decision')),
|
||||
)
|
||||
return result
|
||||
|
||||
def _make_research_decision(self, bull: Dict, bear: Dict, context: Dict) -> str:
|
||||
"""Research manager decision (rule-based, lightweight)."""
|
||||
try:
|
||||
language = context.get("language", "en-US")
|
||||
bull_confidence = bull.get('data', {}).get('confidence', 50)
|
||||
bear_confidence = bear.get('data', {}).get('confidence', 50)
|
||||
|
||||
# Tie-breaker when scores are close (<= 10)
|
||||
score_diff = bull_confidence - bear_confidence
|
||||
|
||||
if abs(score_diff) <= 10:
|
||||
# Use technical + sentiment as a bias signal
|
||||
market_score = context.get('market_report', {}).get('data', {}).get('score', 50)
|
||||
sentiment_score = context.get('sentiment_report', {}).get('data', {}).get('score', 50)
|
||||
|
||||
market_bias = (market_score + sentiment_score) / 2
|
||||
|
||||
if market_bias > 60:
|
||||
return self._t(
|
||||
language,
|
||||
en=(
|
||||
f"Research decision: bull and bear cases are close (bull {bull_confidence}% vs bear {bear_confidence}%), "
|
||||
f"but technical/sentiment are optimistic (avg {market_bias:.1f}), slightly leaning bullish."
|
||||
),
|
||||
zh=(
|
||||
f"研究经理决策:多空论据势均力敌(看涨 {bull_confidence}% vs 看跌 {bear_confidence}%),"
|
||||
f"但鉴于技术面和市场情绪偏乐观(平均分 {market_bias:.1f}),稍微倾向于看涨。"
|
||||
),
|
||||
)
|
||||
elif market_bias < 40:
|
||||
return self._t(
|
||||
language,
|
||||
en=(
|
||||
f"Research decision: bull and bear cases are close (bull {bull_confidence}% vs bear {bear_confidence}%), "
|
||||
f"but technical/sentiment are pessimistic (avg {market_bias:.1f}), slightly leaning bearish."
|
||||
),
|
||||
zh=(
|
||||
f"研究经理决策:多空论据势均力敌(看涨 {bull_confidence}% vs 看跌 {bear_confidence}%),"
|
||||
f"但鉴于技术面和市场情绪偏悲观(平均分 {market_bias:.1f}),稍微倾向于看跌。"
|
||||
),
|
||||
)
|
||||
else:
|
||||
return self._t(
|
||||
language,
|
||||
en=(
|
||||
f"Research decision: bull and bear cases are close (bull {bull_confidence}% vs bear {bear_confidence}%), "
|
||||
"and market bias is unclear. Prefer neutral / wait-and-see."
|
||||
),
|
||||
zh=(
|
||||
f"研究经理决策:多空论据势均力敌(看涨 {bull_confidence}% vs 看跌 {bear_confidence}%),"
|
||||
"且市场情绪不明朗,建议保持中立/观望。"
|
||||
),
|
||||
)
|
||||
|
||||
elif score_diff > 10:
|
||||
return self._t(
|
||||
language,
|
||||
en=(
|
||||
f"Research decision: bullish case (confidence {bull_confidence}%) is clearly stronger than bearish "
|
||||
f"(confidence {bear_confidence}%). Lean bullish."
|
||||
),
|
||||
zh=(
|
||||
f"研究经理决策:基于看涨论据(置信度 {bull_confidence}%)明显强于看跌论据(置信度 {bear_confidence}%),明确倾向于看涨。"
|
||||
),
|
||||
)
|
||||
else: # score_diff < -10
|
||||
return self._t(
|
||||
language,
|
||||
en=(
|
||||
f"Research decision: bearish case (confidence {bear_confidence}%) is clearly stronger than bullish "
|
||||
f"(confidence {bull_confidence}%). Lean bearish."
|
||||
),
|
||||
zh=(
|
||||
f"研究经理决策:基于看跌论据(置信度 {bear_confidence}%)明显强于看涨论据(置信度 {bull_confidence}%),明确倾向于看跌。"
|
||||
),
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Research decision failed: {e}")
|
||||
language = context.get("language", "en-US") if isinstance(context, dict) else "en-US"
|
||||
return self._t(language, en="Research decision: unable to reach a clear conclusion.", zh="研究经理决策:无法做出明确判断。")
|
||||
|
||||
def _make_risk_decision(self, risky: Dict, neutral: Dict, safe: Dict,
|
||||
trader: Dict, context: Dict) -> Dict[str, Any]:
|
||||
"""Risk manager final decision (lightweight)."""
|
||||
try:
|
||||
language = context.get("language", "en-US")
|
||||
trader_decision = trader.get('data', {}).get('decision', 'HOLD')
|
||||
trader_confidence = trader.get('data', {}).get('confidence', 50)
|
||||
|
||||
# Risk debate summary
|
||||
risk_summary = {
|
||||
"risky_view": risky.get('data', {}).get('recommendation', ''),
|
||||
"neutral_view": neutral.get('data', {}).get('recommendation', ''),
|
||||
"safe_view": safe.get('data', {}).get('recommendation', ''),
|
||||
}
|
||||
|
||||
# Final decision (use trader decision + risk debate context)
|
||||
final_decision = {
|
||||
"decision": trader_decision,
|
||||
"confidence": trader_confidence,
|
||||
"reasoning": self._t(
|
||||
language,
|
||||
en=f"Final decision is based on trader analysis ({trader_decision}, confidence {trader_confidence}%) and the risk debate.",
|
||||
zh=f"基于交易员分析({trader_decision},置信度 {trader_confidence}%)和风险辩论,做出最终决策。",
|
||||
),
|
||||
"risk_summary": risk_summary,
|
||||
"recommendation": trader.get('data', {}).get('report', '')
|
||||
}
|
||||
|
||||
return final_decision
|
||||
except Exception as e:
|
||||
logger.error(f"Risk decision failed: {e}")
|
||||
language = context.get("language", "en-US") if isinstance(context, dict) else "en-US"
|
||||
return {
|
||||
"decision": "HOLD",
|
||||
"confidence": 50,
|
||||
"reasoning": self._t(language, en="Risk decision failed", zh="风险决策失败"),
|
||||
"risk_summary": {},
|
||||
"recommendation": ""
|
||||
}
|
||||
|
||||
def _generate_overview(self, context: Dict, final_decision: Dict) -> Dict[str, Any]:
|
||||
"""Generate overview section (lightweight, deterministic)."""
|
||||
try:
|
||||
language = context.get("language", "en-US")
|
||||
# Extract dimension scores
|
||||
technical_data = context.get('market_report', {}).get('data', {})
|
||||
fundamental_data = context.get('fundamental_report', {}).get('data', {})
|
||||
news_data = context.get('news_report', {}).get('data', {})
|
||||
sentiment_data = context.get('sentiment_report', {}).get('data', {})
|
||||
risk_data = context.get('risk_report', {}).get('data', {})
|
||||
|
||||
technical_score = technical_data.get('score', 50)
|
||||
fundamental_score = fundamental_data.get('score', 50)
|
||||
news_score = news_data.get('score', 50)
|
||||
sentiment_score = sentiment_data.get('score', 50)
|
||||
risk_score = risk_data.get('score', 50)
|
||||
|
||||
# Generate an overall score using weighted dimensions + decision/confidence adjustment
|
||||
decision = final_decision.get('decision', 'HOLD')
|
||||
confidence = final_decision.get('confidence', 50)
|
||||
|
||||
# 1) Base score: weighted average (tech 30%, fundamental 25%, news 15%, sentiment 15%, risk 15%)
|
||||
weighted_score = (
|
||||
technical_score * 0.3 +
|
||||
fundamental_score * 0.25 +
|
||||
news_score * 0.15 +
|
||||
sentiment_score * 0.15 +
|
||||
risk_score * 0.15
|
||||
)
|
||||
|
||||
# 2) Decision adjustment: BUY pushes toward 60-100, SELL toward 0-40, HOLD toward 50
|
||||
if decision == 'BUY':
|
||||
target_score = 60 + (confidence / 100 * 40) # Map to 60-100
|
||||
overall_score = (weighted_score * 0.4) + (target_score * 0.6)
|
||||
elif decision == 'SELL':
|
||||
target_score = 40 - (confidence / 100 * 40) # Map to 0-40
|
||||
overall_score = (weighted_score * 0.4) + (target_score * 0.6)
|
||||
else:
|
||||
overall_score = (weighted_score * 0.6) + (50 * 0.4)
|
||||
|
||||
# Clamp to 0..100
|
||||
overall_score = max(0, min(100, int(overall_score)))
|
||||
|
||||
return {
|
||||
"overallScore": overall_score,
|
||||
"recommendation": decision,
|
||||
"confidence": confidence,
|
||||
"dimensionScores": {
|
||||
"fundamental": fundamental_score,
|
||||
"technical": technical_score,
|
||||
"news": news_score,
|
||||
"sentiment": sentiment_score,
|
||||
"risk": risk_score
|
||||
},
|
||||
"report": final_decision.get(
|
||||
'reasoning',
|
||||
self._t(language, en="Overview generated.", zh="综合分析完成"),
|
||||
)
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Generate overview failed: {e}")
|
||||
language = context.get("language", "en-US") if isinstance(context, dict) else "en-US"
|
||||
return {
|
||||
"overallScore": 50,
|
||||
"recommendation": "HOLD",
|
||||
"confidence": 50,
|
||||
"dimensionScores": {
|
||||
"fundamental": 50,
|
||||
"technical": 50,
|
||||
"news": 50,
|
||||
"sentiment": 50,
|
||||
"risk": 50
|
||||
},
|
||||
"report": self._t(language, en="Failed to generate overview.", zh="综合分析生成失败")
|
||||
}
|
||||
|
||||
def reflect_and_learn(self, market: str, symbol: str, decision: str,
|
||||
returns: Optional[float] = None, result: Optional[str] = None):
|
||||
"""
|
||||
Reflection hook: store post-trade outcomes into memory (local-only).
|
||||
|
||||
Args:
|
||||
market: Market
|
||||
symbol: Symbol
|
||||
decision: BUY/SELL/HOLD
|
||||
returns: Return percentage
|
||||
result: Free-text outcome
|
||||
"""
|
||||
if not self.enable_memory:
|
||||
return
|
||||
|
||||
try:
|
||||
situation = f"{market}:{symbol} trading decision"
|
||||
recommendation = f"Decision: {decision}, returns: {returns if returns is not None else 'N/A'}"
|
||||
|
||||
# Update trader memory
|
||||
if 'trader' in self.memories:
|
||||
self.memories['trader'].add_memory(
|
||||
situation,
|
||||
recommendation,
|
||||
result,
|
||||
returns,
|
||||
metadata={
|
||||
"market": market,
|
||||
"symbol": symbol,
|
||||
"timeframe": None,
|
||||
"features": {
|
||||
"source": "manual_reflect",
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
logger.info(f"Reflection completed: {market}:{symbol}, decision={decision}, returns={returns}")
|
||||
except Exception as e:
|
||||
logger.error(f"Reflection failed: {e}")
|
||||
@@ -1,91 +0,0 @@
|
||||
"""
|
||||
Lightweight embedding utilities for local-only deployments.
|
||||
|
||||
We intentionally avoid heavyweight ML deps (torch/sentence-transformers) and external services.
|
||||
This module provides a deterministic "hashed embedding" (similar to feature hashing):
|
||||
- Tokenize text
|
||||
- Hash tokens into a fixed-size dense vector
|
||||
- L2 normalize
|
||||
|
||||
It is not as semantically strong as modern transformer embeddings, but it enables:
|
||||
- Vector storage in SQLite
|
||||
- Cosine similarity retrieval
|
||||
- Recency/return weighted ranking
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
import struct
|
||||
import hashlib
|
||||
from typing import List, Optional
|
||||
|
||||
|
||||
_TOKEN_RE = re.compile(r"[A-Za-z0-9_]+", re.UNICODE)
|
||||
|
||||
|
||||
def _tokenize(text: str) -> List[str]:
|
||||
t = (text or "").lower()
|
||||
return _TOKEN_RE.findall(t)
|
||||
|
||||
|
||||
class EmbeddingService:
|
||||
"""Deterministic local embedding service."""
|
||||
|
||||
def __init__(self, dim: Optional[int] = None):
|
||||
self.dim = int(dim or os.getenv("AGENT_MEMORY_EMBEDDING_DIM", "256") or 256)
|
||||
if self.dim <= 0:
|
||||
self.dim = 256
|
||||
|
||||
def embed(self, text: str) -> List[float]:
|
||||
"""
|
||||
Return a dense, L2-normalized embedding vector.
|
||||
"""
|
||||
vec = [0.0] * self.dim
|
||||
tokens = _tokenize(text)
|
||||
if not tokens:
|
||||
return vec
|
||||
|
||||
# Feature hashing with signed counts
|
||||
for tok in tokens:
|
||||
h = hashlib.blake2b(tok.encode("utf-8"), digest_size=8).digest()
|
||||
# Use first 8 bytes as unsigned int
|
||||
v = int.from_bytes(h, "little", signed=False)
|
||||
idx = v % self.dim
|
||||
sign = -1.0 if ((v >> 63) & 1) else 1.0
|
||||
vec[idx] += sign
|
||||
|
||||
# L2 normalize
|
||||
norm = math.sqrt(sum(x * x for x in vec)) or 1.0
|
||||
return [x / norm for x in vec]
|
||||
|
||||
def to_bytes(self, vec: List[float]) -> bytes:
|
||||
"""
|
||||
Pack float vector into little-endian float32 bytes for SQLite BLOB storage.
|
||||
"""
|
||||
if not vec:
|
||||
return b""
|
||||
return struct.pack("<" + "f" * len(vec), *[float(x) for x in vec])
|
||||
|
||||
def from_bytes(self, blob: bytes) -> List[float]:
|
||||
if not blob:
|
||||
return []
|
||||
n = len(blob) // 4
|
||||
if n <= 0:
|
||||
return []
|
||||
return list(struct.unpack("<" + "f" * n, blob[: n * 4]))
|
||||
|
||||
|
||||
def cosine_sim(a: List[float], b: List[float]) -> float:
|
||||
"""
|
||||
Cosine similarity for L2-normalized vectors.
|
||||
If vectors are not normalized, this becomes a scaled dot product.
|
||||
"""
|
||||
if not a or not b:
|
||||
return 0.0
|
||||
n = min(len(a), len(b))
|
||||
return float(sum(a[i] * b[i] for i in range(n)))
|
||||
|
||||
|
||||
@@ -1,319 +0,0 @@
|
||||
"""
|
||||
Agent memory system (PostgreSQL).
|
||||
|
||||
This module stores agent experiences in PostgreSQL and retrieves relevant past cases
|
||||
to inject into prompts (RAG-style). It does NOT finetune model weights.
|
||||
|
||||
Retrieval (configurable):
|
||||
- Vector similarity via deterministic local embeddings (default)
|
||||
- Fallback to difflib text similarity when embeddings are missing
|
||||
|
||||
Ranking combines:
|
||||
- similarity
|
||||
- recency decay (half-life)
|
||||
- optional returns weight
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import math
|
||||
from typing import List, Dict, Any, Optional
|
||||
from datetime import datetime, timezone
|
||||
import difflib
|
||||
|
||||
from app.utils.logger import get_logger
|
||||
from app.utils.db import get_db_connection
|
||||
from .embedding import EmbeddingService, cosine_sim
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class AgentMemory:
|
||||
"""Agent memory system using PostgreSQL"""
|
||||
|
||||
def __init__(self, agent_name: str, db_path: Optional[str] = None):
|
||||
"""
|
||||
Initialize memory system.
|
||||
|
||||
Args:
|
||||
agent_name: Agent identifier (e.g., 'trader_agent', 'risk_analyst')
|
||||
db_path: Deprecated parameter, kept for backward compatibility
|
||||
"""
|
||||
self.agent_name = agent_name
|
||||
self.embedder = EmbeddingService()
|
||||
self.enable_vector = os.getenv("AGENT_MEMORY_ENABLE_VECTOR", "true").lower() == "true"
|
||||
self.candidate_limit = int(os.getenv("AGENT_MEMORY_CANDIDATE_LIMIT", "500") or 500)
|
||||
self.half_life_days = float(os.getenv("AGENT_MEMORY_HALF_LIFE_DAYS", "30") or 30)
|
||||
self.w_sim = float(os.getenv("AGENT_MEMORY_W_SIM", "0.75") or 0.75)
|
||||
self.w_recency = float(os.getenv("AGENT_MEMORY_W_RECENCY", "0.20") or 0.20)
|
||||
self.w_returns = float(os.getenv("AGENT_MEMORY_W_RETURNS", "0.05") or 0.05)
|
||||
|
||||
def _now_utc(self) -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
def _parse_ts(self, ts_val: Any) -> Optional[datetime]:
|
||||
if ts_val is None:
|
||||
return None
|
||||
if isinstance(ts_val, datetime):
|
||||
return ts_val
|
||||
s = str(ts_val)
|
||||
try:
|
||||
return datetime.fromisoformat(s.replace("Z", ""))
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def _recency_score(self, created_at: Any) -> float:
|
||||
dt = self._parse_ts(created_at)
|
||||
if not dt:
|
||||
return 0.0
|
||||
if dt.tzinfo is None:
|
||||
dt = dt.replace(tzinfo=timezone.utc)
|
||||
age_days = max(0.0, (self._now_utc() - dt).total_seconds() / 86400.0)
|
||||
hl = max(0.1, float(self.half_life_days or 30.0))
|
||||
return float(math.exp(-math.log(2.0) * (age_days / hl)))
|
||||
|
||||
def _returns_score(self, returns: Any) -> float:
|
||||
try:
|
||||
r = float(returns)
|
||||
except Exception:
|
||||
return 0.0
|
||||
return float(math.tanh(r / 10.0))
|
||||
|
||||
def _build_embed_text(self, situation: str, recommendation: str, result: Optional[str], features_json: Optional[str]) -> str:
|
||||
return "\n".join([
|
||||
f"situation: {situation or ''}",
|
||||
f"recommendation: {recommendation or ''}",
|
||||
f"result: {result or ''}",
|
||||
f"features: {features_json or ''}",
|
||||
])
|
||||
|
||||
def add_memory(
|
||||
self,
|
||||
situation: str,
|
||||
recommendation: str,
|
||||
result: Optional[str] = None,
|
||||
returns: Optional[float] = None,
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
"""
|
||||
Add a memory entry.
|
||||
|
||||
Args:
|
||||
situation: Situation description
|
||||
recommendation: Decision/recommendation made
|
||||
result: Outcome description (optional)
|
||||
returns: Return percentage (optional)
|
||||
metadata: Optional structured metadata (market/symbol/timeframe/features...)
|
||||
"""
|
||||
try:
|
||||
meta = metadata or {}
|
||||
market = (meta.get("market") or "").strip() or None
|
||||
symbol = (meta.get("symbol") or "").strip() or None
|
||||
timeframe = (meta.get("timeframe") or "").strip() or None
|
||||
features = meta.get("features") if isinstance(meta, dict) else None
|
||||
try:
|
||||
features_json = json.dumps(features, ensure_ascii=False) if features is not None else None
|
||||
except Exception:
|
||||
features_json = None
|
||||
|
||||
embedding_blob = None
|
||||
if self.enable_vector:
|
||||
text = self._build_embed_text(situation, recommendation, result, features_json)
|
||||
vec = self.embedder.embed(text)
|
||||
embedding_blob = self.embedder.to_bytes(vec)
|
||||
|
||||
with get_db_connection() as conn:
|
||||
cur = conn.cursor()
|
||||
cur.execute(
|
||||
"""
|
||||
INSERT INTO qd_agent_memories
|
||||
(agent_name, situation, recommendation, result, returns, market, symbol, timeframe, features_json, embedding)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(self.agent_name, situation, recommendation, result, returns, market, symbol, timeframe, features_json, embedding_blob)
|
||||
)
|
||||
conn.commit()
|
||||
cur.close()
|
||||
logger.info(f"{self.agent_name} added new memory")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to add memory: {e}")
|
||||
|
||||
def get_memories(self, current_situation: str, n_matches: int = 5, metadata: Optional[Dict[str, Any]] = None) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Retrieve similar memories.
|
||||
|
||||
Args:
|
||||
current_situation: Current situation description
|
||||
n_matches: Number of matches to return
|
||||
metadata: Optional metadata for filtering/weighting
|
||||
|
||||
Returns:
|
||||
List of matching memory entries
|
||||
"""
|
||||
try:
|
||||
with get_db_connection() as conn:
|
||||
cur = conn.cursor()
|
||||
cur.execute(
|
||||
"""
|
||||
SELECT id, situation, recommendation, result, returns, created_at,
|
||||
market, symbol, timeframe, features_json, embedding
|
||||
FROM qd_agent_memories
|
||||
WHERE agent_name = ?
|
||||
ORDER BY created_at DESC
|
||||
LIMIT ?
|
||||
""",
|
||||
(self.agent_name, int(self.candidate_limit))
|
||||
)
|
||||
all_memories = cur.fetchall() or []
|
||||
cur.close()
|
||||
|
||||
if not all_memories:
|
||||
return []
|
||||
|
||||
meta = metadata or {}
|
||||
tf = (meta.get("timeframe") or "").strip()
|
||||
features = meta.get("features") if isinstance(meta, dict) else None
|
||||
try:
|
||||
q_features_json = json.dumps(features, ensure_ascii=False) if features is not None else None
|
||||
except Exception:
|
||||
q_features_json = None
|
||||
|
||||
query_vec = []
|
||||
if self.enable_vector:
|
||||
query_text = self._build_embed_text(current_situation, "", "", q_features_json)
|
||||
query_vec = self.embedder.embed(query_text)
|
||||
|
||||
ranked = []
|
||||
for row in all_memories:
|
||||
mem_id = row['id']
|
||||
situation = row['situation']
|
||||
recommendation = row['recommendation']
|
||||
result = row['result']
|
||||
returns = row['returns']
|
||||
created_at = row['created_at']
|
||||
market = row['market']
|
||||
symbol = row['symbol']
|
||||
timeframe = row['timeframe']
|
||||
features_json = row['features_json']
|
||||
embedding_blob = row['embedding']
|
||||
|
||||
sim = 0.0
|
||||
if self.enable_vector and embedding_blob:
|
||||
try:
|
||||
# Handle memoryview/bytes from PostgreSQL
|
||||
if isinstance(embedding_blob, memoryview):
|
||||
embedding_blob = bytes(embedding_blob)
|
||||
mem_vec = self.embedder.from_bytes(embedding_blob)
|
||||
sim = cosine_sim(query_vec, mem_vec)
|
||||
except Exception:
|
||||
sim = 0.0
|
||||
else:
|
||||
sim = difflib.SequenceMatcher(None, (current_situation or "").lower(), (situation or "").lower()).ratio()
|
||||
|
||||
rec = self._recency_score(created_at)
|
||||
ret = self._returns_score(returns)
|
||||
|
||||
score = (self.w_sim * sim) + (self.w_recency * rec) + (self.w_returns * ret)
|
||||
|
||||
if tf and timeframe and str(timeframe).strip() != tf:
|
||||
score -= 0.15
|
||||
|
||||
ranked.append({
|
||||
'id': mem_id,
|
||||
'matched_situation': situation,
|
||||
'recommendation': recommendation,
|
||||
'result': result,
|
||||
'returns': returns,
|
||||
'created_at': created_at,
|
||||
'market': market,
|
||||
'symbol': symbol,
|
||||
'timeframe': timeframe,
|
||||
'features_json': features_json,
|
||||
'score': float(score),
|
||||
'sim': float(sim),
|
||||
'recency': float(rec),
|
||||
})
|
||||
|
||||
ranked.sort(key=lambda x: x.get('score', 0.0), reverse=True)
|
||||
return ranked[: max(0, int(n_matches or 0))]
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to retrieve memories: {e}")
|
||||
return []
|
||||
|
||||
def update_memory_result(self, memory_id: int, result: str, returns: Optional[float] = None):
|
||||
"""
|
||||
Update memory result.
|
||||
|
||||
Args:
|
||||
memory_id: Memory ID
|
||||
result: Outcome description
|
||||
returns: Return percentage
|
||||
"""
|
||||
try:
|
||||
with get_db_connection() as conn:
|
||||
cur = conn.cursor()
|
||||
cur.execute(
|
||||
"""
|
||||
UPDATE qd_agent_memories
|
||||
SET result = ?, returns = ?, updated_at = NOW()
|
||||
WHERE id = ? AND agent_name = ?
|
||||
""",
|
||||
(result, returns, memory_id, self.agent_name)
|
||||
)
|
||||
conn.commit()
|
||||
cur.close()
|
||||
logger.info(f"{self.agent_name} updated memory {memory_id}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to update memory: {e}")
|
||||
|
||||
def get_statistics(self) -> Dict[str, Any]:
|
||||
"""Get memory statistics for this agent."""
|
||||
try:
|
||||
with get_db_connection() as conn:
|
||||
cur = conn.cursor()
|
||||
|
||||
cur.execute(
|
||||
'SELECT COUNT(*) as cnt FROM qd_agent_memories WHERE agent_name = ?',
|
||||
(self.agent_name,)
|
||||
)
|
||||
total = cur.fetchone()['cnt']
|
||||
|
||||
cur.execute(
|
||||
'SELECT AVG(returns) as avg_ret FROM qd_agent_memories WHERE agent_name = ? AND returns IS NOT NULL',
|
||||
(self.agent_name,)
|
||||
)
|
||||
avg_returns = cur.fetchone()['avg_ret'] or 0
|
||||
|
||||
cur.execute(
|
||||
'SELECT COUNT(*) as cnt FROM qd_agent_memories WHERE agent_name = ? AND returns > 0',
|
||||
(self.agent_name,)
|
||||
)
|
||||
positive = cur.fetchone()['cnt']
|
||||
|
||||
cur.close()
|
||||
|
||||
return {
|
||||
'total_memories': total,
|
||||
'average_returns': round(avg_returns, 2),
|
||||
'positive_decisions': positive,
|
||||
'success_rate': round(positive / total * 100, 2) if total > 0 else 0
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get statistics: {e}")
|
||||
return {}
|
||||
|
||||
def clear_memories(self):
|
||||
"""Clear all memories for this agent (use with caution)."""
|
||||
try:
|
||||
with get_db_connection() as conn:
|
||||
cur = conn.cursor()
|
||||
cur.execute(
|
||||
'DELETE FROM qd_agent_memories WHERE agent_name = ?',
|
||||
(self.agent_name,)
|
||||
)
|
||||
conn.commit()
|
||||
cur.close()
|
||||
logger.warning(f"{self.agent_name} cleared all memories")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to clear memories: {e}")
|
||||
@@ -1,249 +0,0 @@
|
||||
"""
|
||||
Auto-reflection and verification service (PostgreSQL).
|
||||
|
||||
Records analysis predictions and auto-verifies results in the future
|
||||
to achieve closed-loop learning for AI agents.
|
||||
"""
|
||||
|
||||
import os
|
||||
from datetime import datetime, timedelta
|
||||
from typing import List, Dict, Any, Optional
|
||||
|
||||
from app.utils.logger import get_logger
|
||||
from app.utils.db import get_db_connection
|
||||
from .memory import AgentMemory
|
||||
from .tools import AgentTools
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class ReflectionService:
|
||||
"""Reflection service: manages storage and verification of analysis records."""
|
||||
|
||||
def __init__(self, db_path: Optional[str] = None):
|
||||
"""
|
||||
Initialize reflection service.
|
||||
|
||||
Args:
|
||||
db_path: Deprecated parameter, kept for backward compatibility
|
||||
"""
|
||||
self.tools = AgentTools()
|
||||
|
||||
def record_analysis(
|
||||
self,
|
||||
market: str,
|
||||
symbol: str,
|
||||
price: float,
|
||||
decision: str,
|
||||
confidence: int,
|
||||
reasoning: str,
|
||||
check_days: int = 7
|
||||
):
|
||||
"""
|
||||
Record an analysis for future verification.
|
||||
|
||||
Args:
|
||||
market: Market type
|
||||
symbol: Symbol code
|
||||
price: Current price
|
||||
decision: Decision (BUY/SELL/HOLD)
|
||||
confidence: Confidence level (0-100)
|
||||
reasoning: Reasoning text
|
||||
check_days: Days until verification (default 7)
|
||||
"""
|
||||
try:
|
||||
target_date = datetime.now() + timedelta(days=check_days)
|
||||
|
||||
with get_db_connection() as conn:
|
||||
cur = conn.cursor()
|
||||
cur.execute(
|
||||
"""
|
||||
INSERT INTO qd_reflection_records
|
||||
(market, symbol, initial_price, decision, confidence, reasoning, target_check_date)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(market, symbol, price, decision, confidence, reasoning, target_date)
|
||||
)
|
||||
conn.commit()
|
||||
cur.close()
|
||||
logger.info(f"Recorded analysis for reflection: {market}:{symbol}, will verify after {check_days} day(s)")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to record analysis: {e}")
|
||||
|
||||
def run_verification_cycle(self):
|
||||
"""
|
||||
Execute verification cycle: check due records, verify results, and write to memory.
|
||||
"""
|
||||
logger.info("Starting auto-reflection verification cycle...")
|
||||
|
||||
try:
|
||||
with get_db_connection() as conn:
|
||||
cur = conn.cursor()
|
||||
|
||||
# 1. Find all due and pending records
|
||||
cur.execute(
|
||||
"""
|
||||
SELECT id, market, symbol, initial_price, decision, confidence, reasoning, analysis_date
|
||||
FROM qd_reflection_records
|
||||
WHERE status = 'PENDING' AND target_check_date <= NOW()
|
||||
"""
|
||||
)
|
||||
records = cur.fetchall() or []
|
||||
|
||||
if not records:
|
||||
logger.info("No records to verify")
|
||||
cur.close()
|
||||
return
|
||||
|
||||
logger.info(f"Found {len(records)} records to verify")
|
||||
|
||||
# Initialize memory system for writing verification results
|
||||
trader_memory = AgentMemory('trader_agent')
|
||||
|
||||
for record in records:
|
||||
record_id = record['id']
|
||||
market = record['market']
|
||||
symbol = record['symbol']
|
||||
initial_price = record['initial_price']
|
||||
decision = record['decision']
|
||||
confidence = record['confidence']
|
||||
reasoning = record['reasoning']
|
||||
analysis_date = record['analysis_date']
|
||||
|
||||
try:
|
||||
# 2. Get current price
|
||||
current_price_data = self.tools.get_current_price(market, symbol)
|
||||
current_price = current_price_data.get('price')
|
||||
|
||||
if not current_price:
|
||||
logger.warning(f"Cannot get current price for {market}:{symbol}, skipping")
|
||||
continue
|
||||
|
||||
# 3. Calculate return and result
|
||||
if not initial_price or initial_price == 0:
|
||||
actual_return = 0.0
|
||||
else:
|
||||
actual_return = (current_price - initial_price) / initial_price * 100
|
||||
|
||||
# Evaluate result
|
||||
result_desc = ""
|
||||
is_good_prediction = False
|
||||
|
||||
if decision == "BUY":
|
||||
if actual_return > 2.0:
|
||||
result_desc = "Correct: price rose after BUY"
|
||||
is_good_prediction = True
|
||||
elif actual_return < -2.0:
|
||||
result_desc = "Wrong: price fell after BUY"
|
||||
else:
|
||||
result_desc = "Neutral: limited price movement"
|
||||
elif decision == "SELL":
|
||||
if actual_return < -2.0:
|
||||
result_desc = "Correct: price fell after SELL"
|
||||
is_good_prediction = True
|
||||
elif actual_return > 2.0:
|
||||
result_desc = "Wrong: price rose after SELL"
|
||||
else:
|
||||
result_desc = "Neutral: limited price movement"
|
||||
else: # HOLD
|
||||
if -2.0 <= actual_return <= 2.0:
|
||||
result_desc = "Correct: limited movement during HOLD"
|
||||
is_good_prediction = True
|
||||
else:
|
||||
result_desc = f"Deviated: large movement during HOLD ({actual_return:.2f}%)"
|
||||
|
||||
# 4. Write to memory system (agent learning)
|
||||
memory_situation = f"{market}:{symbol} auto-verified (analysis_date: {analysis_date})"
|
||||
memory_recommendation = f"Decision: {decision} (confidence {confidence}), reasoning: {(reasoning or '')[:120]}"
|
||||
memory_result = f"Verification: {result_desc}; return={actual_return:.2f}% (initial {initial_price} -> final {current_price})"
|
||||
|
||||
trader_memory.add_memory(
|
||||
memory_situation,
|
||||
memory_recommendation,
|
||||
memory_result,
|
||||
actual_return,
|
||||
metadata={
|
||||
"market": market,
|
||||
"symbol": symbol,
|
||||
"timeframe": "1D",
|
||||
"features": {
|
||||
"source": "auto_verify",
|
||||
"decision": decision,
|
||||
"confidence": confidence,
|
||||
"initial_price": initial_price,
|
||||
"final_price": current_price,
|
||||
"analysis_date": str(analysis_date),
|
||||
"result_desc": result_desc,
|
||||
"is_good_prediction": bool(is_good_prediction),
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
# 5. Update record status
|
||||
cur.execute(
|
||||
"""
|
||||
UPDATE qd_reflection_records
|
||||
SET status = 'COMPLETED', final_price = ?, actual_return = ?, check_result = ?
|
||||
WHERE id = ?
|
||||
""",
|
||||
(current_price, actual_return, result_desc, record_id)
|
||||
)
|
||||
conn.commit()
|
||||
logger.info(f"Verification completed {market}:{symbol}: {result_desc}")
|
||||
|
||||
except Exception as inner_e:
|
||||
logger.error(f"Failed to process record {record_id}: {inner_e}")
|
||||
# Optionally mark as failed to avoid repeated processing
|
||||
# cur.execute("UPDATE qd_reflection_records SET status = 'FAILED' WHERE id = ?", (record_id,))
|
||||
# conn.commit()
|
||||
|
||||
cur.close()
|
||||
logger.info("Reflection verification cycle completed")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to execute verification cycle: {e}")
|
||||
|
||||
def get_pending_count(self) -> int:
|
||||
"""Get count of pending verification records."""
|
||||
try:
|
||||
with get_db_connection() as conn:
|
||||
cur = conn.cursor()
|
||||
cur.execute("SELECT COUNT(*) as cnt FROM qd_reflection_records WHERE status = 'PENDING'")
|
||||
count = cur.fetchone()['cnt']
|
||||
cur.close()
|
||||
return count
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get pending count: {e}")
|
||||
return 0
|
||||
|
||||
def get_statistics(self) -> Dict[str, Any]:
|
||||
"""Get reflection statistics."""
|
||||
try:
|
||||
with get_db_connection() as conn:
|
||||
cur = conn.cursor()
|
||||
|
||||
cur.execute("SELECT COUNT(*) as cnt FROM qd_reflection_records")
|
||||
total = cur.fetchone()['cnt']
|
||||
|
||||
cur.execute("SELECT COUNT(*) as cnt FROM qd_reflection_records WHERE status = 'PENDING'")
|
||||
pending = cur.fetchone()['cnt']
|
||||
|
||||
cur.execute("SELECT COUNT(*) as cnt FROM qd_reflection_records WHERE status = 'COMPLETED'")
|
||||
completed = cur.fetchone()['cnt']
|
||||
|
||||
cur.execute(
|
||||
"SELECT AVG(actual_return) as avg_ret FROM qd_reflection_records WHERE status = 'COMPLETED' AND actual_return IS NOT NULL"
|
||||
)
|
||||
avg_return = cur.fetchone()['avg_ret'] or 0
|
||||
|
||||
cur.close()
|
||||
|
||||
return {
|
||||
'total_records': total,
|
||||
'pending_records': pending,
|
||||
'completed_records': completed,
|
||||
'average_return': round(avg_return, 2)
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get statistics: {e}")
|
||||
return {}
|
||||
@@ -1,67 +0,0 @@
|
||||
"""
|
||||
Background worker for automated reflection verification.
|
||||
|
||||
This replaces the need for an external cron job in local deployments.
|
||||
It periodically runs ReflectionService.run_verification_cycle().
|
||||
|
||||
Controls (env):
|
||||
- ENABLE_REFLECTION_WORKER=true/false (default: false)
|
||||
- REFLECTION_WORKER_INTERVAL_SEC (default: 86400)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
import threading
|
||||
from typing import Optional
|
||||
|
||||
from app.utils.logger import get_logger
|
||||
from .reflection import ReflectionService
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class ReflectionWorker:
|
||||
def __init__(self):
|
||||
self._thread: Optional[threading.Thread] = None
|
||||
self._stop = threading.Event()
|
||||
|
||||
def start(self):
|
||||
if self._thread and self._thread.is_alive():
|
||||
return
|
||||
|
||||
interval = int(os.getenv("REFLECTION_WORKER_INTERVAL_SEC", "86400") or 86400)
|
||||
interval = max(60, interval) # at least 1 minute
|
||||
|
||||
def _run():
|
||||
logger.info(f"Reflection worker started (interval={interval}s)")
|
||||
svc = ReflectionService()
|
||||
# Initial small delay to avoid fighting startup spikes
|
||||
time.sleep(3)
|
||||
while not self._stop.is_set():
|
||||
try:
|
||||
svc.run_verification_cycle()
|
||||
except Exception as e:
|
||||
logger.error(f"Reflection worker cycle failed: {e}")
|
||||
# Sleep in small steps to react to stop quickly
|
||||
remaining = interval
|
||||
while remaining > 0 and not self._stop.is_set():
|
||||
step = min(5, remaining)
|
||||
time.sleep(step)
|
||||
remaining -= step
|
||||
logger.info("Reflection worker stopped")
|
||||
|
||||
self._thread = threading.Thread(target=_run, name="ReflectionWorker", daemon=True)
|
||||
self._thread.start()
|
||||
|
||||
def stop(self):
|
||||
self._stop.set()
|
||||
t = self._thread
|
||||
if t and t.is_alive():
|
||||
try:
|
||||
t.join(timeout=5)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@@ -1,204 +0,0 @@
|
||||
"""
|
||||
Researcher agents.
|
||||
Includes: bull researcher and bear researcher.
|
||||
"""
|
||||
import json
|
||||
from typing import Dict, Any
|
||||
from .base_agent import BaseAgent
|
||||
from app.services.llm import LLMService
|
||||
|
||||
logger = __import__('app.utils.logger', fromlist=['get_logger']).get_logger(__name__)
|
||||
|
||||
|
||||
class BullResearcher(BaseAgent):
|
||||
"""Bullish researcher."""
|
||||
|
||||
def __init__(self, memory=None):
|
||||
super().__init__("BullResearcher", memory)
|
||||
self.llm_service = LLMService()
|
||||
|
||||
def analyze(self, context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Construct the bull case."""
|
||||
market = context.get('market')
|
||||
symbol = context.get('symbol')
|
||||
language = context.get('language', 'zh-CN')
|
||||
model = context.get('model')
|
||||
|
||||
# Inputs
|
||||
market_report = context.get('market_report', {})
|
||||
fundamental_report = context.get('fundamental_report', {})
|
||||
news_report = context.get('news_report', {})
|
||||
sentiment_report = context.get('sentiment_report', {})
|
||||
|
||||
# Memory
|
||||
situation = f"{market}:{symbol} bull case"
|
||||
memory_meta = {
|
||||
"market": market,
|
||||
"symbol": symbol,
|
||||
"timeframe": context.get("timeframe"),
|
||||
"features": context.get("memory_features") or {},
|
||||
}
|
||||
memories = self.get_memories(situation, n_matches=None, metadata=memory_meta)
|
||||
memory_prompt = self.format_memories_for_prompt(memories)
|
||||
|
||||
lang_instruction = self._get_language_instruction(language)
|
||||
system_prompt = f"""You are a Bullish Analyst, constructing a bullish argument for an investment decision. Your tasks are:
|
||||
{lang_instruction}
|
||||
1. Highlight growth potential, competitive advantages, and positive market indicators.
|
||||
2. Use the provided research and data to build a strong argument.
|
||||
3. Effectively address/counter bearish viewpoints.
|
||||
4. Learn from historical experience: {memory_prompt}
|
||||
5. **Confidence Score**: Evaluate your confidence in the bullish case (0-100). Be realistic. If the data is mixed or weak, lower your confidence. Do NOT default to 75.
|
||||
|
||||
Please return in JSON format as follows:
|
||||
{{
|
||||
"argument": "Detailed bullish argument...",
|
||||
"key_points": ["Point 1", "Point 2", "Point 3"],
|
||||
"confidence": 75
|
||||
}}"""
|
||||
|
||||
user_prompt = f"""Based on the following analysis reports, construct a bullish argument for {symbol} in {market} market:
|
||||
|
||||
**Market Technical Analysis:**
|
||||
{json.dumps(market_report.get('data', {}), ensure_ascii=False, indent=2) if market_report else 'No Data'}
|
||||
|
||||
**Fundamental Analysis:**
|
||||
{json.dumps(fundamental_report.get('data', {}), ensure_ascii=False, indent=2) if fundamental_report else 'No Data'}
|
||||
|
||||
**News Analysis:**
|
||||
{json.dumps(news_report.get('data', {}), ensure_ascii=False, indent=2) if news_report else 'No Data'}
|
||||
|
||||
**Sentiment Analysis:**
|
||||
{json.dumps(sentiment_report.get('data', {}), ensure_ascii=False, indent=2) if sentiment_report else 'No Data'}
|
||||
|
||||
Please construct a strong bullish argument, emphasizing growth potential, competitive advantages, and positive indicators."""
|
||||
|
||||
result = self.llm_service.safe_call_llm(
|
||||
system_prompt,
|
||||
user_prompt,
|
||||
{"argument": "", "key_points": [], "confidence": 50},
|
||||
model=model
|
||||
)
|
||||
|
||||
return {
|
||||
"type": "bull",
|
||||
"data": result
|
||||
}
|
||||
|
||||
def _get_language_instruction(self, language: str) -> str:
|
||||
language_map = {
|
||||
'zh-CN': 'Answer in Simplified Chinese.',
|
||||
'zh-TW': 'Answer in Traditional Chinese.',
|
||||
'en-US': 'Answer in English.',
|
||||
'ja-JP': 'Answer in Japanese.',
|
||||
'ko-KR': 'Answer in Korean.',
|
||||
'vi-VN': 'Answer in Vietnamese.',
|
||||
'th-TH': 'Answer in Thai.',
|
||||
'ar-SA': 'Answer in Arabic.',
|
||||
'fr-FR': 'Answer in French.',
|
||||
'de-DE': 'Answer in German.'
|
||||
}
|
||||
return language_map.get(language, 'Answer in English.')
|
||||
|
||||
|
||||
class BearResearcher(BaseAgent):
|
||||
"""Bearish researcher."""
|
||||
|
||||
def __init__(self, memory=None):
|
||||
super().__init__("BearResearcher", memory)
|
||||
self.llm_service = LLMService()
|
||||
|
||||
def analyze(self, context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Construct the bear case."""
|
||||
market = context.get('market')
|
||||
symbol = context.get('symbol')
|
||||
language = context.get('language', 'zh-CN')
|
||||
model = context.get('model')
|
||||
|
||||
# Inputs
|
||||
market_report = context.get('market_report', {})
|
||||
fundamental_report = context.get('fundamental_report', {})
|
||||
news_report = context.get('news_report', {})
|
||||
sentiment_report = context.get('sentiment_report', {})
|
||||
risk_report = context.get('risk_report', {})
|
||||
|
||||
# Bull argument (if present)
|
||||
bull_argument = context.get('bull_argument', '')
|
||||
|
||||
# Memory
|
||||
situation = f"{market}:{symbol} bear case"
|
||||
memory_meta = {
|
||||
"market": market,
|
||||
"symbol": symbol,
|
||||
"timeframe": context.get("timeframe"),
|
||||
"features": context.get("memory_features") or {},
|
||||
}
|
||||
memories = self.get_memories(situation, n_matches=None, metadata=memory_meta)
|
||||
memory_prompt = self.format_memories_for_prompt(memories)
|
||||
|
||||
lang_instruction = self._get_language_instruction(language)
|
||||
system_prompt = f"""You are a Bearish Analyst, constructing a bearish argument for an investment decision. Your tasks are:
|
||||
{lang_instruction}
|
||||
1. Identify risks, challenges, and negative indicators.
|
||||
2. Use the provided research and data to build a strong argument.
|
||||
3. Effectively address/counter bullish viewpoints.
|
||||
4. Learn from historical experience: {memory_prompt}
|
||||
5. **Confidence Score**: Evaluate your confidence in the bearish case (0-100). Be realistic. If the data is mixed or weak, lower your confidence. Do NOT default to 75.
|
||||
|
||||
Please return in JSON format as follows:
|
||||
{{
|
||||
"argument": "Detailed bearish argument...",
|
||||
"key_points": ["Point 1", "Point 2", "Point 3"],
|
||||
"confidence": 75
|
||||
}}"""
|
||||
|
||||
# Bull argument section (avoid backslashes in f-string expression)
|
||||
bull_argument_section = f"**Bullish Argument (Needs Rebuttal):**\n{bull_argument}" if bull_argument else ""
|
||||
|
||||
user_prompt = f"""Based on the following analysis reports, construct a bearish argument for {symbol} in {market} market:
|
||||
|
||||
**Market Technical Analysis:**
|
||||
{json.dumps(market_report.get('data', {}), ensure_ascii=False, indent=2) if market_report else 'No Data'}
|
||||
|
||||
**Fundamental Analysis:**
|
||||
{json.dumps(fundamental_report.get('data', {}), ensure_ascii=False, indent=2) if fundamental_report else 'No Data'}
|
||||
|
||||
**News Analysis:**
|
||||
{json.dumps(news_report.get('data', {}), ensure_ascii=False, indent=2) if news_report else 'No Data'}
|
||||
|
||||
**Sentiment Analysis:**
|
||||
{json.dumps(sentiment_report.get('data', {}), ensure_ascii=False, indent=2) if sentiment_report else 'No Data'}
|
||||
|
||||
**Risk Analysis:**
|
||||
{json.dumps(risk_report.get('data', {}), ensure_ascii=False, indent=2) if risk_report else 'No Data'}
|
||||
|
||||
{bull_argument_section}
|
||||
|
||||
Please construct a strong bearish argument, emphasizing risks, challenges, and negative indicators."""
|
||||
|
||||
result = self.llm_service.safe_call_llm(
|
||||
system_prompt,
|
||||
user_prompt,
|
||||
{"argument": "", "key_points": [], "confidence": 50},
|
||||
model=model
|
||||
)
|
||||
|
||||
return {
|
||||
"type": "bear",
|
||||
"data": result
|
||||
}
|
||||
|
||||
def _get_language_instruction(self, language: str) -> str:
|
||||
language_map = {
|
||||
'zh-CN': 'Answer in Simplified Chinese.',
|
||||
'zh-TW': 'Answer in Traditional Chinese.',
|
||||
'en-US': 'Answer in English.',
|
||||
'ja-JP': 'Answer in Japanese.',
|
||||
'ko-KR': 'Answer in Korean.',
|
||||
'vi-VN': 'Answer in Vietnamese.',
|
||||
'th-TH': 'Answer in Thai.',
|
||||
'ar-SA': 'Answer in Arabic.',
|
||||
'fr-FR': 'Answer in French.',
|
||||
'de-DE': 'Answer in German.'
|
||||
}
|
||||
return language_map.get(language, 'Answer in English.')
|
||||
@@ -1,206 +0,0 @@
|
||||
"""
|
||||
Risk debate agents.
|
||||
Includes: aggressive / neutral / conservative risk analysts.
|
||||
"""
|
||||
import json
|
||||
from typing import Dict, Any
|
||||
from .base_agent import BaseAgent
|
||||
from app.services.llm import LLMService
|
||||
|
||||
logger = __import__('app.utils.logger', fromlist=['get_logger']).get_logger(__name__)
|
||||
|
||||
|
||||
class RiskyAnalyst(BaseAgent):
|
||||
"""Aggressive risk analyst."""
|
||||
|
||||
def __init__(self, memory=None):
|
||||
super().__init__("RiskyAnalyst", memory)
|
||||
self.llm_service = LLMService()
|
||||
|
||||
def analyze(self, context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Analyze risk from an aggressive perspective."""
|
||||
market = context.get('market')
|
||||
symbol = context.get('symbol')
|
||||
language = context.get('language', 'zh-CN')
|
||||
model = context.get('model')
|
||||
trader_plan = context.get('trader_plan', {})
|
||||
|
||||
lang_instruction = self._get_language_instruction(language)
|
||||
system_prompt = f"""You are an Aggressive Risk Analyst. You tend to:
|
||||
{lang_instruction}
|
||||
1. Emphasize high return potential, even with higher risks.
|
||||
2. Believe current risks are controllable and worth taking.
|
||||
3. Support aggressive trading strategies.
|
||||
|
||||
Please return in JSON format as follows:
|
||||
{{
|
||||
"argument": "Aggressive risk analysis argument...",
|
||||
"risk_assessment": "Risk controllable, high return potential",
|
||||
"recommendation": "Support trading plan"
|
||||
}}"""
|
||||
|
||||
user_prompt = f"""Perform aggressive risk analysis for {symbol} in {market} market.
|
||||
|
||||
**Trading Plan:**
|
||||
{json.dumps(trader_plan, ensure_ascii=False, indent=2) if trader_plan else 'No Data'}
|
||||
|
||||
Please analyze risk from an aggressive perspective, emphasizing return potential."""
|
||||
|
||||
result = self.llm_service.safe_call_llm(
|
||||
system_prompt,
|
||||
user_prompt,
|
||||
{"argument": "", "risk_assessment": "", "recommendation": ""},
|
||||
model=model
|
||||
)
|
||||
|
||||
return {
|
||||
"type": "risky",
|
||||
"data": result
|
||||
}
|
||||
|
||||
def _get_language_instruction(self, language: str) -> str:
|
||||
language_map = {
|
||||
'zh-CN': 'Answer in Simplified Chinese.',
|
||||
'zh-TW': 'Answer in Traditional Chinese.',
|
||||
'en-US': 'Answer in English.',
|
||||
'ja-JP': 'Answer in Japanese.',
|
||||
'ko-KR': 'Answer in Korean.',
|
||||
'vi-VN': 'Answer in Vietnamese.',
|
||||
'th-TH': 'Answer in Thai.',
|
||||
'ar-SA': 'Answer in Arabic.',
|
||||
'fr-FR': 'Answer in French.',
|
||||
'de-DE': 'Answer in German.'
|
||||
}
|
||||
return language_map.get(language, 'Answer in English.')
|
||||
|
||||
|
||||
class NeutralAnalyst(BaseAgent):
|
||||
"""Neutral risk analyst."""
|
||||
|
||||
def __init__(self, memory=None):
|
||||
super().__init__("NeutralAnalyst", memory)
|
||||
self.llm_service = LLMService()
|
||||
|
||||
def analyze(self, context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Analyze risk from a neutral perspective."""
|
||||
market = context.get('market')
|
||||
symbol = context.get('symbol')
|
||||
language = context.get('language', 'zh-CN')
|
||||
model = context.get('model')
|
||||
trader_plan = context.get('trader_plan', {})
|
||||
|
||||
lang_instruction = self._get_language_instruction(language)
|
||||
system_prompt = f"""You are a Neutral Risk Analyst. You tend to:
|
||||
{lang_instruction}
|
||||
1. Balance risk and return.
|
||||
2. Objectively evaluate various possibilities.
|
||||
3. Provide neutral risk advice.
|
||||
|
||||
Please return in JSON format as follows:
|
||||
{{
|
||||
"argument": "Neutral risk analysis argument...",
|
||||
"risk_assessment": "Balance between risk and return",
|
||||
"recommendation": "Cautiously execute trading plan"
|
||||
}}"""
|
||||
|
||||
user_prompt = f"""Perform neutral risk analysis for {symbol} in {market} market.
|
||||
|
||||
**Trading Plan:**
|
||||
{json.dumps(trader_plan, ensure_ascii=False, indent=2) if trader_plan else 'No Data'}
|
||||
|
||||
Please analyze risk from a neutral perspective, balancing risk and return."""
|
||||
|
||||
result = self.llm_service.safe_call_llm(
|
||||
system_prompt,
|
||||
user_prompt,
|
||||
{"argument": "", "risk_assessment": "", "recommendation": ""},
|
||||
model=model
|
||||
)
|
||||
|
||||
return {
|
||||
"type": "neutral",
|
||||
"data": result
|
||||
}
|
||||
|
||||
def _get_language_instruction(self, language: str) -> str:
|
||||
language_map = {
|
||||
'zh-CN': 'Answer in Simplified Chinese.',
|
||||
'zh-TW': 'Answer in Traditional Chinese.',
|
||||
'en-US': 'Answer in English.',
|
||||
'ja-JP': 'Answer in Japanese.',
|
||||
'ko-KR': 'Answer in Korean.',
|
||||
'vi-VN': 'Answer in Vietnamese.',
|
||||
'th-TH': 'Answer in Thai.',
|
||||
'ar-SA': 'Answer in Arabic.',
|
||||
'fr-FR': 'Answer in French.',
|
||||
'de-DE': 'Answer in German.'
|
||||
}
|
||||
return language_map.get(language, 'Answer in English.')
|
||||
|
||||
|
||||
class SafeAnalyst(BaseAgent):
|
||||
"""Conservative risk analyst."""
|
||||
|
||||
def __init__(self, memory=None):
|
||||
super().__init__("SafeAnalyst", memory)
|
||||
self.llm_service = LLMService()
|
||||
|
||||
def analyze(self, context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Analyze risk from a conservative perspective."""
|
||||
market = context.get('market')
|
||||
symbol = context.get('symbol')
|
||||
language = context.get('language', 'zh-CN')
|
||||
model = context.get('model')
|
||||
trader_plan = context.get('trader_plan', {})
|
||||
risk_report = context.get('risk_report', {})
|
||||
|
||||
lang_instruction = self._get_language_instruction(language)
|
||||
system_prompt = f"""You are a Conservative Risk Analyst. You tend to:
|
||||
{lang_instruction}
|
||||
1. Emphasize risk control, prioritizing capital protection.
|
||||
2. Identify potential risk points.
|
||||
3. Suggest cautious or conservative trading strategies.
|
||||
|
||||
Please return in JSON format as follows:
|
||||
{{
|
||||
"argument": "Conservative risk analysis argument...",
|
||||
"risk_assessment": "High risk exists, suggest caution",
|
||||
"recommendation": "Suggest reducing position or suspending trading"
|
||||
}}"""
|
||||
|
||||
user_prompt = f"""Perform conservative risk analysis for {symbol} in {market} market.
|
||||
|
||||
**Trading Plan:**
|
||||
{json.dumps(trader_plan, ensure_ascii=False, indent=2) if trader_plan else 'No Data'}
|
||||
|
||||
**Risk Analysis Report:**
|
||||
{json.dumps(risk_report.get('data', {}), ensure_ascii=False, indent=2) if risk_report else 'No Data'}
|
||||
|
||||
Please analyze risk from a conservative perspective, emphasizing risk control."""
|
||||
|
||||
result = self.llm_service.safe_call_llm(
|
||||
system_prompt,
|
||||
user_prompt,
|
||||
{"argument": "", "risk_assessment": "", "recommendation": ""},
|
||||
model=model
|
||||
)
|
||||
|
||||
return {
|
||||
"type": "safe",
|
||||
"data": result
|
||||
}
|
||||
|
||||
def _get_language_instruction(self, language: str) -> str:
|
||||
language_map = {
|
||||
'zh-CN': 'Answer in Simplified Chinese.',
|
||||
'zh-TW': 'Answer in Traditional Chinese.',
|
||||
'en-US': 'Answer in English.',
|
||||
'ja-JP': 'Answer in Japanese.',
|
||||
'ko-KR': 'Answer in Korean.',
|
||||
'vi-VN': 'Answer in Vietnamese.',
|
||||
'th-TH': 'Answer in Thai.',
|
||||
'ar-SA': 'Answer in Arabic.',
|
||||
'fr-FR': 'Answer in French.',
|
||||
'de-DE': 'Answer in German.'
|
||||
}
|
||||
return language_map.get(language, 'Answer in English.')
|
||||
@@ -1,626 +0,0 @@
|
||||
"""
|
||||
Agent tools.
|
||||
|
||||
Provides data fetching helpers for the multi-agent analysis pipeline.
|
||||
All docstrings/log messages in this module are English. Output language of AI reports
|
||||
is controlled by the `language` value passed through the analysis context.
|
||||
"""
|
||||
from typing import Dict, Any, Optional, List
|
||||
from datetime import datetime, timedelta
|
||||
import os
|
||||
import time
|
||||
import pandas as pd
|
||||
import yfinance as yf
|
||||
import finnhub
|
||||
import ccxt
|
||||
import requests
|
||||
|
||||
from app.utils.logger import get_logger
|
||||
from app.config import APIKeys
|
||||
from app.services.search import SearchService
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class AgentTools:
|
||||
"""A thin wrapper around various public data sources used by agents."""
|
||||
|
||||
def __init__(self):
|
||||
self.search_service = SearchService()
|
||||
self.finnhub_client = None
|
||||
if APIKeys.is_configured('FINNHUB_API_KEY'):
|
||||
try:
|
||||
self.finnhub_client = finnhub.Client(api_key=APIKeys.FINNHUB_API_KEY)
|
||||
except Exception as e:
|
||||
# Safe logging to avoid cascading errors during exception handling
|
||||
try:
|
||||
logger.warning(f"Finnhub init failed: {e}")
|
||||
except Exception:
|
||||
# Fallback to print if logging fails
|
||||
print(f"Warning: Finnhub init failed: {e}")
|
||||
|
||||
# Optional dependency: akshare (A-share fundamentals/company info)
|
||||
try:
|
||||
import akshare as ak # type: ignore
|
||||
self._ak = ak
|
||||
self._has_akshare = True
|
||||
except Exception:
|
||||
self._ak = None
|
||||
self._has_akshare = False
|
||||
|
||||
# AShare spot cache (avoid fetching the full market list repeatedly)
|
||||
self._ashare_spot_cache = None
|
||||
self._ashare_spot_cache_ts = 0
|
||||
self._ashare_spot_cache_ttl = 300 # seconds
|
||||
|
||||
def _get_ashare_spot_df(self):
|
||||
"""Cached AShare spot dataframe via akshare (may be heavy on first load)."""
|
||||
if not self._akshare_required():
|
||||
return None
|
||||
now = int(time.time())
|
||||
if self._ashare_spot_cache is not None and (now - int(self._ashare_spot_cache_ts)) < int(self._ashare_spot_cache_ttl):
|
||||
return self._ashare_spot_cache
|
||||
ak = self._ak
|
||||
if ak is None or not hasattr(ak, "stock_zh_a_spot_em"):
|
||||
return None
|
||||
df = ak.stock_zh_a_spot_em()
|
||||
self._ashare_spot_cache = df
|
||||
self._ashare_spot_cache_ts = now
|
||||
return df
|
||||
|
||||
def _ccxt_exchange(self):
|
||||
"""Create a CCXT exchange client (Binance) with optional proxy support."""
|
||||
cfg: Dict[str, Any] = {'timeout': 5000, 'enableRateLimit': True}
|
||||
# Keep proxy behavior consistent with data sources (.env PROXY_* is supported)
|
||||
from app.config import CCXTConfig
|
||||
proxy = (CCXTConfig.PROXY or '').strip()
|
||||
if proxy:
|
||||
cfg['proxies'] = {'http': proxy, 'https': proxy}
|
||||
return ccxt.binance(cfg)
|
||||
|
||||
def _akshare_required(self) -> bool:
|
||||
"""Whether akshare is available at runtime."""
|
||||
return bool(self._has_akshare and self._ak is not None)
|
||||
|
||||
def get_stock_data(self, market: str, symbol: str, days: int = 30, timeframe: str = "1d") -> Optional[List[Dict[str, Any]]]:
|
||||
"""
|
||||
Get daily Kline data for recent days (best-effort).
|
||||
|
||||
Args:
|
||||
market: Market
|
||||
symbol: Symbol
|
||||
days: Days (for daily) / candle count hint (best-effort for intraday)
|
||||
timeframe: Kline timeframe (best-effort). Common values: 1d, 1h, 4h, 1w
|
||||
|
||||
Returns:
|
||||
List of OHLCV dicts or None
|
||||
"""
|
||||
try:
|
||||
klines = []
|
||||
tf = (timeframe or "1d").strip().lower()
|
||||
# Normalize common UI values
|
||||
tf_map = {
|
||||
"1d": "1d",
|
||||
"1day": "1d",
|
||||
"d": "1d",
|
||||
"1h": "1h",
|
||||
"60m": "1h",
|
||||
"4h": "4h",
|
||||
"240m": "4h",
|
||||
"1w": "1wk",
|
||||
"1wk": "1wk",
|
||||
"w": "1wk",
|
||||
}
|
||||
tf_yf = tf_map.get(tf, "1d")
|
||||
|
||||
if market == 'USStock':
|
||||
end_date = datetime.now().strftime('%Y-%m-%d')
|
||||
start_date = (datetime.now() - timedelta(days=days + 5)).strftime('%Y-%m-%d')
|
||||
|
||||
ticker = yf.Ticker(symbol)
|
||||
# yfinance supports limited intervals. Fallback to 1d when unsupported.
|
||||
interval = tf_yf if tf_yf in ["1d", "1h", "1wk"] else "1d"
|
||||
df = ticker.history(start=start_date, end=end_date, interval=interval)
|
||||
|
||||
if not df.empty:
|
||||
df = df.tail(days).reset_index()
|
||||
for _, row in df.iterrows():
|
||||
klines.append({
|
||||
"time": row['Date'].strftime('%Y-%m-%d'),
|
||||
"open": round(row['Open'], 4),
|
||||
"high": round(row['High'], 4),
|
||||
"low": round(row['Low'], 4),
|
||||
"close": round(row['Close'], 4),
|
||||
"volume": int(row['Volume'])
|
||||
})
|
||||
return klines
|
||||
|
||||
elif market == 'Crypto':
|
||||
exchange = self._ccxt_exchange()
|
||||
# Handle symbol format: ETH/USDT -> ETH/USDT, ETH -> ETH/USDT
|
||||
symbol_pair = symbol if '/' in symbol else f'{symbol}/USDT'
|
||||
start_time = int((datetime.now() - timedelta(days=days)).timestamp())
|
||||
# CCXT timeframes: 1d, 1h, 4h ...
|
||||
ccxt_tf = tf if tf in ["1d", "1h", "4h"] else "1d"
|
||||
ohlcv = exchange.fetch_ohlcv(symbol_pair, ccxt_tf, since=start_time * 1000, limit=days)
|
||||
if ohlcv:
|
||||
for candle in ohlcv:
|
||||
klines.append({
|
||||
"time": datetime.fromtimestamp(candle[0] / 1000).strftime('%Y-%m-%d'),
|
||||
"open": candle[1],
|
||||
"high": candle[2],
|
||||
"low": candle[3],
|
||||
"close": candle[4],
|
||||
"volume": candle[5]
|
||||
})
|
||||
return klines
|
||||
|
||||
# CN/HK stocks
|
||||
if market in ('AShare', 'HShare'):
|
||||
# Prefer akshare for AShare (requested), fall back to yfinance.
|
||||
if market == 'AShare' and self._akshare_required():
|
||||
try:
|
||||
ak = self._ak
|
||||
start_date = (datetime.now() - timedelta(days=days + 10)).strftime('%Y%m%d')
|
||||
end_date = datetime.now().strftime('%Y%m%d')
|
||||
# akshare returns a dataframe with Chinese column names.
|
||||
df = ak.stock_zh_a_hist(symbol=symbol, period="daily", start_date=start_date, end_date=end_date, adjust="qfq")
|
||||
if df is not None and not df.empty:
|
||||
df = df.tail(days)
|
||||
for _, row in df.iterrows():
|
||||
dt = row.get('日期')
|
||||
# dt can be datetime/date/str
|
||||
if hasattr(dt, "strftime"):
|
||||
t = dt.strftime('%Y-%m-%d')
|
||||
else:
|
||||
t = str(dt)[:10]
|
||||
klines.append({
|
||||
"time": t,
|
||||
"open": float(row.get('开盘', 0) or 0),
|
||||
"high": float(row.get('最高', 0) or 0),
|
||||
"low": float(row.get('最低', 0) or 0),
|
||||
"close": float(row.get('收盘', 0) or 0),
|
||||
"volume": float(row.get('成交量', 0) or 0),
|
||||
})
|
||||
return klines
|
||||
except Exception as e:
|
||||
logger.warning(f"akshare AShare kline failed ({symbol}): {e}")
|
||||
|
||||
# yfinance fallback (daily)
|
||||
if market == 'AShare':
|
||||
yf_symbol = f"{symbol}.SS" if symbol.startswith('6') else f"{symbol}.SZ"
|
||||
else:
|
||||
yf_symbol = f"{symbol.zfill(4)}.HK"
|
||||
|
||||
end_date = datetime.now().strftime('%Y-%m-%d')
|
||||
start_date = (datetime.now() - timedelta(days=days + 5)).strftime('%Y-%m-%d')
|
||||
|
||||
ticker = yf.Ticker(yf_symbol)
|
||||
interval = tf_yf if tf_yf in ["1d", "1h", "1wk"] else "1d"
|
||||
df = ticker.history(start=start_date, end=end_date, interval=interval)
|
||||
|
||||
if not df.empty:
|
||||
df = df.tail(days).reset_index()
|
||||
for _, row in df.iterrows():
|
||||
klines.append({
|
||||
"time": row['Date'].strftime('%Y-%m-%d'),
|
||||
"open": round(row['Open'], 4),
|
||||
"high": round(row['High'], 4),
|
||||
"low": round(row['Low'], 4),
|
||||
"close": round(row['Close'], 4),
|
||||
"volume": int(row['Volume'])
|
||||
})
|
||||
return klines
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to fetch kline data {market}:{symbol}: {e}")
|
||||
|
||||
return None
|
||||
|
||||
def get_current_price(self, market: str, symbol: str) -> Optional[Dict[str, Any]]:
|
||||
"""Get current price (best-effort)."""
|
||||
try:
|
||||
if market == 'USStock' and self.finnhub_client:
|
||||
quote = self.finnhub_client.quote(symbol)
|
||||
if quote and quote.get('c'):
|
||||
return {
|
||||
"price": quote.get('c', 0),
|
||||
"change": quote.get('d', 0),
|
||||
"changePercent": quote.get('dp', 0),
|
||||
"high": quote.get('h', 0),
|
||||
"low": quote.get('l', 0),
|
||||
"open": quote.get('o', 0),
|
||||
"previousClose": quote.get('pc', 0)
|
||||
}
|
||||
elif market == 'Crypto':
|
||||
exchange = self._ccxt_exchange()
|
||||
# Handle symbol format: ETH/USDT -> ETH/USDT, ETH -> ETH/USDT
|
||||
symbol_pair = symbol if '/' in symbol else f'{symbol}/USDT'
|
||||
ticker = exchange.fetch_ticker(symbol_pair)
|
||||
if ticker:
|
||||
return {
|
||||
"price": ticker.get('last', 0),
|
||||
"change": ticker.get('change', 0),
|
||||
"changePercent": ticker.get('percentage', 0),
|
||||
"high": ticker.get('high', 0),
|
||||
"low": ticker.get('low', 0),
|
||||
"open": ticker.get('open', 0),
|
||||
"volume": ticker.get('quoteVolume', 0)
|
||||
}
|
||||
|
||||
# CN/HK stocks: prefer akshare for AShare (requested)
|
||||
if market in ('AShare', 'HShare'):
|
||||
if market == 'AShare' and self._akshare_required():
|
||||
try:
|
||||
ak = self._ak
|
||||
df = self._get_ashare_spot_df()
|
||||
if df is not None and not df.empty:
|
||||
row = df[df['代码'] == symbol].iloc[0]
|
||||
price = float(row.get('最新价', 0) or 0)
|
||||
change = float(row.get('涨跌额', 0) or 0)
|
||||
change_pct = float(row.get('涨跌幅', 0) or 0)
|
||||
high = float(row.get('最高', 0) or 0)
|
||||
low = float(row.get('最低', 0) or 0)
|
||||
open_p = float(row.get('今开', 0) or 0)
|
||||
prev_close = float(row.get('昨收', 0) or 0)
|
||||
return {
|
||||
"price": price,
|
||||
"change": change,
|
||||
"changePercent": change_pct,
|
||||
"high": high,
|
||||
"low": low,
|
||||
"open": open_p,
|
||||
"previousClose": prev_close
|
||||
}
|
||||
except Exception as e:
|
||||
logger.warning(f"akshare AShare spot failed ({symbol}): {e}")
|
||||
|
||||
# Do not use Tencent for AShare by default (requested). If akshare is not available,
|
||||
# return None and let the LLM report degrade gracefully.
|
||||
if market == 'AShare':
|
||||
if not self._akshare_required():
|
||||
logger.warning("akshare is not installed; AShare spot price is unavailable.")
|
||||
return None
|
||||
|
||||
# HShare fallback: Tencent quote
|
||||
symbol_code = f'hk{symbol}'
|
||||
|
||||
url = f"http://qt.gtimg.cn/q={symbol_code}"
|
||||
resp = requests.get(url, timeout=10)
|
||||
content = resp.content.decode('gbk', errors='ignore')
|
||||
if '="' in content:
|
||||
data_str = content.split('="')[1].strip('";\n')
|
||||
if data_str:
|
||||
parts = data_str.split('~')
|
||||
if len(parts) > 32:
|
||||
return {
|
||||
"price": float(parts[3]) if parts[3] else 0,
|
||||
"change": float(parts[31]) if parts[31] else 0,
|
||||
"changePercent": float(parts[32]) if parts[32] else 0,
|
||||
"high": float(parts[33]) if len(parts) > 33 and parts[33] else 0,
|
||||
"low": float(parts[34]) if len(parts) > 34 and parts[34] else 0,
|
||||
"open": float(parts[5]) if len(parts) > 5 and parts[5] else 0,
|
||||
"previousClose": float(parts[4]) if parts[4] else 0
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to fetch current price {market}:{symbol}: {e}")
|
||||
|
||||
return None
|
||||
|
||||
def get_fundamental_data(self, market: str, symbol: str) -> Optional[Dict[str, Any]]:
|
||||
"""Get fundamental data (best-effort)."""
|
||||
try:
|
||||
if market == 'USStock' and self.finnhub_client:
|
||||
metrics = self.finnhub_client.company_basic_financials(symbol, 'all')
|
||||
profile = self.finnhub_client.company_profile2(symbol=symbol)
|
||||
|
||||
return {
|
||||
"metrics": metrics.get('metric', {}),
|
||||
"profile_metrics": {
|
||||
"marketCapitalization": profile.get('marketCapitalization', 0),
|
||||
"currency": profile.get('currency', 'USD'),
|
||||
"finnhubIndustry": profile.get('finnhubIndustry', ''),
|
||||
}
|
||||
}
|
||||
|
||||
# AShare fundamentals via akshare (requested)
|
||||
if market == 'AShare' and self._akshare_required():
|
||||
ak = self._ak
|
||||
out: Dict[str, Any] = {"metrics": {}, "profile_metrics": {}}
|
||||
|
||||
# 1) Use spot list (fast) for valuation/market cap
|
||||
try:
|
||||
df = self._get_ashare_spot_df()
|
||||
if df is not None and not df.empty:
|
||||
row = df[df['代码'] == symbol].iloc[0]
|
||||
out["metrics"].update({
|
||||
"pe_ttm": row.get('市盈率-动态'),
|
||||
"pb": row.get('市净率'),
|
||||
"turnoverRate": row.get('换手率'),
|
||||
})
|
||||
out["profile_metrics"].update({
|
||||
"marketCapitalization": row.get('总市值'),
|
||||
"floatMarketCap": row.get('流通市值'),
|
||||
"currency": "CNY",
|
||||
})
|
||||
except Exception as e:
|
||||
logger.debug(f"akshare spot metrics unavailable ({symbol}): {e}")
|
||||
|
||||
# 2) Try akshare indicator endpoints (optional, may be slower / may change)
|
||||
try:
|
||||
if hasattr(ak, "stock_a_lg_indicator"):
|
||||
ind_df = ak.stock_a_lg_indicator(symbol=symbol)
|
||||
if ind_df is not None and not ind_df.empty:
|
||||
last = ind_df.iloc[-1].to_dict()
|
||||
out["metrics"].update(last)
|
||||
except Exception as e:
|
||||
logger.debug(f"akshare indicator fetch failed ({symbol}): {e}")
|
||||
|
||||
return out
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to fetch fundamental data {market}:{symbol}: {e}")
|
||||
|
||||
return None
|
||||
|
||||
def get_company_data(self, market: str, symbol: str, language: str = "en-US") -> Optional[Dict[str, Any]]:
|
||||
"""Get basic company/project info (best-effort)."""
|
||||
try:
|
||||
# 1) Finnhub (mainly for US stocks)
|
||||
if market == 'USStock' and self.finnhub_client:
|
||||
profile = self.finnhub_client.company_profile2(symbol=symbol)
|
||||
if profile:
|
||||
return {
|
||||
"name": profile.get('name', symbol),
|
||||
"ticker": profile.get('ticker', symbol),
|
||||
"exchange": profile.get('exchange', ''),
|
||||
"industry": profile.get('finnhubIndustry', ''),
|
||||
"website": profile.get('weburl', ''),
|
||||
"marketCapitalization": profile.get('marketCapitalization', 0),
|
||||
"description": f"Sector: {profile.get('finnhubIndustry', '')}, Country: {profile.get('country', '')}"
|
||||
}
|
||||
|
||||
# 2) Basic info for AShare / HShare / Crypto
|
||||
elif market in ('AShare', 'HShare', 'Crypto'):
|
||||
name = symbol
|
||||
if market == 'AShare':
|
||||
# Prefer akshare for AShare (requested)
|
||||
if self._akshare_required():
|
||||
try:
|
||||
ak = self._ak
|
||||
# 1) Individual info (more structured)
|
||||
if hasattr(ak, "stock_individual_info_em"):
|
||||
df = ak.stock_individual_info_em(symbol=symbol)
|
||||
if df is not None and not df.empty and 'item' in df.columns and 'value' in df.columns:
|
||||
info = {str(r['item']).strip(): r['value'] for _, r in df.iterrows()}
|
||||
# common keys: 股票简称, 所属行业, 上市时间, 总市值 ...
|
||||
name = str(info.get('股票简称') or info.get('证券简称') or symbol).strip()
|
||||
industry = str(info.get('所属行业') or '').strip()
|
||||
website = str(info.get('公司网址') or '').strip()
|
||||
market_cap = info.get('总市值') or info.get('总市值(元)') or 0
|
||||
return {
|
||||
"name": name or symbol,
|
||||
"ticker": symbol,
|
||||
"market": market,
|
||||
"industry": industry,
|
||||
"website": website,
|
||||
"marketCapitalization": market_cap,
|
||||
"description": f"Industry: {industry}" if industry else ""
|
||||
}
|
||||
# 2) Spot list for name
|
||||
df2 = ak.stock_zh_a_spot_em()
|
||||
if df2 is not None and not df2.empty:
|
||||
row = df2[df2['代码'] == symbol].iloc[0]
|
||||
name = str(row.get('名称') or symbol).strip()
|
||||
except Exception as e:
|
||||
logger.debug(f"akshare company info failed ({symbol}): {e}")
|
||||
|
||||
# Do not use Tencent for AShare by default (requested).
|
||||
if not self._akshare_required():
|
||||
logger.warning("akshare is not installed; AShare company info is limited.")
|
||||
elif market == 'Crypto':
|
||||
name = f"{symbol} Cryptocurrency"
|
||||
|
||||
# Enrich description via web search (best-effort)
|
||||
# Query language should follow UI language when possible.
|
||||
if str(language).lower().startswith('zh'):
|
||||
search_query = f"{name} {symbol} 公司 简介" if market != 'Crypto' else f"{symbol} 加密 项目 介绍"
|
||||
else:
|
||||
search_query = f"{name} {symbol} company profile" if market != 'Crypto' else f"{symbol} crypto project info"
|
||||
search_results = self.search_service.search(search_query, num_results=1)
|
||||
description = ""
|
||||
if search_results:
|
||||
description = search_results[0].get('snippet', '')
|
||||
|
||||
return {
|
||||
"name": name,
|
||||
"ticker": symbol,
|
||||
"market": market,
|
||||
"description": description
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to fetch company data {market}:{symbol}: {e}")
|
||||
|
||||
return None
|
||||
|
||||
def _fetch_page_content(self, url: str) -> str:
|
||||
"""
|
||||
Fetch readable page content via Jina Reader.
|
||||
|
||||
Args:
|
||||
url: Target URL
|
||||
|
||||
Returns:
|
||||
Extracted content (markdown-ish), truncated
|
||||
"""
|
||||
try:
|
||||
jina_url = f"https://r.jina.ai/{url}"
|
||||
# Use a slightly longer timeout for content extraction
|
||||
response = requests.get(jina_url, timeout=15)
|
||||
if response.status_code == 200:
|
||||
content = response.text
|
||||
# Truncate to avoid huge prompts
|
||||
if len(content) > 3000:
|
||||
content = content[:3000] + "..."
|
||||
return content
|
||||
return ""
|
||||
except Exception as e:
|
||||
logger.warning(f"Jina Reader content fetch failed {url}: {e}")
|
||||
return ""
|
||||
|
||||
def get_news(self, market: str, symbol: str, days: int = 7, company_name: str = None) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Get news items (Finnhub + search engine) and optionally enrich via Jina Reader.
|
||||
|
||||
Args:
|
||||
market: Market
|
||||
symbol: Symbol/pair
|
||||
days: Lookback days
|
||||
company_name: Optional company/project name to improve search
|
||||
|
||||
Returns:
|
||||
List of news items
|
||||
"""
|
||||
news_list = []
|
||||
|
||||
# 1) Finnhub news (if available)
|
||||
try:
|
||||
if self.finnhub_client:
|
||||
end_date = datetime.now().strftime('%Y-%m-%d')
|
||||
start_date = (datetime.now() - timedelta(days=days)).strftime('%Y-%m-%d')
|
||||
|
||||
raw_news = []
|
||||
|
||||
if market == 'USStock':
|
||||
raw_news = self.finnhub_client.company_news(symbol, _from=start_date, to=end_date)
|
||||
elif market == 'Crypto':
|
||||
crypto_symbol = symbol.split('/')[0] if '/' in symbol else symbol
|
||||
raw_news = self.finnhub_client.crypto_news(crypto_symbol)
|
||||
else:
|
||||
raw_news = self.finnhub_client.general_news('general', min_id=0)
|
||||
|
||||
if raw_news:
|
||||
for item in raw_news:
|
||||
if not item.get('headline') or not item.get('summary'):
|
||||
continue
|
||||
news_list.append({
|
||||
"id": str(item.get('id', '')),
|
||||
"datetime": datetime.fromtimestamp(item.get('datetime', 0)).strftime('%Y-%m-%d %H:%M'),
|
||||
"headline": item.get('headline', ''),
|
||||
"summary": item.get('summary', ''),
|
||||
"source": f"Finnhub ({item.get('source', '')})",
|
||||
"url": item.get('url', '')
|
||||
})
|
||||
except Exception as e:
|
||||
logger.warning(f"Finnhub news fetch failed: {e}")
|
||||
|
||||
# 2) Supplement with search engine results (useful for non-US markets or specific events)
|
||||
try:
|
||||
# Build search query (use company name to improve relevance)
|
||||
search_query = ""
|
||||
search_name = company_name if company_name else symbol
|
||||
|
||||
# Time restriction for Google CSE
|
||||
date_restrict = f"d{days}"
|
||||
|
||||
if market == 'AShare':
|
||||
# AShare CN keywords
|
||||
search_query = f'{search_name} {symbol} (利好 OR 利空 OR 财报 OR 公告 OR 业绩) after:{datetime.now().year-1}'
|
||||
elif market == 'HShare':
|
||||
search_query = f'{search_name} {symbol} (港股 OR 股价 OR 业绩) after:{datetime.now().year-1}'
|
||||
elif market == 'Crypto':
|
||||
search_query = f'{search_name} {symbol} crypto news analysis'
|
||||
else:
|
||||
search_query = f'{search_name} {symbol} stock news'
|
||||
|
||||
logger.info(f"Running news search: {search_query}")
|
||||
# Google CSE uses `dateRestrict` as a separate param; SearchService supports it.
|
||||
search_results = self.search_service.search(search_query, num_results=10, date_restrict=date_restrict)
|
||||
|
||||
for i, item in enumerate(search_results):
|
||||
# Default: use snippet as summary
|
||||
summary = f"{item.get('snippet', '')} (Source: {item.get('source', '')})"
|
||||
|
||||
# Jina Reader: deep-read only first 2 items to avoid slowdowns
|
||||
if i < 2 and item.get('link'):
|
||||
logger.info(f"Deep reading: {item.get('title')}")
|
||||
full_content = self._fetch_page_content(item.get('link'))
|
||||
if full_content:
|
||||
summary = f"Deep content:\n{full_content}\n(Source: {item.get('source', '')})"
|
||||
|
||||
news_list.append({
|
||||
"id": item.get('link', ''), # Use link as a stable id
|
||||
"datetime": item.get('published', datetime.now().strftime('%Y-%m-%d')), # Fallback to today if missing
|
||||
"headline": item.get('title', ''),
|
||||
"summary": summary,
|
||||
"source": f"Search ({item.get('source', '')})",
|
||||
"url": item.get('link', '')
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Search news failed: {e}")
|
||||
|
||||
# Sort by time desc and keep latest items (best-effort; time formats may vary)
|
||||
news_list.sort(key=lambda x: x.get('datetime', ''), reverse=True)
|
||||
return news_list[:20]
|
||||
|
||||
def calculate_technical_indicators(self, kline_data: List[Dict[str, Any]]) -> Dict[str, Any]:
|
||||
"""
|
||||
Calculate basic technical indicators from kline data.
|
||||
|
||||
Args:
|
||||
kline_data: List of OHLCV dicts
|
||||
|
||||
Returns:
|
||||
Indicators dict
|
||||
"""
|
||||
if not kline_data or len(kline_data) < 20:
|
||||
return {}
|
||||
|
||||
try:
|
||||
df = pd.DataFrame(kline_data)
|
||||
df['close'] = pd.to_numeric(df['close'], errors='coerce')
|
||||
df['high'] = pd.to_numeric(df['high'], errors='coerce')
|
||||
df['low'] = pd.to_numeric(df['low'], errors='coerce')
|
||||
df['volume'] = pd.to_numeric(df['volume'], errors='coerce')
|
||||
|
||||
indicators = {}
|
||||
|
||||
# Moving averages
|
||||
if len(df) >= 20:
|
||||
indicators['MA20'] = round(df['close'].tail(20).mean(), 4)
|
||||
if len(df) >= 50:
|
||||
indicators['MA50'] = round(df['close'].tail(50).mean(), 4)
|
||||
|
||||
# RSI
|
||||
if len(df) >= 14:
|
||||
delta = df['close'].diff()
|
||||
gain = (delta.where(delta > 0, 0)).rolling(window=14).mean()
|
||||
loss = (-delta.where(delta < 0, 0)).rolling(window=14).mean()
|
||||
rs = gain / loss
|
||||
rsi = 100 - (100 / (1 + rs))
|
||||
indicators['RSI'] = round(rsi.iloc[-1], 2) if not rsi.empty else None
|
||||
|
||||
# MACD
|
||||
if len(df) >= 26:
|
||||
exp1 = df['close'].ewm(span=12, adjust=False).mean()
|
||||
exp2 = df['close'].ewm(span=26, adjust=False).mean()
|
||||
macd = exp1 - exp2
|
||||
signal = macd.ewm(span=9, adjust=False).mean()
|
||||
indicators['MACD'] = round(macd.iloc[-1], 4) if not macd.empty else None
|
||||
indicators['MACD_Signal'] = round(signal.iloc[-1], 4) if not signal.empty else None
|
||||
indicators['MACD_Histogram'] = round((macd - signal).iloc[-1], 4) if not (macd - signal).empty else None
|
||||
|
||||
# Bollinger bands
|
||||
if len(df) >= 20:
|
||||
sma = df['close'].rolling(window=20).mean()
|
||||
std = df['close'].rolling(window=20).std()
|
||||
indicators['BB_Upper'] = round((sma + 2 * std).iloc[-1], 4) if not sma.empty else None
|
||||
indicators['BB_Middle'] = round(sma.iloc[-1], 4) if not sma.empty else None
|
||||
indicators['BB_Lower'] = round((sma - 2 * std).iloc[-1], 4) if not sma.empty else None
|
||||
|
||||
return indicators
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to calculate technical indicators: {e}")
|
||||
return {}
|
||||
@@ -1,135 +0,0 @@
|
||||
"""
|
||||
Trader agent.
|
||||
Synthesizes all analysis outputs and produces a final trading decision.
|
||||
"""
|
||||
import json
|
||||
from typing import Dict, Any
|
||||
from .base_agent import BaseAgent
|
||||
from app.services.llm import LLMService
|
||||
|
||||
logger = __import__('app.utils.logger', fromlist=['get_logger']).get_logger(__name__)
|
||||
|
||||
|
||||
class TraderAgent(BaseAgent):
|
||||
"""Trader agent."""
|
||||
|
||||
def __init__(self, memory=None):
|
||||
super().__init__("TraderAgent", memory)
|
||||
self.llm_service = LLMService()
|
||||
|
||||
def analyze(self, context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Make a final trading decision."""
|
||||
market = context.get('market')
|
||||
symbol = context.get('symbol')
|
||||
language = context.get('language', 'zh-CN')
|
||||
model = context.get('model')
|
||||
|
||||
# Inputs
|
||||
market_report = context.get('market_report', {})
|
||||
fundamental_report = context.get('fundamental_report', {})
|
||||
news_report = context.get('news_report', {})
|
||||
sentiment_report = context.get('sentiment_report', {})
|
||||
risk_report = context.get('risk_report', {})
|
||||
|
||||
# Debate outputs
|
||||
bull_argument = context.get('bull_argument', {})
|
||||
bear_argument = context.get('bear_argument', {})
|
||||
research_decision = context.get('research_decision', '')
|
||||
|
||||
# Memory
|
||||
situation = f"{market}:{symbol} trading decision"
|
||||
memory_meta = {
|
||||
"market": market,
|
||||
"symbol": symbol,
|
||||
"timeframe": context.get("timeframe"),
|
||||
"features": context.get("memory_features") or {},
|
||||
}
|
||||
memories = self.get_memories(situation, n_matches=None, metadata=memory_meta)
|
||||
memory_prompt = self.format_memories_for_prompt(memories)
|
||||
|
||||
lang_instruction = self._get_language_instruction(language)
|
||||
system_prompt = f"""You are a Trader, needing to make a final trading decision based on all analysis results.
|
||||
{lang_instruction}
|
||||
|
||||
Your tasks:
|
||||
1. Synthesize analysis results from all dimensions.
|
||||
2. Consider both bullish and bearish arguments.
|
||||
3. Make a clear trading decision: BUY, SELL, or HOLD.
|
||||
4. Provide a detailed trading plan.
|
||||
5. Learn from historical experience: {memory_prompt}
|
||||
6. **Confidence Score**: Evaluate your confidence in the decision (0-100). Be realistic. If the signals are mixed, confidence should be lower (e.g., 40-60). Only use high confidence (>80) for very clear strong signals. Do NOT default to 85.
|
||||
|
||||
Please return in JSON format as follows:
|
||||
{{
|
||||
"decision": "BUY/SELL/HOLD",
|
||||
"confidence": 85,
|
||||
"reasoning": "Reason for decision...",
|
||||
"trading_plan": {{
|
||||
"entry_price": "Suggested entry price",
|
||||
"stop_loss": "Stop loss price",
|
||||
"take_profit": "Take profit price",
|
||||
"position_size": "Suggested position size"
|
||||
}},
|
||||
"report": "Detailed trading plan report..."
|
||||
}}"""
|
||||
|
||||
user_prompt = f"""Based on all the following analyses, make a trading decision for {symbol} in {market} market:
|
||||
|
||||
**Market Technical Analysis:**
|
||||
{json.dumps(market_report.get('data', {}), ensure_ascii=False, indent=2) if market_report else 'No Data'}
|
||||
|
||||
**Fundamental Analysis:**
|
||||
{json.dumps(fundamental_report.get('data', {}), ensure_ascii=False, indent=2) if fundamental_report else 'No Data'}
|
||||
|
||||
**News Analysis:**
|
||||
{json.dumps(news_report.get('data', {}), ensure_ascii=False, indent=2) if news_report else 'No Data'}
|
||||
|
||||
**Sentiment Analysis:**
|
||||
{json.dumps(sentiment_report.get('data', {}), ensure_ascii=False, indent=2) if sentiment_report else 'No Data'}
|
||||
|
||||
**Risk Analysis:**
|
||||
{json.dumps(risk_report.get('data', {}), ensure_ascii=False, indent=2) if risk_report else 'No Data'}
|
||||
|
||||
**Bullish Argument:**
|
||||
{json.dumps(bull_argument.get('data', {}), ensure_ascii=False, indent=2) if bull_argument else 'No Data'}
|
||||
|
||||
**Bearish Argument:**
|
||||
{json.dumps(bear_argument.get('data', {}), ensure_ascii=False, indent=2) if bear_argument else 'No Data'}
|
||||
|
||||
**Research Manager Decision:**
|
||||
{research_decision if research_decision else 'No Data'}
|
||||
|
||||
Please make a clear trading decision (BUY/SELL/HOLD) and provide a detailed trading plan."""
|
||||
|
||||
result = self.llm_service.safe_call_llm(
|
||||
system_prompt,
|
||||
user_prompt,
|
||||
{
|
||||
"decision": "HOLD",
|
||||
"confidence": 50,
|
||||
"reasoning": "",
|
||||
"trading_plan": {},
|
||||
"report": "Failed to parse trader decision"
|
||||
},
|
||||
model=model
|
||||
)
|
||||
|
||||
return {
|
||||
"type": "trader",
|
||||
"data": result
|
||||
}
|
||||
|
||||
def _get_language_instruction(self, language: str) -> str:
|
||||
language_map = {
|
||||
'zh-CN': 'Answer in Simplified Chinese.',
|
||||
'zh-TW': 'Answer in Traditional Chinese.',
|
||||
'en-US': 'Answer in English.',
|
||||
'ja-JP': 'Answer in Japanese.',
|
||||
'ko-KR': 'Answer in Korean.',
|
||||
'vi-VN': 'Answer in Vietnamese.',
|
||||
'th-TH': 'Answer in Thai.',
|
||||
'ar-SA': 'Answer in Arabic.',
|
||||
'fr-FR': 'Answer in French.',
|
||||
'de-DE': 'Answer in German.'
|
||||
}
|
||||
return language_map.get(language, 'Answer in English.')
|
||||
@@ -1,170 +0,0 @@
|
||||
"""
|
||||
Multi-dimensional analysis service.
|
||||
Uses OpenRouter via the internal LLMService and the multi-agent coordinator.
|
||||
Local-only: this project does not implement any paid/credit system itself.
|
||||
"""
|
||||
import json
|
||||
import traceback
|
||||
from typing import Dict, Any, Optional
|
||||
|
||||
from app.utils.logger import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class AnalysisService:
|
||||
"""Multi-dimensional analyzer powered by agent coordinator."""
|
||||
|
||||
# Class-level guard to avoid circular-init recursion
|
||||
_initializing = False
|
||||
|
||||
def __init__(self, use_multi_agent: bool = None):
|
||||
"""
|
||||
Args:
|
||||
use_multi_agent: Deprecated; kept for frontend compatibility
|
||||
"""
|
||||
# Avoid circular-init recursion
|
||||
if AnalysisService._initializing:
|
||||
logger.warning("AnalysisService is initializing; skipping duplicate initialization")
|
||||
self.coordinator = None
|
||||
return
|
||||
|
||||
self.coordinator = None
|
||||
|
||||
try:
|
||||
# Mark initializing
|
||||
AnalysisService._initializing = True
|
||||
|
||||
# Lazy import to avoid circular imports
|
||||
from app.services.agents.coordinator import AgentCoordinator
|
||||
import os
|
||||
enable_memory = os.getenv("ENABLE_AGENT_MEMORY", "true").lower() == "true"
|
||||
self.coordinator = AgentCoordinator(
|
||||
enable_memory=enable_memory,
|
||||
max_debate_rounds=2
|
||||
)
|
||||
logger.info("Multi-agent coordinator initialized")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Coordinator init failed: {e}")
|
||||
logger.error(f"Traceback: {traceback.format_exc()}")
|
||||
self.coordinator = None
|
||||
finally:
|
||||
AnalysisService._initializing = False
|
||||
|
||||
def analyze(self, market: str, symbol: str, language: str = 'en-US', model: str = None, timeframe: str = "1D") -> Dict[str, Any]:
|
||||
"""
|
||||
Args:
|
||||
market: Market (AShare, USStock, HShare, Crypto, Forex, Futures)
|
||||
symbol: Symbol
|
||||
language: Output language tag (e.g. en-US, zh-CN, zh-TW)
|
||||
model: Optional OpenRouter model id
|
||||
Returns:
|
||||
Result dict
|
||||
"""
|
||||
logger.info(f"Starting analysis {market}:{symbol}, language={language}, mode=multi-agent")
|
||||
|
||||
# Default result structure (keeps frontend compatible even when coordinator fails).
|
||||
result = {
|
||||
"overview": {"report": "Initializing..."},
|
||||
"fundamental": {"report": "Initializing..."},
|
||||
"technical": {"report": "Initializing..."},
|
||||
"news": {"report": "Initializing..."},
|
||||
"sentiment": {"report": "Initializing..."},
|
||||
"risk": {"report": "Initializing..."},
|
||||
"error": None
|
||||
}
|
||||
|
||||
if not self.coordinator:
|
||||
result["error"] = "Analysis service is not ready (coordinator init failed)"
|
||||
return result
|
||||
|
||||
try:
|
||||
logger.info(f"Run coordinator: {market}:{symbol}")
|
||||
agent_result = self.coordinator.run_analysis(market, symbol, language, model=model, timeframe=timeframe)
|
||||
|
||||
logger.info(f"Coordinator result keys: {list(agent_result.keys())}")
|
||||
|
||||
# Validate expected keys (defensive)
|
||||
debate = agent_result.get("debate", {})
|
||||
trader_decision = agent_result.get("trader_decision", {})
|
||||
risk_debate = agent_result.get("risk_debate", {})
|
||||
final_decision = agent_result.get("final_decision", {})
|
||||
|
||||
# Keep frontend-compatible shape and fill defaults if empty
|
||||
if "debate" in agent_result and "trader_decision" in agent_result and "risk_debate" in agent_result and "final_decision" in agent_result:
|
||||
if not debate or (isinstance(debate, dict) and len(debate) == 0):
|
||||
logger.warning("debate is empty; using defaults")
|
||||
agent_result["debate"] = {"bull": {}, "bear": {}, "research_decision": "Analyzing..."}
|
||||
if not trader_decision or (isinstance(trader_decision, dict) and len(trader_decision) == 0):
|
||||
logger.warning("trader_decision is empty; using defaults")
|
||||
agent_result["trader_decision"] = {"decision": "HOLD", "confidence": 50, "reasoning": "Analyzing..."}
|
||||
if not risk_debate or (isinstance(risk_debate, dict) and len(risk_debate) == 0):
|
||||
logger.warning("risk_debate is empty; using defaults")
|
||||
agent_result["risk_debate"] = {"risky": {}, "neutral": {}, "safe": {}}
|
||||
if not final_decision or (isinstance(final_decision, dict) and len(final_decision) == 0):
|
||||
logger.warning("final_decision is empty; using defaults")
|
||||
agent_result["final_decision"] = {"decision": "HOLD", "confidence": 50, "reasoning": "Analyzing..."}
|
||||
|
||||
return agent_result
|
||||
else:
|
||||
logger.warning("Coordinator result format is incomplete; filling defaults")
|
||||
return {
|
||||
"overview": agent_result.get("overview", {"report": "Analyzing..."}),
|
||||
"fundamental": agent_result.get("fundamental", {"report": "Analyzing..."}),
|
||||
"technical": agent_result.get("technical", {"report": "Analyzing..."}),
|
||||
"news": agent_result.get("news", {"report": "Analyzing..."}),
|
||||
"sentiment": agent_result.get("sentiment", {"report": "Analyzing..."}),
|
||||
"risk": agent_result.get("risk", {"report": "Analyzing..."}),
|
||||
"debate": agent_result.get("debate", {"bull": {}, "bear": {}, "research_decision": "Analyzing..."}),
|
||||
"trader_decision": agent_result.get("trader_decision", {"decision": "HOLD", "confidence": 50, "reasoning": "Analyzing..."}),
|
||||
"risk_debate": agent_result.get("risk_debate", {"risky": {}, "neutral": {}, "safe": {}}),
|
||||
"final_decision": agent_result.get("final_decision", {"decision": "HOLD", "confidence": 50, "reasoning": "Analyzing..."}),
|
||||
"error": agent_result.get("error")
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
error_msg = str(e)
|
||||
logger.error(f"Analysis failed {market}:{symbol} - {error_msg}")
|
||||
|
||||
# If OpenRouter returns 402, it's an upstream billing/credit issue (not a QuantDinger fee).
|
||||
if "402" in error_msg or "Payment Required" in error_msg:
|
||||
result["error"] = f"OpenRouter returned 402 (billing/credits). Please check your OpenRouter account. Details: {error_msg}"
|
||||
else:
|
||||
result["error"] = f"Analysis failed: {error_msg}"
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def multi_analysis(market: str, symbol: str, language: str = 'en-US', use_multi_agent: bool = None) -> Dict[str, Any]:
|
||||
"""
|
||||
Convenience entrypoint for multi-dimensional analysis.
|
||||
|
||||
Args:
|
||||
market: Market (AShare, USStock, HShare, Crypto, Forex, Futures)
|
||||
symbol: Symbol
|
||||
language: Output language tag
|
||||
use_multi_agent: Deprecated; kept for compatibility
|
||||
"""
|
||||
analyzer = AnalysisService()
|
||||
return analyzer.analyze(market, symbol, language)
|
||||
|
||||
|
||||
def reflect_analysis(market: str, symbol: str, decision: str, returns: float = None, result: str = None):
|
||||
"""
|
||||
Reflection hook: learn from post-trade outcomes (local-only).
|
||||
|
||||
Args:
|
||||
market: Market
|
||||
symbol: Symbol
|
||||
decision: Decision (BUY/SELL/HOLD)
|
||||
returns: Return percentage
|
||||
result: Free-text outcome
|
||||
"""
|
||||
try:
|
||||
analyzer = AnalysisService()
|
||||
if analyzer.coordinator:
|
||||
analyzer.coordinator.reflect_and_learn(market, symbol, decision, returns, result)
|
||||
logger.info(f"Reflection completed: {market}:{symbol}")
|
||||
except Exception as e:
|
||||
logger.error(f"Reflection failed: {e}")
|
||||
@@ -0,0 +1,555 @@
|
||||
"""
|
||||
Analysis Memory System 2.0
|
||||
Simplified memory for fast analysis service.
|
||||
|
||||
Features:
|
||||
1. Store analysis decisions with market context
|
||||
2. Retrieve similar historical patterns
|
||||
3. Track decision outcomes for learning
|
||||
"""
|
||||
import json
|
||||
import time
|
||||
import hashlib
|
||||
from typing import Dict, Any, List, Optional
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from app.utils.logger import get_logger
|
||||
from app.utils.db import get_db_connection
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def _safe_json_parse(val, default=None):
|
||||
"""安全解析 JSON - 处理已是 Python 对象或字符串的情况"""
|
||||
if val is None:
|
||||
return default
|
||||
if isinstance(val, (dict, list)):
|
||||
return val # 已经是 Python 对象 (PostgreSQL JSONB 自动转换)
|
||||
if isinstance(val, str):
|
||||
try:
|
||||
return json.loads(val)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return default
|
||||
return default
|
||||
|
||||
|
||||
class AnalysisMemory:
|
||||
"""
|
||||
Simple but effective memory system for AI analysis.
|
||||
Uses PostgreSQL for persistence.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._ensure_table()
|
||||
|
||||
def _ensure_table(self):
|
||||
"""Create memory table if not exists."""
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
cur.execute("""
|
||||
CREATE TABLE IF NOT EXISTS qd_analysis_memory (
|
||||
id SERIAL PRIMARY KEY,
|
||||
market VARCHAR(50) NOT NULL,
|
||||
symbol VARCHAR(50) NOT NULL,
|
||||
decision VARCHAR(10) NOT NULL,
|
||||
confidence INT DEFAULT 50,
|
||||
price_at_analysis DECIMAL(24, 8),
|
||||
entry_price DECIMAL(24, 8),
|
||||
stop_loss DECIMAL(24, 8),
|
||||
take_profit DECIMAL(24, 8),
|
||||
summary TEXT,
|
||||
reasons JSONB,
|
||||
risks JSONB,
|
||||
scores JSONB,
|
||||
indicators_snapshot JSONB,
|
||||
created_at TIMESTAMP DEFAULT NOW(),
|
||||
validated_at TIMESTAMP,
|
||||
actual_outcome VARCHAR(20),
|
||||
actual_return_pct DECIMAL(10, 4),
|
||||
was_correct BOOLEAN,
|
||||
user_feedback VARCHAR(20),
|
||||
feedback_at TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_analysis_memory_symbol
|
||||
ON qd_analysis_memory(market, symbol);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_analysis_memory_created
|
||||
ON qd_analysis_memory(created_at DESC);
|
||||
""")
|
||||
db.commit()
|
||||
cur.close()
|
||||
except Exception as e:
|
||||
logger.warning(f"Memory table creation skipped: {e}")
|
||||
|
||||
def store(self, analysis_result: Dict[str, Any]) -> Optional[int]:
|
||||
"""
|
||||
Store an analysis result for future reference.
|
||||
|
||||
Args:
|
||||
analysis_result: Result from FastAnalysisService.analyze()
|
||||
|
||||
Returns:
|
||||
Memory ID or None if failed
|
||||
"""
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# 准备数据
|
||||
market = analysis_result.get("market")
|
||||
symbol = analysis_result.get("symbol")
|
||||
decision = analysis_result.get("decision")
|
||||
confidence = analysis_result.get("confidence")
|
||||
price = analysis_result.get("market_data", {}).get("current_price")
|
||||
entry = analysis_result.get("trading_plan", {}).get("entry_price")
|
||||
stop = analysis_result.get("trading_plan", {}).get("stop_loss")
|
||||
take = analysis_result.get("trading_plan", {}).get("take_profit")
|
||||
summary = analysis_result.get("summary")
|
||||
reasons = json.dumps(analysis_result.get("reasons", []))
|
||||
risks = json.dumps(analysis_result.get("risks", []))
|
||||
scores = json.dumps(analysis_result.get("scores", {}))
|
||||
indicators = json.dumps(analysis_result.get("indicators", {}))
|
||||
raw = json.dumps(analysis_result)
|
||||
|
||||
cur.execute("""
|
||||
INSERT INTO qd_analysis_memory (
|
||||
market, symbol, decision, confidence,
|
||||
price_at_analysis, entry_price, stop_loss, take_profit,
|
||||
summary, reasons, risks, scores, indicators_snapshot, raw_result
|
||||
) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
|
||||
RETURNING id
|
||||
""", (market, symbol, decision, confidence, price, entry, stop, take,
|
||||
summary, reasons, risks, scores, indicators, raw))
|
||||
|
||||
# 使用 lastrowid 属性获取 ID(execute 内部已经处理了 RETURNING)
|
||||
memory_id = cur.lastrowid
|
||||
db.commit()
|
||||
cur.close()
|
||||
|
||||
logger.info(f"Stored analysis memory #{memory_id} for {symbol}")
|
||||
return memory_id
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to store analysis memory: {e}", exc_info=True)
|
||||
return None
|
||||
|
||||
def get_recent(self, market: str, symbol: str, days: int = 7, limit: int = 5) -> List[Dict]:
|
||||
"""
|
||||
Get recent analysis history for a symbol.
|
||||
|
||||
Args:
|
||||
market: Market type
|
||||
symbol: Symbol
|
||||
days: Look back period
|
||||
limit: Max results
|
||||
|
||||
Returns:
|
||||
List of historical analyses
|
||||
"""
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
cur.execute(f"""
|
||||
SELECT
|
||||
id, decision, confidence, price_at_analysis,
|
||||
summary, reasons, scores,
|
||||
created_at, validated_at, was_correct, actual_return_pct
|
||||
FROM qd_analysis_memory
|
||||
WHERE market = %s AND symbol = %s
|
||||
AND created_at > NOW() - INTERVAL '{int(days)} days'
|
||||
ORDER BY created_at DESC
|
||||
LIMIT %s
|
||||
""", (market, symbol, limit))
|
||||
|
||||
rows = cur.fetchall() or []
|
||||
cur.close()
|
||||
|
||||
results = []
|
||||
for row in rows:
|
||||
results.append({
|
||||
"id": row['id'],
|
||||
"decision": row['decision'],
|
||||
"confidence": row['confidence'],
|
||||
"price": float(row['price_at_analysis']) if row['price_at_analysis'] else None,
|
||||
"summary": row['summary'],
|
||||
"reasons": _safe_json_parse(row['reasons'], []),
|
||||
"scores": _safe_json_parse(row['scores'], {}),
|
||||
"created_at": row['created_at'].isoformat() if row['created_at'] else None,
|
||||
"was_correct": row['was_correct'],
|
||||
"actual_return_pct": float(row['actual_return_pct']) if row['actual_return_pct'] else None,
|
||||
})
|
||||
|
||||
return results
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get recent memories: {e}")
|
||||
return []
|
||||
|
||||
def get_all_history(self, user_id: int = None, page: int = 1, page_size: int = 20) -> Dict:
|
||||
"""
|
||||
Get all analysis history with pagination.
|
||||
|
||||
Args:
|
||||
user_id: Optional user ID filter (not used currently, for future)
|
||||
page: Page number (1-indexed)
|
||||
page_size: Items per page
|
||||
|
||||
Returns:
|
||||
Dict with items list and total count
|
||||
"""
|
||||
try:
|
||||
offset = (page - 1) * page_size
|
||||
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# Get total count
|
||||
cur.execute("SELECT COUNT(*) as cnt FROM qd_analysis_memory")
|
||||
total_row = cur.fetchone()
|
||||
total = total_row['cnt'] if total_row else 0
|
||||
|
||||
# Get paginated results
|
||||
cur.execute("""
|
||||
SELECT
|
||||
id, market, symbol, decision, confidence, price_at_analysis,
|
||||
summary, reasons, scores, indicators_snapshot, raw_result,
|
||||
created_at, validated_at, was_correct, actual_return_pct
|
||||
FROM qd_analysis_memory
|
||||
ORDER BY created_at DESC
|
||||
LIMIT %s OFFSET %s
|
||||
""", (page_size, offset))
|
||||
|
||||
rows = cur.fetchall() or []
|
||||
cur.close()
|
||||
|
||||
items = []
|
||||
for row in rows:
|
||||
items.append({
|
||||
"id": row['id'],
|
||||
"market": row['market'],
|
||||
"symbol": row['symbol'],
|
||||
"decision": row['decision'],
|
||||
"confidence": row['confidence'],
|
||||
"price": float(row['price_at_analysis']) if row['price_at_analysis'] else None,
|
||||
"summary": row['summary'],
|
||||
"reasons": _safe_json_parse(row['reasons'], []),
|
||||
"scores": _safe_json_parse(row['scores'], {}),
|
||||
"indicators": _safe_json_parse(row['indicators_snapshot'], {}),
|
||||
"full_result": _safe_json_parse(row['raw_result'], None),
|
||||
"created_at": row['created_at'].isoformat() if row['created_at'] else None,
|
||||
"was_correct": row['was_correct'],
|
||||
"actual_return_pct": float(row['actual_return_pct']) if row['actual_return_pct'] else None,
|
||||
})
|
||||
|
||||
return {
|
||||
"items": items,
|
||||
"total": total,
|
||||
"page": page,
|
||||
"page_size": page_size
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get all history: {e}")
|
||||
return {"items": [], "total": 0, "page": page, "page_size": page_size}
|
||||
|
||||
def delete_history(self, memory_id: int) -> bool:
|
||||
"""
|
||||
Delete a history record by ID.
|
||||
|
||||
Args:
|
||||
memory_id: The ID of the analysis memory to delete
|
||||
|
||||
Returns:
|
||||
True if deleted successfully, False otherwise
|
||||
"""
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
cur.execute("DELETE FROM qd_analysis_memory WHERE id = %s", (memory_id,))
|
||||
db.commit()
|
||||
affected = cur.rowcount
|
||||
cur.close()
|
||||
return affected > 0
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to delete memory {memory_id}: {e}")
|
||||
return False
|
||||
|
||||
def get_similar_patterns(self, market: str, symbol: str,
|
||||
current_indicators: Dict, limit: int = 3) -> List[Dict]:
|
||||
"""
|
||||
Find historical analyses with similar technical patterns.
|
||||
|
||||
This is a simplified version - can be enhanced with vector similarity later.
|
||||
Currently matches based on:
|
||||
- Same symbol
|
||||
- Similar RSI range (±10)
|
||||
- Same MACD signal direction
|
||||
- Validated outcomes preferred
|
||||
"""
|
||||
try:
|
||||
rsi = current_indicators.get("rsi", {}).get("value", 50)
|
||||
macd_signal = current_indicators.get("macd", {}).get("signal", "neutral")
|
||||
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# Simple pattern matching query
|
||||
cur.execute("""
|
||||
SELECT
|
||||
id, decision, confidence, price_at_analysis,
|
||||
summary, reasons, indicators_snapshot,
|
||||
created_at, was_correct, actual_return_pct
|
||||
FROM qd_analysis_memory
|
||||
WHERE market = %s AND symbol = %s
|
||||
AND validated_at IS NOT NULL
|
||||
AND was_correct IS NOT NULL
|
||||
ORDER BY
|
||||
CASE WHEN was_correct = true THEN 0 ELSE 1 END,
|
||||
created_at DESC
|
||||
LIMIT %s
|
||||
""", (market, symbol, limit * 2)) # Get more for filtering
|
||||
|
||||
rows = cur.fetchall() or []
|
||||
cur.close()
|
||||
|
||||
results = []
|
||||
for row in rows:
|
||||
indicators = _safe_json_parse(row['indicators_snapshot'], {})
|
||||
hist_rsi = indicators.get("rsi", {}).get("value", 50)
|
||||
hist_macd = indicators.get("macd", {}).get("signal", "neutral")
|
||||
|
||||
# Simple similarity check
|
||||
rsi_similar = abs(hist_rsi - rsi) <= 15
|
||||
macd_similar = hist_macd == macd_signal
|
||||
|
||||
if rsi_similar or macd_similar:
|
||||
results.append({
|
||||
"id": row['id'],
|
||||
"decision": row['decision'],
|
||||
"confidence": row['confidence'],
|
||||
"price": float(row['price_at_analysis']) if row['price_at_analysis'] else None,
|
||||
"summary": row['summary'],
|
||||
"was_correct": row['was_correct'],
|
||||
"actual_return_pct": float(row['actual_return_pct']) if row['actual_return_pct'] else None,
|
||||
"similarity": {
|
||||
"rsi_match": rsi_similar,
|
||||
"macd_match": macd_similar,
|
||||
}
|
||||
})
|
||||
|
||||
if len(results) >= limit:
|
||||
break
|
||||
|
||||
return results
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get similar patterns: {e}")
|
||||
return []
|
||||
|
||||
def record_feedback(self, memory_id: int, feedback: str) -> bool:
|
||||
"""
|
||||
Record user feedback on an analysis.
|
||||
|
||||
Args:
|
||||
memory_id: Analysis memory ID
|
||||
feedback: 'helpful' | 'not_helpful' | 'accurate' | 'inaccurate'
|
||||
"""
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
cur.execute("""
|
||||
UPDATE qd_analysis_memory
|
||||
SET user_feedback = %s, feedback_at = NOW()
|
||||
WHERE id = %s
|
||||
""", (feedback, memory_id))
|
||||
db.commit()
|
||||
cur.close()
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to record feedback: {e}")
|
||||
return False
|
||||
|
||||
def validate_past_decisions(self, days_ago: int = 7) -> Dict[str, Any]:
|
||||
"""
|
||||
Validate historical decisions by comparing with actual price movements.
|
||||
Run this periodically (e.g., daily) to build learning data.
|
||||
|
||||
Args:
|
||||
days_ago: Validate decisions from N days ago
|
||||
|
||||
Returns:
|
||||
Validation statistics
|
||||
"""
|
||||
from app.services.market_data_collector import MarketDataCollector
|
||||
collector = MarketDataCollector()
|
||||
|
||||
stats = {
|
||||
"validated": 0,
|
||||
"correct": 0,
|
||||
"incorrect": 0,
|
||||
"errors": 0,
|
||||
}
|
||||
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# Get unvalidated decisions from N days ago
|
||||
cur.execute(f"""
|
||||
SELECT id, market, symbol, decision, price_at_analysis
|
||||
FROM qd_analysis_memory
|
||||
WHERE validated_at IS NULL
|
||||
AND created_at < NOW() - INTERVAL '{int(days_ago)} days'
|
||||
AND created_at > NOW() - INTERVAL '{int(days_ago + 1)} days'
|
||||
LIMIT 50
|
||||
""")
|
||||
|
||||
rows = cur.fetchall() or []
|
||||
|
||||
for row in rows:
|
||||
try:
|
||||
# Get current price using MarketDataCollector
|
||||
current_price = collector._get_price(row['market'], row['symbol'])
|
||||
if not current_price or current_price <= 0:
|
||||
continue
|
||||
analysis_price = float(row['price_at_analysis'])
|
||||
|
||||
if analysis_price <= 0:
|
||||
continue
|
||||
|
||||
# Calculate return
|
||||
return_pct = ((current_price - analysis_price) / analysis_price) * 100
|
||||
|
||||
# Determine if decision was correct
|
||||
decision = row['decision']
|
||||
was_correct = False
|
||||
|
||||
if decision == 'BUY' and return_pct > 2: # 2% threshold
|
||||
was_correct = True
|
||||
elif decision == 'SELL' and return_pct < -2:
|
||||
was_correct = True
|
||||
elif decision == 'HOLD' and abs(return_pct) <= 5:
|
||||
was_correct = True
|
||||
|
||||
# Update record
|
||||
cur.execute("""
|
||||
UPDATE qd_analysis_memory
|
||||
SET validated_at = NOW(),
|
||||
actual_return_pct = %s,
|
||||
was_correct = %s
|
||||
WHERE id = %s
|
||||
""", (return_pct, was_correct, row['id']))
|
||||
|
||||
stats["validated"] += 1
|
||||
if was_correct:
|
||||
stats["correct"] += 1
|
||||
else:
|
||||
stats["incorrect"] += 1
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to validate memory {row['id']}: {e}")
|
||||
stats["errors"] += 1
|
||||
|
||||
db.commit()
|
||||
cur.close()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Validation batch failed: {e}")
|
||||
|
||||
accuracy = (stats["correct"] / stats["validated"] * 100) if stats["validated"] > 0 else 0
|
||||
stats["accuracy_pct"] = round(accuracy, 2)
|
||||
|
||||
logger.info(f"Validation completed: {stats}")
|
||||
return stats
|
||||
|
||||
def get_performance_stats(self, market: str = None, symbol: str = None,
|
||||
days: int = 30) -> Dict[str, Any]:
|
||||
"""
|
||||
Get AI performance statistics.
|
||||
|
||||
Returns:
|
||||
Performance metrics for display
|
||||
"""
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
where_clauses = ["validated_at IS NOT NULL"]
|
||||
params = []
|
||||
|
||||
if market:
|
||||
where_clauses.append("market = %s")
|
||||
params.append(market)
|
||||
if symbol:
|
||||
where_clauses.append("symbol = %s")
|
||||
params.append(symbol)
|
||||
|
||||
# Use f-string for interval since psycopg2 doesn't support placeholder in INTERVAL
|
||||
where_clauses.append(f"created_at > NOW() - INTERVAL '{int(days)} days'")
|
||||
|
||||
where_sql = " AND ".join(where_clauses)
|
||||
|
||||
cur.execute(f"""
|
||||
SELECT
|
||||
COUNT(*) as total,
|
||||
SUM(CASE WHEN was_correct = true THEN 1 ELSE 0 END) as correct,
|
||||
AVG(actual_return_pct) as avg_return,
|
||||
SUM(CASE WHEN decision = 'BUY' THEN 1 ELSE 0 END) as buy_count,
|
||||
SUM(CASE WHEN decision = 'SELL' THEN 1 ELSE 0 END) as sell_count,
|
||||
SUM(CASE WHEN decision = 'HOLD' THEN 1 ELSE 0 END) as hold_count,
|
||||
SUM(CASE WHEN user_feedback = 'helpful' THEN 1 ELSE 0 END) as helpful_count,
|
||||
SUM(CASE WHEN user_feedback IS NOT NULL THEN 1 ELSE 0 END) as feedback_count
|
||||
FROM qd_analysis_memory
|
||||
WHERE {where_sql}
|
||||
""", tuple(params) if params else None)
|
||||
|
||||
row = cur.fetchone()
|
||||
cur.close()
|
||||
|
||||
if not row or not row['total']:
|
||||
return {
|
||||
"total_analyses": 0,
|
||||
"accuracy_pct": 0,
|
||||
"avg_return_pct": 0,
|
||||
"user_satisfaction_pct": 0,
|
||||
}
|
||||
|
||||
total = row['total']
|
||||
correct = row['correct'] or 0
|
||||
|
||||
return {
|
||||
"total_analyses": total,
|
||||
"accuracy_pct": round((correct / total * 100) if total > 0 else 0, 2),
|
||||
"avg_return_pct": round(float(row['avg_return'] or 0), 2),
|
||||
"decision_distribution": {
|
||||
"buy": row['buy_count'] or 0,
|
||||
"sell": row['sell_count'] or 0,
|
||||
"hold": row['hold_count'] or 0,
|
||||
},
|
||||
"user_satisfaction_pct": round(
|
||||
(row['helpful_count'] / row['feedback_count'] * 100)
|
||||
if row['feedback_count'] and row['feedback_count'] > 0 else 0, 2
|
||||
),
|
||||
"period_days": days,
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get performance stats: {e}")
|
||||
return {
|
||||
"total_analyses": 0,
|
||||
"accuracy_pct": 0,
|
||||
"error": str(e),
|
||||
}
|
||||
|
||||
|
||||
# Singleton
|
||||
_memory_instance = None
|
||||
|
||||
def get_analysis_memory() -> AnalysisMemory:
|
||||
"""Get singleton AnalysisMemory instance."""
|
||||
global _memory_instance
|
||||
if _memory_instance is None:
|
||||
_memory_instance = AnalysisMemory()
|
||||
return _memory_instance
|
||||
@@ -0,0 +1,941 @@
|
||||
"""
|
||||
Community Service - 指标社区服务
|
||||
|
||||
处理指标市场、购买、评论等功能。
|
||||
"""
|
||||
import time
|
||||
from decimal import Decimal
|
||||
from typing import Dict, Any, List, Optional, Tuple
|
||||
|
||||
from app.utils.db import get_db_connection
|
||||
from app.utils.logger import get_logger
|
||||
from app.services.billing_service import get_billing_service
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class CommunityService:
|
||||
"""指标社区服务类"""
|
||||
|
||||
def __init__(self):
|
||||
self.billing = get_billing_service()
|
||||
|
||||
# ==========================================
|
||||
# 指标市场
|
||||
# ==========================================
|
||||
|
||||
def get_market_indicators(
|
||||
self,
|
||||
page: int = 1,
|
||||
page_size: int = 12,
|
||||
keyword: str = None,
|
||||
pricing_type: str = None, # 'free' / 'paid' / None(all)
|
||||
sort_by: str = 'newest', # 'newest' / 'hot' / 'price_asc' / 'price_desc' / 'rating'
|
||||
user_id: int = None # 当前用户ID,用于判断是否已购买
|
||||
) -> Dict[str, Any]:
|
||||
"""获取市场上已发布的指标列表"""
|
||||
offset = (page - 1) * page_size
|
||||
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# 构建查询条件 - 只显示已发布且审核通过的指标
|
||||
where_clauses = ["i.publish_to_community = 1", "(i.review_status = 'approved' OR i.review_status IS NULL)"]
|
||||
params = []
|
||||
|
||||
if keyword and keyword.strip():
|
||||
where_clauses.append("(i.name ILIKE ? OR i.description ILIKE ?)")
|
||||
search_term = f"%{keyword.strip()}%"
|
||||
params.extend([search_term, search_term])
|
||||
|
||||
if pricing_type == 'free':
|
||||
where_clauses.append("(i.pricing_type = 'free' OR i.price <= 0)")
|
||||
elif pricing_type == 'paid':
|
||||
where_clauses.append("(i.pricing_type != 'free' AND i.price > 0)")
|
||||
|
||||
where_sql = " AND ".join(where_clauses)
|
||||
|
||||
# 排序
|
||||
order_map = {
|
||||
'newest': 'i.created_at DESC',
|
||||
'hot': 'i.purchase_count DESC, i.view_count DESC',
|
||||
'price_asc': 'i.price ASC, i.created_at DESC',
|
||||
'price_desc': 'i.price DESC, i.created_at DESC',
|
||||
'rating': 'i.avg_rating DESC, i.rating_count DESC'
|
||||
}
|
||||
order_sql = order_map.get(sort_by, 'i.created_at DESC')
|
||||
|
||||
# 获取总数
|
||||
count_sql = f"""
|
||||
SELECT COUNT(*) as count
|
||||
FROM qd_indicator_codes i
|
||||
WHERE {where_sql}
|
||||
"""
|
||||
cur.execute(count_sql, tuple(params))
|
||||
total = cur.fetchone()['count']
|
||||
|
||||
# 获取列表(联表查询作者信息)
|
||||
query_sql = f"""
|
||||
SELECT
|
||||
i.id, i.name, i.description, i.pricing_type, i.price,
|
||||
i.preview_image, i.purchase_count, i.avg_rating, i.rating_count,
|
||||
i.view_count, i.created_at, i.updated_at,
|
||||
u.id as author_id, u.username as author_username,
|
||||
u.nickname as author_nickname, u.avatar as author_avatar
|
||||
FROM qd_indicator_codes i
|
||||
LEFT JOIN qd_users u ON i.user_id = u.id
|
||||
WHERE {where_sql}
|
||||
ORDER BY {order_sql}
|
||||
LIMIT ? OFFSET ?
|
||||
"""
|
||||
cur.execute(query_sql, tuple(params + [page_size, offset]))
|
||||
rows = cur.fetchall() or []
|
||||
|
||||
# 如果有当前用户,查询已购买的指标
|
||||
purchased_ids = set()
|
||||
if user_id:
|
||||
indicator_ids = [r['id'] for r in rows]
|
||||
if indicator_ids:
|
||||
placeholders = ','.join(['?'] * len(indicator_ids))
|
||||
cur.execute(
|
||||
f"SELECT indicator_id FROM qd_indicator_purchases WHERE buyer_id = ? AND indicator_id IN ({placeholders})",
|
||||
tuple([user_id] + indicator_ids)
|
||||
)
|
||||
purchased_ids = {r['indicator_id'] for r in (cur.fetchall() or [])}
|
||||
|
||||
cur.close()
|
||||
|
||||
# 格式化返回数据
|
||||
items = []
|
||||
for row in rows:
|
||||
items.append({
|
||||
'id': row['id'],
|
||||
'name': row['name'],
|
||||
'description': row['description'][:200] if row['description'] else '',
|
||||
'pricing_type': row['pricing_type'] or 'free',
|
||||
'price': float(row['price'] or 0),
|
||||
'preview_image': row['preview_image'] or '',
|
||||
'purchase_count': row['purchase_count'] or 0,
|
||||
'avg_rating': float(row['avg_rating'] or 0),
|
||||
'rating_count': row['rating_count'] or 0,
|
||||
'view_count': row['view_count'] or 0,
|
||||
'created_at': row['created_at'].isoformat() if row['created_at'] else None,
|
||||
'author': {
|
||||
'id': row['author_id'],
|
||||
'username': row['author_username'],
|
||||
'nickname': row['author_nickname'] or row['author_username'],
|
||||
'avatar': row['author_avatar'] or '/avatar2.jpg'
|
||||
},
|
||||
'is_purchased': row['id'] in purchased_ids,
|
||||
'is_own': row['author_id'] == user_id
|
||||
})
|
||||
|
||||
return {
|
||||
'items': items,
|
||||
'total': total,
|
||||
'page': page,
|
||||
'page_size': page_size,
|
||||
'total_pages': (total + page_size - 1) // page_size if total > 0 else 0
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"get_market_indicators failed: {e}")
|
||||
return {'items': [], 'total': 0, 'page': 1, 'page_size': page_size, 'total_pages': 0}
|
||||
|
||||
def get_indicator_detail(self, indicator_id: int, user_id: int = None) -> Optional[Dict[str, Any]]:
|
||||
"""获取指标详情"""
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# 获取指标信息
|
||||
cur.execute("""
|
||||
SELECT
|
||||
i.id, i.name, i.description, i.pricing_type, i.price,
|
||||
i.preview_image, i.purchase_count, i.avg_rating, i.rating_count,
|
||||
i.view_count, i.publish_to_community, i.created_at, i.updated_at,
|
||||
i.user_id,
|
||||
u.id as author_id, u.username as author_username,
|
||||
u.nickname as author_nickname, u.avatar as author_avatar
|
||||
FROM qd_indicator_codes i
|
||||
LEFT JOIN qd_users u ON i.user_id = u.id
|
||||
WHERE i.id = ?
|
||||
""", (indicator_id,))
|
||||
row = cur.fetchone()
|
||||
|
||||
if not row:
|
||||
cur.close()
|
||||
return None
|
||||
|
||||
# 检查是否已发布到社区(或者是自己的指标)
|
||||
if not row['publish_to_community'] and row['user_id'] != user_id:
|
||||
cur.close()
|
||||
return None
|
||||
|
||||
# 检查是否已购买
|
||||
is_purchased = False
|
||||
if user_id:
|
||||
cur.execute(
|
||||
"SELECT id FROM qd_indicator_purchases WHERE indicator_id = ? AND buyer_id = ?",
|
||||
(indicator_id, user_id)
|
||||
)
|
||||
is_purchased = cur.fetchone() is not None
|
||||
|
||||
# 增加浏览次数
|
||||
cur.execute(
|
||||
"UPDATE qd_indicator_codes SET view_count = COALESCE(view_count, 0) + 1 WHERE id = ?",
|
||||
(indicator_id,)
|
||||
)
|
||||
db.commit()
|
||||
cur.close()
|
||||
|
||||
return {
|
||||
'id': row['id'],
|
||||
'name': row['name'],
|
||||
'description': row['description'] or '',
|
||||
'pricing_type': row['pricing_type'] or 'free',
|
||||
'price': float(row['price'] or 0),
|
||||
'preview_image': row['preview_image'] or '',
|
||||
'purchase_count': row['purchase_count'] or 0,
|
||||
'avg_rating': float(row['avg_rating'] or 0),
|
||||
'rating_count': row['rating_count'] or 0,
|
||||
'view_count': (row['view_count'] or 0) + 1,
|
||||
'created_at': row['created_at'].isoformat() if row['created_at'] else None,
|
||||
'updated_at': row['updated_at'].isoformat() if row['updated_at'] else None,
|
||||
'author': {
|
||||
'id': row['author_id'],
|
||||
'username': row['author_username'],
|
||||
'nickname': row['author_nickname'] or row['author_username'],
|
||||
'avatar': row['author_avatar'] or '/avatar2.jpg'
|
||||
},
|
||||
'is_purchased': is_purchased,
|
||||
'is_own': row['user_id'] == user_id
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"get_indicator_detail failed: {e}")
|
||||
return None
|
||||
|
||||
# ==========================================
|
||||
# 购买功能
|
||||
# ==========================================
|
||||
|
||||
def purchase_indicator(self, buyer_id: int, indicator_id: int) -> Tuple[bool, str, Dict[str, Any]]:
|
||||
"""
|
||||
购买指标
|
||||
|
||||
Returns:
|
||||
(success, message, data)
|
||||
"""
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# 1. 获取指标信息
|
||||
cur.execute("""
|
||||
SELECT id, user_id, name, code, description, pricing_type, price,
|
||||
preview_image, is_encrypted
|
||||
FROM qd_indicator_codes
|
||||
WHERE id = ? AND publish_to_community = 1
|
||||
""", (indicator_id,))
|
||||
indicator = cur.fetchone()
|
||||
|
||||
if not indicator:
|
||||
cur.close()
|
||||
return False, 'indicator_not_found', {}
|
||||
|
||||
seller_id = indicator['user_id']
|
||||
price = float(indicator['price'] or 0)
|
||||
pricing_type = indicator['pricing_type'] or 'free'
|
||||
|
||||
# 2. 检查是否购买自己的指标
|
||||
if seller_id == buyer_id:
|
||||
cur.close()
|
||||
return False, 'cannot_buy_own', {}
|
||||
|
||||
# 3. 检查是否已购买
|
||||
cur.execute(
|
||||
"SELECT id FROM qd_indicator_purchases WHERE indicator_id = ? AND buyer_id = ?",
|
||||
(indicator_id, buyer_id)
|
||||
)
|
||||
if cur.fetchone():
|
||||
cur.close()
|
||||
return False, 'already_purchased', {}
|
||||
|
||||
# 4. 如果是付费指标,检查并扣除积分
|
||||
if pricing_type != 'free' and price > 0:
|
||||
buyer_credits = self.billing.get_user_credits(buyer_id)
|
||||
if buyer_credits < price:
|
||||
cur.close()
|
||||
return False, 'insufficient_credits', {
|
||||
'required': price,
|
||||
'current': float(buyer_credits)
|
||||
}
|
||||
|
||||
# 扣除买家积分
|
||||
new_buyer_balance = buyer_credits - Decimal(str(price))
|
||||
cur.execute(
|
||||
"UPDATE qd_users SET credits = ?, updated_at = NOW() WHERE id = ?",
|
||||
(float(new_buyer_balance), buyer_id)
|
||||
)
|
||||
|
||||
# 记录买家积分日志
|
||||
cur.execute("""
|
||||
INSERT INTO qd_credits_log
|
||||
(user_id, action, amount, balance_after, feature, reference_id, remark, created_at)
|
||||
VALUES (?, 'indicator_purchase', ?, ?, 'indicator_purchase', ?, ?, NOW())
|
||||
""", (buyer_id, -price, float(new_buyer_balance), str(indicator_id),
|
||||
f"购买指标: {indicator['name']}"))
|
||||
|
||||
# 给卖家增加积分(可配置抽成比例,这里先100%给卖家)
|
||||
seller_credits = self.billing.get_user_credits(seller_id)
|
||||
new_seller_balance = seller_credits + Decimal(str(price))
|
||||
cur.execute(
|
||||
"UPDATE qd_users SET credits = ?, updated_at = NOW() WHERE id = ?",
|
||||
(float(new_seller_balance), seller_id)
|
||||
)
|
||||
|
||||
# 记录卖家积分日志
|
||||
cur.execute("""
|
||||
INSERT INTO qd_credits_log
|
||||
(user_id, action, amount, balance_after, feature, reference_id, remark, created_at)
|
||||
VALUES (?, 'indicator_sale', ?, ?, 'indicator_sale', ?, ?, NOW())
|
||||
""", (seller_id, price, float(new_seller_balance), str(indicator_id),
|
||||
f"出售指标: {indicator['name']}"))
|
||||
|
||||
# 5. 创建购买记录
|
||||
cur.execute("""
|
||||
INSERT INTO qd_indicator_purchases
|
||||
(indicator_id, buyer_id, seller_id, price, created_at)
|
||||
VALUES (?, ?, ?, ?, NOW())
|
||||
""", (indicator_id, buyer_id, seller_id, price))
|
||||
|
||||
# 6. 复制指标到买家账户
|
||||
now_ts = int(time.time())
|
||||
cur.execute("""
|
||||
INSERT INTO qd_indicator_codes
|
||||
(user_id, is_buy, end_time, name, code, description,
|
||||
publish_to_community, pricing_type, price, is_encrypted, preview_image,
|
||||
createtime, updatetime, created_at, updated_at)
|
||||
VALUES (?, 1, 0, ?, ?, ?, 0, 'free', 0, ?, ?, ?, ?, NOW(), NOW())
|
||||
""", (
|
||||
buyer_id,
|
||||
indicator['name'],
|
||||
indicator['code'],
|
||||
indicator['description'],
|
||||
indicator['is_encrypted'] or 0,
|
||||
indicator['preview_image'],
|
||||
now_ts, now_ts
|
||||
))
|
||||
|
||||
# 7. 更新指标购买次数
|
||||
cur.execute("""
|
||||
UPDATE qd_indicator_codes
|
||||
SET purchase_count = COALESCE(purchase_count, 0) + 1
|
||||
WHERE id = ?
|
||||
""", (indicator_id,))
|
||||
|
||||
db.commit()
|
||||
cur.close()
|
||||
|
||||
logger.info(f"User {buyer_id} purchased indicator {indicator_id} for {price} credits")
|
||||
return True, 'success', {'indicator_name': indicator['name'], 'price': price}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"purchase_indicator failed: {e}")
|
||||
return False, f'error: {str(e)}', {}
|
||||
|
||||
def get_my_purchases(self, user_id: int, page: int = 1, page_size: int = 20) -> Dict[str, Any]:
|
||||
"""获取用户购买的指标列表"""
|
||||
offset = (page - 1) * page_size
|
||||
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# 获取总数
|
||||
cur.execute(
|
||||
"SELECT COUNT(*) as count FROM qd_indicator_purchases WHERE buyer_id = ?",
|
||||
(user_id,)
|
||||
)
|
||||
total = cur.fetchone()['count']
|
||||
|
||||
# 获取列表
|
||||
cur.execute("""
|
||||
SELECT
|
||||
p.id as purchase_id, p.price as purchase_price, p.created_at as purchase_time,
|
||||
i.id, i.name, i.description, i.preview_image, i.avg_rating,
|
||||
u.nickname as seller_nickname, u.avatar as seller_avatar
|
||||
FROM qd_indicator_purchases p
|
||||
LEFT JOIN qd_indicator_codes i ON p.indicator_id = i.id
|
||||
LEFT JOIN qd_users u ON p.seller_id = u.id
|
||||
WHERE p.buyer_id = ?
|
||||
ORDER BY p.created_at DESC
|
||||
LIMIT ? OFFSET ?
|
||||
""", (user_id, page_size, offset))
|
||||
rows = cur.fetchall() or []
|
||||
cur.close()
|
||||
|
||||
items = []
|
||||
for row in rows:
|
||||
items.append({
|
||||
'purchase_id': row['purchase_id'],
|
||||
'purchase_price': float(row['purchase_price'] or 0),
|
||||
'purchase_time': row['purchase_time'].isoformat() if row['purchase_time'] else None,
|
||||
'indicator': {
|
||||
'id': row['id'],
|
||||
'name': row['name'],
|
||||
'description': row['description'][:100] if row['description'] else '',
|
||||
'preview_image': row['preview_image'] or '',
|
||||
'avg_rating': float(row['avg_rating'] or 0)
|
||||
},
|
||||
'seller': {
|
||||
'nickname': row['seller_nickname'],
|
||||
'avatar': row['seller_avatar'] or '/avatar2.jpg'
|
||||
}
|
||||
})
|
||||
|
||||
return {
|
||||
'items': items,
|
||||
'total': total,
|
||||
'page': page,
|
||||
'page_size': page_size,
|
||||
'total_pages': (total + page_size - 1) // page_size if total > 0 else 0
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"get_my_purchases failed: {e}")
|
||||
return {'items': [], 'total': 0, 'page': 1, 'page_size': page_size, 'total_pages': 0}
|
||||
|
||||
# ==========================================
|
||||
# 评论功能
|
||||
# ==========================================
|
||||
|
||||
def get_comments(self, indicator_id: int, page: int = 1, page_size: int = 20) -> Dict[str, Any]:
|
||||
"""获取指标评论列表"""
|
||||
offset = (page - 1) * page_size
|
||||
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# 获取总数(只统计一级评论)
|
||||
cur.execute("""
|
||||
SELECT COUNT(*) as count FROM qd_indicator_comments
|
||||
WHERE indicator_id = ? AND parent_id IS NULL AND is_deleted = 0
|
||||
""", (indicator_id,))
|
||||
total = cur.fetchone()['count']
|
||||
|
||||
# 获取评论列表
|
||||
cur.execute("""
|
||||
SELECT
|
||||
c.id, c.rating, c.content, c.created_at,
|
||||
u.id as user_id, u.nickname, u.avatar
|
||||
FROM qd_indicator_comments c
|
||||
LEFT JOIN qd_users u ON c.user_id = u.id
|
||||
WHERE c.indicator_id = ? AND c.parent_id IS NULL AND c.is_deleted = 0
|
||||
ORDER BY c.created_at DESC
|
||||
LIMIT ? OFFSET ?
|
||||
""", (indicator_id, page_size, offset))
|
||||
rows = cur.fetchall() or []
|
||||
cur.close()
|
||||
|
||||
items = []
|
||||
for row in rows:
|
||||
items.append({
|
||||
'id': row['id'],
|
||||
'rating': row['rating'],
|
||||
'content': row['content'],
|
||||
'created_at': row['created_at'].isoformat() if row['created_at'] else None,
|
||||
'user': {
|
||||
'id': row['user_id'],
|
||||
'nickname': row['nickname'],
|
||||
'avatar': row['avatar'] or '/avatar2.jpg'
|
||||
}
|
||||
})
|
||||
|
||||
return {
|
||||
'items': items,
|
||||
'total': total,
|
||||
'page': page,
|
||||
'page_size': page_size,
|
||||
'total_pages': (total + page_size - 1) // page_size if total > 0 else 0
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"get_comments failed: {e}")
|
||||
return {'items': [], 'total': 0, 'page': 1, 'page_size': page_size, 'total_pages': 0}
|
||||
|
||||
def add_comment(
|
||||
self,
|
||||
user_id: int,
|
||||
indicator_id: int,
|
||||
rating: int,
|
||||
content: str
|
||||
) -> Tuple[bool, str, Dict[str, Any]]:
|
||||
"""
|
||||
添加评论(只有购买过的用户可以评论,且只能评论一次)
|
||||
"""
|
||||
try:
|
||||
# 验证评分范围
|
||||
rating = max(1, min(5, int(rating)))
|
||||
content = (content or '').strip()[:500] # 限制500字
|
||||
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# 检查指标是否存在
|
||||
cur.execute(
|
||||
"SELECT id, user_id FROM qd_indicator_codes WHERE id = ? AND publish_to_community = 1",
|
||||
(indicator_id,)
|
||||
)
|
||||
indicator = cur.fetchone()
|
||||
if not indicator:
|
||||
cur.close()
|
||||
return False, 'indicator_not_found', {}
|
||||
|
||||
# 不能评论自己的指标
|
||||
if indicator['user_id'] == user_id:
|
||||
cur.close()
|
||||
return False, 'cannot_comment_own', {}
|
||||
|
||||
# 检查是否已购买(免费指标也需要"获取"才能评论)
|
||||
cur.execute(
|
||||
"SELECT id FROM qd_indicator_purchases WHERE indicator_id = ? AND buyer_id = ?",
|
||||
(indicator_id, user_id)
|
||||
)
|
||||
if not cur.fetchone():
|
||||
cur.close()
|
||||
return False, 'not_purchased', {}
|
||||
|
||||
# 检查是否已评论
|
||||
cur.execute(
|
||||
"SELECT id FROM qd_indicator_comments WHERE indicator_id = ? AND user_id = ? AND parent_id IS NULL",
|
||||
(indicator_id, user_id)
|
||||
)
|
||||
if cur.fetchone():
|
||||
cur.close()
|
||||
return False, 'already_commented', {}
|
||||
|
||||
# 添加评论
|
||||
cur.execute("""
|
||||
INSERT INTO qd_indicator_comments
|
||||
(indicator_id, user_id, rating, content, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, NOW(), NOW())
|
||||
""", (indicator_id, user_id, rating, content))
|
||||
comment_id = cur.lastrowid
|
||||
|
||||
# 更新指标的评分统计
|
||||
cur.execute("""
|
||||
UPDATE qd_indicator_codes
|
||||
SET
|
||||
rating_count = COALESCE(rating_count, 0) + 1,
|
||||
avg_rating = (
|
||||
SELECT AVG(rating) FROM qd_indicator_comments
|
||||
WHERE indicator_id = ? AND parent_id IS NULL AND is_deleted = 0
|
||||
)
|
||||
WHERE id = ?
|
||||
""", (indicator_id, indicator_id))
|
||||
|
||||
db.commit()
|
||||
cur.close()
|
||||
|
||||
logger.info(f"User {user_id} commented on indicator {indicator_id} with rating {rating}")
|
||||
return True, 'success', {'comment_id': comment_id}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"add_comment failed: {e}")
|
||||
return False, f'error: {str(e)}', {}
|
||||
|
||||
def update_comment(
|
||||
self,
|
||||
user_id: int,
|
||||
comment_id: int,
|
||||
indicator_id: int,
|
||||
rating: int,
|
||||
content: str
|
||||
) -> Tuple[bool, str, Dict[str, Any]]:
|
||||
"""
|
||||
更新评论(只能修改自己的评论)
|
||||
"""
|
||||
try:
|
||||
rating = max(1, min(5, int(rating)))
|
||||
content = (content or '').strip()[:500]
|
||||
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# 检查评论是否存在且属于当前用户
|
||||
cur.execute("""
|
||||
SELECT id, rating as old_rating FROM qd_indicator_comments
|
||||
WHERE id = ? AND user_id = ? AND indicator_id = ? AND is_deleted = 0
|
||||
""", (comment_id, user_id, indicator_id))
|
||||
comment = cur.fetchone()
|
||||
|
||||
if not comment:
|
||||
cur.close()
|
||||
return False, 'comment_not_found', {}
|
||||
|
||||
old_rating = comment['old_rating']
|
||||
|
||||
# 更新评论
|
||||
cur.execute("""
|
||||
UPDATE qd_indicator_comments
|
||||
SET rating = ?, content = ?, updated_at = NOW()
|
||||
WHERE id = ?
|
||||
""", (rating, content, comment_id))
|
||||
|
||||
# 如果评分变了,更新指标的平均评分
|
||||
if old_rating != rating:
|
||||
cur.execute("""
|
||||
UPDATE qd_indicator_codes
|
||||
SET avg_rating = (
|
||||
SELECT AVG(rating) FROM qd_indicator_comments
|
||||
WHERE indicator_id = ? AND parent_id IS NULL AND is_deleted = 0
|
||||
)
|
||||
WHERE id = ?
|
||||
""", (indicator_id, indicator_id))
|
||||
|
||||
db.commit()
|
||||
cur.close()
|
||||
|
||||
logger.info(f"User {user_id} updated comment {comment_id}")
|
||||
return True, 'success', {'comment_id': comment_id}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"update_comment failed: {e}")
|
||||
return False, f'error: {str(e)}', {}
|
||||
|
||||
def get_user_comment(self, user_id: int, indicator_id: int) -> Optional[Dict[str, Any]]:
|
||||
"""获取用户对某个指标的评论"""
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
cur.execute("""
|
||||
SELECT id, rating, content, created_at, updated_at
|
||||
FROM qd_indicator_comments
|
||||
WHERE user_id = ? AND indicator_id = ? AND parent_id IS NULL AND is_deleted = 0
|
||||
""", (user_id, indicator_id))
|
||||
row = cur.fetchone()
|
||||
cur.close()
|
||||
|
||||
if not row:
|
||||
return None
|
||||
|
||||
return {
|
||||
'id': row['id'],
|
||||
'rating': row['rating'],
|
||||
'content': row['content'],
|
||||
'created_at': row['created_at'].isoformat() if row['created_at'] else None,
|
||||
'updated_at': row['updated_at'].isoformat() if row['updated_at'] else None
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"get_user_comment failed: {e}")
|
||||
return None
|
||||
|
||||
# ==========================================
|
||||
# 管理员审核功能
|
||||
# ==========================================
|
||||
|
||||
def get_pending_indicators(
|
||||
self,
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
review_status: str = 'pending' # 'pending' / 'approved' / 'rejected' / 'all'
|
||||
) -> Dict[str, Any]:
|
||||
"""获取待审核的指标列表(管理员用)"""
|
||||
offset = (page - 1) * page_size
|
||||
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# 构建查询条件
|
||||
where_clauses = ["i.publish_to_community = 1"]
|
||||
params = []
|
||||
|
||||
if review_status and review_status != 'all':
|
||||
where_clauses.append("i.review_status = ?")
|
||||
params.append(review_status)
|
||||
|
||||
where_sql = " AND ".join(where_clauses)
|
||||
|
||||
# 获取总数
|
||||
count_sql = f"""
|
||||
SELECT COUNT(*) as count
|
||||
FROM qd_indicator_codes i
|
||||
WHERE {where_sql}
|
||||
"""
|
||||
cur.execute(count_sql, tuple(params))
|
||||
total = cur.fetchone()['count']
|
||||
|
||||
# 获取列表
|
||||
query_sql = f"""
|
||||
SELECT
|
||||
i.id, i.name, i.description, i.pricing_type, i.price,
|
||||
i.preview_image, i.code, i.review_status, i.review_note,
|
||||
i.reviewed_at, i.reviewed_by, i.created_at,
|
||||
u.id as author_id, u.username as author_username,
|
||||
u.nickname as author_nickname, u.avatar as author_avatar,
|
||||
r.username as reviewer_username
|
||||
FROM qd_indicator_codes i
|
||||
LEFT JOIN qd_users u ON i.user_id = u.id
|
||||
LEFT JOIN qd_users r ON i.reviewed_by = r.id
|
||||
WHERE {where_sql}
|
||||
ORDER BY i.created_at DESC
|
||||
LIMIT ? OFFSET ?
|
||||
"""
|
||||
cur.execute(query_sql, tuple(params + [page_size, offset]))
|
||||
rows = cur.fetchall() or []
|
||||
cur.close()
|
||||
|
||||
items = []
|
||||
for row in rows:
|
||||
items.append({
|
||||
'id': row['id'],
|
||||
'name': row['name'],
|
||||
'description': row['description'][:300] if row['description'] else '',
|
||||
'pricing_type': row['pricing_type'] or 'free',
|
||||
'price': float(row['price'] or 0),
|
||||
'preview_image': row['preview_image'] or '',
|
||||
'code': row['code'] or '', # 管理员可以看代码
|
||||
'review_status': row['review_status'] or 'pending',
|
||||
'review_note': row['review_note'] or '',
|
||||
'reviewed_at': row['reviewed_at'].isoformat() if row['reviewed_at'] else None,
|
||||
'reviewer_username': row['reviewer_username'],
|
||||
'created_at': row['created_at'].isoformat() if row['created_at'] else None,
|
||||
'author': {
|
||||
'id': row['author_id'],
|
||||
'username': row['author_username'],
|
||||
'nickname': row['author_nickname'] or row['author_username'],
|
||||
'avatar': row['author_avatar'] or '/avatar2.jpg'
|
||||
}
|
||||
})
|
||||
|
||||
return {
|
||||
'items': items,
|
||||
'total': total,
|
||||
'page': page,
|
||||
'page_size': page_size,
|
||||
'total_pages': (total + page_size - 1) // page_size if total > 0 else 0
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"get_pending_indicators failed: {e}")
|
||||
return {'items': [], 'total': 0, 'page': 1, 'page_size': page_size, 'total_pages': 0}
|
||||
|
||||
def review_indicator(
|
||||
self,
|
||||
admin_id: int,
|
||||
indicator_id: int,
|
||||
action: str, # 'approve' / 'reject'
|
||||
note: str = ''
|
||||
) -> Tuple[bool, str]:
|
||||
"""审核指标"""
|
||||
try:
|
||||
new_status = 'approved' if action == 'approve' else 'rejected'
|
||||
note = (note or '').strip()[:500]
|
||||
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# 检查指标是否存在且已发布到社区
|
||||
cur.execute("""
|
||||
SELECT id, name, user_id FROM qd_indicator_codes
|
||||
WHERE id = ? AND publish_to_community = 1
|
||||
""", (indicator_id,))
|
||||
indicator = cur.fetchone()
|
||||
|
||||
if not indicator:
|
||||
cur.close()
|
||||
return False, 'indicator_not_found'
|
||||
|
||||
# 更新审核状态
|
||||
cur.execute("""
|
||||
UPDATE qd_indicator_codes
|
||||
SET review_status = ?, review_note = ?, reviewed_at = NOW(), reviewed_by = ?
|
||||
WHERE id = ?
|
||||
""", (new_status, note, admin_id, indicator_id))
|
||||
|
||||
db.commit()
|
||||
cur.close()
|
||||
|
||||
logger.info(f"Admin {admin_id} {action}d indicator {indicator_id}")
|
||||
return True, 'success'
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"review_indicator failed: {e}")
|
||||
return False, f'error: {str(e)}'
|
||||
|
||||
def unpublish_indicator(self, admin_id: int, indicator_id: int, note: str = '') -> Tuple[bool, str]:
|
||||
"""下架指标(取消发布)"""
|
||||
try:
|
||||
note = (note or '').strip()[:500]
|
||||
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# 检查指标是否存在
|
||||
cur.execute("""
|
||||
SELECT id, name FROM qd_indicator_codes WHERE id = ?
|
||||
""", (indicator_id,))
|
||||
indicator = cur.fetchone()
|
||||
|
||||
if not indicator:
|
||||
cur.close()
|
||||
return False, 'indicator_not_found'
|
||||
|
||||
# 下架(取消发布)
|
||||
cur.execute("""
|
||||
UPDATE qd_indicator_codes
|
||||
SET publish_to_community = 0, review_status = 'rejected',
|
||||
review_note = ?, reviewed_at = NOW(), reviewed_by = ?
|
||||
WHERE id = ?
|
||||
""", (f"下架: {note}" if note else "管理员下架", admin_id, indicator_id))
|
||||
|
||||
db.commit()
|
||||
cur.close()
|
||||
|
||||
logger.info(f"Admin {admin_id} unpublished indicator {indicator_id}")
|
||||
return True, 'success'
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"unpublish_indicator failed: {e}")
|
||||
return False, f'error: {str(e)}'
|
||||
|
||||
def admin_delete_indicator(self, admin_id: int, indicator_id: int) -> Tuple[bool, str]:
|
||||
"""管理员删除指标"""
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# 检查指标是否存在
|
||||
cur.execute("SELECT id, name FROM qd_indicator_codes WHERE id = ?", (indicator_id,))
|
||||
indicator = cur.fetchone()
|
||||
|
||||
if not indicator:
|
||||
cur.close()
|
||||
return False, 'indicator_not_found'
|
||||
|
||||
# 删除关联的评论
|
||||
cur.execute("DELETE FROM qd_indicator_comments WHERE indicator_id = ?", (indicator_id,))
|
||||
|
||||
# 删除关联的购买记录
|
||||
cur.execute("DELETE FROM qd_indicator_purchases WHERE indicator_id = ?", (indicator_id,))
|
||||
|
||||
# 删除指标
|
||||
cur.execute("DELETE FROM qd_indicator_codes WHERE id = ?", (indicator_id,))
|
||||
|
||||
db.commit()
|
||||
cur.close()
|
||||
|
||||
logger.info(f"Admin {admin_id} deleted indicator {indicator_id}")
|
||||
return True, 'success'
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"admin_delete_indicator failed: {e}")
|
||||
return False, f'error: {str(e)}'
|
||||
|
||||
def get_review_stats(self) -> Dict[str, int]:
|
||||
"""获取审核统计"""
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
cur.execute("""
|
||||
SELECT
|
||||
COUNT(*) FILTER (WHERE review_status = 'pending' OR review_status IS NULL) as pending_count,
|
||||
COUNT(*) FILTER (WHERE review_status = 'approved') as approved_count,
|
||||
COUNT(*) FILTER (WHERE review_status = 'rejected') as rejected_count
|
||||
FROM qd_indicator_codes
|
||||
WHERE publish_to_community = 1
|
||||
""")
|
||||
row = cur.fetchone()
|
||||
cur.close()
|
||||
|
||||
return {
|
||||
'pending': row['pending_count'] or 0,
|
||||
'approved': row['approved_count'] or 0,
|
||||
'rejected': row['rejected_count'] or 0
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"get_review_stats failed: {e}")
|
||||
return {'pending': 0, 'approved': 0, 'rejected': 0}
|
||||
|
||||
# ==========================================
|
||||
# 实盘表现(聚合回测数据)
|
||||
# ==========================================
|
||||
|
||||
def get_indicator_performance(self, indicator_id: int) -> Dict[str, Any]:
|
||||
"""
|
||||
获取指标的实盘表现统计
|
||||
|
||||
目前基于回测数据统计,未来可扩展为实盘交易数据
|
||||
"""
|
||||
default_result = {
|
||||
'strategy_count': 0,
|
||||
'trade_count': 0,
|
||||
'win_rate': 0,
|
||||
'total_profit': 0,
|
||||
'avg_return': 0,
|
||||
'max_drawdown': 0
|
||||
}
|
||||
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
|
||||
# 首先检查回测记录表是否存在
|
||||
cur.execute("""
|
||||
SELECT COUNT(*) as cnt FROM information_schema.tables
|
||||
WHERE table_name = 'qd_backtest_runs'
|
||||
""")
|
||||
table_exists = cur.fetchone()
|
||||
if not table_exists or table_exists['cnt'] == 0:
|
||||
cur.close()
|
||||
return default_result
|
||||
|
||||
# 从回测记录中统计该指标的表现
|
||||
# 使用 indicator_id 字段匹配
|
||||
cur.execute("""
|
||||
SELECT
|
||||
COUNT(*) as run_count,
|
||||
AVG(CASE WHEN total_return IS NOT NULL THEN total_return ELSE 0 END) as avg_return,
|
||||
AVG(CASE WHEN win_rate IS NOT NULL THEN win_rate ELSE 0 END) as avg_win_rate,
|
||||
AVG(CASE WHEN max_drawdown IS NOT NULL THEN max_drawdown ELSE 0 END) as avg_drawdown,
|
||||
SUM(CASE WHEN trade_count IS NOT NULL THEN trade_count ELSE 0 END) as total_trades
|
||||
FROM qd_backtest_runs
|
||||
WHERE indicator_id = ?
|
||||
""", (indicator_id,))
|
||||
|
||||
row = cur.fetchone()
|
||||
cur.close()
|
||||
|
||||
if not row or row['run_count'] == 0:
|
||||
return default_result
|
||||
|
||||
return {
|
||||
'strategy_count': row['run_count'] or 0,
|
||||
'trade_count': row['total_trades'] or 0,
|
||||
'win_rate': round(float(row['avg_win_rate'] or 0), 2),
|
||||
'total_profit': round(float(row['avg_return'] or 0), 2),
|
||||
'avg_return': round(float(row['avg_return'] or 0), 2),
|
||||
'max_drawdown': round(float(row['avg_drawdown'] or 0), 2)
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"get_indicator_performance failed: {e}")
|
||||
return default_result
|
||||
|
||||
|
||||
# 全局单例
|
||||
_community_service = None
|
||||
|
||||
|
||||
def get_community_service() -> CommunityService:
|
||||
"""获取社区服务单例"""
|
||||
global _community_service
|
||||
if _community_service is None:
|
||||
_community_service = CommunityService()
|
||||
return _community_service
|
||||
@@ -0,0 +1,785 @@
|
||||
"""
|
||||
Fast Analysis Service 3.0
|
||||
系统性重构版本 - 使用统一的数据采集器
|
||||
|
||||
核心改进:
|
||||
1. 数据源统一 - 使用 MarketDataCollector,与K线模块、自选列表完全一致
|
||||
2. 宏观数据 - 新增美元指数、VIX、利率等宏观经济指标
|
||||
3. 多维新闻 - 使用结构化API,无需深度阅读
|
||||
4. 单次LLM调用 - 强约束prompt,输出结构化分析
|
||||
"""
|
||||
import json
|
||||
import time
|
||||
from typing import Dict, Any, Optional, List
|
||||
from decimal import Decimal, ROUND_HALF_UP
|
||||
|
||||
from app.utils.logger import get_logger
|
||||
from app.services.llm import LLMService
|
||||
from app.services.market_data_collector import get_market_data_collector
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class FastAnalysisService:
|
||||
"""
|
||||
快速分析服务 3.0
|
||||
|
||||
架构:
|
||||
1. 数据采集层 - MarketDataCollector (统一数据源)
|
||||
2. 分析层 - 单次LLM调用 (强约束prompt)
|
||||
3. 记忆层 - 分析历史存储和检索
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.llm_service = LLMService()
|
||||
self.data_collector = get_market_data_collector()
|
||||
self._memory_db = None # Lazy init
|
||||
|
||||
# ==================== Data Collection Layer ====================
|
||||
|
||||
def _collect_market_data(self, market: str, symbol: str, timeframe: str = "1D") -> Dict[str, Any]:
|
||||
"""
|
||||
使用统一的数据采集器收集市场数据
|
||||
|
||||
数据层次:
|
||||
1. 核心数据: 价格、K线、技术指标
|
||||
2. 基本面: 公司信息、财务数据
|
||||
3. 宏观数据: DXY、VIX、TNX、黄金等
|
||||
4. 情绪数据: 新闻、市场情绪
|
||||
"""
|
||||
return self.data_collector.collect_all(
|
||||
market=market,
|
||||
symbol=symbol,
|
||||
timeframe=timeframe,
|
||||
include_macro=True,
|
||||
include_news=True,
|
||||
timeout=30
|
||||
)
|
||||
|
||||
def _calculate_indicators(self, kline_data: List[Dict]) -> Dict[str, Any]:
|
||||
"""
|
||||
Calculate technical indicators using rules (no LLM).
|
||||
Returns actionable signals, not raw numbers.
|
||||
"""
|
||||
if not kline_data or len(kline_data) < 5:
|
||||
return {"error": "Insufficient data"}
|
||||
|
||||
try:
|
||||
# Use tools' built-in calculation
|
||||
raw_indicators = self.tools.calculate_technical_indicators(kline_data)
|
||||
|
||||
# Extract key values
|
||||
closes = [float(k.get("close", 0)) for k in kline_data if k.get("close")]
|
||||
if not closes:
|
||||
return {"error": "No close prices"}
|
||||
|
||||
current_price = closes[-1]
|
||||
|
||||
# RSI interpretation
|
||||
rsi = raw_indicators.get("RSI", 50)
|
||||
if rsi < 30:
|
||||
rsi_signal = "oversold"
|
||||
rsi_action = "potential_buy"
|
||||
elif rsi > 70:
|
||||
rsi_signal = "overbought"
|
||||
rsi_action = "potential_sell"
|
||||
else:
|
||||
rsi_signal = "neutral"
|
||||
rsi_action = "hold"
|
||||
|
||||
# MACD interpretation
|
||||
macd = raw_indicators.get("MACD", 0)
|
||||
macd_signal_line = raw_indicators.get("MACD_Signal", 0)
|
||||
macd_hist = raw_indicators.get("MACD_Hist", 0)
|
||||
|
||||
if macd > macd_signal_line and macd_hist > 0:
|
||||
macd_signal = "bullish"
|
||||
macd_trend = "golden_cross" if macd_hist > 0 and len(kline_data) > 1 else "bullish"
|
||||
elif macd < macd_signal_line and macd_hist < 0:
|
||||
macd_signal = "bearish"
|
||||
macd_trend = "death_cross" if macd_hist < 0 and len(kline_data) > 1 else "bearish"
|
||||
else:
|
||||
macd_signal = "neutral"
|
||||
macd_trend = "consolidating"
|
||||
|
||||
# Moving averages
|
||||
ma5 = sum(closes[-5:]) / 5 if len(closes) >= 5 else current_price
|
||||
ma10 = sum(closes[-10:]) / 10 if len(closes) >= 10 else current_price
|
||||
ma20 = sum(closes[-20:]) / 20 if len(closes) >= 20 else current_price
|
||||
|
||||
if current_price > ma5 > ma10 > ma20:
|
||||
ma_trend = "strong_uptrend"
|
||||
elif current_price > ma20:
|
||||
ma_trend = "uptrend"
|
||||
elif current_price < ma5 < ma10 < ma20:
|
||||
ma_trend = "strong_downtrend"
|
||||
elif current_price < ma20:
|
||||
ma_trend = "downtrend"
|
||||
else:
|
||||
ma_trend = "sideways"
|
||||
|
||||
# Support/Resistance (simple: recent highs/lows)
|
||||
recent_highs = [float(k.get("high", 0)) for k in kline_data[-14:] if k.get("high")]
|
||||
recent_lows = [float(k.get("low", 0)) for k in kline_data[-14:] if k.get("low")]
|
||||
|
||||
resistance = max(recent_highs) if recent_highs else current_price * 1.05
|
||||
support = min(recent_lows) if recent_lows else current_price * 0.95
|
||||
|
||||
# Volatility (ATR-like)
|
||||
if len(kline_data) >= 14:
|
||||
ranges = []
|
||||
for k in kline_data[-14:]:
|
||||
h = float(k.get("high", 0))
|
||||
l = float(k.get("low", 0))
|
||||
if h > 0 and l > 0:
|
||||
ranges.append(h - l)
|
||||
atr = sum(ranges) / len(ranges) if ranges else 0
|
||||
volatility_pct = (atr / current_price * 100) if current_price > 0 else 0
|
||||
|
||||
if volatility_pct > 5:
|
||||
volatility = "high"
|
||||
elif volatility_pct > 2:
|
||||
volatility = "medium"
|
||||
else:
|
||||
volatility = "low"
|
||||
else:
|
||||
volatility = "unknown"
|
||||
volatility_pct = 0
|
||||
|
||||
return {
|
||||
"current_price": round(current_price, 6),
|
||||
"rsi": {
|
||||
"value": round(rsi, 2),
|
||||
"signal": rsi_signal,
|
||||
"action": rsi_action,
|
||||
},
|
||||
"macd": {
|
||||
"value": round(macd, 6),
|
||||
"signal_line": round(macd_signal_line, 6),
|
||||
"histogram": round(macd_hist, 6),
|
||||
"signal": macd_signal,
|
||||
"trend": macd_trend,
|
||||
},
|
||||
"moving_averages": {
|
||||
"ma5": round(ma5, 6),
|
||||
"ma10": round(ma10, 6),
|
||||
"ma20": round(ma20, 6),
|
||||
"trend": ma_trend,
|
||||
},
|
||||
"levels": {
|
||||
"support": round(support, 6),
|
||||
"resistance": round(resistance, 6),
|
||||
},
|
||||
"volatility": {
|
||||
"level": volatility,
|
||||
"pct": round(volatility_pct, 2),
|
||||
},
|
||||
"raw": raw_indicators,
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Indicator calculation failed: {e}")
|
||||
return {"error": str(e)}
|
||||
|
||||
def _format_news_summary(self, news_data: List[Dict], max_items: int = 5) -> str:
|
||||
"""Format news into a concise summary for the prompt."""
|
||||
if not news_data:
|
||||
return "No recent news available."
|
||||
|
||||
summaries = []
|
||||
for item in news_data[:max_items]:
|
||||
title = item.get("title", item.get("headline", ""))
|
||||
sentiment = item.get("sentiment", "neutral")
|
||||
date = item.get("date", item.get("datetime", ""))[:10] if item.get("date") or item.get("datetime") else ""
|
||||
|
||||
if title:
|
||||
summaries.append(f"- [{sentiment}] {title} ({date})")
|
||||
|
||||
return "\n".join(summaries) if summaries else "No recent news available."
|
||||
|
||||
# ==================== Memory Layer ====================
|
||||
|
||||
def _get_memory_context(self, market: str, symbol: str, current_indicators: Dict) -> str:
|
||||
"""
|
||||
Retrieve relevant historical analysis for similar market conditions.
|
||||
"""
|
||||
try:
|
||||
from app.services.analysis_memory import get_analysis_memory
|
||||
memory = get_analysis_memory()
|
||||
|
||||
# Get similar patterns
|
||||
patterns = memory.get_similar_patterns(market, symbol, current_indicators, limit=3)
|
||||
|
||||
if not patterns:
|
||||
return "No similar historical patterns found in memory."
|
||||
|
||||
context_lines = ["Historical patterns with similar conditions:"]
|
||||
for p in patterns:
|
||||
outcome = ""
|
||||
if p.get("was_correct") is not None:
|
||||
outcome = f" (Outcome: {'Correct' if p['was_correct'] else 'Incorrect'}"
|
||||
if p.get("actual_return_pct"):
|
||||
outcome += f", Return: {p['actual_return_pct']:.2f}%"
|
||||
outcome += ")"
|
||||
|
||||
context_lines.append(
|
||||
f"- Decision: {p['decision']} at ${p.get('price', 'N/A')}{outcome}"
|
||||
)
|
||||
|
||||
return "\n".join(context_lines)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Memory retrieval failed: {e}")
|
||||
return "Memory retrieval failed."
|
||||
|
||||
# ==================== Prompt Engineering ====================
|
||||
|
||||
def _build_analysis_prompt(self, data: Dict[str, Any], language: str) -> tuple:
|
||||
"""
|
||||
Build the single, comprehensive analysis prompt.
|
||||
Key: Strong constraints to prevent absurd recommendations.
|
||||
"""
|
||||
price_data = data.get("price") or {}
|
||||
current_price = price_data.get("price", 0) if price_data else 0
|
||||
change_24h = price_data.get("changePercent", 0) if price_data else 0
|
||||
|
||||
# Ensure all data fields have safe defaults (may be None from failed fetches)
|
||||
indicators = data.get("indicators") or {}
|
||||
fundamental = data.get("fundamental") or {}
|
||||
company = data.get("company") or {}
|
||||
news_summary = self._format_news_summary(data.get("news") or [])
|
||||
|
||||
# Language instruction - MUST be enforced strictly
|
||||
lang_map = {
|
||||
'zh-CN': '⚠️ 重要:你必须用简体中文回答所有内容,包括summary、key_reasons、risks等所有文本字段。不要使用英文。',
|
||||
'zh-TW': '⚠️ 重要:你必須用繁體中文回答所有內容,包括summary、key_reasons、risks等所有文本字段。不要使用英文。',
|
||||
'en-US': '⚠️ IMPORTANT: You MUST answer ALL content in English, including summary, key_reasons, risks, and all text fields. Do NOT use Chinese.',
|
||||
'ja-JP': '⚠️ 重要:すべての内容を日本語で回答してください。summary、key_reasons、risksなど、すべてのテキストフィールドを日本語で記述してください。',
|
||||
}
|
||||
lang_instruction = lang_map.get(language, '⚠️ IMPORTANT: Answer ALL content in English.')
|
||||
|
||||
# Get pre-calculated trading levels from technical analysis
|
||||
levels = indicators.get("levels", {})
|
||||
trading_levels = indicators.get("trading_levels", {})
|
||||
volatility = indicators.get("volatility", {})
|
||||
|
||||
support = levels.get("support", current_price * 0.95)
|
||||
resistance = levels.get("resistance", current_price * 1.05)
|
||||
pivot = levels.get("pivot", current_price)
|
||||
|
||||
# Use ATR-based suggestions if available, otherwise use percentage
|
||||
atr = volatility.get("atr", current_price * 0.02)
|
||||
suggested_stop_loss = trading_levels.get("suggested_stop_loss", current_price - 2 * atr)
|
||||
suggested_take_profit = trading_levels.get("suggested_take_profit", current_price + 3 * atr)
|
||||
risk_reward_ratio = trading_levels.get("risk_reward_ratio", 1.5)
|
||||
|
||||
# Price bounds (still enforce max 10% deviation)
|
||||
if current_price > 0:
|
||||
price_lower_bound = round(max(suggested_stop_loss, current_price * 0.90), 6)
|
||||
price_upper_bound = round(min(suggested_take_profit, current_price * 1.10), 6)
|
||||
entry_range_low = round(current_price * 0.98, 6)
|
||||
entry_range_high = round(current_price * 1.02, 6)
|
||||
else:
|
||||
price_lower_bound = price_upper_bound = entry_range_low = entry_range_high = 0
|
||||
|
||||
system_prompt = f"""You are QuantDinger's Senior Financial Analyst with 20+ years of experience.
|
||||
Provide professional, detailed analysis like a Wall Street analyst report.
|
||||
|
||||
{lang_instruction}
|
||||
|
||||
📐 TECHNICAL LEVELS (Pre-calculated from chart data):
|
||||
- Support: ${support} | Resistance: ${resistance} | Pivot: ${pivot}
|
||||
- ATR (14-day): ${atr:.4f} ({volatility.get('pct', 0)}% volatility)
|
||||
- Suggested Stop Loss: ${suggested_stop_loss:.4f} (based on 2x ATR below support)
|
||||
- Suggested Take Profit: ${suggested_take_profit:.4f} (based on 3x ATR above resistance)
|
||||
- Risk/Reward Ratio: {risk_reward_ratio}
|
||||
|
||||
⚠️ CRITICAL PRICE RULES:
|
||||
1. Current price: ${current_price}
|
||||
2. Your stop_loss MUST be near ${suggested_stop_loss:.4f} (range: ${price_lower_bound:.4f} ~ ${current_price})
|
||||
3. Your take_profit MUST be near ${suggested_take_profit:.4f} (range: ${current_price} ~ ${price_upper_bound:.4f})
|
||||
4. Entry price: ${entry_range_low:.4f} ~ ${entry_range_high:.4f}
|
||||
5. These levels are based on ATR and support/resistance analysis - use them as reference!
|
||||
|
||||
📊 YOUR ANALYSIS MUST INCLUDE:
|
||||
1. **Technical Analysis**: Interpret the indicators, explain why support/resistance levels matter
|
||||
2. **Fundamental Analysis**: Evaluate valuation, growth if data available
|
||||
3. **Sentiment Analysis**: Assess market mood, news impact, macro factors
|
||||
4. **Risk Assessment**: Explain why the stop loss level is appropriate
|
||||
5. **Clear Recommendation**: BUY/SELL/HOLD with entry, stop loss (near suggested), take profit (near suggested)
|
||||
|
||||
Output ONLY valid JSON (do NOT include word counts or format hints in your actual response):
|
||||
{{
|
||||
"decision": "BUY" | "SELL" | "HOLD",
|
||||
"confidence": 0-100,
|
||||
"summary": "Executive summary in 2-3 sentences",
|
||||
"analysis": {{
|
||||
"technical": "Your detailed technical analysis here - interpret RSI, MACD, MA, support/resistance",
|
||||
"fundamental": "Your fundamental assessment here - valuation, growth, competitive position",
|
||||
"sentiment": "Your market sentiment analysis here - news impact, macro factors, mood"
|
||||
}},
|
||||
"entry_price": number,
|
||||
"stop_loss": number,
|
||||
"take_profit": number,
|
||||
"position_size_pct": 1-100,
|
||||
"timeframe": "short" | "medium" | "long",
|
||||
"key_reasons": ["First key reason for this decision", "Second key reason", "Third key reason"],
|
||||
"risks": ["Primary risk with potential impact", "Secondary risk"],
|
||||
"technical_score": 0-100,
|
||||
"fundamental_score": 0-100,
|
||||
"sentiment_score": 0-100
|
||||
}}
|
||||
|
||||
⚠️ IMPORTANT: The analysis fields should contain your ACTUAL analysis text, NOT the format description above."""
|
||||
|
||||
# Format indicator data for prompt (ensure safe defaults)
|
||||
rsi_data = indicators.get("rsi") or {}
|
||||
macd_data = indicators.get("macd") or {}
|
||||
ma_data = indicators.get("moving_averages") or {}
|
||||
vol_data = indicators.get("volatility") or {}
|
||||
levels = indicators.get("levels") or {}
|
||||
|
||||
# Format macro data
|
||||
macro = data.get("macro") or {}
|
||||
macro_summary = self._format_macro_summary(macro, data.get("market", ""))
|
||||
|
||||
user_prompt = f"""Analyze {data['symbol']} in {data['market']} market.
|
||||
|
||||
📊 REAL-TIME DATA:
|
||||
- Current Price: ${current_price}
|
||||
- 24h Change: {change_24h}%
|
||||
- Support: ${support}
|
||||
- Resistance: ${resistance}
|
||||
|
||||
📈 TECHNICAL INDICATORS:
|
||||
- RSI(14): {rsi_data.get('value', 'N/A')} ({rsi_data.get('signal', 'N/A')})
|
||||
- MACD: {macd_data.get('signal', 'N/A')} ({macd_data.get('trend', 'N/A')})
|
||||
- MA Trend: {ma_data.get('trend', 'N/A')}
|
||||
- Volatility: {vol_data.get('level', 'N/A')} ({vol_data.get('pct', 0)}%)
|
||||
- Trend: {indicators.get('trend', 'N/A')}
|
||||
- Price Position (20d): {indicators.get('price_position', 'N/A')}%
|
||||
|
||||
🌐 MACRO ENVIRONMENT:
|
||||
{macro_summary}
|
||||
|
||||
📰 MARKET NEWS ({len(data.get('news') or [])} items):
|
||||
{news_summary}
|
||||
|
||||
💼 FUNDAMENTALS:
|
||||
- Company: {company.get('name', data['symbol'])}
|
||||
- Industry: {company.get('industry', 'N/A')}
|
||||
- P/E Ratio: {fundamental.get('pe_ratio', 'N/A')}
|
||||
- P/B Ratio: {fundamental.get('pb_ratio', 'N/A')}
|
||||
- Market Cap: {fundamental.get('market_cap', 'N/A')}
|
||||
- 52W High/Low: {fundamental.get('52w_high', 'N/A')} / {fundamental.get('52w_low', 'N/A')}
|
||||
- ROE: {fundamental.get('roe', 'N/A')}
|
||||
|
||||
IMPORTANT: Consider the macro environment (especially DXY, VIX, rates) when making your recommendation.
|
||||
Provide your analysis now. Remember: all prices must be within 10% of ${current_price}."""
|
||||
|
||||
return system_prompt, user_prompt
|
||||
|
||||
def _format_macro_summary(self, macro: Dict[str, Any], market: str) -> str:
|
||||
"""格式化宏观数据摘要"""
|
||||
if not macro:
|
||||
return "宏观数据暂不可用"
|
||||
|
||||
lines = []
|
||||
|
||||
# 美元指数
|
||||
if 'DXY' in macro:
|
||||
dxy = macro['DXY']
|
||||
direction = "↑" if dxy.get('change', 0) > 0 else "↓"
|
||||
lines.append(f"- {dxy.get('name', 'USD Index')}: {dxy.get('price', 'N/A')} ({direction}{abs(dxy.get('changePercent', 0)):.2f}%)")
|
||||
# 美元强弱对不同资产的影响
|
||||
if market == 'Crypto':
|
||||
impact = "利空加密货币" if dxy.get('change', 0) > 0 else "利好加密货币"
|
||||
lines.append(f" ⚠️ 美元{direction} {impact}")
|
||||
elif market == 'Forex':
|
||||
lines.append(f" ⚠️ 美元{direction} 直接影响外汇走势")
|
||||
|
||||
# VIX恐慌指数
|
||||
if 'VIX' in macro:
|
||||
vix = macro['VIX']
|
||||
vix_value = vix.get('price', 0)
|
||||
if vix_value > 30:
|
||||
level = "极度恐慌 (>30)"
|
||||
elif vix_value > 20:
|
||||
level = "较高恐慌 (20-30)"
|
||||
elif vix_value > 15:
|
||||
level = "正常 (15-20)"
|
||||
else:
|
||||
level = "低波动 (<15)"
|
||||
lines.append(f"- {vix.get('name', 'VIX')}: {vix_value:.2f} - {level}")
|
||||
|
||||
# 美债收益率
|
||||
if 'TNX' in macro:
|
||||
tnx = macro['TNX']
|
||||
direction = "↑" if tnx.get('change', 0) > 0 else "↓"
|
||||
lines.append(f"- {tnx.get('name', '10Y Treasury')}: {tnx.get('price', 'N/A'):.3f}% ({direction})")
|
||||
if tnx.get('price', 0) > 4.5:
|
||||
lines.append(" ⚠️ 高利率环境,对估值不利")
|
||||
|
||||
# 黄金
|
||||
if 'GOLD' in macro:
|
||||
gold = macro['GOLD']
|
||||
direction = "↑" if gold.get('change', 0) > 0 else "↓"
|
||||
lines.append(f"- {gold.get('name', 'Gold')}: ${gold.get('price', 'N/A'):.2f} ({direction}{abs(gold.get('changePercent', 0)):.2f}%)")
|
||||
|
||||
# 标普500
|
||||
if 'SPY' in macro:
|
||||
spy = macro['SPY']
|
||||
direction = "↑" if spy.get('change', 0) > 0 else "↓"
|
||||
lines.append(f"- {spy.get('name', 'S&P 500')}: ${spy.get('price', 'N/A'):.2f} ({direction}{abs(spy.get('changePercent', 0)):.2f}%)")
|
||||
|
||||
# 比特币 (作为风险指标)
|
||||
if 'BTC' in macro and market != 'Crypto':
|
||||
btc = macro['BTC']
|
||||
direction = "↑" if btc.get('change', 0) > 0 else "↓"
|
||||
lines.append(f"- {btc.get('name', 'BTC')}: ${btc.get('price', 'N/A'):,.0f} ({direction}{abs(btc.get('changePercent', 0)):.2f}%) [风险偏好指标]")
|
||||
|
||||
return "\n".join(lines) if lines else "宏观数据暂不可用"
|
||||
|
||||
# ==================== Main Analysis ====================
|
||||
|
||||
def analyze(self, market: str, symbol: str, language: str = 'en-US',
|
||||
model: str = None, timeframe: str = "1D") -> Dict[str, Any]:
|
||||
"""
|
||||
Run fast single-call analysis.
|
||||
|
||||
Returns:
|
||||
Complete analysis result with actionable recommendations.
|
||||
"""
|
||||
start_time = time.time()
|
||||
|
||||
result = {
|
||||
"market": market,
|
||||
"symbol": symbol,
|
||||
"language": language,
|
||||
"timeframe": timeframe,
|
||||
"analysis_time_ms": 0,
|
||||
"error": None,
|
||||
}
|
||||
|
||||
try:
|
||||
# Phase 1: Data collection (parallel)
|
||||
logger.info(f"Fast analysis starting: {market}:{symbol}")
|
||||
data = self._collect_market_data(market, symbol, timeframe)
|
||||
|
||||
# Validate we have essential data - with fallback to indicators
|
||||
current_price = None
|
||||
|
||||
# 优先从 price 数据获取
|
||||
if data.get("price") and data["price"].get("price"):
|
||||
current_price = data["price"]["price"]
|
||||
|
||||
# Fallback: 从 indicators 获取 (如果 K 线成功计算了)
|
||||
if not current_price and data.get("indicators"):
|
||||
current_price = data["indicators"].get("current_price")
|
||||
if current_price:
|
||||
logger.info(f"Using price from indicators: ${current_price}")
|
||||
# 构建简化的 price 数据
|
||||
data["price"] = {
|
||||
"price": current_price,
|
||||
"change": 0,
|
||||
"changePercent": 0,
|
||||
"source": "indicators_fallback"
|
||||
}
|
||||
|
||||
# Fallback: 从 kline 最后一根获取
|
||||
if not current_price and data.get("kline"):
|
||||
klines = data["kline"]
|
||||
if klines and len(klines) > 0:
|
||||
current_price = float(klines[-1].get("close", 0))
|
||||
if current_price > 0:
|
||||
logger.info(f"Using price from kline: ${current_price}")
|
||||
prev_close = float(klines[-2].get("close", current_price)) if len(klines) > 1 else current_price
|
||||
change = current_price - prev_close
|
||||
change_pct = (change / prev_close * 100) if prev_close > 0 else 0
|
||||
data["price"] = {
|
||||
"price": current_price,
|
||||
"change": round(change, 6),
|
||||
"changePercent": round(change_pct, 2),
|
||||
"source": "kline_fallback"
|
||||
}
|
||||
|
||||
if not current_price or current_price <= 0:
|
||||
result["error"] = "Failed to fetch current price from all sources"
|
||||
logger.error(f"Price fetch failed for {market}:{symbol}, all sources exhausted")
|
||||
return result
|
||||
|
||||
# Phase 2: Build prompt
|
||||
system_prompt, user_prompt = self._build_analysis_prompt(data, language)
|
||||
|
||||
# Phase 3: Single LLM call
|
||||
logger.info(f"Calling LLM for analysis...")
|
||||
llm_start = time.time()
|
||||
|
||||
analysis = self.llm_service.safe_call_llm(
|
||||
system_prompt,
|
||||
user_prompt,
|
||||
default_structure={
|
||||
"decision": "HOLD",
|
||||
"confidence": 50,
|
||||
"summary": "Analysis failed",
|
||||
"entry_price": current_price,
|
||||
"stop_loss": current_price * 0.95,
|
||||
"take_profit": current_price * 1.05,
|
||||
"position_size_pct": 10,
|
||||
"timeframe": "medium",
|
||||
"key_reasons": ["Unable to analyze"],
|
||||
"risks": ["Analysis error"],
|
||||
"technical_score": 50,
|
||||
"fundamental_score": 50,
|
||||
"sentiment_score": 50,
|
||||
},
|
||||
model=model
|
||||
)
|
||||
|
||||
llm_time = int((time.time() - llm_start) * 1000)
|
||||
logger.info(f"LLM call completed in {llm_time}ms")
|
||||
|
||||
# Phase 4: Validate and constrain output
|
||||
analysis = self._validate_and_constrain(analysis, current_price)
|
||||
|
||||
# Build final result
|
||||
total_time = int((time.time() - start_time) * 1000)
|
||||
|
||||
# Extract detailed analysis sections
|
||||
detailed_analysis = analysis.get("analysis", {})
|
||||
if isinstance(detailed_analysis, str):
|
||||
# If AI returned a string instead of dict, use it as technical analysis
|
||||
detailed_analysis = {"technical": detailed_analysis, "fundamental": "", "sentiment": ""}
|
||||
|
||||
result.update({
|
||||
"decision": analysis.get("decision", "HOLD"),
|
||||
"confidence": analysis.get("confidence", 50),
|
||||
"summary": analysis.get("summary", ""),
|
||||
"detailed_analysis": {
|
||||
"technical": detailed_analysis.get("technical", ""),
|
||||
"fundamental": detailed_analysis.get("fundamental", ""),
|
||||
"sentiment": detailed_analysis.get("sentiment", ""),
|
||||
},
|
||||
"trading_plan": {
|
||||
"entry_price": analysis.get("entry_price"),
|
||||
"stop_loss": analysis.get("stop_loss"),
|
||||
"take_profit": analysis.get("take_profit"),
|
||||
"position_size_pct": analysis.get("position_size_pct", 10),
|
||||
"timeframe": analysis.get("timeframe", "medium"),
|
||||
},
|
||||
"reasons": analysis.get("key_reasons", []),
|
||||
"risks": analysis.get("risks", []),
|
||||
"scores": {
|
||||
"technical": analysis.get("technical_score", 50),
|
||||
"fundamental": analysis.get("fundamental_score", 50),
|
||||
"sentiment": analysis.get("sentiment_score", 50),
|
||||
"overall": self._calculate_overall_score(analysis),
|
||||
},
|
||||
"market_data": {
|
||||
"current_price": current_price,
|
||||
"change_24h": data["price"].get("changePercent", 0),
|
||||
"support": data["indicators"].get("levels", {}).get("support"),
|
||||
"resistance": data["indicators"].get("levels", {}).get("resistance"),
|
||||
},
|
||||
"indicators": data.get("indicators", {}),
|
||||
"analysis_time_ms": total_time,
|
||||
"llm_time_ms": llm_time,
|
||||
"data_collection_time_ms": data.get("collection_time_ms", 0),
|
||||
})
|
||||
|
||||
# Store in memory for future retrieval and get memory_id for feedback
|
||||
memory_id = self._store_analysis_memory(result)
|
||||
if memory_id:
|
||||
result["memory_id"] = memory_id
|
||||
|
||||
logger.info(f"Fast analysis completed in {total_time}ms: {market}:{symbol} -> {result['decision']} (memory_id={memory_id})")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Fast analysis failed: {e}", exc_info=True)
|
||||
result["error"] = str(e)
|
||||
|
||||
return result
|
||||
|
||||
def _validate_and_constrain(self, analysis: Dict, current_price: float) -> Dict:
|
||||
"""
|
||||
Validate LLM output and constrain prices to reasonable ranges.
|
||||
This prevents absurd recommendations like "BTC at 95000, buy at 75000".
|
||||
"""
|
||||
if not current_price or current_price <= 0:
|
||||
return analysis
|
||||
|
||||
# Price bounds
|
||||
min_price = current_price * 0.90
|
||||
max_price = current_price * 1.10
|
||||
|
||||
# Constrain entry price
|
||||
entry = analysis.get("entry_price", current_price)
|
||||
if entry and (entry < min_price or entry > max_price):
|
||||
logger.warning(f"Entry price {entry} out of bounds, constraining to current price {current_price}")
|
||||
analysis["entry_price"] = round(current_price, 6)
|
||||
|
||||
# Constrain stop loss
|
||||
stop_loss = analysis.get("stop_loss", current_price * 0.95)
|
||||
if stop_loss and (stop_loss < min_price or stop_loss > current_price):
|
||||
analysis["stop_loss"] = round(current_price * 0.95, 6)
|
||||
|
||||
# Constrain take profit
|
||||
take_profit = analysis.get("take_profit", current_price * 1.05)
|
||||
if take_profit and (take_profit < current_price or take_profit > max_price):
|
||||
analysis["take_profit"] = round(current_price * 1.05, 6)
|
||||
|
||||
# Constrain confidence
|
||||
confidence = analysis.get("confidence", 50)
|
||||
analysis["confidence"] = max(0, min(100, int(confidence)))
|
||||
|
||||
# Constrain scores
|
||||
for score_key in ["technical_score", "fundamental_score", "sentiment_score"]:
|
||||
score = analysis.get(score_key, 50)
|
||||
analysis[score_key] = max(0, min(100, int(score)))
|
||||
|
||||
# Validate decision
|
||||
decision = str(analysis.get("decision", "HOLD")).upper()
|
||||
if decision not in ["BUY", "SELL", "HOLD"]:
|
||||
analysis["decision"] = "HOLD"
|
||||
else:
|
||||
analysis["decision"] = decision
|
||||
|
||||
return analysis
|
||||
|
||||
def _calculate_overall_score(self, analysis: Dict) -> int:
|
||||
"""Calculate weighted overall score."""
|
||||
tech = analysis.get("technical_score", 50)
|
||||
fund = analysis.get("fundamental_score", 50)
|
||||
sent = analysis.get("sentiment_score", 50)
|
||||
|
||||
# Weights: technical 40%, fundamental 35%, sentiment 25%
|
||||
overall = tech * 0.40 + fund * 0.35 + sent * 0.25
|
||||
|
||||
# Adjust based on decision
|
||||
decision = analysis.get("decision", "HOLD")
|
||||
confidence = analysis.get("confidence", 50)
|
||||
|
||||
if decision == "BUY":
|
||||
overall = overall * 0.6 + (50 + confidence * 0.5) * 0.4
|
||||
elif decision == "SELL":
|
||||
overall = overall * 0.6 + (50 - confidence * 0.5) * 0.4
|
||||
|
||||
return max(0, min(100, int(overall)))
|
||||
|
||||
def _store_analysis_memory(self, result: Dict) -> Optional[int]:
|
||||
"""Store analysis result for future learning. Returns memory_id."""
|
||||
try:
|
||||
from app.services.analysis_memory import get_analysis_memory
|
||||
memory = get_analysis_memory()
|
||||
memory_id = memory.store(result)
|
||||
return memory_id
|
||||
except Exception as e:
|
||||
logger.warning(f"Memory storage failed: {e}")
|
||||
return None
|
||||
|
||||
# ==================== Backward Compatibility ====================
|
||||
|
||||
def analyze_legacy_format(self, market: str, symbol: str, language: str = 'en-US',
|
||||
model: str = None, timeframe: str = "1D") -> Dict[str, Any]:
|
||||
"""
|
||||
Returns analysis in legacy multi-agent format for backward compatibility.
|
||||
"""
|
||||
fast_result = self.analyze(market, symbol, language, model, timeframe)
|
||||
|
||||
if fast_result.get("error"):
|
||||
return {
|
||||
"overview": {"report": f"Analysis failed: {fast_result['error']}"},
|
||||
"fundamental": {"report": "N/A"},
|
||||
"technical": {"report": "N/A"},
|
||||
"news": {"report": "N/A"},
|
||||
"sentiment": {"report": "N/A"},
|
||||
"risk": {"report": "N/A"},
|
||||
"error": fast_result["error"],
|
||||
}
|
||||
|
||||
# Convert to legacy format
|
||||
decision = fast_result.get("decision", "HOLD")
|
||||
confidence = fast_result.get("confidence", 50)
|
||||
scores = fast_result.get("scores", {})
|
||||
|
||||
return {
|
||||
"overview": {
|
||||
"overallScore": scores.get("overall", 50),
|
||||
"recommendation": decision,
|
||||
"confidence": confidence,
|
||||
"dimensionScores": {
|
||||
"fundamental": scores.get("fundamental", 50),
|
||||
"technical": scores.get("technical", 50),
|
||||
"news": scores.get("sentiment", 50),
|
||||
"sentiment": scores.get("sentiment", 50),
|
||||
"risk": 100 - confidence, # Inverse of confidence
|
||||
},
|
||||
"report": fast_result.get("summary", ""),
|
||||
},
|
||||
"fundamental": {
|
||||
"score": scores.get("fundamental", 50),
|
||||
"report": f"Fundamental score: {scores.get('fundamental', 50)}/100",
|
||||
},
|
||||
"technical": {
|
||||
"score": scores.get("technical", 50),
|
||||
"report": f"Technical score: {scores.get('technical', 50)}/100",
|
||||
"indicators": fast_result.get("indicators", {}),
|
||||
},
|
||||
"news": {
|
||||
"score": scores.get("sentiment", 50),
|
||||
"report": "See sentiment analysis",
|
||||
},
|
||||
"sentiment": {
|
||||
"score": scores.get("sentiment", 50),
|
||||
"report": f"Sentiment score: {scores.get('sentiment', 50)}/100",
|
||||
},
|
||||
"risk": {
|
||||
"score": 100 - confidence,
|
||||
"report": "\n".join(fast_result.get("risks", [])),
|
||||
},
|
||||
"debate": {
|
||||
"bull": {"confidence": confidence if decision == "BUY" else 50},
|
||||
"bear": {"confidence": confidence if decision == "SELL" else 50},
|
||||
"research_decision": fast_result.get("summary", ""),
|
||||
},
|
||||
"trader_decision": {
|
||||
"decision": decision,
|
||||
"confidence": confidence,
|
||||
"reasoning": fast_result.get("summary", ""),
|
||||
"trading_plan": fast_result.get("trading_plan", {}),
|
||||
"report": "\n".join(fast_result.get("reasons", [])),
|
||||
},
|
||||
"risk_debate": {
|
||||
"risky": {"recommendation": ""},
|
||||
"neutral": {"recommendation": fast_result.get("summary", "")},
|
||||
"safe": {"recommendation": ""},
|
||||
},
|
||||
"final_decision": {
|
||||
"decision": decision,
|
||||
"confidence": confidence,
|
||||
"reasoning": fast_result.get("summary", ""),
|
||||
"risk_summary": {
|
||||
"risks": fast_result.get("risks", []),
|
||||
},
|
||||
"recommendation": "\n".join(fast_result.get("reasons", [])),
|
||||
},
|
||||
"fast_analysis": fast_result, # Include new format for gradual migration
|
||||
"error": None,
|
||||
}
|
||||
|
||||
|
||||
# Singleton instance
|
||||
_fast_analysis_service = None
|
||||
|
||||
def get_fast_analysis_service() -> FastAnalysisService:
|
||||
"""Get singleton FastAnalysisService instance."""
|
||||
global _fast_analysis_service
|
||||
if _fast_analysis_service is None:
|
||||
_fast_analysis_service = FastAnalysisService()
|
||||
return _fast_analysis_service
|
||||
|
||||
|
||||
def fast_analyze(market: str, symbol: str, language: str = 'en-US',
|
||||
model: str = None, timeframe: str = "1D") -> Dict[str, Any]:
|
||||
"""Convenience function for fast analysis."""
|
||||
service = get_fast_analysis_service()
|
||||
return service.analyze(market, symbol, language, model, timeframe)
|
||||
@@ -155,6 +155,7 @@ class DeepcoinClient(BaseRestClient):
|
||||
"DC-ACCESS-TIMESTAMP": iso_time,
|
||||
"DC-ACCESS-PASSPHRASE": self.passphrase,
|
||||
"Content-Type": "application/json",
|
||||
"appid": "200103",
|
||||
}
|
||||
return headers
|
||||
|
||||
|
||||
@@ -0,0 +1,976 @@
|
||||
"""
|
||||
市场数据采集服务 - AI分析专用
|
||||
|
||||
设计理念:
|
||||
1. 数据为王 - 先把数据获取做好、做稳定
|
||||
2. 统一数据源 - 完全复用 DataSourceFactory 和 kline_service
|
||||
3. 复用全球金融板块 - 宏观数据、情绪数据复用 global_market.py 的缓存
|
||||
4. 快速稳定 - 不依赖慢速外部服务(如Jina Reader)
|
||||
|
||||
数据源映射:
|
||||
- 价格/K线: DataSourceFactory (已验证,与K线模块、自选列表一致)
|
||||
- 宏观数据: 复用 global_market.py (VIX, DXY, TNX, Fear&Greed等,带缓存)
|
||||
- 新闻: Finnhub API (结构化数据,无需深度阅读)
|
||||
- 基本面: Finnhub (美股) / akshare (A股) / 固定描述 (加密)
|
||||
"""
|
||||
|
||||
import time
|
||||
from typing import Dict, List, Any, Optional
|
||||
from datetime import datetime, timedelta
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed, TimeoutError
|
||||
|
||||
import yfinance as yf
|
||||
|
||||
from app.data_sources import DataSourceFactory
|
||||
from app.services.kline import KlineService
|
||||
from app.utils.logger import get_logger
|
||||
from app.config import APIKeys
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class MarketDataCollector:
|
||||
"""
|
||||
市场数据采集器
|
||||
|
||||
职责:为AI分析提供完整、准确、及时的市场数据
|
||||
|
||||
数据层次:
|
||||
1. 核心数据 (必须成功): 价格、K线
|
||||
2. 分析数据 (增强): 技术指标、基本面
|
||||
3. 宏观数据 (可选): 复用 global_market.py (VIX, DXY, TNX, Fear&Greed等)
|
||||
4. 情绪数据 (可选): 新闻、市场情绪
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.kline_service = KlineService()
|
||||
self._finnhub_client = None
|
||||
self._ak = None
|
||||
self._init_clients()
|
||||
|
||||
def _init_clients(self):
|
||||
"""初始化外部API客户端"""
|
||||
# Finnhub
|
||||
finnhub_key = APIKeys.FINNHUB_API_KEY
|
||||
if finnhub_key:
|
||||
try:
|
||||
import finnhub
|
||||
self._finnhub_client = finnhub.Client(api_key=finnhub_key)
|
||||
except Exception as e:
|
||||
logger.warning(f"Finnhub client init failed: {e}")
|
||||
|
||||
# akshare
|
||||
try:
|
||||
import akshare as ak
|
||||
self._ak = ak
|
||||
except ImportError:
|
||||
logger.info("akshare not installed, A-share data will be limited")
|
||||
|
||||
def collect_all(
|
||||
self,
|
||||
market: str,
|
||||
symbol: str,
|
||||
timeframe: str = "1D",
|
||||
include_macro: bool = True,
|
||||
include_news: bool = True,
|
||||
timeout: int = 30
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
采集所有市场数据
|
||||
|
||||
Args:
|
||||
market: 市场类型 (USStock, Crypto, AShare, HShare, Forex, Futures)
|
||||
symbol: 标的代码
|
||||
timeframe: K线周期
|
||||
include_macro: 是否包含宏观数据
|
||||
include_news: 是否包含新闻
|
||||
timeout: 总超时时间(秒)
|
||||
|
||||
Returns:
|
||||
完整的市场数据字典
|
||||
"""
|
||||
start_time = time.time()
|
||||
|
||||
data = {
|
||||
"market": market,
|
||||
"symbol": symbol,
|
||||
"timeframe": timeframe,
|
||||
"collected_at": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
|
||||
# 核心数据
|
||||
"price": None,
|
||||
"kline": None,
|
||||
"indicators": {},
|
||||
# 基本面
|
||||
"fundamental": {},
|
||||
"company": {},
|
||||
# 宏观
|
||||
"macro": {},
|
||||
# 情绪
|
||||
"news": [],
|
||||
"sentiment": {},
|
||||
# 元数据
|
||||
"_meta": {
|
||||
"success_items": [],
|
||||
"failed_items": [],
|
||||
"duration_ms": 0
|
||||
}
|
||||
}
|
||||
|
||||
# === 阶段1: 核心数据 (并行获取) ===
|
||||
with ThreadPoolExecutor(max_workers=4) as executor:
|
||||
core_futures = {
|
||||
executor.submit(self._get_price, market, symbol): "price",
|
||||
executor.submit(self._get_kline, market, symbol, timeframe, 60): "kline",
|
||||
}
|
||||
|
||||
# 如果需要基本面,也并行获取
|
||||
if market in ('USStock', 'AShare', 'HShare'):
|
||||
core_futures[executor.submit(self._get_fundamental, market, symbol)] = "fundamental"
|
||||
core_futures[executor.submit(self._get_company, market, symbol)] = "company"
|
||||
elif market == 'Crypto':
|
||||
# 加密货币的"基本面"是固定描述
|
||||
core_futures[executor.submit(self._get_crypto_info, symbol)] = "fundamental"
|
||||
|
||||
try:
|
||||
for future in as_completed(core_futures, timeout=15):
|
||||
key = core_futures[future]
|
||||
try:
|
||||
result = future.result(timeout=3)
|
||||
if result:
|
||||
data[key] = result
|
||||
data["_meta"]["success_items"].append(key)
|
||||
else:
|
||||
data["_meta"]["failed_items"].append(key)
|
||||
except Exception as e:
|
||||
logger.warning(f"Core data fetch failed ({key}): {e}")
|
||||
data["_meta"]["failed_items"].append(key)
|
||||
except TimeoutError:
|
||||
logger.warning(f"Core data fetch timed out for {market}:{symbol}")
|
||||
|
||||
# 计算技术指标 (本地计算,不需要外部API)
|
||||
if data.get("kline"):
|
||||
data["indicators"] = self._calculate_indicators(data["kline"])
|
||||
data["_meta"]["success_items"].append("indicators")
|
||||
|
||||
# === 阶段2: 宏观数据 (如果需要) ===
|
||||
if include_macro:
|
||||
try:
|
||||
data["macro"] = self._get_macro_data(market, timeout=10)
|
||||
if data["macro"]:
|
||||
data["_meta"]["success_items"].append("macro")
|
||||
except Exception as e:
|
||||
logger.warning(f"Macro data fetch failed: {e}")
|
||||
data["_meta"]["failed_items"].append("macro")
|
||||
|
||||
# === 阶段3: 新闻/情绪 (如果需要) ===
|
||||
if include_news:
|
||||
try:
|
||||
# 获取公司名称以改善搜索
|
||||
company_name = None
|
||||
if data.get("company"):
|
||||
company_name = data["company"].get("name")
|
||||
|
||||
news_result = self._get_news(market, symbol, company_name, timeout=8)
|
||||
data["news"] = news_result.get("news", [])
|
||||
data["sentiment"] = news_result.get("sentiment", {})
|
||||
|
||||
if data["news"]:
|
||||
data["_meta"]["success_items"].append("news")
|
||||
except Exception as e:
|
||||
logger.warning(f"News fetch failed: {e}")
|
||||
data["_meta"]["failed_items"].append("news")
|
||||
|
||||
# 记录总耗时
|
||||
data["_meta"]["duration_ms"] = int((time.time() - start_time) * 1000)
|
||||
logger.info(f"Market data collection completed for {market}:{symbol} in {data['_meta']['duration_ms']}ms")
|
||||
logger.info(f" Success: {data['_meta']['success_items']}")
|
||||
logger.info(f" Failed: {data['_meta']['failed_items']}")
|
||||
|
||||
return data
|
||||
|
||||
# ==================== 核心数据获取 ====================
|
||||
|
||||
def _get_price(self, market: str, symbol: str) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
获取实时价格 - 使用 kline_service (与自选列表一致)
|
||||
"""
|
||||
try:
|
||||
price_data = self.kline_service.get_realtime_price(market, symbol, force_refresh=True)
|
||||
if price_data and price_data.get('price', 0) > 0:
|
||||
# 安全转换为 float,处理 None 值
|
||||
def safe_float(val, default=0.0):
|
||||
if val is None:
|
||||
return default
|
||||
try:
|
||||
return float(val)
|
||||
except (ValueError, TypeError):
|
||||
return default
|
||||
|
||||
price = safe_float(price_data.get('price'))
|
||||
return {
|
||||
"price": price,
|
||||
"change": safe_float(price_data.get('change')),
|
||||
"changePercent": safe_float(price_data.get('changePercent')),
|
||||
"high": safe_float(price_data.get('high'), price),
|
||||
"low": safe_float(price_data.get('low'), price),
|
||||
"open": safe_float(price_data.get('open'), price),
|
||||
"previousClose": safe_float(price_data.get('previousClose'), price),
|
||||
"source": price_data.get('source', 'unknown')
|
||||
}
|
||||
except Exception as e:
|
||||
logger.warning(f"Price fetch failed for {market}:{symbol}: {e}")
|
||||
|
||||
# 如果 kline_service 失败,尝试从 K 线最后一根获取价格
|
||||
try:
|
||||
klines = DataSourceFactory.get_kline(market, symbol, "1D", 2)
|
||||
if klines and len(klines) > 0:
|
||||
latest = klines[-1]
|
||||
price = float(latest.get('close', 0))
|
||||
if price > 0:
|
||||
prev_close = float(klines[-2].get('close', price)) if len(klines) > 1 else price
|
||||
change = price - prev_close
|
||||
change_pct = (change / prev_close * 100) if prev_close > 0 else 0
|
||||
|
||||
logger.info(f"Price fetched from K-line fallback for {market}:{symbol}: ${price}")
|
||||
return {
|
||||
"price": price,
|
||||
"change": round(change, 6),
|
||||
"changePercent": round(change_pct, 2),
|
||||
"high": float(latest.get('high', price)),
|
||||
"low": float(latest.get('low', price)),
|
||||
"open": float(latest.get('open', price)),
|
||||
"previousClose": prev_close,
|
||||
"source": "kline_fallback"
|
||||
}
|
||||
except Exception as e:
|
||||
logger.warning(f"K-line fallback price fetch also failed for {market}:{symbol}: {e}")
|
||||
|
||||
return None
|
||||
|
||||
def _get_kline(
|
||||
self, market: str, symbol: str, timeframe: str, limit: int = 60
|
||||
) -> Optional[List[Dict[str, Any]]]:
|
||||
"""
|
||||
获取K线数据 - 使用 DataSourceFactory (与K线模块一致)
|
||||
"""
|
||||
try:
|
||||
klines = DataSourceFactory.get_kline(market, symbol, timeframe, limit)
|
||||
if klines and len(klines) > 0:
|
||||
return klines
|
||||
except Exception as e:
|
||||
logger.warning(f"Kline fetch failed for {market}:{symbol}: {e}")
|
||||
return None
|
||||
|
||||
def _calculate_indicators(self, klines: List[Dict[str, Any]]) -> Dict[str, Any]:
|
||||
"""
|
||||
计算技术指标 (本地计算,无外部依赖)
|
||||
|
||||
返回格式符合前端 FastAnalysisReport.vue 的期望:
|
||||
{
|
||||
rsi: { value, signal },
|
||||
macd: { signal, trend },
|
||||
moving_averages: { ma5, ma10, ma20, trend },
|
||||
levels: { support, resistance },
|
||||
volatility: { level, pct }
|
||||
}
|
||||
"""
|
||||
if not klines or len(klines) < 5:
|
||||
return {}
|
||||
|
||||
try:
|
||||
closes = [float(k.get('close', 0)) for k in klines]
|
||||
highs = [float(k.get('high', 0)) for k in klines]
|
||||
lows = [float(k.get('low', 0)) for k in klines]
|
||||
volumes = [float(k.get('volume', 0)) for k in klines]
|
||||
|
||||
if not closes:
|
||||
return {}
|
||||
|
||||
current_price = closes[-1]
|
||||
indicators = {}
|
||||
|
||||
# ========== RSI ==========
|
||||
if len(closes) >= 15:
|
||||
rsi_value = self._calc_rsi(closes, 14)
|
||||
if rsi_value < 30:
|
||||
rsi_signal = "oversold"
|
||||
elif rsi_value > 70:
|
||||
rsi_signal = "overbought"
|
||||
else:
|
||||
rsi_signal = "neutral"
|
||||
indicators['rsi'] = {
|
||||
'value': round(rsi_value, 2),
|
||||
'signal': rsi_signal,
|
||||
}
|
||||
|
||||
# ========== MACD ==========
|
||||
if len(closes) >= 26:
|
||||
macd_raw = self._calc_macd(closes)
|
||||
macd_val = macd_raw.get('MACD', 0)
|
||||
macd_sig = macd_raw.get('MACD_signal', 0)
|
||||
macd_hist = macd_raw.get('MACD_histogram', 0)
|
||||
|
||||
if macd_val > macd_sig and macd_hist > 0:
|
||||
macd_signal = "bullish"
|
||||
macd_trend = "golden_cross" if macd_hist > 0 else "bullish"
|
||||
elif macd_val < macd_sig and macd_hist < 0:
|
||||
macd_signal = "bearish"
|
||||
macd_trend = "death_cross" if macd_hist < 0 else "bearish"
|
||||
else:
|
||||
macd_signal = "neutral"
|
||||
macd_trend = "consolidating"
|
||||
|
||||
indicators['macd'] = {
|
||||
'value': round(macd_val, 6),
|
||||
'signal_line': round(macd_sig, 6),
|
||||
'histogram': round(macd_hist, 6),
|
||||
'signal': macd_signal,
|
||||
'trend': macd_trend,
|
||||
}
|
||||
|
||||
# ========== 移动平均线 ==========
|
||||
ma5 = sum(closes[-5:]) / 5 if len(closes) >= 5 else current_price
|
||||
ma10 = sum(closes[-10:]) / 10 if len(closes) >= 10 else current_price
|
||||
ma20 = sum(closes[-20:]) / 20 if len(closes) >= 20 else current_price
|
||||
|
||||
if current_price > ma5 > ma10 > ma20:
|
||||
ma_trend = "strong_uptrend"
|
||||
elif current_price > ma20:
|
||||
ma_trend = "uptrend"
|
||||
elif current_price < ma5 < ma10 < ma20:
|
||||
ma_trend = "strong_downtrend"
|
||||
elif current_price < ma20:
|
||||
ma_trend = "downtrend"
|
||||
else:
|
||||
ma_trend = "sideways"
|
||||
|
||||
indicators['moving_averages'] = {
|
||||
'ma5': round(ma5, 6),
|
||||
'ma10': round(ma10, 6),
|
||||
'ma20': round(ma20, 6),
|
||||
'trend': ma_trend,
|
||||
}
|
||||
|
||||
# ========== 支撑/阻力位 (多种方法综合) ==========
|
||||
# 方法1: 枢轴点 (Pivot Points) - 使用前一日数据
|
||||
if len(klines) >= 2:
|
||||
prev_high = float(klines[-2].get('high', highs[-2]) if len(highs) >= 2 else current_price * 1.02)
|
||||
prev_low = float(klines[-2].get('low', lows[-2]) if len(lows) >= 2 else current_price * 0.98)
|
||||
prev_close = float(klines[-2].get('close', closes[-2]) if len(closes) >= 2 else current_price)
|
||||
|
||||
pivot = (prev_high + prev_low + prev_close) / 3
|
||||
r1 = 2 * pivot - prev_low # 阻力位1
|
||||
s1 = 2 * pivot - prev_high # 支撑位1
|
||||
r2 = pivot + (prev_high - prev_low) # 阻力位2
|
||||
s2 = pivot - (prev_high - prev_low) # 支撑位2
|
||||
else:
|
||||
pivot = current_price
|
||||
r1 = r2 = current_price * 1.02
|
||||
s1 = s2 = current_price * 0.98
|
||||
|
||||
# 方法2: 近期高低点
|
||||
recent_highs = highs[-20:] if len(highs) >= 20 else highs
|
||||
recent_lows = lows[-20:] if len(lows) >= 20 else lows
|
||||
swing_high = max(recent_highs) if recent_highs else current_price * 1.05
|
||||
swing_low = min(recent_lows) if recent_lows else current_price * 0.95
|
||||
|
||||
# 方法3: 布林带中轨上下 (如果有)
|
||||
bb_upper = indicators.get('bollinger', {}).get('upper', swing_high)
|
||||
bb_lower = indicators.get('bollinger', {}).get('lower', swing_low)
|
||||
|
||||
# 综合取值: 取多种方法的平均/加权
|
||||
resistance = round((r1 + swing_high + bb_upper) / 3, 6) if bb_upper else round((r1 + swing_high) / 2, 6)
|
||||
support = round((s1 + swing_low + bb_lower) / 3, 6) if bb_lower else round((s1 + swing_low) / 2, 6)
|
||||
|
||||
indicators['levels'] = {
|
||||
'support': support,
|
||||
'resistance': resistance,
|
||||
'pivot': round(pivot, 6),
|
||||
's1': round(s1, 6),
|
||||
'r1': round(r1, 6),
|
||||
's2': round(s2, 6),
|
||||
'r2': round(r2, 6),
|
||||
'swing_high': round(swing_high, 6),
|
||||
'swing_low': round(swing_low, 6),
|
||||
'method': 'pivot_swing_bb_avg' # 标注计算方法
|
||||
}
|
||||
|
||||
# ========== ATR 和波动率 ==========
|
||||
atr = 0
|
||||
if len(klines) >= 14:
|
||||
# 真实波动幅度 ATR (True Range)
|
||||
true_ranges = []
|
||||
for i in range(-14, 0):
|
||||
h = float(klines[i].get('high', 0))
|
||||
l = float(klines[i].get('low', 0))
|
||||
prev_c = float(klines[i-1].get('close', 0)) if i > -14 else h
|
||||
if h > 0 and l > 0:
|
||||
tr = max(h - l, abs(h - prev_c), abs(l - prev_c))
|
||||
true_ranges.append(tr)
|
||||
|
||||
atr = sum(true_ranges) / len(true_ranges) if true_ranges else 0
|
||||
volatility_pct = (atr / current_price * 100) if current_price > 0 else 0
|
||||
|
||||
if volatility_pct > 5:
|
||||
volatility_level = "high"
|
||||
elif volatility_pct > 2:
|
||||
volatility_level = "medium"
|
||||
else:
|
||||
volatility_level = "low"
|
||||
else:
|
||||
volatility_level = "unknown"
|
||||
volatility_pct = 0
|
||||
|
||||
indicators['volatility'] = {
|
||||
'level': volatility_level,
|
||||
'pct': round(volatility_pct, 2),
|
||||
'atr': round(atr, 6), # 添加 ATR 绝对值
|
||||
}
|
||||
|
||||
# ========== 止盈止损建议 (基于 ATR 和支撑/阻力) ==========
|
||||
# 止损: 基于 2x ATR 或支撑位,取更保守的
|
||||
atr_stop_loss = current_price - (2 * atr) if atr > 0 else current_price * 0.95
|
||||
support_stop = indicators['levels']['support']
|
||||
suggested_stop_loss = max(atr_stop_loss, support_stop * 0.99) # 略低于支撑位
|
||||
|
||||
# 止盈: 基于 3x ATR 或阻力位,考虑风险回报比
|
||||
atr_take_profit = current_price + (3 * atr) if atr > 0 else current_price * 1.05
|
||||
resistance_tp = indicators['levels']['resistance']
|
||||
suggested_take_profit = min(atr_take_profit, resistance_tp * 1.01) # 略高于阻力位
|
||||
|
||||
# 风险回报比
|
||||
risk = current_price - suggested_stop_loss
|
||||
reward = suggested_take_profit - current_price
|
||||
risk_reward_ratio = round(reward / risk, 2) if risk > 0 else 0
|
||||
|
||||
indicators['trading_levels'] = {
|
||||
'suggested_stop_loss': round(suggested_stop_loss, 6),
|
||||
'suggested_take_profit': round(suggested_take_profit, 6),
|
||||
'risk_reward_ratio': risk_reward_ratio,
|
||||
'atr_multiplier_sl': 2.0, # 止损使用 2x ATR
|
||||
'atr_multiplier_tp': 3.0, # 止盈使用 3x ATR
|
||||
'method': 'atr_support_resistance'
|
||||
}
|
||||
|
||||
# ========== 布林带 (附加) ==========
|
||||
if len(closes) >= 20:
|
||||
bb_data = self._calc_bollinger(closes, 20, 2)
|
||||
indicators['bollinger'] = bb_data
|
||||
|
||||
# ========== 成交量 (附加) ==========
|
||||
if len(volumes) >= 20:
|
||||
avg_vol = sum(volumes[-20:]) / 20
|
||||
indicators['volume_ratio'] = round(volumes[-1] / avg_vol, 2) if avg_vol > 0 else 1.0
|
||||
|
||||
# ========== 价格位置 (附加) ==========
|
||||
if len(closes) >= 20:
|
||||
high_20 = max(highs[-20:])
|
||||
low_20 = min(lows[-20:])
|
||||
if high_20 > low_20:
|
||||
indicators['price_position'] = round((current_price - low_20) / (high_20 - low_20) * 100, 1)
|
||||
else:
|
||||
indicators['price_position'] = 50.0
|
||||
|
||||
# ========== 整体趋势 (附加) ==========
|
||||
indicators['trend'] = ma_trend
|
||||
indicators['current_price'] = round(current_price, 6)
|
||||
|
||||
return indicators
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Indicator calculation failed: {e}")
|
||||
return {}
|
||||
|
||||
def _calc_rsi(self, closes: List[float], period: int = 14) -> float:
|
||||
"""计算RSI"""
|
||||
if len(closes) < period + 1:
|
||||
return 50.0
|
||||
|
||||
deltas = [closes[i] - closes[i-1] for i in range(1, len(closes))]
|
||||
gains = [d if d > 0 else 0 for d in deltas]
|
||||
losses = [-d if d < 0 else 0 for d in deltas]
|
||||
|
||||
avg_gain = sum(gains[-period:]) / period
|
||||
avg_loss = sum(losses[-period:]) / period
|
||||
|
||||
if avg_loss == 0:
|
||||
return 100.0
|
||||
|
||||
rs = avg_gain / avg_loss
|
||||
rsi = 100 - (100 / (1 + rs))
|
||||
return round(rsi, 2)
|
||||
|
||||
def _calc_macd(self, closes: List[float]) -> Dict[str, float]:
|
||||
"""计算MACD"""
|
||||
def ema(data, period):
|
||||
multiplier = 2 / (period + 1)
|
||||
ema_values = [data[0]]
|
||||
for i in range(1, len(data)):
|
||||
ema_values.append((data[i] - ema_values[-1]) * multiplier + ema_values[-1])
|
||||
return ema_values
|
||||
|
||||
ema12 = ema(closes, 12)
|
||||
ema26 = ema(closes, 26)
|
||||
|
||||
macd_line = [ema12[i] - ema26[i] for i in range(len(closes))]
|
||||
signal_line = ema(macd_line, 9)
|
||||
histogram = [macd_line[i] - signal_line[i] for i in range(len(closes))]
|
||||
|
||||
return {
|
||||
'MACD': round(macd_line[-1], 4),
|
||||
'MACD_signal': round(signal_line[-1], 4),
|
||||
'MACD_histogram': round(histogram[-1], 4)
|
||||
}
|
||||
|
||||
def _calc_bollinger(self, closes: List[float], period: int = 20, std_dev: int = 2) -> Dict[str, float]:
|
||||
"""计算布林带"""
|
||||
if len(closes) < period:
|
||||
return {}
|
||||
|
||||
recent = closes[-period:]
|
||||
middle = sum(recent) / period
|
||||
|
||||
variance = sum((x - middle) ** 2 for x in recent) / period
|
||||
std = variance ** 0.5
|
||||
|
||||
return {
|
||||
'BB_upper': round(middle + std_dev * std, 4),
|
||||
'BB_middle': round(middle, 4),
|
||||
'BB_lower': round(middle - std_dev * std, 4),
|
||||
'BB_width': round((std_dev * std * 2) / middle * 100, 2) if middle > 0 else 0
|
||||
}
|
||||
|
||||
# ==================== 基本面数据 ====================
|
||||
|
||||
def _get_fundamental(self, market: str, symbol: str) -> Optional[Dict[str, Any]]:
|
||||
"""获取基本面数据"""
|
||||
try:
|
||||
if market == 'USStock':
|
||||
return self._get_us_fundamental(symbol)
|
||||
elif market == 'AShare':
|
||||
return self._get_ashare_fundamental(symbol)
|
||||
elif market == 'HShare':
|
||||
return self._get_hshare_fundamental(symbol)
|
||||
except Exception as e:
|
||||
logger.warning(f"Fundamental data fetch failed for {market}:{symbol}: {e}")
|
||||
return None
|
||||
|
||||
def _get_us_fundamental(self, symbol: str) -> Optional[Dict[str, Any]]:
|
||||
"""美股基本面 - Finnhub + yfinance"""
|
||||
result = {}
|
||||
|
||||
# Finnhub
|
||||
if self._finnhub_client:
|
||||
try:
|
||||
metrics = self._finnhub_client.company_basic_financials(symbol, 'all')
|
||||
if metrics and metrics.get('metric'):
|
||||
m = metrics['metric']
|
||||
result.update({
|
||||
'pe_ratio': m.get('peBasicExclExtraTTM'),
|
||||
'pb_ratio': m.get('pbQuarterly'),
|
||||
'ps_ratio': m.get('psTTM'),
|
||||
'market_cap': m.get('marketCapitalization'),
|
||||
'dividend_yield': m.get('dividendYieldIndicatedAnnual'),
|
||||
'beta': m.get('beta'),
|
||||
'52w_high': m.get('52WeekHigh'),
|
||||
'52w_low': m.get('52WeekLow'),
|
||||
'roe': m.get('roeTTM'),
|
||||
'eps': m.get('epsBasicExclExtraItemsTTM'),
|
||||
'revenue_growth': m.get('revenueGrowthTTMYoy'),
|
||||
})
|
||||
except Exception as e:
|
||||
logger.debug(f"Finnhub fundamental failed for {symbol}: {e}")
|
||||
|
||||
# yfinance 补充
|
||||
if not result:
|
||||
try:
|
||||
ticker = yf.Ticker(symbol)
|
||||
info = ticker.info or {}
|
||||
result.update({
|
||||
'pe_ratio': info.get('trailingPE') or info.get('forwardPE'),
|
||||
'pb_ratio': info.get('priceToBook'),
|
||||
'market_cap': info.get('marketCap'),
|
||||
'dividend_yield': info.get('dividendYield'),
|
||||
'beta': info.get('beta'),
|
||||
'52w_high': info.get('fiftyTwoWeekHigh'),
|
||||
'52w_low': info.get('fiftyTwoWeekLow'),
|
||||
'roe': info.get('returnOnEquity'),
|
||||
'eps': info.get('trailingEps'),
|
||||
})
|
||||
except Exception as e:
|
||||
logger.debug(f"yfinance fundamental failed for {symbol}: {e}")
|
||||
|
||||
return result if result else None
|
||||
|
||||
def _get_ashare_fundamental(self, symbol: str) -> Optional[Dict[str, Any]]:
|
||||
"""A股基本面 - akshare"""
|
||||
if not self._ak:
|
||||
return None
|
||||
|
||||
try:
|
||||
# 个股指标
|
||||
df = self._ak.stock_individual_info_em(symbol=symbol)
|
||||
if df is not None and not df.empty:
|
||||
result = {}
|
||||
for _, row in df.iterrows():
|
||||
item = row.get('item', '')
|
||||
value = row.get('value', '')
|
||||
if '市盈率' in item:
|
||||
result['pe_ratio'] = value
|
||||
elif '市净率' in item:
|
||||
result['pb_ratio'] = value
|
||||
elif '总市值' in item:
|
||||
result['market_cap'] = value
|
||||
elif 'ROE' in item or '净资产收益率' in item:
|
||||
result['roe'] = value
|
||||
elif '每股收益' in item:
|
||||
result['eps'] = value
|
||||
return result if result else None
|
||||
except Exception as e:
|
||||
logger.debug(f"akshare fundamental failed for {symbol}: {e}")
|
||||
|
||||
return None
|
||||
|
||||
def _get_hshare_fundamental(self, symbol: str) -> Optional[Dict[str, Any]]:
|
||||
"""港股基本面 - yfinance"""
|
||||
try:
|
||||
# 港股在yfinance的格式: 0700.HK, 9988.HK
|
||||
yf_symbol = f"{symbol}.HK"
|
||||
ticker = yf.Ticker(yf_symbol)
|
||||
info = ticker.info or {}
|
||||
|
||||
return {
|
||||
'pe_ratio': info.get('trailingPE'),
|
||||
'pb_ratio': info.get('priceToBook'),
|
||||
'market_cap': info.get('marketCap'),
|
||||
'dividend_yield': info.get('dividendYield'),
|
||||
'52w_high': info.get('fiftyTwoWeekHigh'),
|
||||
'52w_low': info.get('fiftyTwoWeekLow'),
|
||||
}
|
||||
except Exception as e:
|
||||
logger.debug(f"yfinance HShare fundamental failed for {symbol}: {e}")
|
||||
|
||||
return None
|
||||
|
||||
def _get_crypto_info(self, symbol: str) -> Optional[Dict[str, Any]]:
|
||||
"""加密货币信息 (固定描述为主)"""
|
||||
# 常见加密货币的描述
|
||||
crypto_info = {
|
||||
'BTC': {
|
||||
'name': 'Bitcoin',
|
||||
'description': '比特币,数字黄金,市值第一的加密货币,作为价值存储和避险资产',
|
||||
'category': 'Store of Value',
|
||||
},
|
||||
'ETH': {
|
||||
'name': 'Ethereum',
|
||||
'description': '以太坊,智能合约平台,DeFi和NFT生态的基础设施',
|
||||
'category': 'Smart Contract Platform',
|
||||
},
|
||||
'BNB': {
|
||||
'name': 'Binance Coin',
|
||||
'description': '币安币,全球最大交易所的平台代币',
|
||||
'category': 'Exchange Token',
|
||||
},
|
||||
'SOL': {
|
||||
'name': 'Solana',
|
||||
'description': '高性能公链,主打高TPS和低Gas费',
|
||||
'category': 'Smart Contract Platform',
|
||||
},
|
||||
'XRP': {
|
||||
'name': 'Ripple',
|
||||
'description': '瑞波币,专注跨境支付解决方案',
|
||||
'category': 'Payment',
|
||||
},
|
||||
'DOGE': {
|
||||
'name': 'Dogecoin',
|
||||
'description': '狗狗币,Meme币代表,社区驱动',
|
||||
'category': 'Meme',
|
||||
},
|
||||
}
|
||||
|
||||
# 提取基础代币名
|
||||
base = symbol.split('/')[0] if '/' in symbol else symbol
|
||||
base = base.upper()
|
||||
|
||||
if base in crypto_info:
|
||||
return crypto_info[base]
|
||||
|
||||
return {
|
||||
'name': base,
|
||||
'description': f'{base} 是一种加密货币',
|
||||
'category': 'Unknown',
|
||||
}
|
||||
|
||||
def _get_company(self, market: str, symbol: str) -> Optional[Dict[str, Any]]:
|
||||
"""获取公司信息"""
|
||||
try:
|
||||
if market == 'USStock' and self._finnhub_client:
|
||||
profile = self._finnhub_client.company_profile2(symbol=symbol)
|
||||
if profile:
|
||||
return {
|
||||
'name': profile.get('name'),
|
||||
'industry': profile.get('finnhubIndustry'),
|
||||
'country': profile.get('country'),
|
||||
'exchange': profile.get('exchange'),
|
||||
'ipo_date': profile.get('ipo'),
|
||||
'market_cap': profile.get('marketCapitalization'),
|
||||
'website': profile.get('weburl'),
|
||||
}
|
||||
|
||||
elif market == 'AShare' and self._ak:
|
||||
df = self._ak.stock_individual_info_em(symbol=symbol)
|
||||
if df is not None and not df.empty:
|
||||
result = {}
|
||||
for _, row in df.iterrows():
|
||||
item = row.get('item', '')
|
||||
value = row.get('value', '')
|
||||
if '名称' in item or '简称' in item:
|
||||
result['name'] = value
|
||||
elif '行业' in item:
|
||||
result['industry'] = value
|
||||
return result if result else None
|
||||
|
||||
except Exception as e:
|
||||
logger.debug(f"Company info fetch failed for {market}:{symbol}: {e}")
|
||||
|
||||
return None
|
||||
|
||||
# ==================== 宏观数据 (复用全球金融板块) ====================
|
||||
|
||||
def _get_macro_data(self, market: str, timeout: int = 10) -> Dict[str, Any]:
|
||||
"""
|
||||
获取宏观经济数据 - 复用 global_market.py 的函数和缓存
|
||||
|
||||
优势:
|
||||
1. 数据与全球金融页面一致
|
||||
2. 复用30秒/5分钟缓存,降低API调用
|
||||
3. 已有完整的数据解读和级别判断
|
||||
"""
|
||||
try:
|
||||
# 复用 global_market.py 的市场情绪数据 (有5分钟缓存)
|
||||
from app.routes.global_market import (
|
||||
_fetch_vix, _fetch_dollar_index, _fetch_yield_curve,
|
||||
_fetch_fear_greed_index,
|
||||
_get_cached, _set_cached
|
||||
)
|
||||
|
||||
result = {}
|
||||
|
||||
# 1) 尝试从缓存获取 (global_market 的缓存, 6小时有效)
|
||||
MACRO_CACHE_TTL = 21600 # 6 hours
|
||||
cached_sentiment = _get_cached("market_sentiment", MACRO_CACHE_TTL)
|
||||
if cached_sentiment:
|
||||
logger.info("Using cached sentiment data from global_market (6h cache)")
|
||||
# 转换格式
|
||||
if cached_sentiment.get('vix'):
|
||||
vix = cached_sentiment['vix']
|
||||
result['VIX'] = {
|
||||
'name': 'VIX恐慌指数',
|
||||
'description': vix.get('interpretation', ''),
|
||||
'price': vix.get('value', 0),
|
||||
'change': vix.get('change', 0),
|
||||
'changePercent': vix.get('change', 0),
|
||||
'level': vix.get('level', 'unknown'),
|
||||
}
|
||||
|
||||
if cached_sentiment.get('dxy'):
|
||||
dxy = cached_sentiment['dxy']
|
||||
result['DXY'] = {
|
||||
'name': '美元指数',
|
||||
'description': dxy.get('interpretation', ''),
|
||||
'price': dxy.get('value', 0),
|
||||
'change': dxy.get('change', 0),
|
||||
'changePercent': dxy.get('change', 0),
|
||||
'level': dxy.get('level', 'unknown'),
|
||||
}
|
||||
|
||||
if cached_sentiment.get('yield_curve'):
|
||||
yc = cached_sentiment['yield_curve']
|
||||
result['TNX'] = {
|
||||
'name': '美债10年收益率',
|
||||
'description': yc.get('interpretation', ''),
|
||||
'price': yc.get('yield_10y', 0),
|
||||
'change': yc.get('change', 0),
|
||||
'changePercent': 0,
|
||||
'spread': yc.get('spread', 0),
|
||||
'level': yc.get('level', 'unknown'),
|
||||
}
|
||||
|
||||
if cached_sentiment.get('fear_greed'):
|
||||
fg = cached_sentiment['fear_greed']
|
||||
result['FEAR_GREED'] = {
|
||||
'name': '恐惧贪婪指数',
|
||||
'description': fg.get('classification', 'Neutral'),
|
||||
'price': fg.get('value', 50),
|
||||
'change': 0,
|
||||
'changePercent': 0,
|
||||
}
|
||||
|
||||
if result:
|
||||
return result
|
||||
|
||||
# 2) 如果没有缓存,快速并行获取关键指标
|
||||
logger.info("Fetching macro data from global_market functions")
|
||||
|
||||
with ThreadPoolExecutor(max_workers=4) as executor:
|
||||
futures = {
|
||||
executor.submit(_fetch_vix): "VIX",
|
||||
executor.submit(_fetch_dollar_index): "DXY",
|
||||
executor.submit(_fetch_yield_curve): "TNX",
|
||||
executor.submit(_fetch_fear_greed_index): "FEAR_GREED",
|
||||
}
|
||||
|
||||
try:
|
||||
for future in as_completed(futures, timeout=timeout):
|
||||
key = futures[future]
|
||||
try:
|
||||
data = future.result(timeout=5)
|
||||
if data:
|
||||
# 转换为统一格式
|
||||
if key == 'VIX':
|
||||
result[key] = {
|
||||
'name': 'VIX恐慌指数',
|
||||
'description': data.get('interpretation', ''),
|
||||
'price': data.get('value', 0),
|
||||
'change': data.get('change', 0),
|
||||
'changePercent': data.get('change', 0),
|
||||
'level': data.get('level', 'unknown'),
|
||||
}
|
||||
elif key == 'DXY':
|
||||
result[key] = {
|
||||
'name': '美元指数',
|
||||
'description': data.get('interpretation', ''),
|
||||
'price': data.get('value', 0),
|
||||
'change': data.get('change', 0),
|
||||
'changePercent': data.get('change', 0),
|
||||
'level': data.get('level', 'unknown'),
|
||||
}
|
||||
elif key == 'TNX':
|
||||
result[key] = {
|
||||
'name': '美债10年收益率',
|
||||
'description': data.get('interpretation', ''),
|
||||
'price': data.get('yield_10y', 0),
|
||||
'change': data.get('change', 0),
|
||||
'changePercent': 0,
|
||||
'spread': data.get('spread', 0),
|
||||
'level': data.get('level', 'unknown'),
|
||||
}
|
||||
elif key == 'FEAR_GREED':
|
||||
result[key] = {
|
||||
'name': '恐惧贪婪指数',
|
||||
'description': data.get('classification', 'Neutral'),
|
||||
'price': data.get('value', 50),
|
||||
'change': 0,
|
||||
'changePercent': 0,
|
||||
}
|
||||
except Exception as e:
|
||||
logger.debug(f"Macro indicator {key} fetch failed: {e}")
|
||||
except TimeoutError:
|
||||
logger.warning("Macro data fetch timed out")
|
||||
|
||||
# 注:黄金等大宗商品数据不再作为宏观指标获取
|
||||
# 原因:1) 如果分析的是黄金,价格已在 _get_price 中获取
|
||||
# 2) 减少 API 调用,提高稳定性
|
||||
pass
|
||||
|
||||
return result
|
||||
|
||||
except ImportError as e:
|
||||
logger.warning(f"Could not import from global_market: {e}")
|
||||
return {}
|
||||
except Exception as e:
|
||||
logger.error(f"_get_macro_data failed: {e}")
|
||||
return {}
|
||||
|
||||
# ==================== 新闻/情绪数据 ====================
|
||||
|
||||
def _get_news(
|
||||
self, market: str, symbol: str, company_name: str = None, timeout: int = 8
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
获取新闻和情绪数据
|
||||
|
||||
策略:
|
||||
1. 使用结构化API (Finnhub) - 无需深度阅读
|
||||
2. 只获取标题和摘要 - 不读取全文
|
||||
3. 多来源聚合 - Finnhub + 市场特定来源
|
||||
"""
|
||||
news_list = []
|
||||
sentiment = {}
|
||||
|
||||
# 1) Finnhub 新闻 (最可靠)
|
||||
if self._finnhub_client:
|
||||
try:
|
||||
end_date = datetime.now().strftime('%Y-%m-%d')
|
||||
start_date = (datetime.now() - timedelta(days=7)).strftime('%Y-%m-%d')
|
||||
|
||||
raw_news = []
|
||||
|
||||
if market == 'USStock':
|
||||
raw_news = self._finnhub_client.company_news(symbol, _from=start_date, to=end_date)
|
||||
else:
|
||||
# 通用新闻
|
||||
raw_news = self._finnhub_client.general_news('general', min_id=0)
|
||||
|
||||
if raw_news:
|
||||
for item in raw_news[:10]: # 最多10条
|
||||
if not item.get('headline'):
|
||||
continue
|
||||
news_list.append({
|
||||
"datetime": datetime.fromtimestamp(item.get('datetime', 0)).strftime('%Y-%m-%d %H:%M'),
|
||||
"headline": item.get('headline', ''),
|
||||
"summary": item.get('summary', '')[:300] if item.get('summary') else '', # 截断摘要
|
||||
"source": item.get('source', 'Finnhub'),
|
||||
"url": item.get('url', ''),
|
||||
"sentiment": item.get('sentiment', 'neutral'), # Finnhub有时提供情绪
|
||||
})
|
||||
except Exception as e:
|
||||
logger.debug(f"Finnhub news fetch failed: {e}")
|
||||
|
||||
# 2) Finnhub 情绪分数 (如果可用)
|
||||
if self._finnhub_client and market == 'USStock':
|
||||
try:
|
||||
# Finnhub 提供社交媒体情绪
|
||||
social = self._finnhub_client.stock_social_sentiment(symbol)
|
||||
if social:
|
||||
sentiment['reddit'] = social.get('reddit', {})
|
||||
sentiment['twitter'] = social.get('twitter', {})
|
||||
except Exception as e:
|
||||
logger.debug(f"Finnhub sentiment fetch failed: {e}")
|
||||
|
||||
# 3) A股特定新闻 (akshare)
|
||||
if market == 'AShare' and self._ak:
|
||||
try:
|
||||
# 个股新闻
|
||||
df = self._ak.stock_news_em(symbol=symbol)
|
||||
if df is not None and not df.empty:
|
||||
for _, row in df.head(10).iterrows():
|
||||
news_list.append({
|
||||
"datetime": str(row.get('发布时间', ''))[:16],
|
||||
"headline": row.get('新闻标题', ''),
|
||||
"summary": row.get('新闻内容', '')[:200] if row.get('新闻内容') else '',
|
||||
"source": row.get('文章来源', 'eastmoney'),
|
||||
"url": row.get('新闻链接', ''),
|
||||
"sentiment": 'neutral',
|
||||
})
|
||||
except Exception as e:
|
||||
logger.debug(f"akshare news fetch failed: {e}")
|
||||
|
||||
# 按时间排序
|
||||
news_list.sort(key=lambda x: x.get('datetime', ''), reverse=True)
|
||||
|
||||
return {
|
||||
"news": news_list[:15], # 最多15条
|
||||
"sentiment": sentiment,
|
||||
}
|
||||
|
||||
|
||||
# 全局实例
|
||||
_collector: Optional[MarketDataCollector] = None
|
||||
|
||||
def get_market_data_collector() -> MarketDataCollector:
|
||||
"""获取市场数据采集器单例"""
|
||||
global _collector
|
||||
if _collector is None:
|
||||
_collector = MarketDataCollector()
|
||||
return _collector
|
||||
@@ -13,7 +13,7 @@ from typing import Any, Dict, List, Optional
|
||||
|
||||
from app.utils.db import get_db_connection
|
||||
from app.utils.logger import get_logger
|
||||
from app.services.analysis import AnalysisService
|
||||
from app.services.fast_analysis import get_fast_analysis_service
|
||||
from app.services.signal_notifier import SignalNotifier
|
||||
from app.services.kline import KlineService
|
||||
|
||||
@@ -156,14 +156,17 @@ def _get_positions_for_monitor(position_ids: List[int] = None, user_id: int = No
|
||||
|
||||
def _run_ai_analysis(positions: List[Dict[str, Any]], config: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Run full multi-agent AI analysis on positions.
|
||||
Uses the same 13-agent analysis flow as the AI Analysis page.
|
||||
Run fast AI analysis on positions.
|
||||
Uses the new FastAnalysisService (single LLM call, faster and more stable).
|
||||
"""
|
||||
try:
|
||||
language = config.get('language', 'en-US')
|
||||
custom_prompt = config.get('prompt', '')
|
||||
|
||||
# Analyze each position using the full agent analysis flow
|
||||
# Get the fast analysis service
|
||||
service = get_fast_analysis_service()
|
||||
|
||||
# Analyze each position
|
||||
position_analyses = []
|
||||
|
||||
for pos in positions:
|
||||
@@ -176,21 +179,24 @@ def _run_ai_analysis(positions: List[Dict[str, Any]], config: Dict[str, Any]) ->
|
||||
continue
|
||||
|
||||
try:
|
||||
logger.info(f"Running multi-agent analysis for {market}:{symbol}")
|
||||
logger.info(f"Running fast AI analysis for {market}:{symbol}")
|
||||
|
||||
# Use the full AnalysisService (13-agent flow)
|
||||
analysis_result = AnalysisService().analyze(
|
||||
# Use the new FastAnalysisService (single LLM call)
|
||||
analysis_result = service.analyze(
|
||||
market=market,
|
||||
symbol=symbol,
|
||||
language=language,
|
||||
timeframe='1D'
|
||||
)
|
||||
|
||||
# Extract key information from the analysis
|
||||
final_decision = analysis_result.get('final_decision', {})
|
||||
trader_decision = analysis_result.get('trader_decision', {})
|
||||
overview = analysis_result.get('overview', {})
|
||||
risk_report = analysis_result.get('risk', {})
|
||||
# Extract information from the new format
|
||||
detailed = analysis_result.get('detailed_analysis', {})
|
||||
trading_plan = analysis_result.get('trading_plan', {})
|
||||
scores = analysis_result.get('scores', {})
|
||||
|
||||
# Build risk report from risks list
|
||||
risks = analysis_result.get('risks', [])
|
||||
risk_report = '\n'.join([f"• {r}" for r in risks]) if risks else ''
|
||||
|
||||
position_analysis = {
|
||||
'market': market,
|
||||
@@ -198,24 +204,35 @@ def _run_ai_analysis(positions: List[Dict[str, Any]], config: Dict[str, Any]) ->
|
||||
'name': name,
|
||||
'group_name': group_name,
|
||||
'entry_price': pos.get('entry_price'),
|
||||
'current_price': pos.get('current_price'),
|
||||
'current_price': pos.get('current_price') or analysis_result.get('market_data', {}).get('current_price'),
|
||||
'pnl': pos.get('pnl'),
|
||||
'pnl_percent': pos.get('pnl_percent'),
|
||||
'quantity': pos.get('quantity'),
|
||||
'side': pos.get('side'),
|
||||
# Multi-agent analysis results
|
||||
'final_decision': final_decision.get('decision', 'HOLD'),
|
||||
'confidence': final_decision.get('confidence', 50),
|
||||
'reasoning': final_decision.get('reasoning', ''),
|
||||
'trader_decision': trader_decision.get('decision', 'HOLD'),
|
||||
'trader_reasoning': trader_decision.get('reasoning', ''),
|
||||
'overview_report': overview.get('report', ''),
|
||||
'risk_report': risk_report.get('report', ''),
|
||||
# New fast analysis results
|
||||
'final_decision': analysis_result.get('decision', 'HOLD'),
|
||||
'confidence': analysis_result.get('confidence', 50),
|
||||
'reasoning': analysis_result.get('summary', ''),
|
||||
'trader_decision': analysis_result.get('decision', 'HOLD'), # Same as final for fast analysis
|
||||
'trader_reasoning': analysis_result.get('summary', ''),
|
||||
'overview_report': detailed.get('technical', ''),
|
||||
'fundamental_report': detailed.get('fundamental', ''),
|
||||
'sentiment_report': detailed.get('sentiment', ''),
|
||||
'risk_report': risk_report,
|
||||
# Trading plan
|
||||
'suggested_entry': trading_plan.get('entry_price'),
|
||||
'suggested_stop_loss': trading_plan.get('stop_loss'),
|
||||
'suggested_take_profit': trading_plan.get('take_profit'),
|
||||
# Scores
|
||||
'technical_score': scores.get('technical', 50),
|
||||
'fundamental_score': scores.get('fundamental', 50),
|
||||
'sentiment_score': scores.get('sentiment', 50),
|
||||
'key_reasons': analysis_result.get('reasons', []),
|
||||
'error': analysis_result.get('error')
|
||||
}
|
||||
|
||||
position_analyses.append(position_analysis)
|
||||
logger.info(f"Analysis completed for {market}:{symbol}: {final_decision.get('decision', 'N/A')}")
|
||||
logger.info(f"Fast analysis completed for {market}:{symbol}: {analysis_result.get('decision', 'N/A')}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to analyze {market}:{symbol}: {e}")
|
||||
@@ -283,7 +300,7 @@ def _build_html_report(
|
||||
# Text translations
|
||||
texts = {
|
||||
'title': '投资组合AI分析报告' if is_zh else 'Portfolio AI Analysis Report',
|
||||
'subtitle': '由 QuantDinger 多智能体分析系统生成' if is_zh else 'Generated by QuantDinger Multi-Agent Analysis System',
|
||||
'subtitle': '由 QuantDinger AI 快速分析引擎生成' if is_zh else 'Generated by QuantDinger Fast AI Analysis Engine',
|
||||
'overview': '组合概览' if is_zh else 'Portfolio Overview',
|
||||
'positions': '持仓数量' if is_zh else 'Positions',
|
||||
'total_value': '总市值' if is_zh else 'Total Value',
|
||||
|
||||
@@ -555,6 +555,11 @@ class TradingExecutor:
|
||||
trade_direction = 'long' # 现货只能做多
|
||||
logger.info(f"Strategy {strategy_id} spot trading; force trade_direction=long")
|
||||
|
||||
# 获取市场类别(Crypto, USStock, Forex, Futures, AShare, HShare)
|
||||
# 这决定了使用哪个数据源来获取价格和K线数据
|
||||
market_category = (strategy.get('market_category') or 'Crypto').strip()
|
||||
logger.info(f"Strategy {strategy_id} market_category: {market_category}")
|
||||
|
||||
# 初始化交易所连接(信号模式下无需真实连接)
|
||||
exchange = None
|
||||
|
||||
@@ -611,7 +616,7 @@ class TradingExecutor:
|
||||
# ============================================
|
||||
# logger.info(f"策略 {strategy_id} 初始化:获取历史K线数据...")
|
||||
history_limit = int(os.getenv('K_LINE_HISTORY_GET_NUMBER', 500))
|
||||
klines = self._fetch_latest_kline(symbol, timeframe, limit=history_limit)
|
||||
klines = self._fetch_latest_kline(symbol, timeframe, limit=history_limit, market_category=market_category)
|
||||
if not klines or len(klines) < 2:
|
||||
logger.error(f"Strategy {strategy_id} failed to fetch K-lines")
|
||||
return
|
||||
@@ -719,16 +724,16 @@ class TradingExecutor:
|
||||
# ============================================
|
||||
# 1. Fetch current price once per tick
|
||||
# ============================================
|
||||
current_price = self._fetch_current_price(exchange, symbol, market_type=market_type)
|
||||
current_price = self._fetch_current_price(exchange, symbol, market_type=market_type, market_category=market_category)
|
||||
if current_price is None:
|
||||
logger.warning(f"Strategy {strategy_id} failed to fetch current price")
|
||||
logger.warning(f"Strategy {strategy_id} failed to fetch current price for {market_category}:{symbol}")
|
||||
continue
|
||||
|
||||
# ============================================
|
||||
# 2. 检查是否需要更新K线(每个K线周期更新一次,从API拉取)
|
||||
# ============================================
|
||||
if current_time - last_kline_update_time >= kline_update_interval:
|
||||
klines = self._fetch_latest_kline(symbol, timeframe, limit=history_limit)
|
||||
klines = self._fetch_latest_kline(symbol, timeframe, limit=history_limit, market_category=market_category)
|
||||
if klines and len(klines) >= 2:
|
||||
df = self._klines_to_dataframe(klines)
|
||||
if len(df) > 0:
|
||||
@@ -978,6 +983,7 @@ class TradingExecutor:
|
||||
leverage=leverage,
|
||||
initial_capital=initial_capital,
|
||||
market_type=market_type,
|
||||
market_category=market_category,
|
||||
execution_mode=execution_mode,
|
||||
notification_config=notification_config,
|
||||
trading_config=trading_config,
|
||||
@@ -1041,7 +1047,8 @@ class TradingExecutor:
|
||||
id, strategy_name, strategy_type, status,
|
||||
initial_capital, leverage, decide_interval,
|
||||
execution_mode, notification_config,
|
||||
indicator_config, exchange_config, trading_config, ai_model_config
|
||||
indicator_config, exchange_config, trading_config, ai_model_config,
|
||||
market_category
|
||||
FROM qd_strategies_trading
|
||||
WHERE id = %s
|
||||
"""
|
||||
@@ -1104,25 +1111,39 @@ class TradingExecutor:
|
||||
"""(Mock) 信号模式不需要真实交易所连接"""
|
||||
return None
|
||||
|
||||
def _fetch_latest_kline(self, symbol: str, timeframe: str, limit: int = 500) -> List[Dict[str, Any]]:
|
||||
"""获取最新K线数据(优先从缓存获取)"""
|
||||
def _fetch_latest_kline(self, symbol: str, timeframe: str, limit: int = 500, market_category: str = 'Crypto') -> List[Dict[str, Any]]:
|
||||
"""获取最新K线数据(优先从缓存获取)
|
||||
|
||||
Args:
|
||||
symbol: 交易对/代码
|
||||
timeframe: 时间周期
|
||||
limit: 数据条数
|
||||
market_category: 市场类型 (Crypto, USStock, Forex, Futures, AShare, HShare)
|
||||
"""
|
||||
try:
|
||||
# 使用 KlineService 获取K线数据(自动处理缓存)
|
||||
return self.kline_service.get_kline(
|
||||
market='Crypto',
|
||||
market=market_category,
|
||||
symbol=symbol,
|
||||
timeframe=timeframe,
|
||||
limit=limit,
|
||||
before_time=int(time.time())
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to fetch K-lines: {str(e)}")
|
||||
logger.error(f"Failed to fetch K-lines for {market_category}:{symbol}: {str(e)}")
|
||||
return []
|
||||
|
||||
def _fetch_current_price(self, exchange: Any, symbol: str, market_type: str = None) -> Optional[float]:
|
||||
"""获取当前价格 (改用 DataSource)"""
|
||||
def _fetch_current_price(self, exchange: Any, symbol: str, market_type: str = None, market_category: str = 'Crypto') -> Optional[float]:
|
||||
"""获取当前价格 (根据 market_category 选择正确的数据源)
|
||||
|
||||
Args:
|
||||
exchange: 交易所实例(信号模式下为 None)
|
||||
symbol: 交易对/代码
|
||||
market_type: 交易类型 (swap/spot)
|
||||
market_category: 市场类型 (Crypto, USStock, Forex, Futures, AShare, HShare)
|
||||
"""
|
||||
# Local in-memory cache first
|
||||
cache_key = (symbol or "").strip().upper()
|
||||
cache_key = f"{market_category}:{(symbol or '').strip().upper()}"
|
||||
if cache_key and self._price_cache_ttl_sec > 0:
|
||||
now = time.time()
|
||||
try:
|
||||
@@ -1138,12 +1159,9 @@ class TradingExecutor:
|
||||
pass
|
||||
|
||||
try:
|
||||
# 默认使用 binance 获取价格 (或者根据配置)
|
||||
# 简单起见,这里硬编码或使用 generic source
|
||||
ds = DataSourceFactory.get_data_source('binance')
|
||||
# normalized symbol handling is tricky without exchange object.
|
||||
# But usually DataSource expects standard 'BTC/USDT'
|
||||
ticker = ds.get_ticker(symbol)
|
||||
# 根据 market_category 选择正确的数据源
|
||||
# 支持: Crypto, USStock, Forex, Futures, AShare, HShare
|
||||
ticker = DataSourceFactory.get_ticker(market_category, symbol)
|
||||
if ticker:
|
||||
price = float(ticker.get('last') or ticker.get('close') or 0)
|
||||
if price > 0:
|
||||
@@ -1155,7 +1173,7 @@ class TradingExecutor:
|
||||
pass
|
||||
return price
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to fetch price: {e}")
|
||||
logger.warning(f"Failed to fetch price for {market_category}:{symbol}: {e}")
|
||||
|
||||
return None
|
||||
|
||||
@@ -1887,6 +1905,7 @@ class TradingExecutor:
|
||||
leverage: int,
|
||||
initial_capital: float,
|
||||
market_type: str = 'swap',
|
||||
market_category: str = 'Crypto',
|
||||
margin_mode: str = 'cross',
|
||||
stop_loss_price: float = None,
|
||||
take_profit_price: float = None,
|
||||
@@ -2013,6 +2032,7 @@ class TradingExecutor:
|
||||
amount=amount,
|
||||
ref_price=float(current_price or 0.0),
|
||||
market_type=market_type,
|
||||
market_category=market_category,
|
||||
leverage=leverage,
|
||||
execution_mode=execution_mode,
|
||||
notification_config=notification_config,
|
||||
@@ -2049,16 +2069,28 @@ class TradingExecutor:
|
||||
)
|
||||
elif sig.startswith("reduce_"):
|
||||
# Partial scale-out: reduce position size, keep entry price unchanged.
|
||||
self._record_trade(
|
||||
strategy_id=strategy_id, symbol=symbol, type=signal_type,
|
||||
price=current_price, amount=amount, value=amount*current_price
|
||||
)
|
||||
# 信号模式下计算部分平仓盈亏
|
||||
side = 'short' if 'short' in signal_type else 'long'
|
||||
old_pos = next((p for p in current_positions if p.get('side') == side), None)
|
||||
if not old_pos:
|
||||
return True
|
||||
old_size = float(old_pos.get('size') or 0.0)
|
||||
old_entry = float(old_pos.get('entry_price') or 0.0)
|
||||
|
||||
# 计算减仓部分的盈亏(信号模式下,不含手续费)
|
||||
reduce_profit = None
|
||||
if old_entry > 0 and amount > 0:
|
||||
if side == 'long':
|
||||
reduce_profit = (current_price - old_entry) * amount
|
||||
else:
|
||||
reduce_profit = (old_entry - current_price) * amount
|
||||
|
||||
self._record_trade(
|
||||
strategy_id=strategy_id, symbol=symbol, type=signal_type,
|
||||
price=current_price, amount=amount, value=amount*current_price,
|
||||
profit=reduce_profit
|
||||
)
|
||||
|
||||
new_size = max(0.0, old_size - float(amount or 0.0))
|
||||
if new_size <= old_size * 0.001:
|
||||
self._close_position(strategy_id, symbol, side)
|
||||
@@ -2068,11 +2100,25 @@ class TradingExecutor:
|
||||
size=new_size, entry_price=old_entry, current_price=current_price
|
||||
)
|
||||
elif 'close' in sig:
|
||||
# 信号模式下计算平仓盈亏
|
||||
side = 'short' if 'short' in signal_type else 'long'
|
||||
old_pos = next((p for p in current_positions if p.get('side') == side), None)
|
||||
|
||||
# 计算盈亏(信号模式下,不含手续费)
|
||||
close_profit = None
|
||||
if old_pos:
|
||||
entry_price = float(old_pos.get('entry_price') or 0)
|
||||
if entry_price > 0 and amount > 0:
|
||||
if side == 'long':
|
||||
close_profit = (current_price - entry_price) * amount
|
||||
else:
|
||||
close_profit = (entry_price - current_price) * amount
|
||||
|
||||
self._record_trade(
|
||||
strategy_id=strategy_id, symbol=symbol, type=signal_type,
|
||||
price=current_price, amount=amount, value=amount*current_price
|
||||
price=current_price, amount=amount, value=amount*current_price,
|
||||
profit=close_profit
|
||||
)
|
||||
side = 'short' if 'short' in signal_type else 'long'
|
||||
self._close_position(strategy_id, symbol, side)
|
||||
|
||||
return True
|
||||
@@ -2145,25 +2191,29 @@ class TradingExecutor:
|
||||
language = str(language or "zh-CN")
|
||||
|
||||
try:
|
||||
# Lazy import to avoid circular deps + heavy init unless the filter is enabled and entry signal happens.
|
||||
from app.services.analysis import AnalysisService
|
||||
# 使用新的 FastAnalysisService (单次LLM调用,更快更稳定)
|
||||
from app.services.fast_analysis import get_fast_analysis_service
|
||||
|
||||
service = AnalysisService()
|
||||
service = get_fast_analysis_service()
|
||||
result = service.analyze(market, symbol, language, model=model)
|
||||
|
||||
if isinstance(result, dict) and result.get("error"):
|
||||
return False, {"ai_decision": "", "reason": "analysis_error", "analysis_error": str(result.get("error") or "")}
|
||||
|
||||
ai_dec = self._extract_ai_trade_decision(result)
|
||||
if not ai_dec:
|
||||
return False, {"ai_decision": "", "reason": "missing_ai_decision"}
|
||||
# FastAnalysisService 直接返回 decision 字段
|
||||
ai_dec = str(result.get("decision", "")).strip().upper()
|
||||
if not ai_dec or ai_dec not in ("BUY", "SELL", "HOLD"):
|
||||
return False, {"ai_decision": ai_dec, "reason": "missing_ai_decision"}
|
||||
|
||||
expected = "BUY" if signal_type == "open_long" else "SELL"
|
||||
confidence = result.get("confidence", 50)
|
||||
summary = result.get("summary", "")
|
||||
|
||||
if ai_dec == expected:
|
||||
return True, {"ai_decision": ai_dec, "reason": "match"}
|
||||
return True, {"ai_decision": ai_dec, "reason": "match", "confidence": confidence, "summary": summary}
|
||||
if ai_dec == "HOLD":
|
||||
return False, {"ai_decision": ai_dec, "reason": "ai_hold"}
|
||||
return False, {"ai_decision": ai_dec, "reason": "direction_mismatch"}
|
||||
return False, {"ai_decision": ai_dec, "reason": "ai_hold", "confidence": confidence, "summary": summary}
|
||||
return False, {"ai_decision": ai_dec, "reason": "direction_mismatch", "confidence": confidence, "summary": summary}
|
||||
except Exception as e:
|
||||
return False, {"ai_decision": "", "reason": "analysis_exception", "analysis_error": str(e)}
|
||||
|
||||
@@ -2262,6 +2312,7 @@ class TradingExecutor:
|
||||
amount: float,
|
||||
ref_price: Optional[float] = None,
|
||||
market_type: str = 'swap',
|
||||
market_category: str = 'Crypto',
|
||||
leverage: float = 1.0,
|
||||
margin_mode: str = 'cross',
|
||||
stop_loss_price: float = None,
|
||||
@@ -2291,7 +2342,7 @@ class TradingExecutor:
|
||||
try:
|
||||
# Reference price at enqueue time: use current tick price if provided to avoid extra fetch.
|
||||
if ref_price is None:
|
||||
ref_price = self._fetch_current_price(None, symbol) or 0.0
|
||||
ref_price = self._fetch_current_price(None, symbol, market_category=market_category) or 0.0
|
||||
ref_price = float(ref_price or 0.0)
|
||||
|
||||
extra_payload = {
|
||||
|
||||
@@ -183,6 +183,72 @@ class UserService:
|
||||
user.pop('password_hash', None)
|
||||
return user
|
||||
|
||||
def get_token_version(self, user_id: int) -> int:
|
||||
"""
|
||||
获取用户当前的 token 版本号。
|
||||
|
||||
Args:
|
||||
user_id: 用户ID
|
||||
|
||||
Returns:
|
||||
当前 token 版本号,默认为 1
|
||||
"""
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
cur.execute(
|
||||
"SELECT token_version FROM qd_users WHERE id = ?",
|
||||
(user_id,)
|
||||
)
|
||||
row = cur.fetchone()
|
||||
cur.close()
|
||||
if row:
|
||||
return int(row.get('token_version') or 1)
|
||||
return 1
|
||||
except Exception as e:
|
||||
logger.error(f"get_token_version failed: {e}")
|
||||
return 1
|
||||
|
||||
def increment_token_version(self, user_id: int) -> int:
|
||||
"""
|
||||
递增用户的 token 版本号,使旧的 token 失效。
|
||||
用于实现单一客户端登录(踢出其他设备)。
|
||||
|
||||
Args:
|
||||
user_id: 用户ID
|
||||
|
||||
Returns:
|
||||
新的 token 版本号
|
||||
"""
|
||||
try:
|
||||
with get_db_connection() as db:
|
||||
cur = db.cursor()
|
||||
# 递增 token_version
|
||||
cur.execute(
|
||||
"""
|
||||
UPDATE qd_users
|
||||
SET token_version = COALESCE(token_version, 0) + 1, updated_at = NOW()
|
||||
WHERE id = ?
|
||||
""",
|
||||
(user_id,)
|
||||
)
|
||||
db.commit()
|
||||
|
||||
# 获取新的 token_version
|
||||
cur.execute(
|
||||
"SELECT token_version FROM qd_users WHERE id = ?",
|
||||
(user_id,)
|
||||
)
|
||||
row = cur.fetchone()
|
||||
cur.close()
|
||||
|
||||
new_version = int(row.get('token_version') or 1) if row else 1
|
||||
logger.info(f"Incremented token_version for user_id={user_id} to {new_version}")
|
||||
return new_version
|
||||
except Exception as e:
|
||||
logger.error(f"increment_token_version failed: {e}")
|
||||
return 1
|
||||
|
||||
def create_user(self, data: Dict[str, Any] = None, **kwargs) -> Optional[int]:
|
||||
"""
|
||||
Create a new user.
|
||||
|
||||
Reference in New Issue
Block a user