This commit is contained in:
Ichinga Samuel
2024-09-30 06:10:19 +01:00
parent e2658425e6
commit 3b37508d4c
28 changed files with 977 additions and 648 deletions
+2 -4
View File
@@ -6,13 +6,11 @@ import logging
from .executor import Executor
from .account import Account
from .core.config import Config
from .symbol import Symbol as _Symbol
from .strategy import Strategy as _Strategy
from .symbol import Symbol as Symbol
from .strategy import Strategy as Strategy
logger = logging.getLogger(__name__)
Strategy = TypeVar("Strategy", bound=_Strategy)
Symbol = TypeVar("Symbol", bound=_Symbol)
class Bot:
+2 -2
View File
@@ -199,8 +199,8 @@ class Candles(Generic[_Candle]):
return self._data[index]
elif isinstance(index, int):
index = index if index >= 0 else len(self) + index
return self.Candle(**self._data.iloc[index], Index=index)
index_ = index if index >= 0 else len(self) + index
return self.Candle(**self._data.iloc[index], Index=index_)
raise TypeError(f"Expected int, slice or str got {type(index)}")
def __setitem__(self, index, value: Series):
+4 -4
View File
@@ -1,8 +1,8 @@
from .meta_tester import MetaTester
from .backtest_engine import BackTestEngine
from .get_data import GetData
from .test_strategy import TestStrategy
from .get_data import GetData, TestData
from .strategy_tester import StrategyTester
from .event_manager import EventManager
from .strategy_tester import StrategyTester, SingleStrategyTester
from .backtester import BackTester
from .test_account import TestAccount
from .types import PositionsManager, OrdersManager, DealsManager
from .trades_manager import PositionsManager, OrdersManager, DealsManager
+239 -171
View File
@@ -14,17 +14,16 @@ from MetaTrader5 import (Tick, SymbolInfo, AccountInfo, TradeOrder, TradePositio
TradeRequest, OrderCheckResult, OrderSendResult, TerminalInfo)
from ...core.meta_trader import MetaTrader
from ...core.constants import (TimeFrame, CopyTicks, OrderType, TradeAction, AccountStopOutMode, PositionReason,
DealType, DealReason, DealEntry, OrderReason)
from ...core.constants import (TimeFrame, OrderType, TradeAction, AccountStopOutMode, PositionReason,
DealType, DealReason, DealEntry, OrderReason, CopyTicks)
from ...core.config import Config
from ...utils import round_down, round_up, error_handler, error_handler_sync, async_cache
from .get_data import TestData, GetData, Cursor
from .test_account import TestAccount
from .types import PositionsManager, OrdersManager, DealsManager
from .trades_manager import PositionsManager, OrdersManager, DealsManager
tz = pytz.timezone('Etc/UTC')
logger = getLogger(__name__)
@@ -41,14 +40,26 @@ class BackTestEngine:
deals: DealsManager
positions: PositionsManager
_account: TestAccount
margins: dict[int, float]
def __init__(self, *, data: TestData = None, speed: int = 1, start: float | datetime = 0,
end: float | datetime = 0, restart: bool = False):
end: float | datetime = 0, restart: bool = False, name: str = ''):
self._data = data or TestData()
self.config = Config(backtest_engine=self)
self.set_up(start=start, end=end, speed=speed, restart=restart)
self.prepare_data()
_name = f"{datetime.fromtimestamp(self.span[0]):%d-%m-%y}_{datetime.fromtimestamp(self.span[-1]):%d-%m-%y}"
self.name = name or _name
def __next__(self) -> Cursor:
try:
index, time = next(self.iter)
self.cursor = Cursor(index=index, time=time)
return self.cursor
except StopIteration:
logger.warning('End of time')
def __repr__(self):
return f"{self.__class__.__name__}()"
def set_up(self, *, start: float | datetime = 0, end: float | datetime = 0, speed: int = 1, restart: bool = False):
span_start = (int(start.timestamp()) if isinstance(start, datetime) else int(start)) or self._data.span.start
@@ -72,27 +83,16 @@ class BackTestEngine:
positions = {}
for ticket, position in self._data.positions.items():
positions[ticket] = TradePosition((position.get(k) for k in TradePosition.__match_args__))
self.positions = PositionsManager(data=positions)
self.positions = PositionsManager(data=positions, open_positions=self._data.open_positions,
margins=self._data.margins)
deals = {}
for ticket, deal in self._data.deals.items():
deals[ticket] = TradeDeal((deal.get(k) for k in TradeDeal.__match_args__))
self.deals = DealsManager(data=deals)
self._account: TestAccount = TestAccount(**self._data.account)
self.margins = self._data.margins
def __next__(self) -> Cursor:
try:
index, time = next(self.iter)
self.cursor = Cursor(index=index, time=time)
return self.cursor
except StopIteration:
logger.warning('End of time')
def __repr__(self):
return f"{self.__class__.__name__}()"
def next(self) -> Cursor:
return next(self)
@@ -122,30 +122,27 @@ class BackTestEngine:
for _ in range(steps):
self.next()
def get_dtype(self, df: DataFrame) -> list[tuple[str, str]]:
@staticmethod
def get_dtype(*, df: DataFrame) -> list[tuple[str, str]]:
return [(c, t) for c, t in zip(df.columns, df.dtypes)]
async def tracker(self):
pos_tasks = [self.check_position(ticket) for ticket in self.positions.open_items]
pos_tasks = [self.check_position(ticket=ticket) for ticket in self.positions._open_positions]
await asyncio.gather(*pos_tasks)
order_tasks = [self.check_order(ticket) for ticket in self.orders.open_items]
await asyncio.gather(*order_tasks)
profit = sum(pos.profit for pos in self.positions.open_positions)
self.update_account(profit=profit)
def save(self):
def wrap_up(self):
try:
if len(self.orders) or len(self.deals):
for symbol in self.orders:
self.history_orders = pd.concat([DataFrame(self.orders[symbol].values()), self.history_orders])
self._data.history_orders = self.history_orders
for symbol in self.deals:
self.history_deals = pd.concat([DataFrame(self.deals[symbol].values()), self.history_deals])
self._data.history_deals = self.history_deals
path = self.config.test_data_dir/f"{self._data.name}.pkl"
GetData.dump_data(data=self._data, name=path, compress=self.config.compress_test_data)
self._data.deals = self.deals.to_dict()
self._data.orders = self.orders.to_dict()
self._data.positions = self.positions.to_dict()
self._data.open_positions = self.positions._open_positions
self._data.margins = self.positions.margins
self._data.account = self._account.asdict()
name = self._data.name or self.name
path = self.config.backtest_dir/f"{name}.pkl"
GetData.pickle_data(data=self._data, name=path)
except Exception as err:
print(err)
@@ -162,52 +159,86 @@ class BackTestEngine:
logger.error(f"Error Getting Price Tick: {exe}")
@error_handler
async def check_order(self, ticket: int):
async def check_order(self, *, ticket: int):
""""
Check if the order has reached its take profit or stop loss levels and close the order if it has.
Checks only **OrderType.BUY** and **OrderType.SELL** orders that have reached their take profit or stop loss levels.
Args:
ticket (int): Order ticket
"""
order = self.orders[ticket]
order_type, symbol = order.type, order.symbol
tick = await self.get_price_tick(symbol, self.cursor.time)
tick = await self.get_price_tick(symbol=symbol, time=self.cursor.time)
tp, sl = order.tp, order.sl
if not (tp and sl):
return
match order_type:
case OrderType.BUY:
if tp >= tick.bid or sl <= tick.bid:
self.close_position(ticket)
self.close_position(ticket=ticket)
case OrderType.SELL:
if tp <= tick.ask or sl >= tick.ask:
self.close_position(ticket)
self.close_position(ticket=ticket)
case _:
...
@error_handler
async def check_position(self, ticket: int, use_terminal=True):
pos = self.positions[ticket]
order_type, symbol, volume, price_open, prev_profit = pos.type, pos.symbol, pos.volume, pos.price_open, pos.profit
tick = await self.get_price_tick(symbol, self.cursor.time)
price_current = tick.bid if order_type == OrderType.BUY else tick.ask
profit = await self.order_calc_profit(order_type, symbol, volume, price_open, price_current, use_terminal)
self.positions.update(ticket=pos.ticket, profit=profit, price_current=price_current, time_update=self.cursor.time)
async def check_position(self, *, ticket: int):
"""
Update the profit of an open position based on the current price of the symbol.
@error_handler(response=False)
Args:
ticket (int): Position ticket
"""
pos = self.positions[ticket]
order_type, symbol, volume, price_open, prev_profit = (pos.type, pos.symbol, pos.volume, pos.price_open,
pos.profit)
tick = await self.get_price_tick(symbol=symbol, time=self.cursor.time)
price_current = tick.bid if order_type == OrderType.BUY else tick.ask
profit = await self.order_calc_profit(action=order_type, symbol=symbol, volume=volume,
price_open=price_open, price_close=price_current)
self.positions.update(ticket=pos.ticket, profit=profit, price_current=price_current,
time_update=self.cursor.time)
await self.check_order(ticket=ticket)
@error_handler_sync(response=False)
def close_position(self, *, ticket: int) -> bool:
position = self.positions.pop(ticket)
margin = self.margins.pop(position.ticket)
del self.orders[ticket]
del self.positions[ticket]
self.orders.update(ticket=ticket, time_done=self.cursor.time)
"""
Close an open position for the trading account using the position ticket.
Args:
ticket: Position ticket
Returns:
bool: True if the position is closed successfully, False otherwise
"""
position = self.positions[ticket]
margin = self.positions.get_margin(ticket=ticket)
self.positions.delete_margin(ticket=ticket)
self.positions.close(ticket=ticket)
self.orders.update(ticket=ticket, time_done=self.cursor.time, time_done_msc=self.cursor.time*1000)
self.update_account(gain=position.profit, margin=-margin)
return True
@error_handler(response=False)
def modify_stops(self, *, ticket: int, sl: int = None, tp: int = None) -> bool:
pos = self.positions[ticket]
order = self.orders[ticket]
sl = sl or pos.sl
tp = tp or pos.tp
self.positions.update(ticket=ticket, sl=sl, tp=tp, time_update=self.cursor.time)
sl = sl or order.sl
tp = tp or order.tp
self.orders.update(ticket=ticket, sl=sl, tp=tp, time_update=self.cursor.time)
def modify_stops(self, *, ticket: int, sl: int, tp: int) -> bool:
"""
Modify the stop loss and take profit levels of an open position.
Args:
ticket: Position ticket
sl: stop loss level
tp: Take profit level
Returns:
bool: True if the stops are modified successfully, False otherwise
"""
self.positions.update(ticket=ticket, sl=sl, tp=tp, time_update=self.cursor.time,
time_update_msc=self.cursor.time*1000)
return True
def update_account(self, *, profit: float = None, margin: float = 0, gain: float = 0):
@@ -288,21 +319,23 @@ class BackTestEngine:
return symbols
@error_handler
async def order_send(self, *, request: dict, use_terminal: bool = True) -> OrderSendResult:
async def order_send(self, *, request: dict) -> OrderSendResult:
osr = {'retcode': 10013, 'comment': 'Invalid request',
'request': TradeRequest(request.get(k, (0 if k != 'comment' else '')) for k in
TradeRequest.__match_args__)}
trade_order = {k: v for k, v in request.items() if k in TradeOrder.__match_args__}
current_tick = await self.get_price_tick(request.get('symbol'), self.cursor.time)
if current_tick is None:
osr['comment'] = 'Market is closed'
osr['retcode'] = 10018
return OrderSendResult((osr.get(k, 0) for k in OrderSendResult.__match_args__))
trade_order = {'external_id': '', 'comment': '',
**{k: v for k, v in request.items() if k in TradeOrder.__match_args__}}
order_type, symbol = request.get('type'), request.get('symbol', '')
action, position_ticket = request.get('action'), request.get('position')
action, position_id = request.get('action'), request.get('position')
sl, tp, volume, symbol = request.get('sl'), request.get('tp'), request.get('volume'), request.get('symbol')
order_type = OrderType(order_type)
current_position = self.positions.get(position_ticket)
current_position = self.positions.get(position_id)
order_ticket = random.randint(100_000_000, 999_999_999)
deal_ticket = random.randint(100_000_000, 999_999_999)
@@ -310,17 +343,21 @@ class BackTestEngine:
if action == TradeAction.DEAL and current_position and order_type.opposite == current_position.type:
res = self.close_position(current_position.ticket)
if res:
trade_order.update({'comment': '', 'position_id': current_position.ticket, 'ticket': order_ticket,
'time_setup': current_tick.time, 'time_expiration': current_tick.time,
'time_setup_msc': current_tick.time_msc, 'time_done': current_tick.time,
'time_done_msc': current_tick.time_msc, 'type': order_type, 'symbol': symbol,
'price_current': current_position.price_current, 'reason': OrderReason.EXPERT,
price_current = current_tick.ask if order_type == OrderType.BUY else current_tick.bid
trade_order.update({'position_id': current_position.ticket, 'ticket': order_ticket,
'time_setup': current_tick.time, 'time_setup_msc': current_tick.time_msc,
'time_done': current_tick.time, 'time_done_msc': current_tick.time_msc,
'type': order_type, 'symbol': symbol, 'sl': current_position.sl,
'tp': current_position.tp,
'price_current': price_current, 'reason': OrderReason.EXPERT,
'volume_initial': current_position.volume})
# TODO: calculate commission and swap if possible or necessary
deal = {'ticket': deal_ticket, 'position_id': current_position.ticket, 'order': order_ticket,
'symbol': symbol, 'commission': 0, 'swap': 0, 'fee': 0, 'time': current_tick.time,
'symbol': symbol, 'time': current_tick.time,
'time_msc': current_tick.time_msc, 'volume': current_position.volume,
'price': current_position.price_current, 'type': DealType(order_type), 'reason': DealReason.EXPERT,
'entry': DealEntry.OUT, 'comment': ''}
'price': price_current, 'type': DealType(order_type), 'reason': DealReason.EXPERT,
'entry': DealEntry.OUT, 'comment': '', 'external_id': ''}
order = TradeOrder((trade_order.get(k, 0) for k in TradeOrder.__match_args__))
self.orders[order.ticket] = order
@@ -331,21 +368,18 @@ class BackTestEngine:
return OrderSendResult((osr.get(k, 0) for k in OrderSendResult.__match_args__))
if action == TradeAction.SLTP and current_position:
check = await self.order_check(position_ticket)
check = await self.order_check(position_id)
if check.retcode != 0:
osr = {'retcode': check.retcode, 'comment': check.comment, 'request': check.request}
return OrderSendResult((osr.get(k, 0) for k in OrderSendResult.__match_args__))
res = self.modify_stops(ticket=position_ticket, sl=sl, tp=tp)
res = self.modify_stops(ticket=position_id, sl=sl, tp=tp)
if res:
# ToDo: Create a deal object here
osr.update({'comment': 'Request completed', 'retcode': 10009, 'order': order_ticket, 'deal': deal_ticket,})
osr.update({'comment': 'Request completed', 'retcode': 10009, 'order': order_ticket,
'deal': deal_ticket})
return OrderSendResult((osr.get(k, 0) for k in OrderSendResult.__match_args__))
if action == TradeAction.DEAL and order_type in (OrderType.BUY, OrderType.SELL):
check = await self.order_check(request=request)
if check.retcode != 0:
osr = {'retcode': check.retcode, 'comment': check.comment, 'request': check.request}
return OrderSendResult((osr.get(k, 0) for k in OrderSendResult.__match_args__))
@@ -376,28 +410,23 @@ class BackTestEngine:
self.orders[order.ticket] = order
osr.update({'order': order_ticket, 'price': price, 'volume': volume, 'bid': current_tick.bid,
'ask': current_tick.ask, 'deal': deal_ticket})
margin = await self.order_calc_margin(action, symbol, volume, price)
self.margins[order_ticket] = margin
margin = await self.order_calc_margin(action=action, symbol=symbol, volume=volume, price=price)
self.positions.set_margin(ticket=order_ticket, margin=margin)
self.update_account(margin=margin)
return OrderSendResult((osr.get(k, 0) for k in OrderSendResult.__match_args__))
@error_handler
async def order_check(self, request: dict) -> OrderCheckResult:
async def order_check(self, *, request: dict) -> OrderCheckResult:
ocr = {'retcode': 10013, 'balance': 0, 'profit': 0, 'margin': 0, 'equity': 0, 'margin_free': 0,
'margin_level': 0, 'comment': 'Invalid request',
'request': TradeRequest(request.get(k, (0 if k != 'comment' else '')) for k in
TradeRequest.__match_args__)}
action, symbol, volume = request.get('action'), request.get('symbol'), request.get('volume')
price, order_type, position_id = request.get('price'), request.get('type'), request.get('position')
price, order_type = request.get('price'), request.get('type')
if price is None and (action is TradeAction.DEAL and order_type in (OrderType.BUY, OrderType.SELL)):
ocr['comment'] = 'Market is closed'
ocr['retcode'] = 10018
return OrderCheckResult((ocr.get(k, 0) for k in OrderCheckResult.__match_args__))
# check margin and confirm order can go through
if action == TradeAction.DEAL and order_type in (OrderType.BUY, OrderType.SELL):
# check margin and confirm order can go through for a deal action and buy or sell order type
if action == TradeAction.DEAL and order_type in (OrderType.BUY, OrderType.SELL) and position_id is None:
margin = await self.order_calc_margin(action, symbol, volume, price)
if margin is None:
return OrderCheckResult((ocr.get(k, 0) for k in OrderCheckResult.__match_args__))
@@ -418,19 +447,19 @@ class BackTestEngine:
# check if the stops level is valid
sym = await self.get_symbol_info(symbol=symbol)
sl, tp = request.get('sl', 0), request.get('tp', 0)
sl, tp = request.get('sl'), request.get('tp')
current_price = price
if tp or sl:
if tp and sl:
if action == TradeAction.SLTP:
pos = self.positions.get(request.get('position'))
sym = sym or await self.get_symbol_info(pos.symbol)
sym = await self.get_symbol_info(pos.symbol)
current_tick = sym or await self.get_price_tick(pos.symbol, self.cursor.time)
current_price = current_tick.bid if pos.type == OrderType.BUY else current_tick.ask
min_sl = min(sl, tp)
dsl = abs(current_price - min_sl) / sym.point
tsl = sym.trade_stops_level + sym.spread
if int(dsl) < int(tsl):
if dsl < tsl:
ocr['retcode'] = 10016
ocr['comment'] = 'Invalid stops'
return OrderCheckResult((ocr.get(k, 0) for k in OrderCheckResult.__match_args__))
@@ -511,11 +540,11 @@ class BackTestEngine:
rates = await self.mt5.copy_rates_from(symbol, timeframe, date_from, count)
return rates
rates = self.rates[symbol][timeframe.name]
rates = self.rates[symbol][timeframe]
start = int(datetime.timestamp(date_from)) if isinstance(date_from, datetime) else int(date_from)
start = round_down(start, timeframe.time)
rates = rates[rates.time <= start].iloc[-count:]
return np.fromiter((tuple(i) for i in rates.iloc), dtype=self.get_dtype(rates))
return np.fromiter((tuple(i) for i in rates.iloc), dtype=self.get_dtype(df=rates))
@error_handler
async def get_rates_from_pos(self, symbol: str, timeframe: TimeFrame, start_pos: int, count: int) -> np.ndarray:
@@ -527,28 +556,30 @@ class BackTestEngine:
res = await self.mt5.copy_rates_from_pos(symbol, timeframe, start_pos, count)
return res
rates = self.rates[symbol][timeframe.name]
rates = self.rates[symbol][timeframe]
end = abs(self.cursor.index - start_pos)
start = abs(end - count)
rates = rates.iloc[start:end]
return np.fromiter((tuple(i) for i in rates.iloc), dtype=self.get_dtype(rates))
return np.fromiter((tuple(i) for i in rates.iloc), dtype=self.get_dtype(df=rates))
@error_handler
async def get_rates_range(self, symbol: str, timeframe: TimeFrame, date_from: datetime | float, date_to: datetime | float) -> np.ndarray:
async def get_rates_range(self, symbol: str, timeframe: TimeFrame, date_from: datetime | float,
date_to: datetime | float) -> np.ndarray:
if self.config.use_terminal_for_backtesting:
rates = await self.mt5.copy_rates_range(symbol, timeframe, date_from, date_to)
return rates
rates = self.rates[symbol][timeframe.name]
rates = self.rates[symbol][timeframe]
start = int(datetime.timestamp(date_from)) if isinstance(date_from, datetime) else int(date_from)
start = round_down(start, timeframe.time)
end = int(datetime.timestamp(date_to)) if isinstance(date_to, datetime) else int(date_to)
end = round_up(end, timeframe.time)
rates = rates[(rates.time >= start) & (rates.time <= end)]
return np.fromiter((tuple(i) for i in rates.iloc), dtype=self.get_dtype(rates))
return np.fromiter((tuple(i) for i in rates.iloc), dtype=self.get_dtype(df=rates))
@error_handler
async def get_ticks_from(self, symbol: str, date_from: datetime | float, count: int, flags: CopyTicks) -> np.ndarray:
async def get_ticks_from(self, *, symbol: str, date_from: datetime | float, count: int,
flags: CopyTicks) -> np.ndarray:
if self.config.use_terminal_for_backtesting:
ticks = await self.mt5.copy_ticks_from(symbol, date_from, count, flags)
return ticks
@@ -556,18 +587,23 @@ class BackTestEngine:
ticks = self.ticks[symbol]
start = int(datetime.timestamp(date_from)) if isinstance(date_from, datetime) else int(date_from)
ticks = ticks[ticks.time <= start].iloc[-count:]
return np.fromiter((tuple(i) for i in ticks.iloc), dtype=self.get_dtype(ticks))
return np.fromiter((tuple(i) for i in ticks.iloc), dtype=self.get_dtype(df=ticks))
@error_handler
async def get_ticks_range(self, symbol: str, date_from: datetime | float, date_to: datetime | float, flags) -> np.ndarray:
async def get_ticks_range(self, *, symbol: str, date_from: datetime | float,
date_to: datetime | float, flags: CopyTicks) -> np.ndarray:
if self.config.use_terminal_for_backtesting:
ticks = await self.mt5.copy_ticks_range(symbol, date_from, date_to, flags)
return ticks
ticks = self.ticks[symbol]
start = int(datetime.timestamp(date_from)) if isinstance(date_from, datetime) else int(date_from)
end = int(datetime.timestamp(date_to)) if isinstance(date_to, datetime) else int(date_to)
ticks = ticks[(ticks.time >= start) & (ticks.time <= end)]
return np.fromiter((tuple(i) for i in ticks.iloc), dtype=self.get_dtype(ticks))
return np.fromiter((tuple(i) for i in ticks.iloc), dtype=self.get_dtype(df=ticks))
@error_handler
async def order_calc_margin(self, action: Literal[OrderType.BUY, OrderType.SELL], symbol: str, volume: float,
async def order_calc_margin(self, *, action: Literal[OrderType.BUY, OrderType.SELL], symbol: str, volume: float,
price: float):
if self.mt5.config.use_terminal_for_backtesting:
return await self.mt5.order_calc_margin(action, symbol, volume, price)
@@ -577,7 +613,7 @@ class BackTestEngine:
return round(margin, self._account.currency_digits)
@error_handler
async def order_calc_profit(self, action: Literal[OrderType.BUY, OrderType.SELL], symbol: str, volume: float,
async def order_calc_profit(self, *, action: Literal[OrderType.BUY, OrderType.SELL], symbol: str, volume: float,
price_open: float, price_close: float):
if self.mt5.config.use_terminal_for_backtesting:
@@ -590,88 +626,120 @@ class BackTestEngine:
@error_handler_sync
def get_orders_total(self) -> int:
return len(self.open_orders)
"""
Get the total number of pending orders.
Returns:
int: Total number of pending orders
"""
return 0
@error_handler_sync
def get_orders(self, symbol: str = '', group: str = '', ticket: int = None) -> tuple[TradeOrder, ...]:
"""
Get pending orders from the terminal history. This has to do with pending orders, which this backtester
doesn't support yet.
if ticket:
order = self.open_orders.get(ticket)
return (order,) if order else ()
elif symbol:
return tuple(order for order in self.orders.get(symbol, ()) if order.ticket in self.open_orders)
elif group:
return tuple(order for order in self.open_orders.values())
else:
return tuple(order for order in self.open_orders.values())
Args:
symbol: Symbol name
group: Group name
ticket: Order ticket
Returns:
tuple[TradeOrder]
"""
if symbol and group and ticket:
return tuple()
return ()
@error_handler_sync
def get_positions_total(self) -> int:
return len(self.open_positions)
"""
Get the total number of open positions.
Returns:
int: Total number of open positions
"""
return self.positions.positions_total()
@error_handler_sync
def get_positions(self, symbol: str = '', group: str = '', ticket: int = None) -> tuple[TradePosition, ...]:
if ticket:
position = self.open_positions.get(ticket)
return (position,) if position else ()
def get_positions(self, *, symbol: str = None, group: str = None, ticket: int = None) -> tuple[TradePosition, ...]:
"""
Get open positions from the terminal history.
elif symbol:
return tuple(position for position in self.positions.get(symbol, ()) if position.ticket in self.open_positions)
Keyword Args:
symbol: The symbol name
group: Group argument to filter by
ticket: Position ticket
elif group:
return tuple(position for position in self.open_positions.values())
else:
return tuple(position for position in self.open_positions.values())
Returns:
tuple[TradePosition]: Open positions
"""
return self.positions.positions_get(ticket=ticket, symbol=symbol, group=group)
@error_handler_sync
def get_history_orders_total(self, date_from: datetime | float, date_to: datetime | float) -> int:
start = int(date_from.timestamp()) if isinstance(date_from, datetime) else int(date_from)
end = int(date_to.timestamp()) if isinstance(date_to, datetime) else int(date_to)
orders = self.history_orders[self.history_orders.time >= start & self.history_orders.time <= end]
return orders.shape[0]
def get_history_orders_total(self, *, date_from: datetime | float, date_to: datetime | float) -> int:
"""
Get the total number of orders in the terminal history.
Args:
date_from: The start date of the history
date_to: The end date of the history
Returns:
int: Total number of orders in the history
"""
return self.orders.history_orders_total(date_from=date_from, date_to=date_to)
@error_handler_sync
def get_history_orders(self, date_from: datetime | float, date_to: datetime | float, group: str = '',
ticket: int = None, position: int = None) -> tuple[TradeOrder, ...]:
start = int(date_from.timestamp()) if isinstance(date_from, datetime) else int(date_from)
end = int(date_to.timestamp()) if isinstance(date_to, datetime) else int(date_to)
orders = self.history_orders[self.history_orders.time >= start & self.history_orders.time <= end]
def get_history_orders(self, *, date_from: datetime | float = None, date_to: datetime | float = None,
group: str = '', ticket: int = None, position: int = None) -> tuple[TradeOrder, ...]:
"""
Get orders from the terminal history.
if ticket:
orders = orders[orders.ticket == ticket]
Keyword Args:
date_from: Date from which to start the history
date_to: Date to which to end the history
group: group keyword to filter by
ticket: ticket id to filter by
position: position id to filter by
elif position:
orders = orders[orders.position == position]
elif group:
...
return tuple(TradeOrder(order) for order in orders.iloc)
Returns:
tuple[TradeOrder]: Orders in the history
"""
return self.orders.history_orders_get(date_from=date_from, date_to=date_to, group=group, ticket=ticket,
position=position)
@error_handler_sync
def get_history_deals_total(self, date_from: datetime | float, date_to: datetime | float) -> int:
start = int(date_from.timestamp()) if isinstance(date_from, datetime) else int(date_from)
end = int(date_to.timestamp()) if isinstance(date_to, datetime) else int(date_to)
deals = self.history_deals[self.history_deals.time >= start & self.history_deals.time <= end]
return deals.shape[0]
def get_history_deals_total(self, *, date_from: datetime | float, date_to: datetime | float) -> int:
"""
Get the total number of deals in the terminal history.
Args:
date_from: Date from which to start the history
date_to: Date to which to end the history
Returns:
int: Total number of deals in the history
"""
return self.deals.history_deals_total(date_from=date_from, date_to=date_to)
@error_handler_sync
def get_history_deals(self, date_from: datetime | float, date_to: datetime | float, group: str = '',
position: int = None, ticket: int = None) -> tuple[TradeDeal, ...]:
start = int(date_from.timestamp()) if isinstance(date_from, datetime) else int(date_from)
end = int(date_to.timestamp()) if isinstance(date_to, datetime) else int(date_to)
deals = self.history_deals[self.history_deals.time >= start & self.history_deals.time <= end]
def get_history_deals(self, *, date_from: datetime | float = None, date_to: datetime | float = None,
group: str = None, position: int = None, ticket: int = None) -> tuple[TradeDeal, ...]:
"""
Get deals from the terminal history.
if ticket:
deals = deals[deals.ticket == ticket]
Keyword Args:
date_from: Date from which to start the history
date_to: Date to which to end the history
group: group keyword to filter by
position: position id to filter by
ticket: ticket id to filter by
elif position:
deals = deals[deals.position == position]
elif group:
...
return tuple(TradeDeal(deal) for deal in deals.iloc)
Returns:
tuple[TradeDeal, ...]: Deals in the history
"""
return self.deals.history_deals_get(date_from=date_from, date_to=date_to, group=group, position=position,
ticket=ticket)
@@ -0,0 +1,32 @@
import asyncio
import signal
from logging import getLogger
from .event_manager import EventManager
from .meta_tester import MetaTester
from .backtest_engine import BackTestEngine
from .strategy_tester import StrategyTester
from ...core import Config
logger = getLogger(__name__)
class BackTester:
def __init__(self, *, strategies: list[StrategyTester] = None, backtest_engine: BackTestEngine = None):
self.strategies = strategies or []
self.event_manager = EventManager()
self.mt5 = MetaTester(backtest_engine=backtest_engine)
signal.signal(signal.SIGINT, self.event_manager.sigint_handler)
async def run(self):
try:
await self.mt5.initialize()
strategies = [strategy for strategy in self.strategies if await strategy.symbol.init()]
print(f"Number of strategies: {len(strategies)}")
self.event_manager.num_main_tasks = len(strategies)
tasks = [*[asyncio.create_task(strategy.test()) for strategy in strategies],
asyncio.create_task(self.event_manager.event_monitor())]
self.event_manager.add_tasks(*tasks)
await asyncio.gather(*tasks, return_exceptions=True) if strategies else ...
except Exception as err:
logger.error(f"Error {err} occurred in StrategyTester")
+17 -6
View File
@@ -1,10 +1,21 @@
class Form:
rest: str
glob = dict()
class Check:
def __init__(self, ty):
self.r = ty
@property
def rest(self):
return 'rest'
def r(self):
print('getting value')
return glob.get('r')
@r.setter
def r(self, value):
print('setting value')
glob['r'] = value
g = Form()
print(g.rest)
f = Check(465)
print(f.r)
f.r = 56
print(f.r)
@@ -25,13 +25,13 @@ class EventManager:
def __init__(self, *, num_tasks: int = 0):
self.num_main_tasks = num_tasks or self.num_main_tasks
def add_task(self, *tasks: Task):
def add_tasks(self, *tasks: Task):
self.tasks.extend(tasks)
def sigint_handler(self, sig, frame):
for task in self.tasks:
task.cancel()
self.config.test_data.save()
self.config.backtest_engine.wrap_up()
async def acquire(self):
await self.condition.acquire()
@@ -44,12 +44,11 @@ class EventManager:
async with self.condition:
if self.task_tracker == self.num_main_tasks:
self.task_tracker = 0
await self.config.test_data.tracker()
self.config.test_data.next()
await self.config.backtest_engine.tracker()
self.config.backtest_engine.next()
self.condition.notify_all()
if ((timestamp := self.config.test_data.cursor.time) % int(60 * 60 * 24)) == 0:
if ((timestamp := self.config.backtest_engine.cursor.time) % int(60 * 60 * 24)) == 0:
print(f"if Time: {datetime.fromtimestamp(timestamp)}")
await asyncio.sleep(0)
async def wait(self):
+10 -28
View File
@@ -1,33 +1,26 @@
from dataclasses import dataclass, field, fields
import pickle
from pathlib import Path
import lzma
from datetime import datetime
from logging import getLogger
from typing import Sequence, ClassVar
from collections import namedtuple
from typing import Sequence, NamedTuple
import MetaTrader5
import pytz
import numpy as np
import pandas as pd
from numpy import ndarray
from pandas import DataFrame
from ...core.meta_trader import MetaTrader
from ...core.config import Config
from ...core.constants import TimeFrame, CopyTicks
from ...core.constants import TimeFrame
from ...core.task_queue import TaskQueue, QueueItem
from ...utils import backoff_decorator
logger = getLogger(__name__)
from MetaTrader5 import TradePosition, TradeOrder, TradeDeal
tof = list(TradeOrder.__match_args__)
tpf = list(TradePosition.__match_args__)
tdf = list(TradeDeal.__match_args__)
Cursor = namedtuple('Cursor', ['index', 'time'])
class Cursor(NamedTuple):
index: int
time: int
@dataclass
@@ -45,23 +38,12 @@ class TestData:
orders: dict[int, dict] = field(default_factory=lambda: {})
deals: dict[int, dict] = field(default_factory=lambda: {})
positions: dict[int, dict] = field(default_factory=lambda: {})
active_orders: tuple[int, ...] = field(default_factory=lambda: ())
open_positions: tuple[int, ...] = field(default_factory=lambda: ())
open_positions: set[int, ...] = field(default_factory=lambda: set())
cursor: Cursor = None
margins: dict[int, float] = field(default_factory=lambda: {})
def __str__(self):
return f"""
Data: {self.name}
Terminal: {str(list(self.terminal.keys())[0:2]) + '...' if len(self.terminal) > 3 else list(self.terminal.keys())}
Version: {self.version}
Account: {str(list(self.account.keys())[0:2]) + '...' if len(self.account) > 3 else list(self.account.keys())}
Symbols: {len(self.symbols)} symbols
Prices: Prices for {len(self.prices)} symbols
Ticks: Ticks for {len(self.ticks)} symbols
Rates: Bars for {len(self.rates)} symbols
Span: {datetime.fromtimestamp(self.span.start)} to {datetime.fromtimestamp(self.span.stop)}
"""
return f"{self.name}"
def __repr__(self):
return f"{self.__class__.__name__}({self.name})"
@@ -191,13 +173,13 @@ class GetData:
@backoff_decorator
async def get_symbol_ticks(self, *, symbol: str):
""""""
res = await self.mt5.copy_ticks_range(symbol, self.start, self.end, CopyTicks.ALL)
res = await self.mt5.copy_ticks_range(symbol, self.start, self.end, MetaTrader5.COPY_TICKS_ALL)
self.data.ticks[symbol] = res
@backoff_decorator
async def get_symbol_prices(self, *, symbol: str):
""""""
res = await self.mt5.copy_ticks_range(symbol, self.start, self.end, CopyTicks.ALL)
res = await self.mt5.copy_ticks_range(symbol, self.start, self.end, MetaTrader5.COPY_TICKS_ALL)
self.data.prices[symbol] = res
@backoff_decorator
+57 -71
View File
@@ -1,12 +1,12 @@
from datetime import datetime
from logging import getLogger
from typing import Literal
from numpy import ndarray
from MetaTrader5 import (Tick, SymbolInfo, AccountInfo, TerminalInfo, TradeOrder, TradePosition, TradeDeal,
OrderCheckResult, OrderSendResult)
from .backtest_engine import BackTestEngine
from .get_data import GetData
from . import BackTestEngine
from ...core.meta_trader import MetaTrader
from ...core.constants import TimeFrame, CopyTicks, OrderType
@@ -17,181 +17,167 @@ logger = getLogger(__name__)
class MetaTester(MetaTrader):
"""A class for testing trading strategies in the MetaTrader 5 terminal. A subclass of MetaTrader."""
backtest_engine: BackTestEngine
def __init__(self, test_data: BackTestEngine = None):
def __init__(self, *, backtest_engine: BackTestEngine = None):
super().__init__()
if test_data is not None:
self.config.test_data = test_data
self.backtest_engine = backtest_engine
@property
def test_data(self) -> BackTestEngine | None:
return self.config.test_data
def backtest_engine(self) -> BackTestEngine:
return self.config.backtest_engine
@test_data.setter
def test_data(self, value: BackTestEngine):
self.config.test_data = value
@backtest_engine.setter
def backtest_engine(self, value: BackTestEngine):
if isinstance(value, BackTestEngine):
self.config.backtest_engine = value
async def last_error(self) -> tuple[int, str]:
return -1, ''
async def initialize(self, path: str = "", login: int = 0, password: str = "", server: str = "",
timeout: int | None = None, portable=False, load_test_data: bool = False,
test_data_file: str = '', use_terminal: bool = True) -> bool:
success = True
async def initialize(self, *, path: str = "", login: int = 0, password: str = "", server: str = "",
timeout: int | None = None, portable=False) -> bool:
if self.config.use_terminal_for_backtesting:
success = await super().initialize(path=path, login=login, password=password, server=server, timeout=timeout)
return await super().initialize(path=path, login=login, password=password, server=server, timeout=timeout)
try:
if load_test_data:
name = f"{self.config.test_data_dir_name}/{test_data_file}"
data = GetData.load_data(name=name, compressed=self.config.compress_test_data)
if data is not None:
self.test_data = BackTestEngine(data)
success = True
except Exception as err:
logger.error(f'{err}: unable to load test data')
success = False
return success
async def login(self, login: int, password: str, server: str, timeout: int = 60000) -> bool:
return await super().login(login, password, server, timeout) if self.config.use_terminal_for_backtesting else True
return True
async def login(self, *, login: int = 0, password: str = '', server: str = '', timeout: int = 60000) -> bool:
if self.config.use_terminal_for_backtesting:
return await super().login(login=login, password=password, server=server, timeout=timeout)
return True
async def shutdown(self) -> None:
await super().shutdown() if self.config.use_terminal_for_backtesting else ...
@error_handler(msg='test data not available', exe=AttributeError)
async def terminal_info(self) -> TerminalInfo:
res = await self.test_data.get_terminal_info()
res = await self.backtest_engine.get_terminal_info()
return res
@error_handler(msg='test data not available', exe=AttributeError)
async def account_info(self) -> AccountInfo:
""""""
return self.test_data.get_account_info()
@error_handler(msg='test data not available', exe=AttributeError)
async def symbol_select(self, symbol: str, enable: bool = True) -> bool:
if self.config.use_terminal_for_backtesting:
res = await super().symbol_select(symbol, enable)
return res
res = (symbol in self.test_data.symbols) and enable
return res
return self.backtest_engine.get_account_info()
@error_handler(msg='test data not available', exe=AttributeError)
async def symbols_total(self) -> int:
tot = await self.test_data.get_symbols_total()
tot = await self.backtest_engine.get_symbols_total()
return tot
@error_handler(msg='test data not available', exe=AttributeError)
async def symbols_get(self, group: str = "") -> tuple[SymbolInfo, ...] | None:
""""""
syms = await self.test_data.get_symbols(group)
syms = await self.backtest_engine.get_symbols(group=group)
return syms
@error_handler(msg='test data not available', exe=AttributeError)
async def symbol_info(self, symbol: str) -> SymbolInfo | None:
sym = await self.test_data.get_symbol_info(symbol)
sym = await self.backtest_engine.get_symbol_info(symbol=symbol)
return sym
@error_handler(msg='test data not available', exe=AttributeError)
async def symbol_info_tick(self, symbol: str) -> Tick | None:
tick = await self.test_data.get_symbol_info_tick(symbol)
tick = await self.backtest_engine.get_symbol_info_tick(symbol=symbol)
return tick
@error_handler(msg='test data not available', exe=AttributeError)
async def copy_rates_from(self, symbol: str, timeframe: TimeFrame, date_from: datetime | float,
count: int) -> ndarray | None:
rates = await self.test_data.get_rates_from(symbol, timeframe, date_from, count)
rates = await self.backtest_engine.get_rates_from(symbol=symbol, timeframe=timeframe, date_from=date_from,
count=count)
return rates
@error_handler(msg='test data not available', exe=AttributeError)
async def copy_rates_from_pos(self, symbol: str, timeframe: TimeFrame, start_pos: int,
count: int) -> ndarray | None:
rates = await self.test_data.get_rates_from_pos(symbol, timeframe, start_pos, count)
rates = await self.backtest_engine.get_rates_from_pos(symbol=symbol, timeframe=timeframe, start_pos=start_pos,
count=count)
return rates
@error_handler(msg='test data not available', exe=AttributeError)
async def copy_rates_range(self, symbol: str, timeframe: TimeFrame, date_from: datetime | float,
date_to: datetime | float) -> ndarray | None:
rates = await self.test_data.get_rates_range(symbol, timeframe, date_from, date_to)
rates = await self.backtest_engine.get_rates_range(symbol=symbol, timeframe=timeframe, date_from=date_from,
date_to=date_to)
return rates
@error_handler(msg='test data not available', exe=AttributeError)
async def copy_ticks_from(self, symbol: str, date_from: datetime | float, count: int,
flags: CopyTicks) -> ndarray | None:
ticks = await self.test_data.get_ticks_from(symbol, date_from, count, flags)
ticks = await self.backtest_engine.get_ticks_from(symbol=symbol, date_from=date_from, count=count, flags=flags)
return ticks
@error_handler(msg='test data not available', exe=AttributeError)
async def copy_ticks_range(self, symbol: str, date_from: datetime | float, date_to: datetime | float,
flags: CopyTicks) -> ndarray | None:
ticks = await self.test_data.get_ticks_range(symbol, date_from, date_to, flags)
ticks = await self.backtest_engine.get_ticks_range(symbol=symbol, date_from=date_from, date_to=date_to,
flags=flags)
return ticks
@error_handler(msg='test data not available', exe=AttributeError)
async def orders_total(self) -> int:
return self.test_data.get_orders_total()
return self.backtest_engine.get_orders_total()
@error_handler(msg='test data not available', exe=AttributeError)
async def orders_get(self, group: str = "", ticket: int = 0, symbol: str = "") -> tuple[TradeOrder, ...] | None:
kwargs = {key: value for key, value in (('group', group), ('ticket', ticket), ('symbol', symbol)) if value}
return self.test_data.get_orders(**kwargs)
return self.backtest_engine.get_orders(**kwargs)
@error_handler(msg='test data not available', exe=AttributeError)
async def order_calc_margin(self, action: OrderType, symbol: str, volume: float, price: float) -> float | None:
res = await self.test_data.order_calc_margin(action, symbol, volume, price)
res = await self.backtest_engine.order_calc_margin(action=action, symbol=symbol, volume=volume, price=price)
return res
@error_handler(msg='test data not available', exe=AttributeError)
async def order_calc_profit(self, action: OrderType, symbol: str, volume: float, price_open: float,
async def order_calc_profit(self, action: Literal[0, 1], symbol: str, volume: float, price_open: float,
price_close: float) -> float | None:
profit = await self.test_data.order_calc_profit(action, symbol, volume,
price_open, price_close)
profit = await self.backtest_engine.order_calc_profit(action=action, symbol=symbol, volume=volume,
price_open=price_open, price_close=price_close)
return profit
@error_handler(msg='test data not available', exe=AttributeError)
async def order_check(self, request: dict) -> OrderCheckResult:
ocr = await self.test_data.order_check(request)
ocr = await self.backtest_engine.order_check(request=request)
return ocr
async def order_send(self, request: dict) -> OrderSendResult:
osr = await self.test_data.order_send(request)
osr = await self.backtest_engine.order_send(request=request)
return osr
@error_handler(msg='test data not available', exe=AttributeError)
async def positions_total(self) -> int:
return self.test_data.get_positions_total()
return self.backtest_engine.get_positions_total()
@error_handler(msg='test data not available', exe=AttributeError)
async def positions_get(self, group: str = "", ticket: int = None,
symbol: str = "") -> tuple[TradePosition, ...] | None:
kwargs = {key: value for key, value in (('group', group), ('ticket', ticket), ('symbol', symbol)) if value}
return self.test_data.get_positions(**kwargs)
return self.backtest_engine.get_positions(**kwargs)
@error_handler(msg='test data not available', exe=AttributeError)
async def history_orders_total(self, date_from: datetime | float, date_to: datetime | float) -> int:
return self.test_data.get_history_orders_total(date_from, date_to)
async def history_orders_total(self, date_from: datetime | float, date_to: datetime | float) -> int | None:
return self.backtest_engine.get_history_orders_total(date_from=date_from, date_to=date_to)
@error_handler(msg='test data not available', exe=AttributeError)
async def history_orders_get(self, date_from: datetime | float = None, date_to: datetime | float = None,
group: str = '', ticket: int = None,
position: int = None) -> tuple[TradeOrder, ...] | None:
kwargs = {key: value for key, value in (('group', group), ('ticket', ticket), ('position', position)) if value}
args = tuple(arg for arg in (date_from, date_to) if arg)
return self.test_data.get_history_orders(*args, **kwargs)
args = (('date_from', date_from), ('date_to', date_to), ('group', group), ('ticket', ticket),
('position', position))
kwargs = {key: value for key, value in args if value}
return self.backtest_engine.get_history_orders(**kwargs)
@error_handler(msg='test data not available', exe=AttributeError)
async def history_deals_total(self, date_from: datetime | float, date_to: datetime | float) -> int:
return self.test_data.get_history_deals_total(date_from, date_to)
async def history_deals_total(self, date_from: datetime | float, date_to: datetime | float) -> int | None:
return self.backtest_engine.get_history_deals_total(date_from=date_from, date_to=date_to)
@error_handler(msg='test data not available', exe=AttributeError)
async def history_deals_get(self, date_from: datetime | float = None, date_to: datetime | float = None,
group: str = '', ticket: int = None,
position: int = None) -> tuple[TradeDeal, ...] | None:
kwargs = {key: value for key, value in (('group', group), ('ticket', ticket), ('position', position)) if value}
args = tuple(arg for arg in (date_from, date_to) if arg)
return self.test_data.get_history_deals(*args, **kwargs)
args = (('date_from', date_from), ('date_to', date_to), ('group', group), ('ticket', ticket),
('position', position))
kwargs = {key: value for key, value in args if value}
return self.backtest_engine.get_history_deals(**kwargs)
@@ -1,55 +1,27 @@
import asyncio
import signal
from logging import getLogger
from .event_manager import EventManager
from .test_data import TestData
from .meta_tester import MetaTester
from .test_strategy import TestStrategy, TestSingleStrategy
from ...core import Config
logger = getLogger(__name__)
from ...core.config import Config
class StrategyTester:
def __init__(self, *, strategies: list[TestStrategy] = None):
self.strategies = strategies or []
self.event_manager = EventManager(num_tasks=len(self.strategies))
signal.signal(signal.SIGINT, self.event_manager.sigint_handler)
event_manager: EventManager
config: Config
async def run(self, test_data: TestData):
try:
config = Config()
config.test_data = test_data
mt5 = MetaTester()
acc = config.account_info()
await mt5.initialize(**acc)
await mt5.login(**acc)
strategies = [strategy for strategy in self.strategies if await strategy.symbol.init()]
print(f"Number of strategies: {len(strategies)}")
self.event_manager.num_main_tasks = len(strategies)
tasks = [*[asyncio.create_task(strategy.test()) for strategy in strategies],
asyncio.create_task(self.event_manager.event_monitor())]
self.event_manager.add_task(*tasks)
await asyncio.gather(*tasks, return_exceptions=True) if strategies else ...
except Exception as err:
logger.error(f"Error {err} occurred in StrategyTester")
def set_up(self):
self.event_manager = EventManager()
async def sleep(self, secs: float):
time = self.config.backtest_engine.cursor.time
mod = time % secs
secs = secs - mod if mod != 0 else mod
if self.event_manager.num_main_tasks == 1:
self.config.backtest_engine.fast_forward(secs)
await self.event_manager.wait()
elif self.event_manager.num_main_tasks > 1:
time = self.config.backtest_engine.cursor.time + secs
while time > self.config.backtest_engine.cursor.time:
await self.event_manager.wait()
else:
...
class SingleStrategyTester:
def __init__(self, *, strategy: TestSingleStrategy):
self.strategy = strategy
async def run(self, test_data: TestData):
try:
config = Config()
config.test_data = test_data
mt5 = MetaTester()
acc = config.account_info()
await mt5.initialize(**acc)
await mt5.login(**acc)
sym = self.strategy.symbol
res = await sym.init()
await self.strategy.test() if res else ...
except Exception as err:
logger.error(f"Error {err} occurred in SingleStrategyTester")
def test(self):
raise NotImplementedError("Implement this method in your subclass")
@@ -1,2 +0,0 @@
class FingerTrapTest:
...
@@ -1,40 +0,0 @@
from typing import TypeVar
from .event_manager import EventManager
from ...core.config import Config
Symbol = TypeVar("Symbol")
class TestStrategy:
event_manager: EventManager
config: Config
symbol: Symbol
def set_up(self):
self.event_manager = EventManager()
async def sleep(self, secs: float):
time = self.config.test_data.cursor.time
mod = time % secs
secs = secs - mod if mod != 0 else mod
time = self.config.test_data.cursor.time + secs
while time > self.config.test_data.cursor.time:
await self.event_manager.wait()
def test(self):
raise NotImplementedError("Implement this method in your subclass")
class TestSingleStrategy:
config: Config
symbol: Symbol
async def sleep(self, secs: float):
time = self.config.test_data.cursor.time
mod = time % secs
secs = secs - mod if mod != 0 else mod
self.config.test_data.fast_forward(secs)
def test(self):
raise NotImplementedError("Implement this method in your subclass")
@@ -0,0 +1,180 @@
from datetime import datetime
from typing import TypeVar, Generic
from MetaTrader5 import TradePosition, TradeOrder, TradeDeal
from aiomql.utils import logger
TradeData = TypeVar('TradeData', bound=TradePosition | TradeOrder | TradeDeal)
class TradeManager(Generic[TradeData]):
_data: dict[int, TradeData]
def __init__(self, *, data: dict = None):
self._data = data or {}
def __iter__(self):
return iter(self._data)
def __len__(self):
return len(self._data)
def __contains__(self, item: TradeData):
return item.ticket in self._data
def __getitem__(self, item):
return self._data[item]
def __setitem__(self, key, value: TradeData):
self._data[key] = value
def __delitem__(self, key):
del self._data[key]
def get(self, key, default=None) -> TradeData | None:
return self._data.get(key, default)
def update(self, *, ticket: int, **kwargs):
try:
res = self[ticket]
klass = type(res)
res = res._asdict()
res.update(**kwargs)
res = klass(res.get(v) for v in klass.__match_args__)
self[res.ticket] = res
return res
except KeyError:
logger.error(f"Update Operation Failed: Could Not Find Ticket")
def values(self) -> tuple[TradeData, ...]:
return tuple(value for value in self._data.values())
def keys(self) -> tuple[int, ...]:
return tuple(key for key in self._data.keys())
def items(self) -> tuple[tuple[int, TradeData], ...]:
return tuple((key, value) for key, value in self._data.items())
def to_dict(self):
return {key: value._asdict() for key, value in self._data.items()}
class PositionsManager(TradeManager):
_data: dict[int, TradePosition]
_open_positions: set[int]
margins: dict[int, float]
def __init__(self, *, data: dict = None, open_positions: set = None, margins: dict = None):
super().__init__(data=data)
self._open_positions = open_positions or {trade.ticket for trade in self._data.values()}
self.margins: dict[int, float] = margins or dict()
def __len__(self):
return len(self._open_positions)
def __contains__(self, item: TradePosition):
return item.ticket in self._open_positions
def __getitem__(self, item):
if item in self:
return super().__getitem__(item)
raise KeyError('Position not found')
def __setitem__(self, key, value: TradeData):
self._open_positions.add(value.ticket)
self._data[key] = value
def __delitem__(self, key):
self._open_positions.discard(key)
del self._data[key]
def close(self, *, ticket: int) -> bool:
is_open = ticket in self._open_positions
self._open_positions.discard(ticket)
return is_open
def get_margin(self, *, ticket: int) -> float:
return self.margins.get(ticket, 0.0)
def delete_margin(self, *, ticket: int):
return self.margins.pop(ticket, 0)
def set_margin(self, *, ticket: int, margin: float):
self.margins[ticket] = margin
def positions_get(self, *, ticket: int = None, symbol: str = None, group: None) -> tuple[TradePosition, ...]:
if ticket:
return tuple(position for position in self.open_positions if position.ticket == ticket)
if symbol:
return tuple(position for position in self.open_positions if position.symbol == symbol)
if group:
return self.open_positions
if ticket == group == symbol is None:
return self.open_positions
return tuple()
def positions_total(self) -> int:
return len(self)
@property
def open_positions(self) -> tuple[TradePosition, ...]:
return tuple(position for position in self.values() if position.ticket in self._open_positions)
class OrdersManager(TradeManager):
_data = dict[int, TradeOrder]
def get_orders_range(self, *, date_from: float, date_to: float) -> tuple[TradeData, ...]:
start = date_from.timestamp() if isinstance(date_from, datetime) else date_from
end = date_to.timestamp() if isinstance(date_to, datetime) else date_to
return tuple(order for order in self.values() if start <= order.time_setup <= end)
def history_orders_get(self, *, date_from: float | datetime = None, date_to: float | datetime = None,
group: str = '', ticket: int = None, position: int = None) -> tuple[TradeOrder, ...]:
if date_from and date_to:
orders = self.get_orders_range(date_from=date_from, date_to=date_to)
if group:
orders = orders
return orders
if ticket:
return tuple(order for order in self.values() if order.ticket == ticket)
if position:
return tuple(order for order in self.values() if order.position == position)
return ()
def history_orders_total(self, *, date_from: datetime | float, date_to: datetime | float) -> int:
return len(self.get_orders_range(date_from=date_from, date_to=date_to))
class DealsManager(TradeManager):
_data = dict[int, TradeDeal]
def get_deals_range(self, *, date_from: float, date_to: float) -> tuple[TradeData, ...]:
start = date_from.timestamp() if isinstance(date_from, datetime) else date_from
end = date_to.timestamp() if isinstance(date_to, datetime) else date_to
return tuple(deal for deal in self.values() if start <= deal.time <= end)
def history_deals_get(self, *, date_from: float | datetime = None, date_to: float | datetime = None,
group: str = '', ticket: int = None, position: int = None) -> tuple[TradeDeal, ...]:
if date_from and date_to:
deals = self.get_deals_range(date_from=date_from, date_to=date_to)
if group:
deals = deals
return deals
if ticket:
return tuple(deal for deal in self.values() if deal.ticket == ticket)
if position and (ticket is None):
return tuple(deal for deal in self.values() if deal.position == position)
return ()
def history_deals_total(self, *, date_from: datetime | float, date_to: datetime | float) -> int:
return len(self.get_deals_range(date_from=date_from, date_to=date_to))
-131
View File
@@ -1,131 +0,0 @@
from typing import TypeVar, Generic
from MetaTrader5 import TradePosition, TradeOrder, TradeDeal
from aiomql.utils import logger
TradeData = TypeVar('TradeData', bound=TradePosition | TradeOrder | TradeDeal)
class TradingData(Generic[TradeData]):
_data: dict[int, TradeData]
def __init__(self, open_items: set[int] = None, data: dict = None):
self._data = data or {}
def __len__(self):
return len(self._data)
def __getitem__(self, item):
return self._data[item]
def __setitem__(self, key, value: TradeData):
self._data[key] = value
def __contains__(self, item: int):
return item in self._open_items
def __iter__(self):
return iter(self._data)
def get(self, key, default=None) -> TradeData | None:
return self._data.get(key, default)
def update(self, *, ticket: int, **kwargs):
try:
res = self[ticket]
klass = type(res)
res = res._asdict()
res.update(**kwargs)
res = klass(res.get(v) for v in klass.__match_args__)
self[res.ticket] = res
return res
except KeyError:
logger.error(f"Update Operation Failed: Could Not Find Ticket")
def get_by_symbol(self, *, symbol: str) -> tuple[TradeData, ...]:
return tuple(position for position in self._data.values() if position.symbol == symbol)
def get_by_ticket(self, *, ticket: int) -> tuple[TradeData, ...]:
return tuple(position for position in self._data.values() if position.ticket == ticket)
class PositionsManager(TradingData):
_data: dict[int, TradePosition]
def __init__(self, open_items: set[int] = None, data: dict = None):
super().__init__(data=data)
self._open_items = open_items or {trade.ticket for trade in self._data.values()}
def __len__(self):
return len(self._open_items)
def __setitem__(self, key, value: TradeData):
self._open_items.add(value.ticket)
self._data[key] = value
def __delitem__(self, key):
try:
self._open_items.discard(key)
except KeyError:
logger.warning(f'{key} not found')
def get(self, key, default=None) -> TradeData | None:
return self._data.get(key, default) if key in self._open_items else default
def pop(self, key, default=None) -> TradeData | None:
self._open_items.discard(key)
return self._data.get(key, default)
def positions_get(self, *, ticket: int = None, symbol: str = None, group: None) -> tuple[TradePosition, ...]:
if ticket and (symbol == group == None):
return self.get_by_ticket(ticket=ticket)
if symbol and (ticket == group == None):
return self.get_by_symbol(symbol=symbol)
if group and (ticket == symbol == None):
return tuple(position for position in self._data.values())
if ticket == group == symbol == None:
return tuple(position for position in self._data.values())
return tuple()
def positions_total(self) -> int:
return len(self)
@property
def open_positions(self) -> tuple[TradePosition, ...]:
return tuple(position for position in self._data.values() if position.ticket in self.open_items)
class OrdersManager(TradingData):
_data = dict[int, TradeOrder]
def get_orders_range(self, *, date_from: float, date_to: float) -> tuple[TradeData, ...]:
start = date_from.timestamp() if isinstance(date_from, datetime) else date_from
end = date_to.timestamp() if isinstance(date_to, datetime) else date_to
return tuple(v for v in self._data.values() if start <= v.time_setup <= end)
def history_orders_get(self, *, date_from: float | datetime, date_to: float | datetime,
group: str = '', ticket: int = None, position: int = None) -> tuple[TradeOrder, ...]:
orders = self.get_orders_range(date_from=date_from, date_to=date_to)
if ticket and (position == None):
return tuple(order for order in orders if order.ticket == ticket)
if position and (ticket == None):
return tuple(order for order in orders if order.position == position)
return orders
class DealsManager(TradingData):
_data = dict[int, TradeDeal]
def get_deals_range(self, *, date_from: float, date_to: float) -> tuple[TradeData, ...]:
start = date_from.timestamp() if isinstance(date_from, datetime) else date_from
end = date_to.timestamp() if isinstance(date_to, datetime) else date_to
return tuple(v for v in self._data.values() if start <= v.time <= end)
def history_deals_get(self, date_from: float | datetime, date_to: float | datetime,
group: str = '', ticket: int = None, position: int = None) -> tuple[TradeDeal, ...]:
deals = self.get_deals_range(date_from, date_to)
if ticket and (position == None):
return tuple(deal for deal in deals if deal.ticket == ticket)
if position and (ticket == None):
return tuple(deal for deal in deals if deal.position == position)
return deals
-24
View File
@@ -1,24 +0,0 @@
from MetaTrader5 import TradePosition, TradeOrder, TradeDeal
# from .get_data import Data
class LiveDesc:
"""A Descriptor for live trading data"""
def __set_name__(self, owner, name):
self.access_name = name
def __get__(self, instance, owner):
return instance.__dict__.get(self.access_name, {})
def __set__(self, instance, value: tuple[int, str]):
prop = instance.__dict__.setdefault(self.access_name, {})
prop[value[0]] = value[1]
class Data:
pos = LiveDesc()
ords = LiveDesc()
dd = Data()
dd.pos = (1, 'EURUSD')
print(dd.pos)
+356
View File
@@ -0,0 +1,356 @@
from typing import Callable
import MetaTrader5
from MetaTrader5 import (Tick, SymbolInfo, AccountInfo, TerminalInfo, TradeOrder, TradePosition, TradeDeal,
OrderCheckResult, OrderSendResult, BookInfo, TradeRequest)
constants = ('TIMEFRAME_M1', 'TIMEFRAME_M2', 'TIMEFRAME_M3', 'TIMEFRAME_M4', 'TIMEFRAME_M5', 'TIMEFRAME_M6',
'TIMEFRAME_M10', 'TIMEFRAME_M12', 'TIMEFRAME_M15', 'TIMEFRAME_M20', 'TIMEFRAME_M30', 'TIMEFRAME_H1',
'TIMEFRAME_H2', 'TIMEFRAME_H4', 'TIMEFRAME_H3', 'TIMEFRAME_H6', 'TIMEFRAME_H8', 'TIMEFRAME_H12',
'TIMEFRAME_D1', 'TIMEFRAME_W1', 'TIMEFRAME_MN1', 'COPY_TICKS_ALL', 'COPY_TICKS_INFO', 'COPY_TICKS_TRADE',
'TICK_FLAG_BID', 'TICK_FLAG_ASK', 'TICK_FLAG_LAST', 'TICK_FLAG_VOLUME', 'TICK_FLAG_BUY', 'TICK_FLAG_SELL',
'POSITION_TYPE_BUY', 'POSITION_TYPE_SELL', 'POSITION_REASON_CLIENT', 'POSITION_REASON_MOBILE',
'POSITION_REASON_WEB', 'POSITION_REASON_EXPERT', 'ORDER_TYPE_BUY', 'ORDER_TYPE_SELL',
'ORDER_TYPE_BUY_LIMIT', 'ORDER_TYPE_SELL_LIMIT', 'ORDER_TYPE_BUY_STOP', 'ORDER_TYPE_SELL_STOP',
'ORDER_TYPE_BUY_STOP_LIMIT', 'ORDER_TYPE_SELL_STOP_LIMIT', 'ORDER_TYPE_CLOSE_BY', 'ORDER_STATE_STARTED',
'ORDER_STATE_PLACED', 'ORDER_STATE_CANCELED', 'ORDER_STATE_PARTIAL', 'ORDER_STATE_FILLED',
'ORDER_STATE_REJECTED', 'ORDER_STATE_EXPIRED', 'ORDER_STATE_REQUEST_ADD', 'ORDER_STATE_REQUEST_MODIFY',
'ORDER_STATE_REQUEST_CANCEL', 'ORDER_FILLING_FOK', 'ORDER_FILLING_IOC', 'ORDER_FILLING_RETURN',
'ORDER_FILLING_BOC', 'ORDER_TIME_GTC', 'ORDER_TIME_DAY', 'ORDER_TIME_SPECIFIED',
'ORDER_TIME_SPECIFIED_DAY', 'ORDER_REASON_CLIENT', 'ORDER_REASON_MOBILE', 'ORDER_REASON_WEB',
'ORDER_REASON_EXPERT', 'ORDER_REASON_SL', 'ORDER_REASON_TP', 'ORDER_REASON_SO', 'DEAL_TYPE_BUY',
'DEAL_TYPE_SELL', 'DEAL_TYPE_BALANCE', 'DEAL_TYPE_CREDIT', 'DEAL_TYPE_CHARGE', 'DEAL_TYPE_CORRECTION',
'DEAL_TYPE_BONUS', 'DEAL_TYPE_COMMISSION', 'DEAL_TYPE_COMMISSION_DAILY', 'DEAL_TYPE_COMMISSION_MONTHLY',
'DEAL_TYPE_COMMISSION_AGENT_DAILY', 'DEAL_TYPE_COMMISSION_AGENT_MONTHLY', 'DEAL_TYPE_INTEREST',
'DEAL_TYPE_BUY_CANCELED', 'DEAL_TYPE_SELL_CANCELED', 'DEAL_DIVIDEND', 'DEAL_DIVIDEND_FRANKED', 'DEAL_TAX',
'DEAL_ENTRY_IN', 'DEAL_ENTRY_OUT', 'DEAL_ENTRY_INOUT', 'DEAL_ENTRY_OUT_BY', 'DEAL_REASON_CLIENT',
'DEAL_REASON_MOBILE', 'DEAL_REASON_WEB', 'DEAL_REASON_EXPERT', 'DEAL_REASON_SL', 'DEAL_REASON_TP',
'DEAL_REASON_SO', 'DEAL_REASON_ROLLOVER', 'DEAL_REASON_VMARGIN', 'DEAL_REASON_SPLIT', 'TRADE_ACTION_DEAL',
'TRADE_ACTION_PENDING', 'TRADE_ACTION_SLTP', 'TRADE_ACTION_MODIFY', 'TRADE_ACTION_REMOVE',
'TRADE_ACTION_CLOSE_BY', 'SYMBOL_CHART_MODE_BID', 'SYMBOL_CHART_MODE_LAST', 'SYMBOL_CALC_MODE_FOREX',
'SYMBOL_CALC_MODE_FUTURES', 'SYMBOL_CALC_MODE_CFD', 'SYMBOL_CALC_MODE_CFDINDEX',
'SYMBOL_CALC_MODE_CFDLEVERAGE', 'SYMBOL_CALC_MODE_FOREX_NO_LEVERAGE', 'SYMBOL_CALC_MODE_EXCH_STOCKS',
'SYMBOL_CALC_MODE_EXCH_FUTURES', 'SYMBOL_CALC_MODE_EXCH_OPTIONS', 'SYMBOL_CALC_MODE_EXCH_OPTIONS_MARGIN',
'SYMBOL_CALC_MODE_EXCH_BONDS', 'SYMBOL_CALC_MODE_EXCH_STOCKS_MOEX', 'SYMBOL_CALC_MODE_EXCH_BONDS_MOEX',
'SYMBOL_CALC_MODE_SERV_COLLATERAL', 'SYMBOL_TRADE_MODE_DISABLED', 'SYMBOL_TRADE_MODE_LONGONLY',
'SYMBOL_TRADE_MODE_SHORTONLY', 'SYMBOL_TRADE_MODE_CLOSEONLY', 'SYMBOL_TRADE_MODE_FULL',
'SYMBOL_TRADE_EXECUTION_REQUEST', 'SYMBOL_TRADE_EXECUTION_INSTANT', 'SYMBOL_TRADE_EXECUTION_MARKET',
'SYMBOL_TRADE_EXECUTION_EXCHANGE', 'SYMBOL_SWAP_MODE_DISABLED', 'SYMBOL_SWAP_MODE_POINTS',
'SYMBOL_SWAP_MODE_CURRENCY_SYMBOL', 'SYMBOL_SWAP_MODE_CURRENCY_MARGIN',
'SYMBOL_SWAP_MODE_CURRENCY_DEPOSIT', 'SYMBOL_SWAP_MODE_INTEREST_CURRENT', 'SYMBOL_SWAP_MODE_INTEREST_OPEN',
'SYMBOL_SWAP_MODE_REOPEN_CURRENT', 'SYMBOL_SWAP_MODE_REOPEN_BID', 'DAY_OF_WEEK_SUNDAY',
'DAY_OF_WEEK_MONDAY', 'DAY_OF_WEEK_TUESDAY', 'DAY_OF_WEEK_WEDNESDAY', 'DAY_OF_WEEK_THURSDAY',
'DAY_OF_WEEK_FRIDAY', 'DAY_OF_WEEK_SATURDAY', 'SYMBOL_ORDERS_GTC', 'SYMBOL_ORDERS_DAILY',
'SYMBOL_ORDERS_DAILY_NO_STOPS', 'SYMBOL_OPTION_RIGHT_CALL', 'SYMBOL_OPTION_RIGHT_PUT',
'SYMBOL_OPTION_MODE_EUROPEAN', 'SYMBOL_OPTION_MODE_AMERICAN', 'ACCOUNT_TRADE_MODE_DEMO',
'ACCOUNT_TRADE_MODE_CONTEST', 'ACCOUNT_TRADE_MODE_REAL', 'ACCOUNT_STOPOUT_MODE_PERCENT',
'ACCOUNT_STOPOUT_MODE_MONEY', 'ACCOUNT_MARGIN_MODE_RETAIL_NETTING', 'ACCOUNT_MARGIN_MODE_EXCHANGE',
'ACCOUNT_MARGIN_MODE_RETAIL_HEDGING', 'BOOK_TYPE_SELL', 'BOOK_TYPE_BUY', 'BOOK_TYPE_SELL_MARKET',
'BOOK_TYPE_BUY_MARKET', 'TRADE_RETCODE_REQUOTE', 'TRADE_RETCODE_REJECT', 'TRADE_RETCODE_CANCEL',
'TRADE_RETCODE_PLACED', 'TRADE_RETCODE_DONE', 'TRADE_RETCODE_DONE_PARTIAL', 'TRADE_RETCODE_ERROR',
'TRADE_RETCODE_TIMEOUT', 'TRADE_RETCODE_INVALID', 'TRADE_RETCODE_INVALID_VOLUME',
'TRADE_RETCODE_INVALID_PRICE', 'TRADE_RETCODE_INVALID_STOPS', 'TRADE_RETCODE_TRADE_DISABLED',
'TRADE_RETCODE_MARKET_CLOSED', 'TRADE_RETCODE_NO_MONEY', 'TRADE_RETCODE_PRICE_CHANGED',
'TRADE_RETCODE_PRICE_OFF', 'TRADE_RETCODE_INVALID_EXPIRATION', 'TRADE_RETCODE_ORDER_CHANGED',
'TRADE_RETCODE_TOO_MANY_REQUESTS', 'TRADE_RETCODE_NO_CHANGES', 'TRADE_RETCODE_SERVER_DISABLES_AT',
'TRADE_RETCODE_CLIENT_DISABLES_AT', 'TRADE_RETCODE_LOCKED', 'TRADE_RETCODE_FROZEN',
'TRADE_RETCODE_INVALID_FILL', 'TRADE_RETCODE_CONNECTION', 'TRADE_RETCODE_ONLY_REAL',
'TRADE_RETCODE_LIMIT_ORDERS', 'TRADE_RETCODE_LIMIT_VOLUME', 'TRADE_RETCODE_INVALID_ORDER',
'TRADE_RETCODE_POSITION_CLOSED', 'TRADE_RETCODE_INVALID_CLOSE_VOLUME', 'TRADE_RETCODE_CLOSE_ORDER_EXIST',
'TRADE_RETCODE_LIMIT_POSITIONS', 'TRADE_RETCODE_REJECT_CANCEL', 'TRADE_RETCODE_LONG_ONLY',
'TRADE_RETCODE_SHORT_ONLY', 'TRADE_RETCODE_CLOSE_ONLY', 'TRADE_RETCODE_FIFO_CLOSE', 'RES_S_OK',
'RES_E_FAIL', 'RES_E_INVALID_PARAMS', 'RES_E_NO_MEMORY', 'RES_E_NOT_FOUND', 'RES_E_INVALID_VERSION',
'RES_E_AUTH_FAILED', 'RES_E_UNSUPPORTED', 'RES_E_AUTO_TRADING_DISABLED', 'RES_E_INTERNAL_FAIL',
'RES_E_INTERNAL_FAIL_SEND', 'RES_E_INTERNAL_FAIL_RECEIVE', 'RES_E_INTERNAL_FAIL_INIT',
'RES_E_INTERNAL_FAIL_CONNECT', 'RES_E_INTERNAL_FAIL_TIMEOUT')
core_mt5_functions = ('initialize', 'shutdown', 'login', 'version', 'terminal_info', 'account_info', 'copy_ticks_from',
'copy_ticks_range', 'copy_rates_from', 'copy_rates_from_pos', 'copy_rates_range', 'positions_total',
'positions_get', 'orders_total', 'orders_get', 'history_orders_total', 'history_orders_get',
'history_deals_total', 'history_deals_get', 'order_check', 'order_send', 'order_calc_margin',
'order_calc_profit', 'symbol_info', 'symbol_info_tick', 'symbol_select', 'symbols_total', 'symbols_get',
'market_book_add', 'market_book_release', 'market_book_get', 'last_error')
types = ('TradePosition', 'TradeOrder', 'TradeDeal', 'TradeRequest', 'OrderSendResult', 'OrderCheckResult', 'Tick',
'TerminalInfo', 'SymbolInfo', 'AccountInfo', 'BookInfo')
class BaseMeta(type):
def __new__(mcs, cls_name, bases, cls_dict):
defaults: dict = getattr(MetaTrader5, '__dict__', {})
callables = {f"_{key}": value for key in core_mt5_functions if (value := defaults.get(key, None)) is not None}
consts = {key: value for key in constants if (value := defaults.get(key, None)) is not None}
types_ = {key: value for key in types if (value := defaults.get(key, None)) is not None}
cls_dict |= callables
cls_dict |= consts
cls_dict |= types_
return super().__new__(mcs, cls_name, bases, cls_dict)
class MetaCore(metaclass=BaseMeta):
TIMEFRAME_M1: int
TIMEFRAME_M2: int
TIMEFRAME_M3: int
TIMEFRAME_M4: int
TIMEFRAME_M5: int
TIMEFRAME_M6: int
TIMEFRAME_M10: int
TIMEFRAME_M12: int
TIMEFRAME_M15: int
TIMEFRAME_M20: int
TIMEFRAME_M30: int
TIMEFRAME_H1: int
TIMEFRAME_H2: int
TIMEFRAME_H4: int
TIMEFRAME_H3: int
TIMEFRAME_H6: int
TIMEFRAME_H8: int
TIMEFRAME_H12: int
TIMEFRAME_D1: int
TIMEFRAME_W1: int
TIMEFRAME_MN1: int
COPY_TICKS_ALL: int
COPY_TICKS_INFO: int
COPY_TICKS_TRADE: int
TICK_FLAG_BID: int
TICK_FLAG_ASK: int
TICK_FLAG_LAST: int
TICK_FLAG_VOLUME: int
TICK_FLAG_BUY: int
TICK_FLAG_SELL: int
POSITION_TYPE_BUY: int
POSITION_TYPE_SELL: int
POSITION_REASON_CLIENT: int
POSITION_REASON_MOBILE: int
POSITION_REASON_WEB: int
POSITION_REASON_EXPERT: int
ORDER_TYPE_BUY: int
ORDER_TYPE_SELL: int
ORDER_TYPE_BUY_LIMIT: int
ORDER_TYPE_SELL_LIMIT: int
ORDER_TYPE_BUY_STOP: int
ORDER_TYPE_SELL_STOP: int
ORDER_TYPE_BUY_STOP_LIMIT: int
ORDER_TYPE_SELL_STOP_LIMIT: int
ORDER_TYPE_CLOSE_BY: int
ORDER_STATE_STARTED: int
ORDER_STATE_PLACED: int
ORDER_STATE_CANCELED: int
ORDER_STATE_PARTIAL: int
ORDER_STATE_FILLED: int
ORDER_STATE_REJECTED: int
ORDER_STATE_EXPIRED: int
ORDER_STATE_REQUEST_ADD: int
ORDER_STATE_REQUEST_MODIFY: int
ORDER_STATE_REQUEST_CANCEL: int
ORDER_FILLING_FOK: int
ORDER_FILLING_IOC: int
ORDER_FILLING_RETURN: int
ORDER_FILLING_BOC: int
ORDER_TIME_GTC: int
ORDER_TIME_DAY: int
ORDER_TIME_SPECIFIED: int
ORDER_TIME_SPECIFIED_DAY: int
ORDER_REASON_CLIENT: int
ORDER_REASON_MOBILE: int
ORDER_REASON_WEB: int
ORDER_REASON_EXPERT: int
ORDER_REASON_SL: int
ORDER_REASON_TP: int
ORDER_REASON_SO: int
DEAL_TYPE_BUY: int
DEAL_TYPE_SELL: int
DEAL_TYPE_BALANCE: int
DEAL_TYPE_CREDIT: int
DEAL_TYPE_CHARGE: int
DEAL_TYPE_CORRECTION: int
DEAL_TYPE_BONUS: int
DEAL_TYPE_COMMISSION: int
DEAL_TYPE_COMMISSION_DAILY: int
DEAL_TYPE_COMMISSION_MONTHLY: int
DEAL_TYPE_COMMISSION_AGENT_DAILY: int
DEAL_TYPE_COMMISSION_AGENT_MONTHLY: int
DEAL_TYPE_INTEREST: int
DEAL_TYPE_BUY_CANCELED: int
DEAL_TYPE_SELL_CANCELED: int
DEAL_DIVIDEND: int
DEAL_DIVIDEND_FRANKED: int
DEAL_TAX: int
DEAL_ENTRY_IN: int
DEAL_ENTRY_OUT: int
DEAL_ENTRY_INOUT: int
DEAL_ENTRY_OUT_BY: int
DEAL_REASON_CLIENT: int
DEAL_REASON_MOBILE: int
DEAL_REASON_WEB: int
DEAL_REASON_EXPERT: int
DEAL_REASON_SL: int
DEAL_REASON_TP: int
DEAL_REASON_SO: int
DEAL_REASON_ROLLOVER: int
DEAL_REASON_VMARGIN: int
DEAL_REASON_SPLIT: int
TRADE_ACTION_DEAL: int
TRADE_ACTION_PENDING: int
TRADE_ACTION_SLTP: int
TRADE_ACTION_MODIFY: int
TRADE_ACTION_REMOVE: int
TRADE_ACTION_CLOSE_BY: int
SYMBOL_CHART_MODE_BID: int
SYMBOL_CHART_MODE_LAST: int
SYMBOL_CALC_MODE_FOREX: int
SYMBOL_CALC_MODE_FUTURES: int
SYMBOL_CALC_MODE_CFD: int
SYMBOL_CALC_MODE_CFDINDEX: int
SYMBOL_CALC_MODE_CFDLEVERAGE: int
SYMBOL_CALC_MODE_FOREX_NO_LEVERAGE: int
SYMBOL_CALC_MODE_EXCH_STOCKS: int
SYMBOL_CALC_MODE_EXCH_FUTURES: int
SYMBOL_CALC_MODE_EXCH_OPTIONS: int
SYMBOL_CALC_MODE_EXCH_OPTIONS_MARGIN: int
SYMBOL_CALC_MODE_EXCH_BONDS: int
SYMBOL_CALC_MODE_EXCH_STOCKS_MOEX: int
SYMBOL_CALC_MODE_EXCH_BONDS_MOEX: int
SYMBOL_CALC_MODE_SERV_COLLATERAL: int
SYMBOL_TRADE_MODE_DISABLED: int
SYMBOL_TRADE_MODE_LONGONLY: int
SYMBOL_TRADE_MODE_SHORTONLY: int
SYMBOL_TRADE_MODE_CLOSEONLY: int
SYMBOL_TRADE_MODE_FULL: int
SYMBOL_TRADE_EXECUTION_REQUEST: int
SYMBOL_TRADE_EXECUTION_INSTANT: int
SYMBOL_TRADE_EXECUTION_MARKET: int
SYMBOL_TRADE_EXECUTION_EXCHANGE: int
SYMBOL_SWAP_MODE_DISABLED: int
SYMBOL_SWAP_MODE_POINTS: int
SYMBOL_SWAP_MODE_CURRENCY_SYMBOL: int
SYMBOL_SWAP_MODE_CURRENCY_MARGIN: int
SYMBOL_SWAP_MODE_CURRENCY_DEPOSIT: int
SYMBOL_SWAP_MODE_INTEREST_CURRENT: int
SYMBOL_SWAP_MODE_INTEREST_OPEN: int
SYMBOL_SWAP_MODE_REOPEN_CURRENT: int
SYMBOL_SWAP_MODE_REOPEN_BID: int
DAY_OF_WEEK_SUNDAY: int
DAY_OF_WEEK_MONDAY: int
DAY_OF_WEEK_TUESDAY: int
DAY_OF_WEEK_WEDNESDAY: int
DAY_OF_WEEK_THURSDAY: int
DAY_OF_WEEK_FRIDAY: int
DAY_OF_WEEK_SATURDAY: int
SYMBOL_ORDERS_GTC: int
SYMBOL_ORDERS_DAILY: int
SYMBOL_ORDERS_DAILY_NO_STOPS: int
SYMBOL_OPTION_RIGHT_CALL: int
SYMBOL_OPTION_RIGHT_PUT: int
SYMBOL_OPTION_MODE_EUROPEAN: int
SYMBOL_OPTION_MODE_AMERICAN: int
ACCOUNT_TRADE_MODE_DEMO: int
ACCOUNT_TRADE_MODE_CONTEST: int
ACCOUNT_TRADE_MODE_REAL: int
ACCOUNT_STOPOUT_MODE_PERCENT: int
ACCOUNT_STOPOUT_MODE_MONEY: int
ACCOUNT_MARGIN_MODE_RETAIL_NETTING: int
ACCOUNT_MARGIN_MODE_EXCHANGE: int
ACCOUNT_MARGIN_MODE_RETAIL_HEDGING: int
BOOK_TYPE_SELL: int
BOOK_TYPE_BUY: int
BOOK_TYPE_SELL_MARKET: int
BOOK_TYPE_BUY_MARKET: int
TRADE_RETCODE_REQUOTE: int
TRADE_RETCODE_REJECT: int
TRADE_RETCODE_CANCEL: int
TRADE_RETCODE_PLACED: int
TRADE_RETCODE_DONE: int
TRADE_RETCODE_DONE_PARTIAL: int
TRADE_RETCODE_ERROR: int
TRADE_RETCODE_TIMEOUT: int
TRADE_RETCODE_INVALID: int
TRADE_RETCODE_INVALID_VOLUME: int
TRADE_RETCODE_INVALID_PRICE: int
TRADE_RETCODE_INVALID_STOPS: int
TRADE_RETCODE_TRADE_DISABLED: int
TRADE_RETCODE_MARKET_CLOSED: int
TRADE_RETCODE_NO_MONEY: int
TRADE_RETCODE_PRICE_CHANGED: int
TRADE_RETCODE_PRICE_OFF: int
TRADE_RETCODE_INVALID_EXPIRATION: int
TRADE_RETCODE_ORDER_CHANGED: int
TRADE_RETCODE_TOO_MANY_REQUESTS: int
TRADE_RETCODE_NO_CHANGES: int
TRADE_RETCODE_SERVER_DISABLES_AT: int
TRADE_RETCODE_CLIENT_DISABLES_AT: int
TRADE_RETCODE_LOCKED: int
TRADE_RETCODE_FROZEN: int
TRADE_RETCODE_INVALID_FILL: int
TRADE_RETCODE_CONNECTION: int
TRADE_RETCODE_ONLY_REAL: int
TRADE_RETCODE_LIMIT_ORDERS: int
TRADE_RETCODE_LIMIT_VOLUME: int
TRADE_RETCODE_INVALID_ORDER: int
TRADE_RETCODE_POSITION_CLOSED: int
TRADE_RETCODE_INVALID_CLOSE_VOLUME: int
TRADE_RETCODE_CLOSE_ORDER_EXIST: int
TRADE_RETCODE_LIMIT_POSITIONS: int
TRADE_RETCODE_REJECT_CANCEL: int
TRADE_RETCODE_LONG_ONLY: int
TRADE_RETCODE_SHORT_ONLY: int
TRADE_RETCODE_CLOSE_ONLY: int
TRADE_RETCODE_FIFO_CLOSE: int
RES_S_OK: int
RES_E_FAIL: int
RES_E_INVALID_PARAMS: int
RES_E_NO_MEMORY: int
RES_E_NOT_FOUND: int
RES_E_INVALID_VERSION: int
RES_E_AUTH_FAILED: int
RES_E_UNSUPPORTED: int
RES_E_AUTO_TRADING_DISABLED: int
RES_E_INTERNAL_FAIL: int
RES_E_INTERNAL_FAIL_SEND: int
RES_E_INTERNAL_FAIL_RECEIVE: int
RES_E_INTERNAL_FAIL_INIT: int
RES_E_INTERNAL_FAIL_CONNECT: int
RES_E_INTERNAL_FAIL_TIMEOUT: int
_account_info: Callable
_copy_rates_from: Callable
_copy_rates_from_pos: Callable
_copy_rates_range: Callable
_copy_ticks_from: Callable
_copy_ticks_range: Callable
_history_deals_get: Callable
_history_deals_total: Callable
_history_orders_get: Callable
_history_orders_total: Callable
_initialize: Callable
_last_error: Callable
_login: Callable
_market_book_add: Callable
_market_book_get: Callable
_market_book_release: Callable
_order_calc_margin: Callable
_order_calc_profit: Callable
_order_check: Callable
_order_send: Callable
_orders_get: Callable
_orders_total: Callable
_positions_get: Callable
_positions_total: Callable
_shutdown: Callable
_symbol_info: Callable
_symbol_info_tick: Callable
_symbol_select: Callable
_symbols_get: Callable
_symbols_total: Callable
_terminal_info: Callable
_version: Callable
config: Config
AccountInfo: AccountInfo
TradePosition: TradePosition
TradeOrder: TradeOrder
TradeDeal: TradeDeal
TradeRequest: TradeRequest
OrderSendResult: OrderSendResult
OrderCheckResult: OrderCheckResult
Tick: Tick
TerminalInfo: TerminalInfo
SymbolInfo: SymbolInfo
BookInfo: BookInfo
+3 -2
View File
@@ -4,7 +4,8 @@ from logging import getLogger
from .config import Config
from .meta_trader import MetaTrader
# from ..contrib.backtester import MetaTester
from ..contrib.backtester import MetaTester
logger = getLogger(__name__)
@@ -22,7 +23,7 @@ class Base:
**kwargs: Set instance attributes with keyword arguments. Only if they are annotated on the class body.
"""
self.config = Config()
self.mt5 = MetaTrader() #if self.config.mode == 'live' else MetaTester()
self.mt5 = MetaTrader() if self.config.mode == 'live' else MetaTester()
self.exclude = {'mt5', "config", 'exclude', 'include', 'annotations', 'class_vars', 'dict'}
self.include = set()
self.set_attributes(**kwargs)
+18 -19
View File
@@ -7,8 +7,9 @@ from logging import getLogger
from .task_queue import TaskQueue
logger = getLogger(__name__)
Bot = TypeVar("Bot")
TestData = TypeVar("TestData")
BackTestEngine = TypeVar("BackTestEngine")
class Config:
@@ -45,20 +46,17 @@ class Config:
record_trades: bool
records_dir: Path
records_dir_name: str
compress_test_data: bool
test_data_dir: Path
test_data_dir_name: str
backtest_dir: Path
backtest_dir_name: str
task_queue: TaskQueue
_test_data: TestData
_backtest_engine: BackTestEngine
bot: Bot
_instance: 'Config'
mode: Literal['backtest', 'live']
use_terminal_for_backtesting: bool
test_data_file: str
_defaults = {"timeout": 60000, "record_trades": True, "trade_record_mode": "csv", "mode": "live",
'filename': "aiomql.json", "records_dir_name": "trade_records", "test_data_dir_name": "test_data",
"use_terminal_for_backtesting": True, 'path': '', 'login': 0, 'password': '', 'server': '',
"compress_test_data": False, 'test_data_file': ''}
'filename': "aiomql.json", "records_dir_name": "trade_records", "backtest_dir_name": "backtester",
"use_terminal_for_backtesting": True, 'path': '', 'login': 0, 'password': '', 'server': ''}
def __new__(cls, *args, **kwargs):
if not hasattr(cls, "_instance"):
@@ -66,7 +64,7 @@ class Config:
cls._instance.state = {}
cls._instance.task_queue = TaskQueue()
cls._instance.set_attributes(**cls._defaults)
cls._instance._test_data = None
cls._instance._backtest_engine = None
cls._instance.load_config(**kwargs)
return cls._instance
@@ -74,12 +72,12 @@ class Config:
self.set_attributes(**kwargs)
@property
def test_data(self):
return self._test_data
def backtest_engine(self):
return self._backtest_engine
@test_data.setter
def test_data(self, value: TestData):
self._test_data = value
@backtest_engine.setter
def backtest_engine(self, value: BackTestEngine):
self._backtest_engine = value
def set_attributes(self, **kwargs):
"""Set keyword arguments as object attributes
@@ -111,7 +109,8 @@ class Config:
if os.path.isfile(check_path):
return check_path
return None
except Exception as _:
except Exception as err:
logger.debug(f"Error finding config file: {err}")
return
def load_config(self, *, file: str | Path = None, filename: str = None, root: str | Path = None, **kwargs):
@@ -156,9 +155,9 @@ class Config:
self.records_dir = self.root / self.records_dir_name
self.records_dir.mkdir(parents=True, exist_ok=True)
if self.mode == "backtest" and not hasattr(self, "test_data_dir"):
self.test_data_dir = self.root / self.test_data_dir_name
self.test_data_dir.mkdir(parents=True, exist_ok=True)
if self.mode == "backtest" and not hasattr(self, "backtest_dir"):
self.backtest_dir = self.root / self.backtest_dir_name
self.backtest_dir.mkdir(parents=True, exist_ok=True)
def account_info(self) -> dict[str, int | str]:
"""Returns Account login details as found in the config object if available
+1 -1
View File
@@ -17,7 +17,7 @@ class Repr:
__enum_name__ = ""
name: str
def __repr__(self):
def __str__(self):
return f"{self.__enum_name__}_{self.name}"
+1 -1
View File
@@ -20,7 +20,7 @@ class Error:
-10005: 'internal timeout',
}
conn_errors = (-10000, -10001, -10002, -10003, -10004, -10005)
conn_errors = (-10000, -10001, -10002, -10003, -10004, -10005, -6)
def __init__(self, code: int, description: str = ''):
self.code = code
+14 -54
View File
@@ -1,63 +1,22 @@
from datetime import datetime
import asyncio
from logging import getLogger
from typing import Callable
from typing import Literal
import MetaTrader5
import numpy as np
from MetaTrader5 import BookInfo, SymbolInfo, AccountInfo, Tick, TerminalInfo, TradeOrder, TradeDeal, \
TradePosition, OrderSendResult, OrderCheckResult
from MetaTrader5 import (BookInfo, SymbolInfo, AccountInfo, Tick, TerminalInfo, TradeOrder, TradeDeal,
TradePosition, OrderSendResult, OrderCheckResult)
from .constants import TimeFrame, CopyTicks, OrderType
from .constants import OrderType, CopyTicks
from ._core import MetaCore
from .errors import Error
from .config import Config
logger = getLogger()
class BaseMeta(type):
def __new__(mcs, cls_name, bases, cls_dict):
defaults = MetaTrader5.__dict__
defaults = {f'_{key}': value for key, value in defaults.items() if not key.startswith('_')}
cls_dict |= defaults
return super().__new__(mcs, cls_name, bases, cls_dict)
class MetaTrader(metaclass=BaseMeta):
_account_info: Callable
_copy_rates_from: Callable
_copy_rates_from_pos: Callable
_copy_rates_range: Callable
_copy_ticks_from: Callable
_copy_ticks_range: Callable
_history_deals_get: Callable
_history_deals_total: Callable
_history_orders_get: Callable
_history_orders_total: Callable
_initialize: Callable
_last_error: Callable
_login: Callable
_market_book_add: Callable
_market_book_get: Callable
_market_book_release: Callable
_order_calc_margin: Callable
_order_calc_profit: Callable
_order_check: Callable
_order_send: Callable
_orders_get: Callable
_orders_total: Callable
_positions_get: Callable
_positions_total: Callable
_shutdown: Callable
_symbol_info: Callable
_symbol_info_tick: Callable
_symbol_select: Callable
_symbols_get: Callable
_symbols_total: Callable
_terminal_info: Callable
_version: Callable
config: Config
class MetaTrader(MetaCore):
def __init__(self):
self.config = Config()
self.error: Error = Error(1)
@@ -225,26 +184,27 @@ class MetaTrader(metaclass=BaseMeta):
res = await self._handler(api)
return res
async def copy_rates_from(self, symbol: str, timeframe: TimeFrame, date_from: datetime | float, count: int) -> np.ndarray | None:
async def copy_rates_from(self, symbol: str, timeframe: int, date_from: datetime | float, count: int) -> np.ndarray | None:
api = {'func': self._copy_rates_from, 'args': (symbol, timeframe, date_from, count),
'error_msg': f'Error in obtaining rates for {symbol}'}
res = await self._handler(api)
return res
async def copy_rates_from_pos(self, symbol: str, timeframe: TimeFrame, start_pos: int, count: int) -> np.ndarray | None:
async def copy_rates_from_pos(self, symbol: str, timeframe: int, start_pos: int, count: int) -> np.ndarray | None:
api = {'func': self._copy_rates_from_pos, 'args': (symbol, timeframe, start_pos, count),
'error_msg': f'Error in obtaining rates for {symbol}'}
res = await self._handler(api)
return res
async def copy_rates_range(self, symbol: str, timeframe: TimeFrame, date_from: datetime | float,
async def copy_rates_range(self, symbol: str, timeframe: int, date_from: datetime | float,
date_to: datetime | float) -> np.ndarray | None:
api = {'func': self._copy_rates_range, 'args': (symbol, timeframe, date_from, date_to),
'error_msg': f'Error in obtaining rates for {symbol}'}
res = await self._handler(api)
return res
async def copy_ticks_from(self, symbol: str, date_from: datetime | float, count: int, flags: CopyTicks) -> np.ndarray | None:
async def copy_ticks_from(self, symbol: str, date_from: datetime | float, count: int,
flags: CopyTicks) -> np.ndarray | None:
api = {'func': self._copy_ticks_from, 'args': (symbol, date_from, count, flags),
'error_msg': f'Error in obtaining ticks for {symbol}'}
res = await self._handler(api)
@@ -274,8 +234,8 @@ class MetaTrader(metaclass=BaseMeta):
res = await self._handler(api)
return res
async def order_calc_profit(self, action: OrderType, symbol: str, volume: float, price_open: float,
price_close: float) -> float | None:
async def order_calc_profit(self, action: Literal[OrderType.BUY, OrderType.SELL], symbol: str,
volume: float, price_open: float, price_close: float) -> float | None:
api = {'func': self._order_calc_profit, 'args': (action, symbol, volume, price_open, price_close),
'error_msg': 'Error in calculating profit.'}
res = await self._handler(api)
+1 -1
View File
@@ -4,7 +4,7 @@ from .constants import BookType, TradeAction, OrderType, OrderTime, OrderFilling
DealReason, SymbolChartMode, SymbolTradeMode, SymbolCalcMode, SymbolOptionMode, SymbolOrderGTCMode, \
SymbolOptionRight, \
SymbolTradeExecution, SymbolSwapMode, DayOfWeek, AccountTradeMode, AccountStopOutMode, AccountMarginMode, \
OrderReason, TickFlag
OrderReason
from .base import Base
+2 -1
View File
@@ -13,6 +13,7 @@ class QueueItem:
self.task_item = task_item
self.args = args
self.kwargs = kwargs
self.must_complete = False
self.time = asyncio.get_event_loop().time()
def __hash__(self):
@@ -55,7 +56,7 @@ class TaskQueue:
self.queue.put_nowait(item)
except asyncio.QueueFull:
logger.error(f"Queue is full, could not add {item}")
logger.error(f"Queue is full")
async def worker(self):
while True:
+5 -5
View File
@@ -1,6 +1,6 @@
import asyncio
from concurrent.futures import ThreadPoolExecutor
from typing import Sequence, Coroutine, Callable
from typing import Coroutine, Callable
from logging import getLogger
from .strategy import Strategy
@@ -20,7 +20,7 @@ class Executor:
def __init__(self):
self.executor = ThreadPoolExecutor
self.workers: list[type(Strategy)] = []
self.workers: list[Strategy] = []
self.coroutines: dict[Coroutine | Callable: dict] = {}
self.functions: dict[Callable: dict] = {}
@@ -30,7 +30,7 @@ class Executor:
def add_coroutine(self, coro: Coroutine, kwargs: dict):
self.coroutines[coro] = kwargs
def add_workers(self, strategies: Sequence[type(Strategy)]):
def add_workers(self, strategies: tuple[Strategy]):
"""Add multiple strategies at once
Args:
@@ -42,7 +42,7 @@ class Executor:
"""Removes any worker running on a symbol not successfully initialized."""
self.workers = [worker for worker in self.workers if worker.symbol in symbols]
def add_worker(self, strategy: type(Strategy)):
def add_worker(self, strategy: Strategy):
"""Add a strategy instance to the list of workers
Args:
@@ -51,7 +51,7 @@ class Executor:
self.workers.append(strategy)
@staticmethod
def trade(strategy: type(Strategy)):
def trade(strategy: Strategy):
"""Wraps the coroutine trade method of each strategy with 'asyncio.run'.
Args:
@@ -1,27 +1,8 @@
from .finger_trap import FingerTrap
from ...contrib.backtester.test_strategy import TestStrategy, TestSingleStrategy
from ...contrib.backtester.strategy_tester import StrategyTester
class FingerTrapSingleTest(TestSingleStrategy, FingerTrap):
async def test(self):
print(f"Backtesting {self.symbol}")
while True:
try:
await self.watch_market()
if not self.tracker.new:
continue
if self.tracker.order_type is not None:
await self.trader.place_trade(order_type=self.tracker.order_type, parameters=self.parameters,
sl=self.tracker.sl)
await self.sleep(self.tracker.snooze)
except Exception as err:
print(f"{err} For {self.symbol} in {self.__class__.__name__}.trade")
await self.sleep(self.ttf.time)
class FingerTrapTest(TestStrategy, FingerTrap):
class FingerTrapTest(StrategyTester, FingerTrap):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.set_up()
+2 -2
View File
@@ -104,13 +104,13 @@ class Positions:
async def close_by(self, pos: TradePosition):
"""Close an open position for the trading account."""
order = Order(position=pos.ticket, symbol=pos.symbol, volume=pos.volume, type=pos.type.opposite,
price=pos.price_current)
price=pos.price_current, action=TradeAction.DEAL)
return await order.send()
async def close_position(self, *, position: TradePosition):
"""Close an open position for the trading account. Using a position object."""
order = Order(position=position.ticket, symbol=position.symbol, volume=position.volume,
type=position.type.opposite, price=position.price_current)
type=position.type.opposite, price=position.price_current, action=TradeAction.DEAL)
return await order.send()
async def close_all(self, symbol: str = '', group: str = '') -> int:
+2 -2
View File
@@ -8,7 +8,7 @@ from datetime import time as dtime
from .core.meta_trader import MetaTrader
from .symbol import Symbol as _Symbol
from .core import Config
# from .contrib.backtester.meta_tester import MetaTester
from .contrib.backtester.meta_tester import MetaTester
from .sessions import Sessions, Session
Symbol = TypeVar("Symbol", bound=_Symbol)
@@ -48,7 +48,7 @@ class Strategy(ABC):
self.parameters["name"] = self.name
self.sessions = sessions or Sessions(Session(start=0, end=dtime(hour=23, minute=59, second=59)))
self.config = Config()
self.mt5 = MetaTrader() #if self.config.mode == 'live' else MetaTester()
self.mt5 = MetaTrader() if self.config.mode == 'live' else MetaTester()
def __repr__(self):
return f"{self.name}({self.symbol!r})"
+2 -2
View File
@@ -61,9 +61,9 @@ def error_handler(func=None, *, msg='', exe = Exception, response=None):
return wrapper
def error_handler_sync(func=None, *, msg='', exe=Exception):
def error_handler_sync(func=None, *, msg='', exe=Exception, response=None):
if func is None:
return partial(error_handler, msg=msg, exe=exe)
return partial(error_handler, msg=msg, exe=exe, response=response)
@wraps(func)
def wrapper(*args, **kwargs):