Files
AI-Trader/market/pivot_detector.py
T
guaiwoluo2020 51b2f30748 feat: 添加前端界面和市场分析模块
- 新增 Vue 3 + Vuetify 前端界面
- 新增市场分析模块 (market/)
- 更新主服务器和路由
- 更新 MT5 EA 文件
- 添加 .gitignore 排除临时文件
2026-03-10 17:38:13 +08:00

414 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
转折点检测模块
识别K线的高点和低点(分型识别)
"""
from collections import defaultdict
from datetime import datetime
from typing import List, Dict, Optional, Tuple
import threading
from .store import KlineData, normalize_symbol
class PivotPoint:
"""转折点数据结构"""
def __init__(self, symbol: str, period: str, timestamp, price: float,
direction: str, strength: int = 3):
self.symbol = normalize_symbol(symbol)
self.period = period
self.timestamp = timestamp
self.price = price
self.direction = direction # "high" 或 "low"
self.strength = strength # 转折强度(左右各N根K线)
def to_dict(self) -> Dict:
"""转换为字典"""
ts = self.timestamp
if isinstance(ts, datetime):
ts_str = ts.strftime("%Y-%m-%d %H:%M:%S")
else:
ts_str = str(ts)
return {
"symbol": self.symbol,
"period": self.period,
"timestamp": ts_str,
"price": self.price,
"direction": self.direction,
"strength": self.strength
}
class PivotDetector:
"""转折点检测器"""
# 各周期接近阈值(千分比)
THRESHOLDS = {
'H4': 0.0015, # 千分之1.5
'H1': 0.0015, # 千分之1.5
'M15': 0.0015, # 千分之1.5
'M5': 0.0005, # 千分之0.5
'M1': 0.0002 # 千分之0.2
}
def __init__(self):
# 存储转折点: {SYMBOL: {PERIOD: [PivotPoint, ...]}}
self._pivots = defaultdict(lambda: defaultdict(list))
self._lock = threading.RLock()
# 默认转折强度(左右各N根K线)
self.default_strength = 3
print("[PivotDetector] 转折点检测器已初始化")
def detect_pivots(self, symbol: str, period: str, klines: List[KlineData],
strength: int = None) -> List[PivotPoint]:
"""
检测转折点
Args:
symbol: 交易品种
period: 周期
klines: K线数据列表
strength: 转折强度(左右各N根K线)
Returns:
检测到的转折点列表
"""
if strength is None:
strength = self.default_strength
if len(klines) < 2 * strength + 1:
return []
pivots = []
# 遍历K线,检测分型
for i in range(strength, len(klines) - strength):
current = klines[i]
# 检查是否为高点(顶分型)
is_high = True
for j in range(1, strength + 1):
if klines[i - j].high >= current.high or klines[i + j].high >= current.high:
is_high = False
break
if is_high:
pivot = PivotPoint(
symbol=symbol,
period=period,
timestamp=current.timestamp,
price=current.high,
direction="high",
strength=strength
)
pivots.append(pivot)
# 检查是否为低点(底分型)
is_low = True
for j in range(1, strength + 1):
if klines[i - j].low <= current.low or klines[i + j].low <= current.low:
is_low = False
break
if is_low:
pivot = PivotPoint(
symbol=symbol,
period=period,
timestamp=current.timestamp,
price=current.low,
direction="low",
strength=strength
)
pivots.append(pivot)
return pivots
def _merge_pivots(self, pivots: List[PivotPoint], klines: List[KlineData]) -> List[PivotPoint]:
"""
合并相近的转折点
合并规则:
- K线距离小于26根
- 价格相差在万分之三范围内
- 高点合并:取较高的价格
- 低点合并:取较低的价格
Args:
pivots: 原始转折点列表
klines: K线数据(用于计算K线索引)
Returns:
合并后的转折点列表
"""
if len(pivots) < 2:
return pivots
# 建立K线时间戳到索引的映射
kline_index = {str(k.timestamp): i for i, k in enumerate(klines)}
# 按时间排序
pivots = sorted(pivots, key=lambda p: str(p.timestamp))
# 分开处理高点和低点
high_pivots = [p for p in pivots if p.direction == "high"]
low_pivots = [p for p in pivots if p.direction == "low"]
# 合并高点
merged_highs = self._merge_same_direction(
high_pivots, kline_index, "high"
)
# 合并低点
merged_lows = self._merge_same_direction(
low_pivots, kline_index, "low"
)
# 合并结果
result = merged_highs + merged_lows
return result
def _merge_same_direction(self, pivots: List[PivotPoint],
kline_index: Dict[str, int],
direction: str) -> List[PivotPoint]:
"""
合并同方向的转折点
"""
if len(pivots) < 2:
return pivots
merged = []
i = 0
while i < len(pivots):
current = pivots[i]
current_idx = kline_index.get(str(current.timestamp), -1)
if current_idx < 0:
i += 1
continue
# 查找需要合并的转折点
group = [current]
j = i + 1
while j < len(pivots):
next_pivot = pivots[j]
next_idx = kline_index.get(str(next_pivot.timestamp), -1)
if next_idx < 0:
j += 1
continue
# 检查K线距离
kline_distance = abs(next_idx - current_idx)
if kline_distance >= 26:
break
# 检查价格差距(万分之三)
if current.price > 0:
price_diff_pct = abs(next_pivot.price - current.price) / current.price
if price_diff_pct <= 0.0003: # 万分之三
group.append(next_pivot)
j += 1
continue
break
# 从组中选择代表性转折点
if direction == "high":
# 高点:取价格最高的
best = max(group, key=lambda p: p.price)
else:
# 低点:取价格最低的
best = min(group, key=lambda p: p.price)
merged.append(best)
i = j
return merged
def update_pivots(self, symbol: str, period: str, klines: List[KlineData],
strength: int = None) -> int:
"""
更新转折点数据
Returns:
更新后的转折点数量
"""
symbol = normalize_symbol(symbol)
pivots = self.detect_pivots(symbol, period, klines, strength)
# 合并相近的转折点
merged_pivots = self._merge_pivots(pivots, klines)
with self._lock:
self._pivots[symbol][period] = merged_pivots
count = len(merged_pivots)
original_count = len(pivots)
if original_count != count:
print(f"[PivotDetector] {symbol} {period} 检测到 {original_count} 个转折点,合并后 {count} 个")
else:
print(f"[PivotDetector] {symbol} {period} 检测到 {count} 个转折点")
return count
def get_pivots(self, symbol: str, period: str, direction: str = None,
count: int = 50) -> List[Dict]:
"""
获取转折点数据
Args:
symbol: 交易品种
period: 周期
direction: "high" 或 "low"None表示全部
count: 返回数量
Returns:
转折点列表
"""
symbol = normalize_symbol(symbol)
with self._lock:
pivots = self._pivots[symbol][period]
if direction:
pivots = [p for p in pivots if p.direction == direction]
# 按时间排序,返回最新的
pivots = sorted(pivots, key=lambda x: str(x.timestamp), reverse=True)[:count]
return [p.to_dict() for p in pivots]
def get_recent_pivots(self, symbol: str, period: str, count: int = 10) -> List[Dict]:
"""获取最近的转折点(按时间倒序)"""
symbol = normalize_symbol(symbol)
with self._lock:
pivots = self._pivots[symbol][period]
pivots = sorted(pivots, key=lambda x: str(x.timestamp), reverse=True)[:count]
return [p.to_dict() for p in pivots]
def check_near_pivot(self, symbol: str, current_price: float) -> List[Dict]:
"""
检查当前价格是否接近某个转折点
Args:
symbol: 交易品种
current_price: 当前价格
Returns:
接近的转折点列表,包含距离信息
预警逻辑:
- 接近高点:当前价格 < 高点价格 且 距离在阈值范围内
- 接近低点:当前价格 > 低点价格 且 距离在阈值范围内
- 突破高点:当前价格超过高点价格的万分之一点二(基于实时价格)
- 突破低点:当前价格低于低点价格的万分之一点二(基于实时价格)
- 超过千分之一不再提示
"""
symbol = normalize_symbol(symbol)
near_pivots = []
# 突破阈值:万分之一点二
BREAKTHROUGH_THRESHOLD = 0.00012
# 最大提示范围:千分之一
MAX_ALERT_THRESHOLD = 0.001
with self._lock:
for period in self._pivots[symbol]:
pivots = self._pivots[symbol][period]
threshold = self.THRESHOLDS.get(period, 0.001)
for pivot in pivots:
if pivot.price == 0 or current_price == 0:
continue
# 基于实时价格计算阈值
breakthrough_value = current_price * BREAKTHROUGH_THRESHOLD # 万分之一点二
max_alert_value = current_price * MAX_ALERT_THRESHOLD # 千分之一
# 判断是接近还是突破
is_near = False
is_breakthrough = False
alert_type = ""
if pivot.direction == "high":
# 高点转折
if current_price > pivot.price:
# 当前价格高于高点,判断是否突破
# 突破:超过高点的距离在万分之一点二到千分之一之间
distance = current_price - pivot.price
if distance >= breakthrough_value and distance < max_alert_value:
is_breakthrough = True
alert_type = "breakthrough_high"
# 超过千分之一不再提示
else:
# 当前价格低于高点
distance_pct = (pivot.price - current_price) / current_price
if distance_pct <= threshold:
is_near = True
alert_type = "near_high"
elif pivot.direction == "low":
# 低点转折
if current_price < pivot.price:
# 当前价格低于低点,判断是否突破
# 突破:低于低点的距离在万分之一点二到千分之一之间
distance = pivot.price - current_price
if distance >= breakthrough_value and distance < max_alert_value:
is_breakthrough = True
alert_type = "breakthrough_low"
# 超过千分之一不再提示
else:
# 当前价格高于低点
distance_pct = (current_price - pivot.price) / current_price
if distance_pct <= threshold:
is_near = True
alert_type = "near_low"
if is_near or is_breakthrough:
distance_pct = abs(current_price - pivot.price) / current_price
near_pivots.append({
**pivot.to_dict(),
"current_price": current_price,
"distance_pct": round(distance_pct * 100, 4),
"threshold_pct": round(threshold * 100, 4),
"distance": round(current_price - pivot.price, 2),
"alert_type": alert_type,
"is_breakthrough": is_breakthrough
})
# 按距离排序,最近的优先
near_pivots.sort(key=lambda x: x['distance_pct'])
return near_pivots
def get_threshold(self, period: str) -> float:
"""获取某个周期的接近阈值"""
return self.THRESHOLDS.get(period, 0.001)
def clear_symbol(self, symbol: str):
"""清除某个Symbol的转折点数据"""
symbol = normalize_symbol(symbol)
with self._lock:
if symbol in self._pivots:
del self._pivots[symbol]
def get_status(self) -> Dict:
"""获取状态"""
with self._lock:
status = {}
for symbol in self._pivots:
status[symbol] = {}
for period in self._pivots[symbol]:
count = len(self._pivots[symbol][period])
status[symbol][period] = {"pivot_count": count}
return status