This commit is contained in:
Ichinga Samuel
2025-07-20 17:07:55 +01:00
parent 4178de2a0b
commit 675dbc95a2
55 changed files with 1413 additions and 234 deletions
+1
View File
@@ -72,6 +72,7 @@ target/
.pypirc .pypirc
.vscode/ .vscode/
*.sqlite3
# config files # config files
aiomql.json aiomql.json
View File
+21
View File
@@ -0,0 +1,21 @@
import logging
from aiomql import Bot, ForexSymbol, auto_commit, OpenPositionsTracker
from .emaxover import EMAXOver
logging.basicConfig(level=logging.INFO)
def x_bot():
syms = ["LTCUSD", "ETHUSD", "SOLUSD", "BTCUSD"]
symbols = [ForexSymbol(name=sym) for sym in syms]
strategies = [EMAXOver(symbol=symbol) for symbol in symbols]
bot = Bot()
bot.add_strategies(strategies=strategies)
bot.add_coroutine(coroutine=OpenPositionsTracker(autocommit=True).track, on_separate_thread=True)
# bot.add_coroutine(coroutine=auto_commit, on_separate_thread=True)
bot.execute()
x_bot()
@@ -1,5 +1,6 @@
from aiomql import Strategy, ForexSymbol, TimeFrame, Tracker, OrderType, Sessions, Trader, ScalpTrader from aiomql import Strategy, ForexSymbol, TimeFrame, Tracker, OrderType, Sessions, Trader, ScalpTrader
from .traders import TestTrader
class EMAXOver(Strategy): class EMAXOver(Strategy):
ttf: TimeFrame # time frame for the strategy ttf: TimeFrame # time frame for the strategy
@@ -12,18 +13,18 @@ class EMAXOver(Strategy):
# default parameters for the strategy # default parameters for the strategy
# they are set as attributes. You can override them in the constructor via the params argument. # they are set as attributes. You can override them in the constructor via the params argument.
parameters = {'ttf': TimeFrame.H1, 'tcc': 3000, 'fast_ema': 34, 'slow_ema': 55, 'interval': TimeFrame.M15, parameters = {'ttf': TimeFrame.M10, 'tcc': 3000, 'fast_ema': 34, 'slow_ema': 55, 'interval': TimeFrame.M5,
'timeout': 3 * 60 * 60} 'timeout': 120}
def __init__(self, *, symbol: ForexSymbol, params: dict | None = None, trader: Trader = None, def __init__(self, *, symbol: ForexSymbol, params: dict | None = None, trader: Trader = None,
sessions: Sessions = None, name: str = "EMAXOver"): sessions: Sessions = None, name: str = "EMAXOver"):
super().__init__(symbol=symbol, params=params, sessions=sessions, name=name) super().__init__(symbol=symbol, params=params, sessions=sessions, name=name)
self.tracker = Tracker(snooze=self.interval.seconds) self.tracker = Tracker(snooze=self.interval.seconds)
self.trader = trader or ScalpTrader(symbol=self.symbol) self.trader = trader or TestTrader(symbol=self.symbol)
async def find_entry(self): async def find_entry(self):
# get the candles # get the candles
candles = await self.symbol.copy_rates_from_pos(timeframe=self.ttf, start_position=0, count=self.tcc) candles = await self.symbol.copy_rates_from_pos(timeframe=self.ttf, count=self.tcc)
# get the fast moving average # get the fast moving average
candles.ta.ema(length=self.fast_ema, append=True) candles.ta.ema(length=self.fast_ema, append=True)
@@ -34,9 +35,9 @@ class EMAXOver(Strategy):
# check for crossovers # check for crossovers
# fast above slow # fast above slow
fas = candles.ta_lib.cross(candles.fast_ema, candles.slow_ema, above=True) fas = candles.ta_lib.above(candles.fast_ema, candles.slow_ema)
# fast below slow # fast below slow
fbs = candles.ta_lib.cross(candles.fast_ema, candles.slow_ema, above=False) fbs = candles.ta_lib.below(candles.fast_ema, candles.slow_ema)
## check for entry signals in the current candle ## check for entry signals in the current candle
if fas.iloc[-1]: if fas.iloc[-1]:
+1
View File
@@ -0,0 +1 @@
from .track import close_after
+18
View File
@@ -0,0 +1,18 @@
from logging import getLogger
from datetime import datetime
from aiomql import OpenPosition
logger = getLogger(__name__)
async def close_after(open_pos: OpenPosition, /, *, duration: int, start: float = 0):
if not await open_pos.update_position():
return
pos = open_pos.position
start = start or pos.time
diff = datetime.now().timestamp() - start
if diff > duration:
_, res = await open_pos.close_position()
if res.retcode == 10009:
logger.info("%s, %d closed", pos.symbol, pos.ticket)
+1
View File
@@ -0,0 +1 @@
from .test_trader import TestTrader
@@ -0,0 +1,42 @@
from logging import getLogger
from datetime import datetime
from aiomql import Trader, OrderType, OpenPosition, Positions, PositionTracker, Store
from ..trackers import close_after
logger = getLogger(__name__)
class TestTrader(Trader):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.positions = Positions()
self.store = Store()
async def place_trade(self, *, order_type: OrderType, volume: float = None, parameters: dict = None):
"""Places a trade based on the order_type and volume. The volume is optional. If not provided, the minimum volume
for the symbol will be used. This trade is placed without a stop_loss or take_profit. The trade is recorded in the
trade_record file.
Args:
order_type (OrderType): The order_type
volume (float): The volume to trade
parameters (dict): Parameters associated with the trade
"""
try:
self.parameters |= parameters or {}
volume = volume or self.symbol.volume_min
await self.create_order_no_stops(order_type=order_type, volume=volume)
if not await self.check_order():
return
self.order.comment = self.parameters.get("name", self.__class__.__name__)
res = await self.send_order()
if res is not None and res.retcode == 10009:
position = await self.positions.get_position_by_ticket(ticket=res.order)
open_position = OpenPosition(ticket=res.order, symbol=self.symbol, position=position)
kwargs = {"duration": 40, "start": datetime.now().timestamp()}
ca = PositionTracker(close_after, **kwargs)
open_position.add_tracker(tracker=ca)
await self.record_trade(result=res, parameters=self.parameters)
except Exception as err:
logger.error(f"{err} in {self.__class__.__name__}.place_trade for {self.symbol.name}")
-19
View File
@@ -1,19 +0,0 @@
import logging
from aiomql import Bot, ForexSymbol
from emaxover import EMAXOver
logging.basicConfig(level=logging.INFO)
def x_bot():
syms = ["EURUSD", "GBPUSD", "USDJPY"]
symbols = [ForexSymbol(name=sym) for sym in syms]
strategies = [EMAXOver(symbol=symbol) for symbol in symbols]
bot = Bot()
bot.add_strategies(strategies=strategies)
bot.execute()
x_bot()
+1 -1
View File
@@ -16,7 +16,7 @@ classifiers = [
keywords = ["MetaTrader5", "Asynchronous", "Algorithmic Trading", "Trading Bot", "Backtesting", keywords = ["MetaTrader5", "Asynchronous", "Algorithmic Trading", "Trading Bot", "Backtesting",
"Technical Analysis", "Forex", "Stocks", "Cryptocurrency", "Futures", "Options", "Crypto", "Algo Trading"] "Technical Analysis", "Forex", "Stocks", "Cryptocurrency", "Futures", "Options", "Crypto", "Algo Trading"]
dependencies = ["MetaTrader5>=5.0.37", "pandas>=1.5.0", "pandas-ta>=0.3.14b0"] dependencies = ["MetaTrader5>=5.0.37", "pandas>=1.5.0", "pandas-ta>=0.3.14b0", "mplfinance>=0.10.1"]
authors = [{name = "Ichinga Samuel", email = "ichingasamuel@gmail.com"}] authors = [{name = "Ichinga Samuel", email = "ichingasamuel@gmail.com"}]
+1 -2
View File
@@ -1,5 +1,4 @@
from .core import * from .core import *
from .lib import * from .lib import *
from .contrib import * from .contrib import *
from ._utils import * from .utils import *
from .process_pool import process_pool
+2
View File
@@ -3,3 +3,5 @@ from .candle_patterns import *
from .symbols import * from .symbols import *
from .utils import * from .utils import *
from .traders import * from .traders import *
from .quants import *
from .trackers import *
+173 -11
View File
@@ -1,15 +1,177 @@
from ...lib.candle import Candle, Candles from aiomql import Candle, Candles
from ..quants.change import percentage_difference
def find_bearish_fractal(candles: Candles) -> Candle | None: def is_bullish_fractal(candles: Candles) -> tuple[bool, Candle|None]:
"""Given a candles object, find the most recent bearish fractal.""" """Check if the candles form a five-candle bullish fractal"""
for i in range(len(candles) - 3, 1, -1): if len(candles) < 5:
if candles[i].high > max(candles[i - 1].high, candles[i + 1].high, candles[i - 2].high, candles[i + 2].high): return False, None
return candles[i]
# find the two candles by the left and right of the middle candle
# from left to right; oldest to most recent
first_left, second_left = candles[-5], candles[-4]
first_right, second_right = candles[-2], candles[-1]
center = candles[-3]
if not ((first_left.low > second_left.low > center.low < first_right.low < second_right.low)
and center.is_bearish() and first_right.is_bullish()):
return False, None
return True, center
def find_bullish_fractal(candles: Candles) -> Candle | None: def is_half_bullish_fractal(candles: Candles) -> tuple[bool, Candle|None]:
"""Given a candles object, find the most recent bullish fractal.""" """Check for a partially formed bullish fractal"""
for i in range(len(candles) - 3, 1, -1): if len(candles) < 4:
if candles[i].low < min(candles[i - 1].low, candles[i + 1].low, candles[i - 2].low, candles[i + 2].low): return False, None
return candles[i]
# find the two candles by the left and right of the current candle
# from left to right; oldest to most recent
first_left, second_left = candles[-4], candles[-3]
first_right = candles[-1]
center = candles[-3]
if not ((first_left.low > second_left.low > center.low < first_right.low) and center.is_bearish() and first_right.is_bullish()):
return False, None
return True, center
def is_half_bearish_fractal(candles: Candles) -> tuple[bool, Candle|None]:
if len(candles) < 4:
return False, None
# find the two candles by the left and right of the current candle
# from left to right; oldest to most recent
first_left, second_left = candles[-4], candles[-3]
first_right = candles[-1]
center = candles[-2]
if not ((first_left.high < second_left.high < center.high > first_right.high) and center.is_bullish() and first_right.is_bearish()):
return False, None
return True, center
def is_bearish_fractal(candles: Candles) -> tuple[bool, Candle|None]:
"""Check if the candles form a bearish fractal"""
if len(candles) < 5:
return False, None
# find the two candles by the left and right of the middle candle
# from left to right; oldest to most recent
first_left, second_left = candles[-5], candles[-4]
first_right, second_right = candles[-2], candles[-1]
center = candles[-3]
if not ((first_left.high < second_left.high < center.high > first_right.high > second_right.high) and
center.is_bullish() and first_right.is_bearish()):
return False, None
return True, center
def is_double_bullish_fractal(candles: Candles, tolerance: float = 1) -> tuple[bool, Candle|None]:
"""Check if the candles form a double bullish fractal"""
if len(candles) < 4:
return False, None
# find the two candles on the left and right
# from left to right; oldest to most recent
left, right = candles[-4], candles[-1]
center_left, center_right = candles[-3], candles[-2]
if not (percentage_difference(center_left.low, center_right.low) <= tolerance):
return False, None
if not (center_left.is_bearish() and center_right.is_bullish()):
return False, None
if not (left.low > center_left.low and center_right.low < right.low):
return False, None
return True, min(center_left, center_right, key=lambda c: c.low)
def is_double_bearish_fractal(candles: Candles, tolerance: float = 1) -> tuple[bool, Candle|None]:
"""Check if the candles form a double bearish fractal"""
if len(candles) < 4:
return False, None
# find the two candles on the left and right
# from left to right; oldest to most recent
left, right = candles[-4], candles[-1]
center_left, center_right = candles[-3], candles[-2]
if not (percentage_difference(center_left.high, center_right.high) <= tolerance):
return False, None
if not (center_left.is_bullish() and center_right.is_bearish()):
return False, None
if not (left.high < center_left.high and center_right.high > right.high):
return False, None
return True, max(center_left, center_right, key=lambda c: c.high)
def find_bullish_fractal(candles: Candles, swing_number: int = 1, min_price: float = None) -> tuple[Candle, Candles] | None:
"""Find a bullish fractal pattern"""
fractal_candles: Candles
swing_candle: Candle | None = None
min_price: float = min_price or candles[-1].low
for c in reversed(candles):
fractal_candles = candles[c.Index: c.Index + 5]
ok, candle = is_bullish_fractal(fractal_candles)
if ok and ((swing_candle is not None and candle.low < swing_candle.low) or swing_candle is None) and candle.low < min_price:
swing_number -= 1
swing_candle = candle
if swing_number <= 0:
return swing_candle, fractal_candles
continue
return None
def find_bearish_fractal(candles: Candles, swing_number: int = 1, max_price: float = None) -> tuple[Candle, Candles] | None:
"""Find a bearish fractal pattern"""
fractal_candles: Candles
swing_candle: Candle | None = None
max_price: float = max_price or candles[-1].high
for c in reversed(candles):
fractal_candles = candles[c.Index: c.Index + 5]
ok, candle = is_bearish_fractal(fractal_candles)
if (ok and ((swing_candle is not None and candle.high > swing_candle.high) or swing_candle is None) and
candle.high > max_price):
swing_number -= 1
swing_candle = candle
if swing_number <= 0:
return swing_candle, fractal_candles
continue
return None
def find_double_bearish_fractal(candles: Candles, swing_number: int = 1,
tolerance: float = 0.05, max_price: float = None) -> tuple[Candle, Candles] | None:
"""Find double bearish Fractals"""
fractal_candles: Candles
swing_candle: Candle | None = None
max_price: float = max_price or candles[-1].high
for c in reversed(candles):
fractal_candles = candles[c.Index: c.Index + 4]
ok, candle = is_double_bearish_fractal(fractal_candles, tolerance=tolerance)
if (ok and ((swing_candle is not None and candle.high > swing_candle.high) or swing_candle is None)
and candle.high > max_price):
swing_number -= 1
swing_candle = candle
if swing_number <= 0:
return swing_candle, fractal_candles
continue
return None
def find_double_bullish_fractal(candles: Candles, swing_number: int = 1,
tolerance: float = 0.05, min_price: float = None) -> tuple[Candle, Candles] | None:
"""Find double bullish Fractals"""
fractal_candles: Candles
swing_candle: Candle | None = None
min_price: float = min_price or candles[-1].low
for c in reversed(candles):
fractal_candles = candles[c.Index: c.Index + 4]
ok, candle = is_double_bullish_fractal(fractal_candles, tolerance=tolerance)
if (ok and ((swing_candle is not None and candle.low < swing_candle.low) or swing_candle is None)
and candle.low < min_price):
swing_number -= 1
swing_candle = candle
if swing_number <= 0:
return swing_candle, fractal_candles
continue
return None
+1
View File
@@ -0,0 +1 @@
from .change import *
+76
View File
@@ -0,0 +1,76 @@
def relative_difference(first: float, second: float) -> float:
"""Find the relative difference between two values"""
diff = abs(first - second)
div = (first + second) / 2
return (diff / div) * 100
def percentage_difference(first: float, second: float) -> float:
"""Find the percentage difference between two values"""
diff = abs(first - second)
div = (first + second) / 2
return (diff / div) * 100
def relative_position(start: float, end: float, value: float):
"""Find the relative position of a value between two values"""
return (value - start) / (end - start)
def percentage_position(start: float, end: float, value: float):
"""Find the percentage position of a value between two values"""
return ((value - start) / (end - start)) * 100
def get_relative_position(start: float, end: float, factor: float):
"""Find position within an interval based on a factor"""
span = end - start
position = start + (factor * span)
return position
def get_percentage_position(start: float, end: float, rate: float):
"""Find position within an interval based on a percentage"""
span = end - start
position = start + ((rate/100) * span)
return position
def extend_interval_by_factor(start: float, end: float, factor: float):
"""Extend the end of an interval by a factor"""
span = end - start
span *= factor
return end + span
def extend_interval_by_percentage(start: float, end: float, rate: float):
"""Extend the end of an interval by a factor"""
span = end - start
span *= (rate/100)
return end + span
def relative_change(start_value: float, end_value: float) -> float:
"""Find the relative change from one value to another"""
return (end_value - start_value) / start_value
def percentage_change(start_value: float, end_value: float) -> float:
"""Find the percentage change from one value to another"""
return ((end_value - start_value) / start_value) * 100
def relative_increase(value: float, factor: float) -> float:
"""Increase the value by a factor"""
return value * (1 + factor)
def percentage_increase(value: float, rate: float) -> float:
"""Increase the value by a factor"""
return value * ((100 + rate) / 100)
def relative_decrease(value: float, factor: float) -> float:
"""Increase the value by a factor"""
return value * (1 - factor)
def percentage_decrease(value: float, rate: float) -> float:
"""Increase the value by a factor"""
return value * ((100 - rate) / 100)
+1 -1
View File
@@ -8,7 +8,7 @@ from ...core.constants import TimeFrame, OrderType
from ...lib.sessions import Sessions from ...lib.sessions import Sessions
from ..candle_patterns import find_bearish_fractal, find_bullish_fractal from ..candle_patterns import find_bearish_fractal, find_bullish_fractal
from ..traders import SimpleTrader from ..traders import SimpleTrader
from ..utils.tracker import Tracker from ..utils import Tracker
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
+5
View File
@@ -0,0 +1,5 @@
from .position_tracker import PositionTracker
from .open_position import OpenPosition
from .positions_tracker import OpenPositionsTracker
from .positions_tracker import OpenPositionsTracker
from .position_tracking_functions import *
@@ -0,0 +1,164 @@
from dataclasses import dataclass, field
from typing import ClassVar
from logging import getLogger
from ...lib import Symbol, Positions, Order
from ...core.models import TradePosition, TradeAction, OrderSendResult
from ...core.constants import OrderType
from ...core.config import Config
from .position_tracker import PositionTracker
logger = getLogger()
@dataclass
class OpenPosition:
symbol: Symbol
ticket: int
position: TradePosition
is_open: bool = True
trackers: list[PositionTracker] = field(default_factory=list)
use_checkpoint: bool = False
checkpoint: float = 0
positions: ClassVar[Positions]
state_key: ClassVar[str] = "tracked_positions"
config: ClassVar[Config]
def __new__(cls, *args, **kwargs):
if not hasattr(cls, "config"):
cls.config = Config()
if not hasattr(cls, "positions"):
cls.positions = Positions()
return super().__new__(cls)
def __post_init__(self):
self.config.state.setdefault(self.state_key, {}).setdefault(self.ticket, self)
async def remove_closed(self):
try:
await self.update_position()
if not self.is_open:
self.remove_from_state()
except Exception as exe:
logger.error("%s: Unable to remove closed position from state", exe)
def add_tracker(self, *, tracker: PositionTracker, number: int = None):
number = number or len(self.trackers)
tracker.set_position(self, number)
self.trackers.append(tracker)
self.trackers.sort(key=lambda x: x.number)
async def update_position(self) -> bool:
pos = await self.positions.get_position_by_ticket(ticket=self.ticket)
if pos is not None:
self.position = pos
self.is_open = True
else:
self.is_open = False
return self.is_open
async def modify_stops(self, *, sl: float = None, tp: float = None,
use_stop_levels=False) -> tuple[bool, OrderSendResult | None]:
try:
tick = await self.symbol.info_tick()
# modify stop_loss
if sl is not None and use_stop_levels is True:
min_stops_value = (self.symbol.trade_stops_level + self.symbol.spread) * self.symbol.point
if self.position.type == OrderType.BUY:
sl = min(sl, tick.ask - min_stops_value)
elif self.position.type == OrderType.SELL:
sl = max(sl, tick.bid + min_stops_value)
else:
raise TypeError("Invalid OrderType %s: In %s.modify_stops", self.position.type, self.__class__.__name__)
elif sl is not None and use_stop_levels is False:
sl = sl or self.position.sl
# modify take_profit
if tp is not None and use_stop_levels is True:
min_stops_value = (self.symbol.trade_stops_level + self.symbol.spread) * self.symbol.point
if self.position.type == OrderType.BUY:
tp = max(tp, tick.ask + min_stops_value)
elif self.position.type == OrderType.SELL:
tp = min(tp, tick.bid - min_stops_value)
else:
raise TypeError("Invalid OrderType %s: In %s.modify_stops", self.position.type, self.__class__.__name__)
elif tp is not None and use_stop_levels is False:
tp = tp or self.position.tp
# send order
order = Order(position=self.ticket, sl=sl, tp=tp, action=TradeAction.SLTP)
res = await order.send()
if res and res.retcode == 10009:
await self.update_position()
return True, res
else:
return False, res
except Exception as exe:
logger.error("%s: Error occurred in %s.modify_stops for %s:%d",
exe, self.__class__.__name__, self.symbol.name, self.ticket)
return False, None
def remove_from_state(self):
try:
self.config.state.get(self.state_key, {}).pop(self.ticket, None)
except KeyError as err:
logger.error("%s: Unable to remove closed position from state in", err)
async def close_position(self, remove_from_state: bool = True) -> tuple[bool, OrderSendResult | None]:
try:
res = await self.positions.close_position(position=self.position)
if res.retcode == 10009:
self.is_open = False
if remove_from_state:
self.remove_from_state()
return True, res
else:
return False, res
except Exception as exe:
logger.error("%s: Error occurred in %s.close_position for %s:%d", exe, self.__class__.__name__,
self.symbol.name, self.ticket)
return False, None
async def track(self):
try:
for tracker in self.trackers:
await tracker()
except Exception as exe:
logger.error("%s: Error occurred in track method of Open Position for %d:%s",
exe, self.symbol.name, self.ticket)
async def hedge_position(self, *, hedge_params: dict = None) -> tuple[bool, OrderSendResult | None]:
try:
# hedge_params for customizing the hedge
hedge_params = hedge_params or {}
volume = hedge_params.get("volume", self.position.volume)
order_type = hedge_params.get("type") or OrderType(self.position.type).opposite
tick = await self.symbol.info_tick()
price = hedge_params.get("price") or (tick.ask if order_type == OrderType.BUY else tick.bid)
comment = hedge_params.get("comment", f"Hedge_{self.ticket}")
hedge_order = Order(type=order_type, symbol=self.symbol.name, volume=volume, price=price, comment=comment)
# if position has stops get them from the sl and tp of the position
if self.position.sl and self.position.tp:
osl = abs(self.position.sl - self.position.price_open)
otp = abs(self.position.tp - self.position.price_open)
if order_type == OrderType.BUY:
sl = hedge_params.get("sl") or tick.ask - osl
tp = hedge_params.get("tp") or tick.ask + otp
hedge_order.set_attributes(tp=tp, sl=sl)
if order_type == OrderType.SELL:
sl = hedge_params.get("sl") or tick.bid + osl
tp = hedge_params.get("sl") or tick.bid - otp
hedge_order.set_attributes(tp=tp, sl=sl)
res = await hedge_order.send()
if res.retcode == 10009:
return True, res
return False, res
except Exception as exe:
logger.error("%s: Error occurred in %s.hedge_position for %s:%d", exe, self.__class__.__name__,
self.symbol.name, self.ticket)
return False, None
@@ -0,0 +1,33 @@
from collections.abc import Callable
from typing import Any, TypeVar
from logging import getLogger
logger = getLogger(__name__)
OpenPosition = TypeVar('OpenPosition')
class PositionTracker:
params: dict[str, Any]
function: Callable
number: int
open_position: OpenPosition
def __init__(self, function: Callable, /, **kwargs) -> None:
self.function = function
self.kwargs = kwargs
self.number = 0
async def __call__(self, **kwargs):
try:
kwargs = self.kwargs if not kwargs else (self.kwargs | kwargs)
await self.function(self.open_position, **kwargs)
except Exception as exe:
logger.error("%s: Error occurred in %s for %s:%d", exe, self.function.__name__,
self.open_position.symbol.name, self.open_position.ticket)
def set_position(self, open_position: "OpenPosition", number: int = None):
self.open_position = open_position
if number is not None:
self.number = number
self.open_position = open_position
@@ -0,0 +1,86 @@
from logging import getLogger
from aiomql import OrderType
from .open_position import OpenPosition
from ..quants import extend_interval_by_percentage, get_percentage_position, percentage_position
logger = getLogger(__name__)
async def exit_at_price(*, pos: OpenPosition, tp: float = None, sl: float = None):
"""Exit a trade at a particular price"""
if not await pos.update_position():
return
if (tp is not None and pos.position.profit >= tp) or (sl is not None and pos.position.profit <= sl):
ok, res = await pos.close_position()
if not ok:
logger.warning("Unable to close %s:%d due to %s in exit_at_price", pos.symbol, pos.ticket, res.comment)
async def extend_take_profit(*, pos: OpenPosition, increase: float = 20, start: float = 80,
use_stop_levels: bool = True):
is_open = await pos.update_position()
if is_open is False or pos.position.profit < 0:
return
position = pos.position
if percentage_position(position.price_open, position.tp, position.price_current) >= start:
new_tp = extend_interval_by_percentage(position.price_open, position.tp, increase)
ok, res = await pos.modify_stops(tp=new_tp, use_stop_levels=use_stop_levels)
if ok:
logger.info("%s:%d take_profit extended by extend_take_profit", position.symbol, position.ticket)
else:
logger.warning("Unable to extend take profit of %s:%d due to %s in extend_take_profit",
position.symbol, position.ticket, res.comment)
async def extend_stop_loss(*, pos: OpenPosition, increase: float = 20, start: float = 80,
use_stop_levels: bool = True):
is_open = await pos.update_position()
if is_open is False or pos.position.profit > 0:
return
position = pos.position
if percentage_position(position.price_open, position.tp, position.price_current) >= start:
new_sl = extend_interval_by_percentage(position.price_open, position.sl, increase)
ok, res = await pos.modify_stops(sl=new_sl, use_stop_levels=use_stop_levels)
if ok:
logger.info("%s:%d extend_stop_loss extended by extend_stop_loss",
position.symbol, position.ticket)
else:
logger.warning("Unable to extend stop loss of %s:%d due to %s in extend_stop_loss",
position.symbol, position.ticket, res.comment)
async def exit_at_checkpoint(*, pos: OpenPosition, start: float = 80, trail: float = 15):
is_open = await pos.update_position()
if is_open is False or pos.position.profit < 0:
return
position = pos.position
if percentage_position(position.price_open, position.tp, position.price_current) >= start:
new_checkpoint = get_percentage_position(position.price_open, position.price_current, 100-trail)
change_checkpoint = False
if position.type == OrderType.BUY and new_checkpoint > pos.checkpoint:
change_checkpoint = True
elif position.type == OrderType.SELL and new_checkpoint < pos.checkpoint:
change_checkpoint = True
if change_checkpoint:
pos.checkpoint = new_checkpoint
pos.use_checkpoint = True
logger.info("New checkpoint created for %s:%d at %f:%f",
position.symbol, position.ticket, new_checkpoint, position.profit)
close = False
if position.type == OrderType.BUY and position.price_current <= pos.checkpoint and pos.use_checkpoint:
close = True
elif position.type == OrderType.SELL and position.price_current >= pos.checkpoint and pos.use_checkpoint:
close = True
if close:
ok, res = await pos.close_position()
if not ok:
logger.warning("Unable to close %s:%d due to %s in checkpoint", position.symbol,
position.ticket, res.comment)
else:
logger.info("Closed %s:%d in checkpoint at %f:%f",
position.symbol, position.ticket, position.price_current, position.profit)
@@ -0,0 +1,50 @@
import asyncio
from logging import getLogger
from typing import ClassVar
from ...core import Config, State, Store, sleep
from ...lib import Positions
logger = getLogger(__name__)
class OpenPositionsTracker:
state: ClassVar[State]
store: ClassVar[Store]
config: ClassVar[Config]
def __new__(cls, *args, **kwargs):
if not hasattr(cls, "config"):
cls.config = Config()
if not hasattr(cls, "state"):
cls.state = cls.config.state
if not hasattr(cls, "positions"):
cls.positions = Positions()
if not hasattr(cls, "store"):
cls.store = cls.config.store
return super().__new__(cls)
def __init__(self, interval: int = 10, state_key: str = "tracked_positions", autocommit=False):
self.interval = interval
self.state_key = state_key
self.autocommit = autocommit
async def track(self):
conn = self.config.state.conn
while self.config.shutdown is False:
try:
await sleep(self.interval)
tracked_positions = self.state.get("tracked_positions", {})
await asyncio.gather(*(pos.track() for pos in tracked_positions.values()))
await self.remove_closed_positions(tracked_positions)
if self.autocommit:
await self.state.acommit(conn=conn, close=False)
except Exception as exe:
logger.error("%s: Error occurred in main position tracker", exe)
conn.close()
async def remove_closed_positions(self, tracked_positions):
all_pos = await self.positions.get_positions()
all_pos = {pos.ticket for pos in all_pos}
self.state[self.state_key] = {pos.ticket: pos for pos in tracked_positions.values() if pos.ticket in all_pos}
+1 -1
View File
@@ -1 +1 @@
from .tracker import Tracker from .strategy_tracker import StrategyTracker as Tracker
@@ -1,3 +1,4 @@
from datetime import datetime
from dataclasses import dataclass from dataclasses import dataclass
from typing import Literal from typing import Literal
@@ -5,9 +6,8 @@ from ...core.constants import OrderType
@dataclass @dataclass
class Tracker: class StrategyTracker:
"""Keeps track of a strategy's data and state""" """Keeps track of a strategy's data and state"""
trend: Literal["ranging", "bullish", "bearish"] = "ranging" trend: Literal["ranging", "bullish", "bearish"] = "ranging"
bullish: bool = False bullish: bool = False
bearish: bool = False bearish: bool = False
@@ -22,6 +22,14 @@ class Tracker:
sl: float = 0 sl: float = 0
tp: float = 0 tp: float = 0
@property
def time(self):
return datetime.now()
@property
def timestamp(self):
return self.time.timestamp()
def update(self, **kwargs): def update(self, **kwargs):
"""Updates the tracker with the given kwargs""" """Updates the tracker with the given kwargs"""
fields = self.__dict__ fields = self.__dict__
+4
View File
@@ -8,3 +8,7 @@ from .errors import Error
from .exceptions import * from .exceptions import *
from .task_queue import TaskQueue from .task_queue import TaskQueue
from .backtesting import * from .backtesting import *
from .utils import *
from .db import DB
from .state import State
from .store import Store
@@ -38,7 +38,7 @@ from ..constants import (
CopyTicks, CopyTicks,
) )
from ..._utils import round_down, round_up, error_handler, error_handler_sync, async_cache from ...utils import round_down, round_up, error_handler, error_handler_sync, async_cache
from .get_data import BackTestData, GetData, Cursor from .get_data import BackTestData, GetData, Cursor
from .backtest_account import BackTestAccount from .backtest_account import BackTestAccount
+1 -1
View File
@@ -12,7 +12,7 @@ from ..meta_trader import MetaTrader
from ..config import Config from ..config import Config
from ..constants import TimeFrame from ..constants import TimeFrame
from ..task_queue import TaskQueue, QueueItem from ..task_queue import TaskQueue, QueueItem
from ..._utils import backoff_decorator from ...utils import backoff_decorator
logger = getLogger(__name__) logger = getLogger(__name__)
+4 -4
View File
@@ -19,7 +19,6 @@ class Base:
include (set [str]): A set of attributes to be included when retrieving attributes include (set [str]): A set of attributes to be included when retrieving attributes
using the get_dict and dict method. using the get_dict and dict method.
""" """
exclude: set[str] exclude: set[str]
include: set[str] include: set[str]
@@ -137,8 +136,9 @@ class _Base(Base):
"""Base class that provides access to the MetaTrader and Config classes as well as the MetaBackTester class for """Base class that provides access to the MetaTrader and Config classes as well as the MetaBackTester class for
backtesting mode. backtesting mode.
""" """
mt5: MetaTrader | MetaBackTester
config: Config
def __init__(self, **kwargs): def __init__(self, **kwargs):
self.config = Config() self.__class__.config = Config()
self.mt5 = MetaTrader() if self.config.mode != "backtest" else MetaBackTester() self.__class__.mt5 = MetaTrader() if self.config.mode != "backtest" else MetaBackTester()
super().__init__(**kwargs) super().__init__(**kwargs)
+59 -18
View File
@@ -1,10 +1,13 @@
import os import os
import json
from pathlib import Path from pathlib import Path
from typing import Literal, TypeVar, Self from typing import Literal, TypeVar, Self
import json
from logging import getLogger from logging import getLogger
from threading import Lock
from .task_queue import TaskQueue from .task_queue import TaskQueue
from .state import State
from .store import Store
logger = getLogger(__name__) logger = getLogger(__name__)
Bot = TypeVar("Bot") Bot = TypeVar("Bot")
@@ -14,14 +17,15 @@ BackTestController = TypeVar("BackTestController")
class Config: class Config:
login: int login: int
trade_record_mode: Literal["csv", "json"] trade_record_mode: Literal["csv", "json", "sql"]
password: str password: str
config_file: str | Path config_file: str | Path
server: str server: str
path: str | Path path: str | Path
timeout: int timeout: int
filename: str filename: str
state: dict _state: State
_store: Store
root: Path root: Path
record_trades: bool record_trades: bool
records_dir: Path records_dir: Path
@@ -30,6 +34,7 @@ class Config:
records_dir_name: str records_dir_name: str
plots_dir_name: str plots_dir_name: str
backtest_dir_name: str #Todo: add to docs backtest_dir_name: str #Todo: add to docs
db_name: str
task_queue: TaskQueue task_queue: TaskQueue
_backtest_engine: BackTestEngine _backtest_engine: BackTestEngine
bot: Bot bot: Bot
@@ -39,40 +44,46 @@ class Config:
use_terminal_for_backtesting: bool use_terminal_for_backtesting: bool
shutdown: bool shutdown: bool
force_shutdown: bool force_shutdown: bool
db_commit_interval: float
auto_commit: bool = False
lock: Lock
_defaults = { _defaults = {
"timeout": 60000, "timeout": 60000,
"record_trades": True, "record_trades": True,
"records_dir_name": "trade_records", "records_dir_name": "trade_records",
"backtest_dir_name": "backtesting", "backtest_dir_name": "backtesting",
"config_file": None, "config_file": None,
"trade_record_mode": "csv", "trade_record_mode": "sql",
"mode": "live", "mode": "live",
"filename": "aiomql.json", "filename": "aiomql.json",
"use_terminal_for_backtesting": True, "use_terminal_for_backtesting": True,
"db_name": "",
"path": "", "path": "",
"login": 0, "login": None,
"password": "", "password": "",
"server": "", "server": "",
"shutdown": False, "shutdown": False,
"force_shutdown": False, "force_shutdown": False,
"root": None, "root": None,
"plots_dir_name": "plots" "plots_dir_name": "plots",
"db_commit_interval": 30,
"auto_commit": False
} }
def __new__(cls, *args, **kwargs): def __new__(cls, *args, **kwargs):
if not hasattr(cls, "_instance"): with (lock := Lock()) as _:
cls._instance = super().__new__(cls) if not hasattr(cls, "_instance"):
cls._instance.state = {} cls._lock = lock
cls._instance.task_queue = TaskQueue(mode='infinite', workers=10) cls._instance = super().__new__(cls)
cls._instance.set_attributes(**cls._defaults) cls._instance.task_queue = TaskQueue(mode='infinite', workers=10)
cls._instance._backtest_engine = None cls._instance.set_attributes(**cls._defaults)
cls._instance.bot = None cls._instance._backtest_engine = None
cls._instance.backtest_controller = None cls._instance.bot = None
cls._instance.backtest_controller = None
return cls._instance return cls._instance
def __init__(self, **kwargs): def __init__(self, **kwargs):
"""Initialize the Config object. The root directory can be set here or in the load_config method.""" """Initialize the Config object. The root directory can be set here or in the load_config method."""
root = kwargs.pop("root", None) root = kwargs.pop("root", None)
config_file = kwargs.pop("config_file", None) config_file = kwargs.pop("config_file", None)
if self.root is None or root is not None or config_file is not None: if self.root is None or root is not None or config_file is not None:
@@ -123,7 +134,6 @@ class Config:
config_file = dirname / self.filename config_file = dirname / self.filename
if config_file.exists(): if config_file.exists():
return config_file return config_file
if self.root == dirname: if self.root == dirname:
break break
return None return None
@@ -176,11 +186,15 @@ class Config:
else: else:
fh = open(self.config_file, mode="r") fh = open(self.config_file, mode="r")
file_config = json.load(fh) file_config = json.load(fh)
# print(file_config)
fh.close() fh.close()
data = file_config | kwargs data = file_config | kwargs
self.set_attributes(**data) self.set_attributes(**data)
self.db_name = self.db_name or (f"db_{self.login}.sqlite3" if self.login else "db.sqlite3")
os.environ["DB_NAME"] = self.db_name
self.init_state()
self.init_store()
try: try:
if self.path: if self.path:
self.path = self.root / self.path if not Path(self.path).resolve().exists() else self.path self.path = self.root / self.path if not Path(self.path).resolve().exists() else self.path
@@ -189,6 +203,32 @@ class Config:
self.path = "" self.path = ""
return self return self
@property
def state(self):
if not hasattr(self, "_state"):
self.init_state()
return self._state
@state.setter
def state(self, value: State):
self._state = value
@property
def store(self):
if not hasattr(self, "_store"):
self.init_store()
return self._store
@store.setter
def store(self, value: Store):
self._store = value
def init_state(self):
self.state = State(db_name=self.db_name)
def init_store(self):
self.store = Store(db_name=self.db_name)
@property @property
def records_dir(self): def records_dir(self):
rec_dir = self.root / self.records_dir_name or 'trade_records' rec_dir = self.root / self.records_dir_name or 'trade_records'
@@ -207,6 +247,7 @@ class Config:
p_dir.mkdir(parents=True, exist_ok=True) if p_dir.exists() is False else ... p_dir.mkdir(parents=True, exist_ok=True) if p_dir.exists() is False else ...
return p_dir return p_dir
@property
def account_info(self) -> dict[str, int | str]: def account_info(self) -> dict[str, int | str]:
"""Returns Account login details as found in the config object if available """Returns Account login details as found in the config object if available
@@ -226,7 +267,7 @@ Config.__doc__ = """A class for handling configuration settings for the aiomql p
timeout (int): The timeout argument for the terminal timeout (int): The timeout argument for the terminal
filename (str): The filename of the config file filename (str): The filename of the config file
config_file (Path): The config file path config_file (Path): The config file path
state (dict): The state (State): A key-value database
root (Path): The root directory of the project root (Path): The root directory of the project
record_trades (bool): To record trades or not. Default is True record_trades (bool): To record trades or not. Default is True
records_dir (Path): The directory to store trade records, relative to the root directory records_dir (Path): The directory to store trade records, relative to the root directory
+167
View File
@@ -0,0 +1,167 @@
import os
import sqlite3
import re
from logging import getLogger
from dataclasses import Field, fields, asdict, MISSING, is_dataclass
from typing import ClassVar
from .config import Config
logger = getLogger(__name__)
class DB:
_table: ClassVar[str] = ""
TYPES: ClassVar[dict] = {str: "TEXT", float: "REAL", int: "INTEGER", bool: "BOOLEAN", None: "NULL", bytes: "BLOB"}
config: ClassVar[Config]
conn: sqlite3.Connection
cursor: sqlite3.Cursor
def __new__(cls, *args, **kwargs):
cls.config = Config()
return super().__new__(cls)
def __post_init__(self):
self.init_db()
def init_db(self):
db_name = getattr(self.config, "db_name", os.getenv("DB_NAME", "db.sqlite3"))
self.conn = sqlite3.connect(db_name)
self.conn.row_factory = self.dict_factory()
self.cursor = self.conn.cursor()
self.create_table(self.conn)
@classmethod
def get_connection(cls):
db_name = os.getenv("DB_NAME", "db.sqlite3")
conn = sqlite3.connect(db_name)
conn.row_factory = cls.dict_factory()
return conn
@classmethod
def create_table(cls, conn: sqlite3.Connection):
try:
if not is_dataclass(cls):
return
cls._table = cls._table or cls.__name__.lower()
columns = cls.get_columns()
q = f"""CREATE TABLE IF NOT EXISTS '{cls._table}' ({columns})"""
conn.execute(f"""CREATE TABLE IF NOT EXISTS '{cls._table}' ({columns})""")
conn.commit()
except Exception as e:
logger.error("%s: Failed to create table", e)
@classmethod
def dict_factory(cls):
def dict_factory(cursor, row):
cols = [column[0] for column in cursor.description]
kw = {key: value for key, value in zip(cols, row)}
print(type(cls), cls.__name__)
return cls(**kw)
return dict_factory
@staticmethod
def sanitize(identifier):
if re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]*", identifier):
return f'"{identifier}"'
raise ValueError("Invalid table name")
@classmethod
def types(cls, key):
return cls.TYPES.get(key, "TEXT")
@staticmethod
def get_default(col: Field):
if not isinstance(col.default, type(MISSING)):
return f"DEFAULT {col.default}"
elif not isinstance(col.default_factory, type(MISSING)):
return f"DEFAULT {col.default_factory()}"
else:
return ""
@staticmethod
def get_metadata(col: Field):
return f"{' '.join(meta for meta, value in col.metadata.items() if value)}"
@classmethod
def get_columns(cls):
cols = fields(cls)
cols = [f"'{col.name}' {cls.types(col.type)} {cls.get_metadata(col)} {cls.get_default(col)}" for col in
cols]
return ",".join(cols)
def commit(self):
self.conn.commit()
def close(self):
self.conn.close()
def asdict(self):
return asdict(self)
def get_data(self):
return self.asdict()
@classmethod
def clear(cls):
conn = cls.get_connection()
conn.execute(f"DELETE FROM {cls._table}")
conn.commit()
conn.close()
def save(self, commit=True):
data = self.get_data()
columns = ", ".join(self.sanitize(key) for key in data.keys())
placeholders = ", ".join(["?"] * len(data))
values = tuple(data.values())
table = self.sanitize(self._table)
query = f"INSERT INTO {table} ({columns}) VALUES ({placeholders})"
self.cursor.execute(query, values)
if commit:
self.commit()
@classmethod
def get(cls, **kwargs):
conn = cls.get_connection()
_query = " AND ".join([f'"{key}" = "{value}"' for key, value in kwargs.items()])
query = f"""SELECT * FROM {cls.sanitize(cls._table)} WHERE {_query}"""
res = conn.execute(query).fetchone()
conn.close()
return res
@classmethod
def filter(cls, **kwargs):
conn = cls.get_connection()
_query = " AND ".join([f'"{key}" = "{value}"' for key, value in kwargs.items()])
# _query = " AND ".join([f'{key} = {value}' for key, value in kwargs.items()])
_query = f"WHERE {_query}" if _query else ""
query = f"SELECT * FROM {cls.sanitize(cls._table)} {_query}"
res = conn.execute(query).fetchall()
conn.close()
return res
@classmethod
def update(cls, data=None, /, **kwargs):
conn = cls.get_connection()
update = data or {}
_query = " AND ".join([f'"{key}" = "{value}"' for key, value in kwargs.items()])
update = ", ".join([f'"{key}" = "{value}"' for key, value in update.items()])
table = cls.sanitize(cls._table)
_query = f"WHERE {_query}" if _query else ""
query = f"""UPDATE {table} SET {update} {_query}"""
conn.execute(query)
conn.commit()
conn.close()
return True
@classmethod
def fields(cls) -> list[str]:
fs = fields(cls)
return [f.name for f in fs]
@classmethod
def drop_table(cls):
conn = cls.get_connection()
conn.execute(f"DROP TABLE IF EXISTS {cls._table}")
conn.commit()
conn.close()
+3 -3
View File
@@ -3,7 +3,7 @@ class Error:
descriptions = { descriptions = {
# common errors # common errors
1: "Successful", 1: "successful",
-1: "generic fail", -1: "generic fail",
-2: "invalid arguments/parameters", -2: "invalid arguments/parameters",
-3: "no memory condition", -3: "no memory condition",
@@ -23,9 +23,9 @@ class Error:
conn_errors = (-10000, -10001, -10002, -10003, -10004, -10005, -6) conn_errors = (-10000, -10001, -10002, -10003, -10004, -10005, -6)
def __init__(self, code: int, description: str = ""): def __init__(self, code: int = 1, description: str = ""):
self.code = code self.code = code
self.description = description or self.descriptions.get(code, "unknown error") self.description = self.descriptions.get(code, description or "unknown error")
def is_connection_error(self): def is_connection_error(self):
return self.code in self.conn_errors return self.code in self.conn_errors
+1 -1
View File
@@ -17,7 +17,7 @@ from MetaTrader5 import (
from .meta_trader import MetaTrader from .meta_trader import MetaTrader
from .constants import TimeFrame, CopyTicks, OrderType from .constants import TimeFrame, CopyTicks, OrderType
from .._utils import error_handler from ..utils import error_handler
logger = getLogger(__name__) logger = getLogger(__name__)
+53 -90
View File
@@ -30,7 +30,8 @@ logger = getLogger()
class MetaTrader(MetaCore): class MetaTrader(MetaCore):
def __init__(self): def __init__(self):
self.config = Config() self.config = Config()
self.error: Error = Error(1) self.error: Error = Error(code=1)
self._lock = asyncio.Lock()
async def __aenter__(self) -> Self: async def __aenter__(self) -> Self:
""" """
@@ -50,28 +51,24 @@ class MetaTrader(MetaCore):
""" """
await self.shutdown() await self.shutdown()
async def _handler(self, api: dict): async def _handler(self, api: dict, retries=3):
func = api["func"] func = api["func"]
args = api.get("args", ()) args = api.get("args", ())
kwargs = api.get("kwargs", {}) kwargs = api.get("kwargs", {})
error_msg = api.get("error_msg", "An error occurred") error_msg = api.get("error_msg", "An error occurred")
res = await asyncio.to_thread(func, *args, **kwargs) res = await asyncio.to_thread(func, *args, **kwargs)
if res is None:
if res is not None:
return res
if res is None and self.error.is_connection_error() and retries > 0:
await self.initialize()
await self.login()
return await self._handler(api, retries=retries - 1)
else:
err = await self.last_error() err = await self.last_error()
self.error = Error(*err) self.error = Error(*err)
logger.warning(f"{error_msg}:{self.error.description}")
if self.error.is_connection_error():
await self.initialize()
await self.login()
res = await asyncio.to_thread(func, *args, **kwargs)
if res is None:
err = await self.last_error()
self.error = Error(*err)
logger.warning(f"{error_msg}:{self.error.description}")
else:
logger.warning(f"{error_msg}:{self.error.description}")
return res return res
async def login(self, *, login: int = 0, password: str = "", server: str = "", timeout: int = 60000) -> bool: async def login(self, *, login: int = 0, password: str = "", server: str = "", timeout: int = 60000) -> bool:
@@ -87,7 +84,7 @@ class MetaTrader(MetaCore):
Returns: Returns:
bool: True if successful, False otherwise. bool: True if successful, False otherwise.
""" """
acc_details = self.config.account_info() acc_details = self.config.account_info
login = login or acc_details.get("login", 0) login = login or acc_details.get("login", 0)
password = password or acc_details.get("password", "") password = password or acc_details.get("password", "")
server = server or acc_details.get("server", "") server = server or acc_details.get("server", "")
@@ -106,70 +103,15 @@ class MetaTrader(MetaCore):
Returns: Returns:
bool: True if successful, False otherwise. bool: True if successful, False otherwise.
""" """
acc_details = self.config.account_info() acc_details = self.config.account_info
login = login or acc_details.get("login", 0) login = login or acc_details.get("login", 0)
password = password or acc_details.get("password", "") password = password or acc_details.get("password", "")
server = server or acc_details.get("server", "") server = server or acc_details.get("server", "")
res = self._login(login, password=password, server=server, timeout=timeout) res = self._login(login, password=password, server=server, timeout=timeout)
return res return res
async def initialize( async def initialize(self, path: str = None, login: int = 0, password: str = "", server: str = "",
self, timeout: int | None = None, portable=False) -> bool:
path: str = None,
login: int = 0,
password: str = "",
server: str = "",
timeout: int | None = None,
portable=False,
) -> bool:
"""
Initializes the connection to the MetaTrader terminal. All parameters are optional.
Keyword Args:
path (str): The path to the MetaTrader terminal executable.
login (int): The trading account number.
password (str): The trading account password.
server (str): The trading server name.
timeout (int): The timeout for the connection in milliseconds.
portable (bool): If True, the terminal will be launched in portable mode.
Returns:
bool: True if successful, False otherwise.
"""
async with asyncio.Lock() as _:
path = self.config.path if path is None else path
path = "" if Path(path).exists() is False else path
args = (str(path),) if path else ()
acc = self.config.account_info()
kwargs = {
key: value
for key, value in (
("login", login or acc.get("login")),
("password", password or acc.get("password")),
("server", server or acc.get("server")),
("timeout", timeout or 60000),
("portable", portable),
)
if key is not None
}
res = await asyncio.to_thread(self._initialize, *args, **kwargs)
if res is False:
await self.shutdown()
res = await asyncio.to_thread(self._initialize, *args, **kwargs)
if not res:
err = await self.last_error()
self.error = Error(*err)
return res
def initialize_sync(
self,
path: str = None,
login: int = 0,
password: str = "",
server: str = "",
timeout: int | None = None,
portable=False,
) -> bool:
""" """
Initializes the connection to the MetaTrader terminal. All parameters are optional. Initializes the connection to the MetaTrader terminal. All parameters are optional.
@@ -187,21 +129,42 @@ class MetaTrader(MetaCore):
path = self.config.path if path is None else path path = self.config.path if path is None else path
path = "" if Path(path).exists() is False else path path = "" if Path(path).exists() is False else path
args = (str(path),) if path else () args = (str(path),) if path else ()
acc = self.config.account_info() acc = self.config.account_info
kwargs = { kwargs = {key: value for key, value in (("login", login or acc.get("login")),
key: value ("password", password or acc.get("password")),
for key, value in ( ("server", server or acc.get("server")), ("timeout", timeout or 45000),
("login", login or acc.get("login")), ("portable", portable)) if key is not None}
("password", password or acc.get("password")), res = await asyncio.to_thread(self._initialize, *args, **kwargs)
("server", server or acc.get("server")), if not res:
("timeout", timeout or 60000), err = await self.last_error()
("portable", portable), self.error = Error(*err)
) return res
if key is not None
} def initialize_sync(self, path: str = None, login: int = 0, password: str = "", server: str = "",
timeout: int | None = None, portable=False) -> bool:
"""
Initializes the connection to the MetaTrader terminal. All parameters are optional.
Keyword Args:
path (str): The path to the MetaTrader terminal executable.
login (int): The trading account number.
password (str): The trading account password.
server (str): The trading server name.
timeout (int): The timeout for the connection in milliseconds.
portable (bool): If True, the terminal will be launched in portable mode.
Returns:
bool: True if successful, False otherwise.
"""
path = self.config.path if path is None else path
path = "" if Path(path).exists() is False else path
args = (str(path),) if path else ()
acc = self.config.account_info
kwargs = {key: value for key, value in (("login", login or acc.get("login")),
("password", password or acc.get("password")),
("server", server or acc.get("server")), ("timeout", timeout or 45000),
("portable", portable)) if key is not None}
res = self._initialize(*args, **kwargs) res = self._initialize(*args, **kwargs)
if res is False:
self._shutdown()
if not res: if not res:
err = self._last_error() err = self._last_error()
self.error = Error(*err) self.error = Error(*err)
@@ -216,7 +179,7 @@ class MetaTrader(MetaCore):
res = await asyncio.to_thread(self._last_error) res = await asyncio.to_thread(self._last_error)
return res return res
except Exception as err: except Exception as err:
logger.warning(f"Error in obtaining last error.") logger.warning("%s: Error in obtaining last error.", err)
return -1, str(err) return -1, str(err)
async def version(self) -> tuple[int, int, str] | None: async def version(self) -> tuple[int, int, str] | None:
-3
View File
@@ -412,7 +412,6 @@ class TradeOrder(Base):
comment: str comment: str
external_id: str external_id: str
""" """
ticket: int ticket: int
time_setup: int time_setup: int
time_setup_msc: int time_setup_msc: int
@@ -529,7 +528,6 @@ class OrderSendResult(Base):
request_id: int request_id: int
retcode_external: int retcode_external: int
""" """
retcode: int retcode: int
deal: int deal: int
order: int order: int
@@ -569,7 +567,6 @@ class TradePosition(Base):
comment: str comment: str
external_id: str external_id: str
""" """
ticket: int ticket: int
time: int time: int
time_msc: int time_msc: int
+150
View File
@@ -0,0 +1,150 @@
import os
import pickle
import sqlite3
from threading import Lock
from typing import Self
from typing import MutableMapping, Iterable, Any, ClassVar
from logging import getLogger
SENTINEL = object()
logger = getLogger(__name__)
sqlite3.register_converter("pickle", pickle.loads)
sqlite3.register_adapter(dict, pickle.dumps)
class State(MutableMapping):
_data: ClassVar
_instance: Self
_lock: Lock
db_name: str
autocommit: bool
_initialized = False
def __new__(cls, *args, **kwargs):
with (lock := Lock()) as _:
if not hasattr(cls, '_instance'):
cls._data = {}
cls._lock = lock
cls._instance = super().__new__(cls)
return cls._instance
def __init__(self, db_name: str = "", data: dict = None, flush: bool = False, autocommit: bool = False):
with self._lock:
self.autocommit = autocommit
self.init(data=data, flush=flush, db_name=db_name)
@property
def data(self):
return self.__class__._data
@data.setter
def data(self, value):
assert isinstance(value, dict)
self.__class__._data = value
def __repr__(self):
return repr(self.data)
def __iter__(self):
return iter(self.data)
def __len__(self):
return len(self.data)
def __contains__(self, key):
return key in self.data
def __setitem__(self, key, value):
self.data[key] = value
if self.autocommit:
self.commit()
def __getitem__(self, key):
value = self.data[key]
return value
def __delitem__(self, key):
del self.data[key]
if self.autocommit:
self.commit()
def pop(self, key, default=SENTINEL):
if default is SENTINEL:
value = self.data.pop(key)
else:
value = self.data.pop(key, default)
if self.autocommit:
self.commit()
return value
def get(self, key, default=None):
return self.data.get(key, default)
def update(self, data: MutableMapping | Iterable[Iterable[Any]] = None, /, **kwargs):
self.data.update(data, **kwargs)
if self.autocommit:
self.commit()
def setdefault(self, key, default = None, /):
value = self.data.setdefault(key, default)
if self.autocommit:
self.commit()
return value
def keys(self):
return self.data.keys()
def values(self):
return self.data.values()
def items(self):
return self.data.items()
def load(self, *, conn = None, db_name: str = "", data: dict = None):
try:
db_name = db_name or self.db_name
conn = conn or self.conn
res = conn.execute("SELECT value FROM state where key = 'data'").fetchone()
db_data = pickle.loads(res[0]) if res else {}
db_data |= (data or {})
self.update(db_data)
except Exception as err:
logger.error("%s: Failed to load data from database", err)
def init(self, data: dict = None, flush: bool = False, db_name: str = ""):
if not self._initialized:
self.db_name = db_name or os.environ.get("DB_NAME", "db.sqlite3")
conn = self.conn
conn.execute("CREATE TABLE IF NOT EXISTS state (key text unique, value blob)")
self._initialized = True
if flush:
self.data = data or {}
self.commit(conn=conn)
else:
self.load(conn=conn, data=data)
conn.close()
if flush:
self.data = data or {}
self.commit()
def commit(self, *, conn: sqlite3.Connection = None, close: bool = True):
value = pickle.dumps(self.data, protocol=pickle.HIGHEST_PROTOCOL)
value = sqlite3.Binary(value)
conn = conn or self.conn
conn.execute("REPLACE INTO state (key, value) VALUES ('data', ?)", (value,))
conn.commit()
if close:
conn.close()
@property
def conn(self):
return sqlite3.connect(self.db_name)
async def acommit(self, conn=None, close=True):
conn = conn or self.conn
value = pickle.dumps(self.data, protocol=pickle.HIGHEST_PROTOCOL)
value = sqlite3.Binary(value)
conn.execute("REPLACE INTO state (key, value) VALUES ('data', ?)", (value,))
conn.commit()
if close:
conn.close()
+128
View File
@@ -0,0 +1,128 @@
import os
import sqlite3
from typing import MutableMapping, Iterable, Any
from logging import getLogger
SENTINEL = object()
logger = getLogger(__name__)
class Store(MutableMapping):
autocommit: bool
db_name: str
cursor: sqlite3.Cursor
conn: sqlite3.Connection
def __init__(self, db_name: str = "", data: dict = None, flush: bool = False, autocommit: bool = True):
self.autocommit = autocommit
self.db_name = db_name or os.environ.get("DB_NAME", "db.sqlite3")
self.conn = sqlite3.connect(self.db_name, check_same_thread=False)
self.cursor = self.conn.cursor()
self.conn.execute("CREATE TABLE IF NOT EXISTS store (key unique, value)")
if flush:
self.conn.execute("DELETE FROM store")
self.commit()
if data:
self.conn.executemany("REPLACE INTO store VALUES(?, ?)", data.items())
self.commit()
def __len__(self):
rows = self.cursor.execute('SELECT COUNT(*) FROM store').fetchone()[0]
return rows if rows is not None else 0
def __contains__(self, key):
return self.cursor.execute('SELECT 1 FROM store WHERE key = ?', (key,)).fetchone() is not None
def __getitem__(self, key):
item = self.cursor.execute('SELECT value FROM store WHERE key = ?', (key,)).fetchone()
if item is None:
raise KeyError(key)
return item[0]
def __setitem__(self, key, value):
self.cursor.execute('REPLACE INTO store (key, value) VALUES (?,?)', (key, value))
if self.autocommit:
self.conn.commit()
def __delitem__(self, key):
if key not in self:
raise KeyError(key)
self.cursor.execute('DELETE FROM store WHERE key = ?', (key,))
if self.autocommit:
self.commit()
def __iter__(self):
return self.iterkeys()
def __repr__(self):
return f"{self.__class__.__name__}()"
def iterkeys(self):
for row in self.cursor.execute('SELECT key FROM store'):
yield row[0]
def itervalues(self):
for row in self.cursor.execute('SELECT value FROM store'):
yield row[0]
def iteritems(self):
for row in self.cursor.execute('SELECT key, value FROM store'):
yield row[0], row[1]
def keys(self):
return list(self.iterkeys())
def values(self):
return list(self.itervalues())
def items(self):
return list(self.iteritems())
def update(self, data: MutableMapping | Iterable[Iterable[Any]] = None, /, **kwargs):
data = (dict(data if data is not None else {}) or {}) | kwargs
self.cursor.executemany("REPLACE INTO store VALUES(?, ?)", data.items())
if self.autocommit:
self.commit()
def setdefault(self, key, default = None, /):
if key in self:
return self[key]
self[key] = default
return self[key]
def get(self, key, /, default=None):
try:
value = self[key]
return value
except KeyError:
return default
def clear(self):
self.cursor.execute("DELETE FROM store")
if self.autocommit:
self.commit()
def pop(self, key, /, default=SENTINEL):
try:
value = self[key]
del self[key]
return value
except KeyError as err:
if default is SENTINEL:
raise err
return default
def commit(self, conn: sqlite3.Connection = None, close: bool = False):
conn = conn or self.conn
conn.commit()
if close:
conn.close()
async def acommit(self, conn: sqlite3.Connection = None, close: bool = False):
self.commit(conn=conn, close=close)
@property
def data(self):
res = self.cursor.execute("SELECT * FROM store").fetchall()
return {key: value for key, value in res}
+11 -3
View File
@@ -22,12 +22,12 @@ class QueueItem:
- `time` (int): The time the task was added to the queue. - `time` (int): The time the task was added to the queue.
""" """
def __init__(self, task_item: Callable | Coroutine, *args, **kwargs): def __init__(self, task_item: Callable | Coroutine, *args, on_separate_thread=True, **kwargs):
self.task_item = task_item self.task_item = task_item
self.args = args self.args = args
self.kwargs = kwargs self.kwargs = kwargs
self.must_complete = False
self.time = time.time_ns() self.time = time.time_ns()
self.on_separate_thread = on_separate_thread
def __hash__(self): def __hash__(self):
return self.time return self.time
@@ -45,14 +45,18 @@ class QueueItem:
try: try:
if asyncio.iscoroutinefunction(self.task_item): if asyncio.iscoroutinefunction(self.task_item):
return await self.task_item(*self.args, **self.kwargs) return await self.task_item(*self.args, **self.kwargs)
else: elif self.on_separate_thread:
return await asyncio.to_thread(self.task_item, *self.args, **self.kwargs) return await asyncio.to_thread(self.task_item, *self.args, **self.kwargs)
else:
return self.task_item(*self.args, **self.kwargs)
except asyncio.CancelledError: except asyncio.CancelledError:
logger.debug("Task %s with args %s and %s was cancelled", logger.debug("Task %s with args %s and %s was cancelled",
self.task_item.__name__, self.args, self.kwargs) self.task_item.__name__, self.args, self.kwargs)
return None
except Exception as err: except Exception as err:
logger.error("Error %s occurred in %s with args %s and %s", logger.error("Error %s occurred in %s with args %s and %s",
err, self.task_item.__name__, self.args, self.kwargs) err, self.task_item.__name__, self.args, self.kwargs)
return None
class TaskQueue: class TaskQueue:
@@ -75,6 +79,10 @@ class TaskQueue:
self.start_time = time.perf_counter() self.start_time = time.perf_counter()
signal(SIGINT, self.sigint_handle) signal(SIGINT, self.sigint_handle)
def add_task(self, task_item: Callable | Coroutine, *args, on_separate_thread=True, must_complete=True, priority=3, **kwargs):
task_item = QueueItem(task_item, *args, on_separate_thread=on_separate_thread, **kwargs)
self.add(item=task_item, priority=priority, must_complete=must_complete)
def add(self, *, item: QueueItem, priority=3, must_complete=False): def add(self, *, item: QueueItem, priority=3, must_complete=False):
"""Add a task to the queue. """Add a task to the queue.
+45
View File
@@ -0,0 +1,45 @@
import asyncio
import time
from logging import getLogger
from .config import Config
logger = getLogger(__name__)
async def auto_commit():
config = Config()
try:
with config.state.conn as conn:
while config.shutdown is False:
await config.state.acommit(conn=conn, close=False)
await sleep(config.db_commit_interval)
except Exception as err:
logger.error("%s: Error occurred in auto_commit", err)
async def backtest_sleep(secs):
"""An async sleep function for use during backtesting."""
config = Config()
secs = config.backtest_engine.cursor.time + secs
while secs > config.backtest_engine.cursor.time:
await asyncio.sleep(0)
async def sleep(secs):
if Config.mode == "backtest":
await backtest_sleep(secs)
else:
await asyncio.sleep(secs)
def sleep_sync(secs):
if Config.mode == "backtest":
backtest_sleep_sync(secs)
else:
time.sleep(secs)
def backtest_sleep_sync(secs):
"""A sleep function for use during backtesting."""
config = Config()
secs = config.backtest_engine.cursor.time + secs
while secs > config.backtest_engine.cursor.time:
time.sleep(0)
+1
View File
@@ -15,3 +15,4 @@ from .sessions import Sessions, Session
from .trade_records import TradeRecords from .trade_records import TradeRecords
from .terminal import Terminal from .terminal import Terminal
from .backtester import BackTester from .backtester import BackTester
from .result_db import ResultDB
+8 -5
View File
@@ -1,3 +1,4 @@
from threading import Lock
from logging import getLogger from logging import getLogger
from typing import Self from typing import Self
@@ -15,14 +16,16 @@ class Account(_Base, AccountInfo):
Attributes: Attributes:
connected (bool): Status of connection to MetaTrader 5 Terminal connected (bool): Status of connection to MetaTrader 5 Terminal
""" """
_instance: Self _instance: Self
_lock: Lock
connected: bool connected: bool
def __new__(cls, *args, **kwargs): def __new__(cls, *args, **kwargs):
if not hasattr(cls, "_instance"): with (lock := Lock()) as _:
cls._instance = super().__new__(cls) if not hasattr(cls, "_instance"):
cls._instance.connected = False cls._lock = lock
cls._instance = super().__new__(cls)
cls._instance.connected = False
return cls._instance return cls._instance
async def __aenter__(self) -> Self: async def __aenter__(self) -> Self:
@@ -48,7 +51,7 @@ class Account(_Base, AccountInfo):
async def refresh(self): async def refresh(self):
"""Refreshes the account instance with the latest account details from the MetaTrader 5 terminal""" """Refreshes the account instance with the latest account details from the MetaTrader 5 terminal"""
account_info = await self.mt5.account_info() account_info = await self.mt5.account_info
acc = account_info._asdict() acc = account_info._asdict()
self.connected = True self.connected = True
self.set_attributes(**acc) self.set_attributes(**acc)
+1 -1
View File
@@ -34,7 +34,7 @@ class Bot:
@classmethod @classmethod
def process_pool(cls, processes: dict[Callable:dict] = None, num_workers: int = None): def process_pool(cls, processes: dict[Callable:dict] = None, num_workers: int = None):
"""Run multiple processes in parallel using a ProcessPoolExecutor. Each bot should be a callable that accepts """Run multiple processes in parallel using a ProcessPoolExecutor. Each process should be a callable that accepts
keyword arguments only. keyword arguments only.
Args: Args:
+1 -4
View File
@@ -8,7 +8,6 @@ from ..core.config import Config
from ..core.meta_trader import MetaTrader from ..core.meta_trader import MetaTrader
from ..core.models import TradeDeal, TradeOrder from ..core.models import TradeDeal, TradeOrder
from ..core.meta_backtester import MetaBackTester from ..core.meta_backtester import MetaBackTester
from .._utils import backoff_decorator
logger = getLogger(__name__) logger = getLogger(__name__)
@@ -45,7 +44,7 @@ class History:
date_to (datetime, float): Date up to which the orders are requested. Set by the 'datetime' object or as a date_to (datetime, float): Date up to which the orders are requested. Set by the 'datetime' object or as a
number of seconds elapsed since 1970.01.01. Defaults to the current time in "utc" number of seconds elapsed since 1970.01.01. Defaults to the current time in "utc"
group (str): Filter for selecting history by symbols. Defaults to an empty string group (str): Filter for selecting history by symbols. This defaults to an empty string
""" """
self.config = Config() self.config = Config()
self.mt5 = MetaTrader() if self.config.mode != "backtest" else MetaBackTester() self.mt5 = MetaTrader() if self.config.mode != "backtest" else MetaBackTester()
@@ -67,7 +66,6 @@ class History:
self.total_deals = len(self.deals) self.total_deals = len(self.deals)
self.total_orders = len(self.orders) self.total_orders = len(self.orders)
@backoff_decorator
async def get_deals(self) -> tuple[TradeDeal, ...]: async def get_deals(self) -> tuple[TradeDeal, ...]:
"""Get deals from trading history using the parameters set in the constructor. """Get deals from trading history using the parameters set in the constructor.
@@ -103,7 +101,6 @@ class History:
""" """
return tuple(sorted((deal for deal in self.deals if deal.position_id == position), key=lambda x: x.time_msc)) return tuple(sorted((deal for deal in self.deals if deal.position_id == position), key=lambda x: x.time_msc))
@backoff_decorator
async def get_orders(self) -> tuple[TradeOrder, ...]: async def get_orders(self) -> tuple[TradeOrder, ...]:
"""Get orders from trading history using the parameters set in the constructor or the method arguments. """Get orders from trading history using the parameters set in the constructor or the method arguments.
+1 -2
View File
@@ -4,7 +4,7 @@ from ..core.models import TradeRequest, TradeOrder, OrderCheckResult, OrderSendR
from ..core.constants import TradeAction, OrderTime, OrderFilling from ..core.constants import TradeAction, OrderTime, OrderFilling
from ..core.exceptions import OrderError from ..core.exceptions import OrderError
from ..core.base import _Base from ..core.base import _Base
from .._utils import backoff_decorator, error_handler from ..utils import error_handler
logger = getLogger(__name__) logger = getLogger(__name__)
@@ -77,7 +77,6 @@ class Order(_Base, TradeRequest):
return tuple(TradeOrder(**order._asdict()) for order in orders) return tuple(TradeOrder(**order._asdict()) for order in orders)
return tuple() return tuple()
@backoff_decorator
async def check(self, **kwargs) -> OrderCheckResult: async def check(self, **kwargs) -> OrderCheckResult:
"""Check funds sufficiency for performing a required trading operation and the possibility of executing it. """Check funds sufficiency for performing a required trading operation and the possibility of executing it.
+1 -3
View File
@@ -7,7 +7,6 @@ from ..core.meta_trader import MetaTrader
from ..core.models import TradePosition, OrderSendResult from ..core.models import TradePosition, OrderSendResult
from ..core.constants import OrderType, TradeAction from ..core.constants import OrderType, TradeAction
from ..core.config import Config from ..core.config import Config
from .._utils import backoff_decorator
from ..core.meta_backtester import MetaBackTester from ..core.meta_backtester import MetaBackTester
from .order import Order from .order import Order
@@ -33,7 +32,6 @@ class Positions:
self.positions = () self.positions = ()
self.total_positions = 0 self.total_positions = 0
@backoff_decorator
async def get_positions(self, *, symbol: str = None, ticket: int = None, group: str = None) -> tuple[TradePosition, ...]: async def get_positions(self, *, symbol: str = None, ticket: int = None, group: str = None) -> tuple[TradePosition, ...]:
"""Get open positions with the ability to filter by symbol, ticket or group of symbols. """Get open positions with the ability to filter by symbol, ticket or group of symbols.
Args: Args:
@@ -147,7 +145,7 @@ class Positions:
return len([res for res in results if (isinstance(res, OrderSendResult) and res.retcode == 10009)]) return len([res for res in results if (isinstance(res, OrderSendResult) and res.retcode == 10009)])
async def get_total_positions(self) -> int: async def get_total_positions(self) -> int:
"""Get total number of open positions.""" """Get the total number of open positions."""
total = await self.mt5.positions_total() total = await self.mt5.positions_total()
self.total_positions = total or self.total_positions self.total_positions = total or self.total_positions
return self.total_positions return self.total_positions
+20 -3
View File
@@ -6,6 +6,7 @@ from threading import Lock
from ..core.config import Config from ..core.config import Config
from ..core.models import OrderSendResult from ..core.models import OrderSendResult
from .result_db import ResultDB
logger = getLogger(__name__) logger = getLogger(__name__)
@@ -17,11 +18,10 @@ class Result:
config (Config): The configuration object config (Config): The configuration object
name: Any desired name for the result file object name: Any desired name for the result file object
""" """
config: Config config: Config
lock = Lock() lock = Lock()
def __init__(self, *, result: OrderSendResult, parameters: dict = None, name: str = ""): def __init__(self, *, result: OrderSendResult, parameters: dict = None, name: str = "", **kwargs):
""" """
Prepare result data Prepare result data
Args: Args:
@@ -33,13 +33,28 @@ class Result:
self.parameters = parameters or {} self.parameters = parameters or {}
self.result = result self.result = result
self.name = name or self.parameters.get("name", "Trades") self.name = name or self.parameters.get("name", "Trades")
self.extra_params = kwargs
def to_sql(self):
try:
res = self.result.get_dict(include=set(ResultDB.fields()))
res["parameters"] = self.parameters
res["name"] = self.name
res["date"] = self.extra_params.get("date", "")
res["symbol"] = self.parameters["symbol"]
db = ResultDB(**res)
db.save(commit=True)
db.close()
except Exception as err:
logger.error("%s: Error occurred while saving", err)
def get_data(self) -> dict: def get_data(self) -> dict:
res = self.result.get_dict(exclude={"retcode", "comment", "retcode_external", "request_id", "request"}) res = self.result.get_dict(exclude={"retcode", "comment", "retcode_external", "request_id", "request"})
return self.parameters | res | {"actual_profit": 0, "closed": False, "win": False} return self.parameters | res | {"actual_profit": 0, "closed": False, "win": False} | self.extra_params
async def save(self, *, trade_record_mode: Literal["csv", "json"] = None): async def save(self, *, trade_record_mode: Literal["csv", "json"] = None):
"""Record trade results as a csv or json file """Record trade results as a csv or json file
Args: Args:
trade_record_mode (Literal['csv'|'json']): Mode of saving trade records trade_record_mode (Literal['csv'|'json']): Mode of saving trade records
""" """
@@ -49,6 +64,8 @@ class Result:
await self.to_csv() await self.to_csv()
elif trade_record_mode == "json": elif trade_record_mode == "json":
await self.to_json() await self.to_json()
elif trade_record_mode == "sql":
self.to_sql()
else: else:
logger.error(f"Invalid trade record mode: {trade_record_mode}") logger.error(f"Invalid trade record mode: {trade_record_mode}")
+40
View File
@@ -0,0 +1,40 @@
import pickle
from dataclasses import dataclass, field
from typing import ClassVar
from ..core.db import DB
@dataclass(kw_only=True)
class ResultDB(DB):
_table: ClassVar[str] = "result"
deal: int = field(metadata={"NOT NULL": True})
order: int = field(metadata={"PRIMARY KEY": True, "UNIQUE": True, "NOT NULL": True})
name: str = field(metadata={"NOT NULL": True})
symbol: str = field(metadata={"NOT NULL": True})
date: str = field(metadata={"NOT NULL": True})
volume: float
price: float
bid: float
ask: float
tp: float = 0
sl: float = 0
actual_profit: float = 0
expected_profit: float = 0
win: bool = False
closed: bool = False
profit: float = 0
loss: float = 0
comment: str = field(default="''")
parameters: dict|bytes|str = field(default="''")
def __post_init__(self):
super().__post_init__()
if isinstance(self.parameters, bytes):
self.parameters = pickle.loads(self.parameters)
def get_data(self):
data = self.asdict()
if not isinstance(params:=data["parameters"], bytes):
data["parameters"] = pickle.dumps(params, protocol=pickle.HIGHEST_PROTOCOL)
return data
+3 -10
View File
@@ -6,7 +6,7 @@ from logging import getLogger
from ..core.constants import TimeFrame, CopyTicks from ..core.constants import TimeFrame, CopyTicks
from ..core.base import _Base from ..core.base import _Base
from ..core.models import SymbolInfo, BookInfo from ..core.models import SymbolInfo, BookInfo
from .._utils import round_off, backoff_decorator from ..utils import round_off
from .ticks import Tick from .ticks import Tick
from .account import Account from .account import Account
from .candle import Candles from .candle import Candles
@@ -42,7 +42,6 @@ class Symbol(_Base, SymbolInfo):
self.account = Account() self.account = Account()
self.initialized = False self.initialized = False
@backoff_decorator
async def info_tick(self, *, name: str = "") -> Tick | None: async def info_tick(self, *, name: str = "") -> Tick | None:
"""Get the current price tick of a financial instrument. """Get the current price tick of a financial instrument.
@@ -250,7 +249,6 @@ class Symbol(_Base, SymbolInfo):
logger.warning(f"{err}: Currency conversion failed: Unable to convert {amount} in {quote} to {base}") logger.warning(f"{err}: Currency conversion failed: Unable to convert {amount} in {quote} to {base}")
return None return None
@backoff_decorator
async def copy_rates_from(self, *, timeframe: TimeFrame, date_from: datetime | int, count: int = 500) -> Candles: async def copy_rates_from(self, *, timeframe: TimeFrame, date_from: datetime | int, count: int = 500) -> Candles:
""" """
Get bars from the MetaTrader 5 terminal starting from the specified date. Get bars from the MetaTrader 5 terminal starting from the specified date.
@@ -274,7 +272,6 @@ class Symbol(_Base, SymbolInfo):
return Candles(data=rates) return Candles(data=rates)
raise ValueError(f"Could not get rates for {self.name}.") raise ValueError(f"Could not get rates for {self.name}.")
@backoff_decorator
async def copy_rates_from_pos(self, *, timeframe: TimeFrame, count: int = 500, start_position: int = 0) -> Candles: async def copy_rates_from_pos(self, *, timeframe: TimeFrame, count: int = 500, start_position: int = 0) -> Candles:
"""Get bars from the MetaTrader 5 terminal starting from the specified index. """Get bars from the MetaTrader 5 terminal starting from the specified index.
@@ -297,7 +294,6 @@ class Symbol(_Base, SymbolInfo):
return Candles(data=rates) return Candles(data=rates)
raise ValueError(f"Could not get rates for {self.name}.") raise ValueError(f"Could not get rates for {self.name}.")
@backoff_decorator
async def copy_rates_range( async def copy_rates_range(
self, *, timeframe: TimeFrame, date_from: datetime | int, date_to: datetime | int self, *, timeframe: TimeFrame, date_from: datetime | int, date_to: datetime | int
) -> Candles: ) -> Candles:
@@ -327,10 +323,8 @@ class Symbol(_Base, SymbolInfo):
return Candles(data=rates) return Candles(data=rates)
raise ValueError(f"Could not get rates for {self.name}.") raise ValueError(f"Could not get rates for {self.name}.")
@backoff_decorator async def copy_ticks_from(self, *, date_from: datetime | int, count: int = 100,
async def copy_ticks_from( flags: CopyTicks = CopyTicks.ALL) -> Ticks:
self, *, date_from: datetime | int, count: int = 100, flags: CopyTicks = CopyTicks.ALL
) -> Ticks:
""" """
Get ticks from the MetaTrader 5 terminal starting from the specified date. Get ticks from the MetaTrader 5 terminal starting from the specified date.
@@ -352,7 +346,6 @@ class Symbol(_Base, SymbolInfo):
return Ticks(data=ticks) return Ticks(data=ticks)
raise ValueError(f"Could not get ticks for {self.name}.") raise ValueError(f"Could not get ticks for {self.name}.")
@backoff_decorator
async def copy_ticks_range( async def copy_ticks_range(
self, *, date_from: datetime | int, date_to: datetime | int, flags: CopyTicks = CopyTicks.ALL self, *, date_from: datetime | int, date_to: datetime | int, flags: CopyTicks = CopyTicks.ALL
) -> Ticks: ) -> Ticks:
+1 -1
View File
@@ -17,7 +17,7 @@ logger = logging.getLogger(__name__)
class TradeRecords: class TradeRecords:
"""This utility class read trade records from csv files, and update them based on their closing positions. """This utility class read trade records from the csv and json files, and update them based on their closing positions.
Attributes: Attributes:
config: Config object config: Config object
+6 -5
View File
@@ -11,7 +11,8 @@ from .result import Result
from .order import Order from .order import Order
from .symbol import Symbol as _Symbol from .symbol import Symbol as _Symbol
from .ram import RAM from .ram import RAM
from .._utils import error_handler from ..utils import error_handler
# from .result_db import ResultDB
logger = getLogger(__name__) logger = getLogger(__name__)
Symbol = TypeVar("Symbol", bound=_Symbol) Symbol = TypeVar("Symbol", bound=_Symbol)
@@ -200,14 +201,14 @@ class Trader(ABC):
return return
params = {**parameters} if isinstance(parameters, dict) else {} params = {**parameters} if isinstance(parameters, dict) else {}
profit = await self.order.calc_profit() profit = await self.order.calc_profit()
params["expected_profit"] = profit # params["expected_profit"] = profit
date = ( date = (
datetime.now(tz=UTC) datetime.now()
if self.config.mode == "live" if self.config.mode == "live"
else datetime.fromtimestamp(self.config.backtest_engine.cursor.time, tz=UTC) else datetime.fromtimestamp(self.config.backtest_engine.cursor.time, tz=UTC)
) )
params["date"] = date.strftime("%Y-%m-%d %H:%M:%S.%f") # params["date"] = date.strftime("%Y-%m-%d %H:%M:%S.%f")
res = Result(result=result, parameters=params, name=name) res = Result(result=result, parameters=params, name=name, date=date, expected=profit)
self.config.task_queue.add(item=QueueItem(res.save), must_complete=True) self.config.task_queue.add(item=QueueItem(res.save), must_complete=True)
@abstractmethod @abstractmethod
+2
View File
@@ -0,0 +1,2 @@
from .utils import *
from .process_pool import *
@@ -6,38 +6,10 @@ from functools import wraps, partial
import asyncio import asyncio
from threading import RLock from threading import RLock
from logging import getLogger from logging import getLogger
from .core.config import Config from ..core.config import Config
logger = getLogger(__name__) logger = getLogger(__name__)
async def backtest_sleep(secs):
"""An async sleep function for use during backtesting."""
config = Config()
sleep = config.backtest_engine.cursor.time + secs
while sleep > config.backtest_engine.cursor.time:
await asyncio.sleep(0)
# async def backtest_sleep(secs: float):
# config = Config()
# btc = config.backtest_controller
# try:
# if btc.parties == 2:
# steps = int(secs) // config.backtest_engine.speed
# steps = max(steps, 1)
# config.backtest_engine.fast_forward(steps=steps)
# btc.wait()
#
# elif btc.parties > 2:
# _time = config.backtest_engine.cursor.time + secs
# while _time > config.backtest_engine.cursor.time:
# btc.wait()
# else:
# btc.wait()
# except Exception as err:
# btc.wait()
# logger.error("Error: %s in backtest_sleep", err)
def dict_to_string(data: dict, multi=False) -> str: def dict_to_string(data: dict, multi=False) -> str:
"""Convert a dict to a string. Useful for logging. """Convert a dict to a string. Useful for logging.
@@ -104,7 +76,8 @@ def error_handler(func=None, *, msg="", exe=Exception, response=None, log_error_
return res return res
except exe as err: except exe as err:
if log_error_msg: if log_error_msg:
logger.error(f"Error in {func.__name__}: {msg or err}") err_msg = msg or f"Error in {func.__name__}: {err}"
logger.error(err_msg)
return response return response
return wrapper return wrapper
+1 -1
View File
@@ -15,6 +15,6 @@ class TestAccount:
assert self.account.connected is True assert self.account.connected is True
async def test_account_info(self): async def test_account_info(self):
acc_info = await self.account.mt5.account_info() acc_info = await self.account.mt5.account_info
assert acc_info.login == self.account.login assert acc_info.login == self.account.login
assert acc_info.server == self.account.server assert acc_info.server == self.account.server
+1 -1
View File
@@ -3,7 +3,7 @@ from math import ceil
from aiomql import TimeFrame from aiomql import TimeFrame
from aiomql.core.backtesting import BackTestEngine from aiomql.core.backtesting import BackTestEngine
from aiomql.core.backtesting.get_data import GetData from aiomql.core.backtesting.get_data import GetData
from aiomql._utils import round_down from aiomql.utils import round_down
from aiomql.core.constants import OrderType, TradeAction from aiomql.core.constants import OrderType, TradeAction
import pytest import pytest
+1 -1
View File
@@ -18,7 +18,7 @@ class TestConfig:
assert config.backtest_engine is engine assert config.backtest_engine is engine
def test_account_info(self, config): def test_account_info(self, config):
account_info = config.account_info() account_info = config.account_info
assert isinstance(account_info, dict) assert isinstance(account_info, dict)
assert "login" in account_info assert "login" in account_info
assert "password" in account_info assert "password" in account_info