Add files via upload

This commit is contained in:
xiaochuan
2025-11-14 22:56:44 +00:00
committed by GitHub
parent cf1944a688
commit ef6e1e278d
96 changed files with 3300 additions and 0 deletions
+10
View File
@@ -0,0 +1,10 @@
"""
QuantTrader core package.
Expose runtime modules (data, execution, risk, strategy).
"""
from . import data # noqa: F401
from . import execution # noqa: F401
from . import risk # noqa: F401
from . import strategy # noqa: F401
Binary file not shown.
Binary file not shown.
+3
View File
@@ -0,0 +1,3 @@
# core.data package init
from . import base
from . import oanda
+136
View File
@@ -0,0 +1,136 @@
from abc import ABC, abstractmethod
from typing import Dict, List, Optional, Any
from datetime import datetime
import pandas as pd
class DataFeed(ABC):
"""
数据源的基础抽象类
定义了获取市场数据的标准接口
"""
def __init__(self, instrument: str, timeframe: str):
"""
初始化数据源
Args:
instrument: 交易品种 (例如: "EUR_USD")
timeframe: 时间周期 (例如: "H1", "D")
"""
self.instrument = instrument
self.timeframe = timeframe
@abstractmethod
async def get_historical_data(
self,
start: datetime,
end: datetime,
**kwargs
) -> pd.DataFrame:
"""
获取历史数据
Args:
start: 开始时间
end: 结束时间
**kwargs: 额外参数
Returns:
包含历史数据的DataFrame,至少应该包含以下列:
- datetime: 时间戳
- open: 开盘价
- high: 最高价
- low: 最低价
- close: 收盘价
- volume: 成交量(如果可用)
"""
pass
@abstractmethod
async def get_latest_data(self) -> Dict[str, Any]:
"""
获取最新的市场数据
Returns:
包含最新市场数据的字典
"""
pass
@abstractmethod
async def subscribe(self, callback) -> None:
"""
订阅实时数据更新
Args:
callback: 处理实时数据的回调函数
"""
pass
@abstractmethod
async def unsubscribe(self) -> None:
"""
取消订阅实时数据
"""
pass
class DataProcessor:
"""
数据处理器基类
用于对原始市场数据进行预处理和计算指标
"""
def __init__(self):
self.indicators = {}
def add_indicator(self, name: str, func, **params):
"""
添加技术指标计算
Args:
name: 指标名称
func: 计算指标的函数
**params: 指标参数
"""
self.indicators[name] = {
'function': func,
'params': params
}
def process_data(self, data: pd.DataFrame) -> pd.DataFrame:
"""
处理数据并计算所有已注册的指标
Args:
data: 原始市场数据
Returns:
添加了技术指标的DataFrame
"""
result = data.copy()
for name, indicator in self.indicators.items():
try:
result[name] = indicator['function'](
data,
**indicator['params']
)
except Exception as e:
print(f"计算指标 {name} 时发生错误: {str(e)}")
return result
class MarketDataEvent:
"""
市场数据事件类
用于在系统各层之间传递市场数据更新
"""
def __init__(
self,
instrument: str,
timestamp: datetime,
data: Dict[str, Any],
event_type: str = "MARKET_DATA"
):
self.event_type = event_type
self.instrument = instrument
self.timestamp = timestamp
self.data = data
+257
View File
@@ -0,0 +1,257 @@
from __future__ import annotations
import asyncio
import time
from datetime import datetime
from typing import Dict, Any, Optional, Sequence
import pandas as pd
from queue import Queue
import threading
from loguru import logger
from oandapyV20 import API
try:
from oandapyV20.endpoints.pricing import PricingStream as OandaPricingStream
except ImportError: # pragma: no cover
OandaPricingStream = None # type: ignore
from .base import DataFeed, MarketDataEvent
def _normalize_instrument(symbol: str) -> str:
"""规范化交易品种名称"""
s = symbol.upper().replace(" ", "").replace("/", "_").replace("-", "_")
if "_" in s and len(s) == 7:
return s
stripped = s.replace("_", "")
if len(stripped) == 6:
return f"{stripped[:3]}_{stripped[3:]}"
return s
class OANDADataFeed(DataFeed):
"""
OANDA数据源实现
提供实时和历史市场数据
"""
def __init__(
self,
instrument: str,
timeframe: str,
account_id: str,
access_token: str,
environment: str = "practice",
reconnect_wait: float = 5.0,
log_heartbeat: bool = False
):
super().__init__(instrument, timeframe)
self.account_id = account_id
self.access_token = access_token
self.environment = environment
self.reconnect_wait = reconnect_wait
self.log_heartbeat = log_heartbeat
self.client = API(access_token=access_token, environment=environment)
self._stop = threading.Event()
self._thread: Optional[threading.Thread] = None
self._callback = None
self._latest_data = None
self._loop: Optional[asyncio.AbstractEventLoop] = None
async def get_historical_data(
self,
start: datetime,
end: datetime,
**kwargs
) -> pd.DataFrame:
"""获取历史数据
Args:
start: 开始时间
end: 结束时间
**kwargs: 额外参数,支持:
- granularity: str, 时间周期 (如 "H1", "D")
- count: int, 返回的K线数量
Returns:
DataFrame包含以下列:datetime, open, high, low, close, volume
"""
from oandapyV20 import API
import oandapyV20.endpoints.instruments as instruments
granularity = kwargs.get("granularity", self.timeframe)
price = kwargs.get("price", "M")
# OANDA 的单次请求有最大返回数限制(例如 5000 candles),对长区间需要分段请求
MAX_CANDLES = 5000
# 估算每个candle的时间长度(秒),支持常见的granularity
def granularity_seconds(g: str) -> int:
if g.endswith('H'):
return int(g[:-1]) * 3600
if g.endswith('D'):
return int(g[:-1]) * 86400 if g[:-1].isdigit() else 86400
if g.endswith('M') and len(g) > 1 and g[0].isdigit():
# 例如 M1, M5 (分钟)
return int(g[1:]) * 60 if g[0] == 'M' else 30
# 默认按小时处理
return 3600
step_seconds = granularity_seconds(granularity) * MAX_CANDLES
all_data = []
current_start = pd.to_datetime(start)
end_ts = pd.to_datetime(end)
while current_start < end_ts:
current_end = current_start + pd.Timedelta(seconds=step_seconds)
if current_end > end_ts:
current_end = end_ts
params = {
"from": current_start.strftime("%Y-%m-%dT%H:%M:%S.000000Z"),
"to": current_end.strftime("%Y-%m-%dT%H:%M:%S.000000Z"),
"granularity": granularity,
"price": price
}
request = instruments.InstrumentsCandles(
instrument=self.instrument,
params=params
)
try:
response = self.client.request(request)
candles = response.get("candles", [])
for candle in candles:
if candle.get("complete"):
all_data.append({
"datetime": pd.to_datetime(candle["time"]),
"open": float(candle["mid"]["o"]),
"high": float(candle["mid"]["h"]),
"low": float(candle["mid"]["l"]),
"close": float(candle["mid"]["c"]),
"volume": int(candle.get("volume", 0))
})
except Exception as e:
logger.error(f"获取历史数据失败: {str(e)}")
raise
# 推进起点
current_start = current_end
if not all_data:
return pd.DataFrame()
df = pd.DataFrame(all_data)
df.drop_duplicates(subset=["datetime"], inplace=True)
df.sort_values(by="datetime", inplace=True)
df.set_index("datetime", inplace=True)
return df
async def get_latest_data(self) -> Dict[str, Any]:
"""获取最新数据"""
return self._latest_data if self._latest_data else {}
async def subscribe(self, callback) -> None:
"""订阅实时数据"""
self._callback = callback
self._loop = asyncio.get_running_loop()
if not self._thread or not self._thread.is_alive():
self._start_stream()
async def unsubscribe(self) -> None:
"""取消订阅"""
self._stop.set()
if self._thread:
self._thread.join(timeout=2.0)
self._loop = None
logger.info("[OANDA] Pricing stream stopped.")
def _start_stream(self) -> None:
"""启动价格流"""
if self._thread and self._thread.is_alive():
return
self._stop.clear()
self._thread = threading.Thread(target=self._run_stream, daemon=True)
self._thread.start()
logger.info(
f"[OANDA] Pricing stream started for {self.instrument} (account={self.account_id})"
)
def _run_stream(self) -> None:
"""运行价格流"""
if OandaPricingStream is None:
raise RuntimeError(
"oandapyV20.endpoints.pricing.PricingStream is unavailable. "
"Ensure oandapyV20 is installed."
)
params = {"instruments": self.instrument}
while not self._stop.is_set():
request = OandaPricingStream(
accountID=self.account_id,
params=params
)
try:
for msg in self.client.request(request):
if self._stop.is_set():
break
self._handle_msg(msg)
except Exception as exc:
if self._stop.is_set():
break
logger.warning(
f"[OANDA] Pricing stream error: {exc}. "
f"Reconnecting in {self.reconnect_wait}s"
)
time.sleep(self.reconnect_wait)
def _handle_msg(self, msg: dict) -> None:
"""处理价格消息"""
msg_type = msg.get("type")
if msg_type == "HEARTBEAT":
if self.log_heartbeat:
logger.debug(f"[OANDA] Heartbeat {msg.get('time')}")
return
if msg_type != "PRICE":
logger.debug(f"[OANDA] Skip message type={msg_type}")
return
try:
bids = msg.get("bids")
asks = msg.get("asks")
if not bids or not asks:
return
bid = float(bids[0]["price"])
ask = float(asks[0]["price"])
timestamp = pd.to_datetime(msg["time"]).to_pydatetime()
self._latest_data = {
"bid": bid,
"ask": ask,
"timestamp": timestamp,
"mid": (bid + ask) / 2
}
if self._callback and self._loop:
event = MarketDataEvent(
instrument=self.instrument,
timestamp=timestamp,
data=self._latest_data
)
try:
coro = self._callback(event)
if asyncio.iscoroutine(coro):
asyncio.run_coroutine_threadsafe(coro, self._loop)
else:
self._loop.call_soon_threadsafe(self._callback, event)
except RuntimeError as exc:
logger.warning(f"[OANDA] Failed to dispatch callback: {exc}")
elif self._callback and not self._loop:
logger.warning("[OANDA] Callback set but event loop missing; dropping tick")
except Exception as exc:
logger.warning(f"[OANDA] Malformed price message: {msg} ({exc})")
+38
View File
@@ -0,0 +1,38 @@
# fx_backtest/core/events.py
import os
import sys
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from dataclasses import dataclass
from datetime import datetime
from typing import Literal, Optional
@dataclass(frozen=True)
class TickEvent:
ts: datetime
symbol: str
bid: float
ask: float
@dataclass(frozen=True)
class SignalEvent:
ts: datetime
symbol: str
direction: Literal["LONG", "SHORT", "EXIT"]
size: float # 单位:合约单位/手,随你定义
@dataclass(frozen=True)
class OrderEvent:
ts: datetime
symbol: str
side: Literal["BUY", "SELL"]
size: float
price: Optional[float] = None # 市价可为 None;限价时填价格
@dataclass(frozen=True)
class FillEvent:
ts: datetime
symbol: str
side: Literal["BUY", "SELL"]
size: float
price: float
commission: float
+14
View File
@@ -0,0 +1,14 @@
# 执行层包初始化文件
from .base import ExecutionHandler, OrderEvent, FillEvent
try:
from .oanda_handler import OANDAExecutionHandler
except Exception:
# optional: OANDA handler may not be available if dependencies missing
OANDAExecutionHandler = None
__all__ = [
"ExecutionHandler",
"OrderEvent",
"FillEvent",
"OANDAExecutionHandler",
]
+182
View File
@@ -0,0 +1,182 @@
from abc import ABC, abstractmethod
from typing import Dict, List, Optional, Any
from datetime import datetime
import asyncio
from loguru import logger
from ..strategy.base import SignalEvent, Position
class OrderEvent:
"""订单事件"""
def __init__(
self,
instrument: str,
order_type: str, # "MARKET", "LIMIT", "STOP"
direction: str, # "BUY", "SELL"
quantity: float,
timestamp: datetime,
price: Optional[float] = None,
stop_loss: Optional[float] = None,
take_profit: Optional[float] = None,
order_id: Optional[str] = None
):
self.event_type = "ORDER"
self.instrument = instrument
self.order_type = order_type
self.direction = direction
self.quantity = quantity
self.timestamp = timestamp
self.price = price
self.stop_loss = stop_loss
self.take_profit = take_profit
self.order_id = order_id
self.status = "CREATED" # CREATED, SUBMITTED, FILLED, CANCELLED, REJECTED
class FillEvent:
"""成交事件"""
def __init__(
self,
instrument: str,
direction: str,
quantity: float,
price: float,
timestamp: datetime,
commission: float = 0.0,
order_id: Optional[str] = None
):
self.event_type = "FILL"
self.instrument = instrument
self.direction = direction
self.quantity = quantity
self.price = price
self.timestamp = timestamp
self.commission = commission
self.order_id = order_id
class ExecutionHandler(ABC):
"""
执行处理器基类
负责订单执行和管理
"""
def __init__(self):
self.orders: Dict[str, OrderEvent] = {} # order_id -> OrderEvent
self.positions: Dict[str, Position] = {} # instrument -> Position
self.fills: List[FillEvent] = []
self._order_callbacks = []
self._fill_callbacks = []
async def process_signal(self, signal: SignalEvent, price: Optional[float] = None) -> OrderEvent:
"""
处理交易信号并创建订单
Args:
signal: 交易信号
Returns:
创建的订单事件
"""
order = OrderEvent(
instrument=signal.instrument,
order_type="MARKET", # 默认为市价单
direction="BUY" if signal.signal_type == "LONG" else "SELL",
quantity=abs(signal.strength),
timestamp=signal.timestamp,
stop_loss=signal.stop_loss,
take_profit=signal.take_profit,
price=price
)
# 生成订单ID
order.order_id = f"{order.instrument}_{order.timestamp.strftime('%Y%m%d_%H%M%S')}"
self.orders[order.order_id] = order
# 执行订单
try:
await self.execute_order(order)
except Exception as e:
logger.error(f"订单执行失败: {str(e)}")
order.status = "REJECTED"
return order
@abstractmethod
async def execute_order(self, order: OrderEvent) -> None:
"""
执行订单
Args:
order: 要执行的订单
"""
pass
@abstractmethod
async def cancel_order(self, order_id: str) -> bool:
"""
取消订单
Args:
order_id: 要取消的订单ID
Returns:
是否成功取消
"""
pass
def add_order_callback(self, callback):
"""添加订单状态更新回调"""
self._order_callbacks.append(callback)
def add_fill_callback(self, callback):
"""添加成交更新回调"""
self._fill_callbacks.append(callback)
async def _notify_order(self, order: OrderEvent):
"""通知订单状态更新"""
for callback in self._order_callbacks:
await callback(order)
async def _notify_fill(self, fill: FillEvent):
"""通知成交更新"""
for callback in self._fill_callbacks:
await callback(fill)
def get_position(self, instrument: str) -> Optional[Position]:
"""获取某个品种的持仓"""
return self.positions.get(instrument)
def get_all_positions(self) -> List[Position]:
"""获取所有持仓"""
return list(self.positions.values())
def update_position(self, fill: FillEvent) -> None:
"""根据成交更新持仓"""
instrument = fill.instrument
position = self.positions.get(instrument)
if position is None:
# 新建仓位
position = Position(
instrument=instrument,
direction=fill.direction,
size=fill.quantity,
entry_price=fill.price,
entry_time=fill.timestamp
)
self.positions[instrument] = position
else:
# 更新现有仓位
if fill.direction == position.direction:
# 同向加仓
new_size = position.size + fill.quantity
position.entry_price = (position.entry_price * position.size +
fill.price * fill.quantity) / new_size
position.size = new_size
else:
# 反向减仓
position.size -= fill.quantity
if position.size <= 0:
# 清仓
del self.positions[instrument]
+117
View File
@@ -0,0 +1,117 @@
from typing import Optional
from datetime import datetime
from loguru import logger
from oandapyV20 import API
import oandapyV20.endpoints.orders as orders
import csv
from pathlib import Path
from .base import ExecutionHandler, OrderEvent, FillEvent
class OANDAExecutionHandler(ExecutionHandler):
"""实盘环境下的 OANDA 执行实现。"""
def __init__(self, account_id: str, access_token: str, environment: str = "practice", fills_path: Optional[str] = None):
super().__init__()
self.client = API(access_token=access_token, environment=environment)
self.account_id = account_id
default_dir = Path("results/execution/live") if environment == "live" else Path("results/execution/paper")
default_dir.mkdir(parents=True, exist_ok=True)
self._fills_csv = Path(fills_path) if fills_path else default_dir / "fills.csv"
if not self._fills_csv.exists():
self._init_csv()
def _init_csv(self) -> None:
with self._fills_csv.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(
f,
fieldnames=["order_id", "ts", "symbol", "pnl", "adapter_latency_ms", "direction", "price", "quantity"],
)
writer.writeheader()
async def execute_order(self, order: OrderEvent) -> None:
# 将 OrderEvent 转换为 OANDA 下单请求
instrument = order.instrument.replace("_", "_")
units = int(order.quantity) if order.direction.upper() in ("BUY", "LONG") else -int(order.quantity)
order_body = {
"instrument": instrument,
"units": str(units),
"timeInForce": "FOK",
"positionFill": "DEFAULT",
}
if order.price:
order_body["type"] = "LIMIT"
order_body["price"] = f"{order.price:.5f}"
order_body["timeInForce"] = "GTC"
else:
order_body["type"] = "MARKET"
payload = {"order": order_body}
req = orders.OrderCreate(accountID=self.account_id, data=payload)
try:
resp = self.client.request(req)
except Exception as exc:
logger.error(f"[OANDA] Order submission failed for {instrument}: {exc}")
order.status = "REJECTED"
await self._notify_order(order)
return
order.status = "SUBMITTED"
await self._notify_order(order)
fill_txn = resp.get("orderFillTransaction")
if not fill_txn:
logger.warning(f"[OANDA] Order accepted but no fill: {resp}")
return
try:
price = float(fill_txn["price"])
filled_units = abs(float(fill_txn["units"]))
commission = float(fill_txn.get("commission", 0))
ts = datetime.fromisoformat(fill_txn["time"].replace("Z", "+00:00"))
side = "BUY" if float(fill_txn["units"]) > 0 else "SELL"
except Exception as exc:
logger.error(f"[OANDA] Unable to parse fill transaction: {fill_txn} ({exc})")
return
fill = FillEvent(
instrument=order.instrument,
direction=side,
quantity=filled_units,
price=price,
timestamp=ts,
commission=abs(commission),
order_id=order.order_id
)
order.status = "FILLED"
await self._notify_order(order)
await self._notify_fill(fill)
self.fills.append(fill)
self.update_position(fill)
self._append_fill_csv(fill)
def _append_fill_csv(self, fill: FillEvent) -> None:
record = {
"order_id": fill.order_id or "",
"ts": fill.timestamp.isoformat(),
"symbol": fill.instrument,
"pnl": 0.0,
"adapter_latency_ms": None,
"direction": fill.direction,
"price": fill.price,
"quantity": fill.quantity,
}
with self._fills_csv.open("a", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=record.keys())
writer.writerow(record)
async def cancel_order(self, order_id: str) -> bool:
# OANDA 取消需要调用 OrderCancel 或交易 API;简单实现为更新状态
if order_id in self.orders:
order = self.orders[order_id]
order.status = "CANCELLED"
await self._notify_order(order)
return True
return False
+93
View File
@@ -0,0 +1,93 @@
"""
OANDA 执行适配器:把 OrderEvent 转换为 OANDA API 下单,并回写 FillEvent。
"""
from __future__ import annotations
from queue import Queue
import pandas as pd
from loguru import logger
from oandapyV20 import API
import oandapyV20.endpoints.orders as orders
from .events import FillEvent, OrderEvent
def _normalize_instrument(symbol: str) -> str:
s = symbol.upper().replace(" ", "").replace("/", "_").replace("-", "_")
if "_" in s and len(s) == 7:
return s
stripped = s.replace("_", "")
if len(stripped) == 6:
return f"{stripped[:3]}_{stripped[3:]}"
return s
class OandaExecution:
"""
把 OrderEvent 翻译为 OANDA 订单;成交后投递 FillEvent。
"""
def __init__(
self,
q: Queue,
account_id: str,
access_token: str,
environment: str = "practice",
) -> None:
self.q = q
self.account_id = account_id
self.client = API(access_token=access_token, environment=environment)
def on_event(self, ev) -> None:
if not isinstance(ev, OrderEvent):
return
instrument = _normalize_instrument(ev.symbol)
units = ev.size if ev.side == "BUY" else -ev.size
order_body = {
"instrument": instrument,
"units": str(int(units)),
"timeInForce": "FOK",
"positionFill": "DEFAULT",
}
if ev.price is None:
order_body["type"] = "MARKET"
else:
order_body["type"] = "LIMIT"
order_body["timeInForce"] = "GTC"
order_body["price"] = f"{ev.price:.5f}"
payload = {"order": order_body}
req = orders.OrderCreate(accountID=self.account_id, data=payload)
try:
resp = self.client.request(req)
except Exception as exc:
logger.error(f"[OANDA] Order submission failed for {instrument}: {exc}")
return
fill_txn = resp.get("orderFillTransaction")
if not fill_txn:
logger.warning(f"[OANDA] Order accepted but no fill: {resp}")
return
try:
price = float(fill_txn["price"])
filled_units = abs(float(fill_txn["units"]))
commission = float(fill_txn.get("commission", 0))
ts = pd.to_datetime(fill_txn["time"]).to_pydatetime()
side = "BUY" if float(fill_txn["units"]) > 0 else "SELL"
except Exception as exc:
logger.error(f"[OANDA] Unable to parse fill transaction: {fill_txn} ({exc})")
return
self.q.put(
FillEvent(
ts=ts,
symbol=ev.symbol,
side=side,
size=filled_units,
price=price,
commission=abs(commission),
)
)
+2
View File
@@ -0,0 +1,2 @@
# core.risk package init
from . import base
+161
View File
@@ -0,0 +1,161 @@
from abc import ABC, abstractmethod
from typing import Dict, List, Optional, Any
from datetime import datetime
import pandas as pd
from ..strategy.base import Position, SignalEvent
class RiskEvent:
"""风险事件"""
def __init__(
self,
event_type: str, # "RISK_LIMIT", "STOP_LOSS", "MARGIN_CALL" etc.
instrument: str,
timestamp: datetime,
message: str,
severity: str = "WARNING", # "INFO", "WARNING", "CRITICAL"
data: Optional[Dict[str, Any]] = None
):
self.event_type = event_type
self.instrument = instrument
self.timestamp = timestamp
self.message = message
self.severity = severity
self.data = data or {}
class PositionSizer(ABC):
"""
仓位管理器基类
负责计算每笔交易的具体仓位大小
"""
@abstractmethod
def calculate_position_size(
self,
signal: SignalEvent,
portfolio_value: float,
risk_per_trade: float
) -> float:
"""
计算交易仓位大小
Args:
signal: 交易信号
portfolio_value: 当前组合总价值
risk_per_trade: 每笔交易的风险比例
Returns:
建议的仓位大小
"""
pass
class RiskManager(ABC):
"""
风险管理器基类
负责风险控制和监控
"""
def __init__(
self,
max_position_size: float,
max_portfolio_risk: float,
max_drawdown: float
):
self.max_position_size = max_position_size
self.max_portfolio_risk = max_portfolio_risk
self.max_drawdown = max_drawdown
self.current_drawdown = 0.0
self.peak_value = 0.0
@abstractmethod
async def check_signal(self, signal: SignalEvent) -> bool:
"""
检查交易信号是否符合风险控制要求
Args:
signal: 交易信号
Returns:
True if signal is acceptable, False otherwise
"""
pass
@abstractmethod
async def check_position(self, position: Position) -> List[RiskEvent]:
"""
检查持仓的风险状况
Args:
position: 当前持仓
Returns:
风险事件列表
"""
pass
def update_drawdown(self, portfolio_value: float) -> Optional[RiskEvent]:
"""
更新和检查回撤状况
Args:
portfolio_value: 当前组合价值
Returns:
如果超过最大回撤限制,返回风险事件
"""
if portfolio_value > self.peak_value:
self.peak_value = portfolio_value
self.current_drawdown = 0.0
else:
self.current_drawdown = (self.peak_value - portfolio_value) / self.peak_value
if self.current_drawdown > self.max_drawdown:
return RiskEvent(
event_type="MAX_DRAWDOWN_BREACH",
instrument="PORTFOLIO",
timestamp=datetime.now(),
message=f"Maximum drawdown breached: {self.current_drawdown:.2%}",
severity="CRITICAL",
data={"drawdown": self.current_drawdown}
)
return None
class SimpleRiskManager(RiskManager):
"""
简单风险管理器实现
实现基本的风险控制功能
"""
async def check_signal(self, signal: SignalEvent) -> bool:
"""检查交易信号"""
# 实现基本的信号检查逻辑
if not signal.stop_loss:
return False # 要求必须有止损
return True
async def check_position(self, position: Position) -> List[RiskEvent]:
"""检查持仓风险"""
events = []
# 检查持仓规模
if abs(position.size) > self.max_position_size:
events.append(RiskEvent(
event_type="POSITION_SIZE_LIMIT",
instrument=position.instrument,
timestamp=datetime.now(),
message=f"Position size {position.size} exceeds limit {self.max_position_size}",
severity="WARNING"
))
# 检查止损
if not position.stop_loss:
events.append(RiskEvent(
event_type="MISSING_STOP_LOSS",
instrument=position.instrument,
timestamp=datetime.now(),
message="Position has no stop loss",
severity="WARNING"
))
return events
+70
View File
@@ -0,0 +1,70 @@
"""Lightweight risk engine enforcing exposure, leverage, and loss caps."""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Dict, Tuple
@dataclass
class RiskLimits:
max_position_notional: float
max_gross_leverage: float
max_daily_loss: float
max_drawdown: float
@dataclass
class RiskState:
equity: float = 0.0
peak_equity: float = 0.0
min_equity: float = float("inf")
realized_pnl: float = 0.0
gross_notional: float = 0.0
exposures: Dict[str, float] = field(default_factory=dict)
class RiskViolation(Exception):
"""Raised when orders violate limits."""
class RiskEngine:
def __init__(self, limits: RiskLimits, starting_equity: float):
self.limits = limits
self.state = RiskState(equity=starting_equity, peak_equity=starting_equity, min_equity=starting_equity)
def evaluate_order(self, symbol: str, side: str, notional: float) -> Tuple[bool, str]:
exposure = self.state.exposures.get(symbol, 0.0)
proposed = exposure + (notional if side.lower() == "buy" else -notional)
if abs(proposed) > self.limits.max_position_notional:
return False, f"symbol_exposure_limit:{symbol}"
gross = self.state.gross_notional + abs(notional)
leverage = gross / self.state.equity if self.state.equity else float("inf")
if leverage > self.limits.max_gross_leverage:
return False, "gross_leverage_limit"
return True, "ok"
def record_fill(self, symbol: str, side: str, notional: float, pnl: float) -> None:
delta = notional if side.lower() == "buy" else -notional
self.state.exposures[symbol] = self.state.exposures.get(symbol, 0.0) + delta
self.state.gross_notional = sum(abs(v) for v in self.state.exposures.values())
self.state.realized_pnl += pnl
self.state.equity += pnl
self.state.peak_equity = max(self.state.peak_equity, self.state.equity)
self.state.min_equity = min(self.state.min_equity, self.state.equity)
def check_loss_limits(self) -> Tuple[bool, str]:
if -self.state.realized_pnl > self.limits.max_daily_loss:
return False, "daily_loss_limit"
drawdown = (self.state.equity - self.state.peak_equity) / self.state.peak_equity if self.state.peak_equity else 0.0
if drawdown < -self.limits.max_drawdown:
return False, "drawdown_limit"
return True, "ok"
def max_drawdown_pct(self) -> float:
if not self.state.peak_equity:
return 0.0
trough = self.state.min_equity if self.state.min_equity != float("inf") else self.state.equity
return abs((trough - self.state.peak_equity) / self.state.peak_equity)
+3
View File
@@ -0,0 +1,3 @@
# core.strategy package init
from . import base
from . import rsi_mean_reversion
+122
View File
@@ -0,0 +1,122 @@
from abc import ABC, abstractmethod
from typing import Dict, List, Optional, Any
from datetime import datetime
import pandas as pd
from ..data.base import MarketDataEvent
class Position:
"""持仓类,表示当前市场头寸"""
def __init__(
self,
instrument: str,
direction: str, # "LONG" or "SHORT"
size: float,
entry_price: float,
entry_time: datetime,
stop_loss: Optional[float] = None,
take_profit: Optional[float] = None
):
self.instrument = instrument
self.direction = direction
self.size = size
self.entry_price = entry_price
self.entry_time = entry_time
self.stop_loss = stop_loss
self.take_profit = take_profit
self.unrealized_pnl = 0.0
self.realized_pnl = 0.0
class SignalEvent:
"""交易信号事件"""
def __init__(
self,
instrument: str,
timestamp: datetime,
signal_type: str, # "LONG", "SHORT", "EXIT"
direction: str,
strength: float = 1.0,
stop_loss: Optional[float] = None,
take_profit: Optional[float] = None
):
self.event_type = "SIGNAL"
self.instrument = instrument
self.timestamp = timestamp
self.signal_type = signal_type
self.direction = direction
self.strength = strength
self.stop_loss = stop_loss
self.take_profit = take_profit
class Strategy(ABC):
"""
策略基类
定义了策略开发的标准接口
"""
def __init__(
self,
instrument: str,
position_size: float = 1.0,
max_positions: int = 1
):
self.instrument = instrument
self.position_size = position_size
self.max_positions = max_positions
self.positions: List[Position] = []
self.historical_data: Optional[pd.DataFrame] = None
@abstractmethod
async def on_data(self, event: MarketDataEvent) -> Optional[SignalEvent]:
"""
处理市场数据更新
Args:
event: 市场数据事件
Returns:
如果产生交易信号,返回SignalEvent;否则返回None
"""
pass
@abstractmethod
async def calculate_signals(self, data: pd.DataFrame) -> List[SignalEvent]:
"""
基于历史数据计算交易信号
Args:
data: 历史市场数据
Returns:
交易信号列表
"""
pass
def update_position(self, position: Position, current_price: float) -> None:
"""
更新持仓的未实现盈亏
Args:
position: 需要更新的持仓
current_price: 当前市场价格
"""
if position.direction == "LONG":
position.unrealized_pnl = (current_price - position.entry_price) * position.size
else:
position.unrealized_pnl = (position.entry_price - current_price) * position.size
def can_open_position(self) -> bool:
"""检查是否可以开新仓位"""
return len(self.positions) < self.max_positions
def get_position_value(self) -> float:
"""获取当前持仓的总价值"""
return sum(abs(pos.unrealized_pnl) for pos in self.positions)
def get_total_pnl(self) -> float:
"""获取总盈亏(已实现 + 未实现)"""
unrealized = sum(pos.unrealized_pnl for pos in self.positions)
realized = sum(pos.realized_pnl for pos in self.positions)
return realized + unrealized
+65
View File
@@ -0,0 +1,65 @@
from typing import List, Optional
import pandas as pd
from datetime import datetime
from .base import Strategy, SignalEvent
from .base import Position
from ..data.base import MarketDataEvent
class MACrossoverStrategy(Strategy):
"""
简单移动平均交叉策略
快线上穿慢线做多,下穿做空
"""
def __init__(self, instrument: str, fast_period: int = 50, slow_period: int = 200, position_size: float = 1.0):
super().__init__(instrument, position_size)
self.fast_period = fast_period
self.slow_period = slow_period
async def on_data(self, event: MarketDataEvent) -> Optional[SignalEvent]:
if self.historical_data is None:
return None
# 更新收盘价
close = event.data.get('close') or event.data.get('mid')
self.historical_data.loc[event.timestamp] = {
'open': event.data.get('open', close),
'high': event.data.get('high', close),
'low': event.data.get('low', close),
'close': close
}
if len(self.historical_data) < self.slow_period:
return None
fast = self.historical_data['close'].rolling(self.fast_period).mean()
slow = self.historical_data['close'].rolling(self.slow_period).mean()
if fast.iloc[-2] <= slow.iloc[-2] and fast.iloc[-1] > slow.iloc[-1]:
return SignalEvent(
instrument=self.instrument,
timestamp=event.timestamp,
signal_type="LONG",
direction="BUY",
strength=self.position_size
)
if fast.iloc[-2] >= slow.iloc[-2] and fast.iloc[-1] < slow.iloc[-1]:
return SignalEvent(
instrument=self.instrument,
timestamp=event.timestamp,
signal_type="SHORT",
direction="SELL",
strength=self.position_size
)
return None
async def calculate_signals(self, data: pd.DataFrame) -> List[SignalEvent]:
signals = []
self.historical_data = data.copy()
fast = data['close'].rolling(self.fast_period).mean()
slow = data['close'].rolling(self.slow_period).mean()
for i in range(self.slow_period, len(data)):
ts = data.index[i]
if fast.iloc[i-1] <= slow.iloc[i-1] and fast.iloc[i] > slow.iloc[i]:
signals.append(SignalEvent(instrument=self.instrument, timestamp=ts, signal_type="LONG", direction="BUY", strength=self.position_size))
if fast.iloc[i-1] >= slow.iloc[i-1] and fast.iloc[i] < slow.iloc[i]:
signals.append(SignalEvent(instrument=self.instrument, timestamp=ts, signal_type="SHORT", direction="SELL", strength=self.position_size))
return signals
+46
View File
@@ -0,0 +1,46 @@
from typing import List, Optional
import pandas as pd
from datetime import datetime
from .base import Strategy, SignalEvent
from ..data.base import MarketDataEvent
class MomentumStrategy(Strategy):
"""
简单动量策略:当价格高于N日均线时做多,低于时做空
"""
def __init__(self, instrument: str, lookback: int = 20, position_size: float = 1.0):
super().__init__(instrument, position_size)
self.lookback = lookback
async def on_data(self, event: MarketDataEvent) -> Optional[SignalEvent]:
if self.historical_data is None:
return None
close = event.data.get('close') or event.data.get('mid')
self.historical_data.loc[event.timestamp] = {
'open': event.data.get('open', close),
'high': event.data.get('high', close),
'low': event.data.get('low', close),
'close': close
}
if len(self.historical_data) < self.lookback:
return None
ma = self.historical_data['close'].rolling(self.lookback).mean()
if close > ma.iloc[-1] and self.can_open_position():
return SignalEvent(self.instrument, event.timestamp, "LONG", "BUY", strength=self.position_size)
if close < ma.iloc[-1] and self.can_open_position():
return SignalEvent(self.instrument, event.timestamp, "SHORT", "SELL", strength=self.position_size)
return None
async def calculate_signals(self, data: pd.DataFrame) -> List[SignalEvent]:
signals = []
self.historical_data = data.copy()
ma = data['close'].rolling(self.lookback).mean()
for i in range(self.lookback, len(data)):
ts = data.index[i]
price = data['close'].iloc[i]
if price > ma.iloc[i-1]:
signals.append(SignalEvent(self.instrument, ts, "LONG", "BUY", strength=self.position_size))
elif price < ma.iloc[i-1]:
signals.append(SignalEvent(self.instrument, ts, "SHORT", "SELL", strength=self.position_size))
return signals
@@ -0,0 +1,150 @@
from typing import List, Optional
import pandas as pd
import numpy as np
from datetime import datetime
from ..data.base import MarketDataEvent
from .base import Strategy, SignalEvent
class RSIMeanReversionStrategy(Strategy):
"""
RSI均值回归策略
当RSI超买时做空,超卖时做多
"""
def __init__(
self,
instrument: str,
position_size: float = 1.0,
max_positions: int = 1,
rsi_period: int = 14,
overbought: float = 70.0,
oversold: float = 30.0,
stop_loss_atr: float = 2.0,
atr_period: int = 14
):
super().__init__(instrument, position_size, max_positions)
self.rsi_period = rsi_period
self.overbought = overbought
self.oversold = oversold
self.stop_loss_atr = stop_loss_atr
self.atr_period = atr_period
self.last_rsi = None
self.last_atr = None
@staticmethod
def calculate_rsi(data: pd.Series, period: int = 14) -> pd.Series:
"""计算RSI指标"""
delta = data.diff()
gain = (delta.where(delta > 0, 0)).rolling(window=period).mean()
loss = (-delta.where(delta < 0, 0)).rolling(window=period).mean()
rs = gain / loss
return 100 - (100 / (1 + rs))
@staticmethod
def calculate_atr(data: pd.DataFrame, period: int = 14) -> pd.Series:
"""计算ATR指标"""
high = data['high']
low = data['low']
close = data['close']
tr1 = high - low
tr2 = abs(high - close.shift())
tr3 = abs(low - close.shift())
tr = pd.concat([tr1, tr2, tr3], axis=1).max(axis=1)
return tr.rolling(window=period).mean()
async def on_data(self, event: MarketDataEvent) -> Optional[SignalEvent]:
"""
处理实时市场数据
Args:
event: 市场数据事件
Returns:
如果触发信号则返回SignalEvent,否则返回None
"""
if self.historical_data is None:
return None
# 更新数据
current_price = event.data['mid']
self.historical_data.loc[event.timestamp] = current_price
# 计算指标
close_prices = self.historical_data['close']
rsi = self.calculate_rsi(close_prices, self.rsi_period).iloc[-1]
atr = self.calculate_atr(self.historical_data, self.atr_period).iloc[-1]
self.last_rsi = rsi
self.last_atr = atr
# 生成信号
if self.can_open_position():
if rsi > self.overbought:
return SignalEvent(
instrument=self.instrument,
timestamp=event.timestamp,
signal_type="SHORT",
direction="SELL",
strength=self.position_size,
stop_loss=current_price + self.stop_loss_atr * atr
)
elif rsi < self.oversold:
return SignalEvent(
instrument=self.instrument,
timestamp=event.timestamp,
signal_type="LONG",
direction="BUY",
strength=self.position_size,
stop_loss=current_price - self.stop_loss_atr * atr
)
return None
async def calculate_signals(self, data: pd.DataFrame) -> List[SignalEvent]:
"""
基于历史数据计算交易信号
Args:
data: 历史市场数据
Returns:
交易信号列表
"""
signals = []
self.historical_data = data.copy()
# 计算指标
close_prices = data['close']
rsi = self.calculate_rsi(close_prices, self.rsi_period)
atr = self.calculate_atr(data, self.atr_period)
# 生成信号
for i in range(self.rsi_period, len(data)):
timestamp = data.index[i]
current_price = close_prices[i]
current_rsi = rsi[i]
current_atr = atr[i]
if current_rsi > self.overbought:
signals.append(SignalEvent(
instrument=self.instrument,
timestamp=timestamp,
signal_type="SHORT",
direction="SELL",
strength=self.position_size,
stop_loss=current_price + self.stop_loss_atr * current_atr
))
elif current_rsi < self.oversold:
signals.append(SignalEvent(
instrument=self.instrument,
timestamp=timestamp,
signal_type="LONG",
direction="BUY",
strength=self.position_size,
stop_loss=current_price - self.stop_loss_atr * current_atr
))
return signals