reorganize the library into three folders lib, contrib and core

Write unittests with pytest
Make all functions and method signatures as keyword only arguments
This commit is contained in:
Ichinga Samuel
2024-10-06 23:20:59 +01:00
parent 9b55214d1a
commit 76a947640f
68 changed files with 1439 additions and 817 deletions
+4 -5
View File
@@ -48,7 +48,6 @@ nosetests.xml
coverage.xml coverage.xml
*,cover *,cover
.hypothesis/ .hypothesis/
scrap.py
# Translations # Translations
*.mo *.mo
@@ -74,9 +73,9 @@ target/
.vscode/ .vscode/
# config files # config files
config.json
aiomql.json aiomql.json
config/
test_data/ # development
test.json terminals/
+1 -19
View File
@@ -1,21 +1,3 @@
from .core import * from .core import *
from .account import Account
from .ram import RAM
from .symbol import Symbol
from .strategy import Strategy
from .bot_builder import Bot
from .result import Result
from .records import Records
from .trade_records import TradeRecords
from .candle import Candle, Candles
from .positions import Positions
from .executor import Executor
from .order import Order
from .ticks import Tick, Ticks
from .history import History
from .trader import Trader
from .terminal import Terminal
from .sessions import Session, Sessions
from .utils import dict_to_string, round_off, backoff_decorator, error_handler, error_handler_sync, round_up, round_down
from .lib import * from .lib import *
# from .contrib import * from .contrib import *
+6 -2
View File
@@ -22,7 +22,7 @@ def dict_to_string(data: dict, multi=False) -> str:
return f"{sep}".join(f"{key}: {value}" for key, value in data.items()) return f"{sep}".join(f"{key}: {value}" for key, value in data.items())
def backoff_decorator(func=None, *, max_retries: int = 5, retries: int = 0, error='') -> callable: def backoff_decorator(func=None, *, max_retries: int = 3, retries: int = 0, error='') -> callable:
if func is None: if func is None:
return partial(backoff_decorator, max_retries=max_retries, retries=retries, error=error) return partial(backoff_decorator, max_retries=max_retries, retries=retries, error=error)
@@ -37,7 +37,9 @@ def backoff_decorator(func=None, *, max_retries: int = 5, retries: int = 0, erro
res = await func(*args, **kwargs) res = await func(*args, **kwargs)
if error != '' and res == error: if error != '' and res == error:
raise TypeError('Invalid return type') raise TypeError('Invalid return type')
return res else:
retries = 0
return res
except Exception as err: except Exception as err:
logger.error(f'Error in {func.__name__}: {err}') logger.error(f'Error in {func.__name__}: {err}')
await asyncio.sleep(retries + random.randint(1, max_retries)) await asyncio.sleep(retries + random.randint(1, max_retries))
@@ -72,6 +74,7 @@ def error_handler_sync(func=None, *, msg='', exe=Exception, response=None):
return res return res
except exe as err: except exe as err:
logger.error(f'Error in {func.__name__}: {msg or err}') logger.error(f'Error in {func.__name__}: {msg or err}')
return response
return wrapper return wrapper
@@ -83,6 +86,7 @@ def round_up(value: int, base: int) -> int:
return value if value % base == 0 else value + base - (value % base) return value if value % base == 0 else value + base - (value % base)
# noinspection PyShadowingNames
def round_off(value: float, step: float, round_down: bool = False) -> float: def round_off(value: float, step: float, round_down: bool = False) -> float:
"""Round off a number to the nearest step.""" """Round off a number to the nearest step."""
with decimal.localcontext() as ctx: with decimal.localcontext() as ctx:
+4 -1
View File
@@ -1 +1,4 @@
# from .backtester import * from .backtesting import *
from .strategies import *
from .candle_patterns import *
from .symbols import *
@@ -1,8 +0,0 @@
from .meta_tester import MetaTester
from .backtest_engine import BackTestEngine
from .get_data import GetData, TestData
from .strategy_tester import StrategyTester
from .event_manager import EventManager
from .backtester import BackTester
from .test_account import TestAccount
from .trades_manager import PositionsManager, OrdersManager, DealsManager
@@ -1,27 +0,0 @@
from .event_manager import EventManager
from ...core.config import Config
class StrategyTester:
event_manager: EventManager
config: Config
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:
...
def test(self):
raise NotImplementedError("Implement this method in your subclass")
@@ -0,0 +1,4 @@
from .get_data import GetData, TestData
from .backtest_engine import BackTestEngine
from .backtest_account import BackTestAccount
from .trades_manager import PositionsManager, OrdersManager, DealsManager
@@ -5,7 +5,7 @@ from ...core.constants import AccountTradeMode, AccountMarginMode, AccountStopOu
@dataclass @dataclass
class TestAccount: class BackTestAccount:
login: int = 0 login: int = 0
trade_mode: AccountTradeMode = AccountTradeMode.DEMO trade_mode: AccountTradeMode = AccountTradeMode.DEMO
leverage: float = 0 leverage: float = 0
@@ -15,20 +15,19 @@ from MetaTrader5 import (Tick, SymbolInfo, AccountInfo, TradeOrder, TradePositio
from ...core.meta_trader import MetaTrader from ...core.meta_trader import MetaTrader
from ...core.constants import (TimeFrame, OrderType, TradeAction, AccountStopOutMode, PositionReason, from ...core.constants import (TimeFrame, OrderType, TradeAction, AccountStopOutMode, PositionReason,
DealType, DealReason, DealEntry, OrderReason, CopyTicks) 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 ..._utils import round_down, round_up, error_handler, error_handler_sync, async_cache
from .get_data import TestData, GetData, Cursor from .get_data import TestData, GetData, Cursor
from .test_account import TestAccount from .backtest_account import BackTestAccount
from .trades_manager import PositionsManager, OrdersManager, DealsManager from .trades_manager import PositionsManager, OrdersManager, DealsManager
tz = pytz.timezone('Etc/UTC')
logger = getLogger(__name__) logger = getLogger(__name__)
class BackTestEngine: class BackTestEngine:
mt5: MetaTrader = MetaTrader() mt5: MetaTrader
span: range span: range
range: range range: range
cursor: Cursor cursor: Cursor
@@ -39,15 +38,19 @@ class BackTestEngine:
orders: OrdersManager orders: OrdersManager
deals: DealsManager deals: DealsManager
positions: PositionsManager positions: PositionsManager
_account: TestAccount _account: BackTestAccount
def __init__(self, *, data: TestData = None, speed: int = 1, start: float | datetime = 0, def __init__(self, *, data: TestData = None, speed: int = 1, start: float | datetime = 0,
end: float | datetime = 0, restart: bool = False, name: str = ''): end: float | datetime = 0, restart: bool = False, name: str = ''):
self._data = data or TestData() self._data = data or TestData()
self.config = Config(backtest_engine=self) self.mt5 = MetaTrader()
self.config = self.mt5.config
self.config.backtest_engine = self
self.set_up(start=start, end=end, speed=speed, restart=restart) self.set_up(start=start, end=end, speed=speed, restart=restart)
self.prepare_data() self.prepare_data()
_name = f"{datetime.fromtimestamp(self.span[0]):%d-%m-%y}_{datetime.fromtimestamp(self.span[-1]):%d-%m-%y}" start, end = (self.span[0], self.span[-1]) if self.span else ((now := datetime.now(pytz.UTC).timestamp()), now)
_name = f"{datetime.fromtimestamp(start):%d-%m-%y}_{datetime.fromtimestamp(end):%d-%m-%y}"
self.name = name or _name self.name = name or _name
def __next__(self) -> Cursor: def __next__(self) -> Cursor:
@@ -91,7 +94,7 @@ class BackTestEngine:
deals[ticket] = TradeDeal((deal.get(k) for k in TradeDeal.__match_args__)) deals[ticket] = TradeDeal((deal.get(k) for k in TradeDeal.__match_args__))
self.deals = DealsManager(data=deals) self.deals = DealsManager(data=deals)
self._account: TestAccount = TestAccount(**self._data.account) self._account: BackTestAccount = BackTestAccount(**self._data.account)
def next(self) -> Cursor: def next(self) -> Cursor:
return next(self) return next(self)
@@ -99,7 +102,7 @@ class BackTestEngine:
@property @property
def data(self): def data(self):
return self._data return self._data
def reset(self): def reset(self):
self.iter = zip_longest(self.range, self.span) self.iter = zip_longest(self.range, self.span)
self.cursor = Cursor(index=self.range.start, time=self.span.start) self.cursor = Cursor(index=self.range.start, time=self.span.start)
@@ -107,7 +110,7 @@ class BackTestEngine:
def go_to(self, *, time: datetime | float): def go_to(self, *, time: datetime | float):
time = int(time.timestamp()) if isinstance(time, datetime) else int(time) time = int(time.timestamp()) if isinstance(time, datetime) else int(time)
steps = time - self.cursor.time steps = time - self.cursor.time
if steps > 0: if steps > 0:
self.fast_forward(steps=steps) self.fast_forward(steps=steps)
return return
@@ -451,11 +454,11 @@ class BackTestEngine:
current_price = price current_price = price
if tp and sl: if tp and sl:
if action == TradeAction.SLTP: if action == TradeAction.SLTP:
pos = self.positions.get(request.get('position')) pos = self.positions.get(request.get('position'))
sym = 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_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 current_price = current_tick.bid if pos.type == OrderType.BUY else current_tick.ask
min_sl = min(sl, tp) min_sl = min(sl, tp)
dsl = abs(current_price - min_sl) / sym.point dsl = abs(current_price - min_sl) / sym.point
tsl = sym.trade_stops_level + sym.spread tsl = sym.trade_stops_level + sym.spread
@@ -549,7 +552,7 @@ class BackTestEngine:
@error_handler @error_handler
async def get_rates_from_pos(self, symbol: str, timeframe: TimeFrame, start_pos: int, count: int) -> np.ndarray: async def get_rates_from_pos(self, symbol: str, timeframe: TimeFrame, start_pos: int, count: int) -> np.ndarray:
if self.config.use_terminal_for_backtesting: if self.config.use_terminal_for_backtesting:
now = datetime.now(tz=tz) now = datetime.now(tz=pytz.UTC)
b_now = self.cursor.time b_now = self.cursor.time
diff = (now.timestamp() - b_now) // timeframe.time diff = (now.timestamp() - b_now) // timeframe.time
start_pos = int(diff + start_pos) start_pos = int(diff + start_pos)
@@ -13,7 +13,7 @@ from ...core.meta_trader import MetaTrader
from ...core.config import Config from ...core.config import Config
from ...core.constants import TimeFrame from ...core.constants import TimeFrame
from ...core.task_queue import TaskQueue, QueueItem from ...core.task_queue import TaskQueue, QueueItem
from ...utils import backoff_decorator from ..._utils import backoff_decorator
logger = getLogger(__name__) logger = getLogger(__name__)
@@ -1,8 +1,10 @@
from datetime import datetime from datetime import datetime
from typing import TypeVar, Generic from typing import TypeVar, Generic
from logging import getLogger
from MetaTrader5 import TradePosition, TradeOrder, TradeDeal from MetaTrader5 import TradePosition, TradeOrder, TradeDeal
from aiomql.utils import logger logger = getLogger(__name__)
TradeData = TypeVar('TradeData', bound=TradePosition | TradeOrder | TradeDeal) TradeData = TypeVar('TradeData', bound=TradePosition | TradeOrder | TradeDeal)
@@ -1,4 +1,4 @@
from ...candle import Candle, Candles from ...lib.candle import Candle, Candles
def find_bearish_fractal(candles: Candles) -> Candle | None: def find_bearish_fractal(candles: Candles) -> Candle | None:
@@ -0,0 +1,2 @@
from .finger_trap import FingerTrap
from .tracker import Tracker
@@ -1,15 +1,15 @@
import asyncio import asyncio
import logging import logging
from .tracker import Tracker from ...lib.symbol import Symbol
from ..traders import SimpleTrader from ...lib.trader import Trader
from ...symbol import Symbol from ...lib.candle import Candles
from ...trader import Trader from ...lib.strategy import Strategy
from ...candle import Candles from ...core.constants import TimeFrame, OrderType
from ...strategy import Strategy from ...lib.sessions import Sessions
from ...core import TimeFrame, OrderType
from ...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 .tracker import Tracker
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -1,4 +1,4 @@
from ...symbol import Symbol from ...lib.symbol import Symbol
from ...core.exceptions import VolumeError from ...core.exceptions import VolumeError
@@ -6,6 +6,16 @@ class ForexSymbol(Symbol):
"""Subclass of Symbol for Forex Symbols. Handles the conversion of currency and the computation of stop loss, """Subclass of Symbol for Forex Symbols. Handles the conversion of currency and the computation of stop loss,
take profit and volume. take profit and volume.
""" """
@property
def pip(self):
"""Returns the pip value of the symbol. This is ten times the point value for forex symbols.
Returns:
float: The pip value of the symbol.
"""
return self.point * 10
def compute_points(self, *, amount: float, volume) -> float: def compute_points(self, *, amount: float, volume) -> float:
"""Compute the number of points required for a trade. Given the amount and the volume of the trade. """Compute the number of points required for a trade. Given the amount and the volume of the trade.
Args: Args:
@@ -1,9 +1,9 @@
from logging import getLogger from logging import getLogger
from ..symbols import ForexSymbol from ...lib.ram import RAM
from ...ram import RAM
from ...core.models import OrderType from ...core.models import OrderType
from ...trader import Trader from ...lib.trader import Trader
from ..symbols import ForexSymbol
logger = getLogger(__name__) logger = getLogger(__name__)
+1 -1
View File
@@ -2,7 +2,7 @@ from .meta_trader import MetaTrader
from .config import Config from .config import Config
from .models import * from .models import *
from .constants import * from .constants import *
from .base import Base from .base import Base, _Base
from .errors import Error from .errors import Error
from .exceptions import * from .exceptions import *
from .task_queue import TaskQueue from .task_queue import TaskQueue
+20 -14
View File
@@ -4,7 +4,7 @@ from logging import getLogger
from .config import Config from .config import Config
from .meta_trader import MetaTrader from .meta_trader import MetaTrader
from ..contrib.backtester import MetaTester from .meta_backtester import MetaBackTester
logger = getLogger(__name__) logger = getLogger(__name__)
@@ -13,18 +13,17 @@ class Base:
"""A base class for all data structure classes in the aiomql package. This class provides a set of common methods """A base class for all data structure classes in the aiomql package. This class provides a set of common methods
and attributes for handling data. and attributes for handling data.
""" """
mt5: MetaTrader exclude: set
config: Config include: set
def __init__(self, **kwargs): def __init__(self, **kwargs):
""" """
Initialize a new instance of the Base class Initialize a new instance of the Base class
Args: Args:
**kwargs: Set instance attributes with keyword arguments. Only if they are annotated on the class body. **kwargs: Set instance attributes with keyword arguments. Only if they are annotated on the class body.
""" """
self.config = Config() self.exclude = {'mt5', "config", 'exclude', 'include', 'annotations', 'class_vars', 'dict', '_instance'}
self.mt5 = MetaTrader() if self.config.mode == 'live' else MetaTester()
self.exclude = {'mt5', "config", 'exclude', 'include', 'annotations', 'class_vars', 'dict'}
self.include = set() self.include = set()
self.set_attributes(**kwargs) self.set_attributes(**kwargs)
@@ -37,7 +36,7 @@ class Base:
def set_attributes(self, **kwargs): def set_attributes(self, **kwargs):
"""Set keyword arguments as object attributes """Set keyword arguments as object attributes
Keyword Args: Keyword Args:
**kwargs: Object attributes and values as keyword arguments **kwargs: Object attributes and values as keyword arguments
@@ -51,15 +50,15 @@ class Base:
try: try:
setattr(self, i, self.annotations[i](j)) setattr(self, i, self.annotations[i](j))
except KeyError: except KeyError:
logger.warning(f"Attribute {i} does not belong to class {self.__class__.__name__}") logger.debug(f"Attribute {i} does not belong to class {self.__class__.__name__}")
continue continue
except ValueError: except (ValueError, TypeError):
logger.warning(f'Cannot covert object of type {type(j)} to type {self.annotations[i]}') logger.debug(f'Cannot covert object of type {type(j)} to type {self.annotations[i]}')
continue setattr(self, i, j)
except Exception as exe: except Exception as exe:
logger.warning(f'Did not set attribute {i} on class {self.__class__.__name__} due to {exe}') logger.debug(f'Did not set attribute {i} on class {self.__class__.__name__} due to {exe}')
continue continue
@property @property
@@ -71,7 +70,7 @@ class Base:
dict: A dictionary of class annotations dict: A dictionary of class annotations
""" """
annots = {} annots = {}
for base in self.__class__.__mro__[-3::-1]: for base in self.__class__.__mro__[::-1]:
annots |= getattr(base, '__annotations__', {}) annots |= getattr(base, '__annotations__', {})
return annots return annots
@@ -100,7 +99,7 @@ class Base:
Returns: Returns:
dict: A dictionary of available class attributes in all ancestor classes and the current class. dict: A dictionary of available class attributes in all ancestor classes and the current class.
""" """
clss = self.__class__.__mro__[-3::-1] clss = self.__class__.__mro__[::-1]
cls_dict = {} cls_dict = {}
for cls in clss: for cls in clss:
cls_dict |= cls.__dict__ cls_dict |= cls.__dict__
@@ -119,3 +118,10 @@ class Base:
key not in _filter} key not in _filter}
except Exception as err: except Exception as err:
logger.warning(err) logger.warning(err)
class _Base(Base):
def __init__(self, **kwargs):
self.config = Config()
self.mt5 = MetaTrader() if self.config.mode != 'backtest' else MetaBackTester()
super().__init__(**kwargs)
+17 -5
View File
@@ -1,3 +1,4 @@
import inspect
import os import os
from pathlib import Path from pathlib import Path
from typing import Iterator, Literal, TypeVar, Self from typing import Iterator, Literal, TypeVar, Self
@@ -7,10 +8,15 @@ from logging import getLogger
from .task_queue import TaskQueue from .task_queue import TaskQueue
logger = getLogger(__name__) logger = getLogger(__name__)
Bot = TypeVar("Bot") Bot = TypeVar("Bot")
BackTestEngine = TypeVar("BackTestEngine") BackTestEngine = TypeVar("BackTestEngine")
def func():
stack = inspect.stack()
calling_context = next(context for context in stack if context.filename != __file__)
print(calling_context.filename)
return calling_context.filename
class Config: class Config:
"""A class for handling configuration settings for the aiomql package. """A class for handling configuration settings for the aiomql package.
@@ -26,7 +32,7 @@ class Config:
path (str): Path to terminal file path (str): Path to terminal file
timeout (int): Timeout for terminal connection timeout (int): Timeout for terminal connection
state (dict): A global state dictionary for storing data across the framework state (dict): A global state dictionary for storing data across the framework
root (str): Root directory of the project root (Path): Root directory of the project
Notes: Notes:
By default, the config class looks for a file named aiomql.json. By default, the config class looks for a file named aiomql.json.
@@ -80,11 +86,13 @@ class Config:
self._backtest_engine = value self._backtest_engine = value
def set_attributes(self, **kwargs): def set_attributes(self, **kwargs):
"""Set keyword arguments as object attributes """Set keyword arguments as object attributes, The root folder attribute can't be set here.
Keyword Args: Keyword Args:
**kwargs: Object attributes and values as keyword arguments **kwargs: Object attributes and values as keyword arguments
""" """
if kwargs.pop('root', None) is not None:
logger.warning('Tried setting root from set_attributes. Use load_config to change project root')
[setattr(self, key, value) for key, value in kwargs.items()] [setattr(self, key, value) for key, value in kwargs.items()]
@staticmethod @staticmethod
@@ -113,7 +121,7 @@ class Config:
logger.debug(f"Error finding config file: {err}") logger.debug(f"Error finding config file: {err}")
return return
def load_config(self, *, file: str | Path = None, filename: str = None, root: str | Path = None, **kwargs): def load_config(self, *, file: str | Path = None, filename: str = None, root: str | Path = None, **kwargs) -> Self:
"""Load configuration settings from a file. """Load configuration settings from a file.
Keyword Args: Keyword Args:
@@ -128,7 +136,6 @@ class Config:
self.root = root self.root = root
else: else:
self.root = self.root if hasattr(self, 'root') else Path.cwd() self.root = self.root if hasattr(self, 'root') else Path.cwd()
if file is not None: if file is not None:
file = Path(file).resolve() file = Path(file).resolve()
if not file.exists(): if not file.exists():
@@ -151,6 +158,9 @@ class Config:
data = file_config | kwargs data = file_config | kwargs
self.set_attributes(**data) self.set_attributes(**data)
if self.path:
self.path = self.root / self.path if not Path(self.path).resolve().exists() else self.path
if self.record_trades and not hasattr(self, "records_dir"): if self.record_trades and not hasattr(self, "records_dir"):
self.records_dir = self.root / self.records_dir_name self.records_dir = self.root / self.records_dir_name
self.records_dir.mkdir(parents=True, exist_ok=True) self.records_dir.mkdir(parents=True, exist_ok=True)
@@ -159,6 +169,8 @@ class Config:
self.backtest_dir = self.root / self.backtest_dir_name self.backtest_dir = self.root / self.backtest_dir_name
self.backtest_dir.mkdir(parents=True, exist_ok=True) self.backtest_dir.mkdir(parents=True, exist_ok=True)
return self
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
@@ -2,7 +2,7 @@ import asyncio
from asyncio import Condition, Task from asyncio import Condition, Task
from typing import Self from typing import Self
from datetime import datetime from datetime import datetime
from ...core import Config from .config import Config
class EventManager: class EventManager:
@@ -42,7 +42,7 @@ class EventManager:
async def event_monitor(self): async def event_monitor(self):
while True: while True:
async with self.condition: async with self.condition:
if self.task_tracker == self.num_main_tasks: if self.task_tracker == self.num_main_tasks: # all main tasks have been completed in the current cycle
self.task_tracker = 0 self.task_tracker = 0
await self.config.backtest_engine.tracker() await self.config.backtest_engine.tracker()
self.config.backtest_engine.next() self.config.backtest_engine.next()
@@ -1,21 +1,20 @@
from datetime import datetime from datetime import datetime
from logging import getLogger from logging import getLogger
from typing import Literal from typing import Literal, TypeVar
from numpy import ndarray from numpy import ndarray
from MetaTrader5 import (Tick, SymbolInfo, AccountInfo, TerminalInfo, TradeOrder, TradePosition, TradeDeal, from MetaTrader5 import (Tick, SymbolInfo, AccountInfo, TerminalInfo, TradeOrder, TradePosition, TradeDeal,
OrderCheckResult, OrderSendResult) OrderCheckResult, OrderSendResult)
from . import BackTestEngine from .meta_trader import MetaTrader
from .constants import TimeFrame, CopyTicks, OrderType
from ...core.meta_trader import MetaTrader from .._utils import error_handler
from ...core.constants import TimeFrame, CopyTicks, OrderType
from ...utils import error_handler
logger = getLogger(__name__) logger = getLogger(__name__)
BackTestEngine = TypeVar('BackTestEngine')
class MetaTester(MetaTrader): class MetaBackTester(MetaTrader):
"""A class for testing trading strategies in the MetaTrader 5 terminal. A subclass of MetaTrader.""" """A class for testing trading strategies in the MetaTrader 5 terminal. A subclass of MetaTrader."""
backtest_engine: BackTestEngine backtest_engine: BackTestEngine
@@ -29,7 +28,7 @@ class MetaTester(MetaTrader):
@backtest_engine.setter @backtest_engine.setter
def backtest_engine(self, value: BackTestEngine): def backtest_engine(self, value: BackTestEngine):
if isinstance(value, BackTestEngine): if BackTestEngine is not None:
self.config.backtest_engine = value self.config.backtest_engine = value
async def last_error(self) -> tuple[int, str]: async def last_error(self) -> tuple[int, str]:
+25 -15
View File
@@ -1,5 +1,5 @@
from datetime import datetime
import asyncio import asyncio
from datetime import datetime
from logging import getLogger from logging import getLogger
from typing import Literal from typing import Literal
@@ -17,6 +17,7 @@ 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(1)
@@ -29,7 +30,8 @@ class MetaTrader(MetaCore):
Returns: Returns:
MetaTrader: An instance of the MetaTrader class. MetaTrader: An instance of the MetaTrader class.
""" """
await self.initialize(**Config().account_info()) await self.initialize()
await self.login()
return self return self
async def __aexit__(self, exc_type, exc_val, exc_tb): async def __aexit__(self, exc_type, exc_val, exc_tb):
@@ -50,7 +52,7 @@ class MetaTrader(MetaCore):
self.error = Error(*err) self.error = Error(*err)
if self.error.is_connection_error(): if self.error.is_connection_error():
await self.initialize(path=self.config.path) await self.initialize()
await self.login() await self.login()
res = await asyncio.to_thread(func, *args, **kwargs) res = await asyncio.to_thread(func, *args, **kwargs)
@@ -81,7 +83,7 @@ class MetaTrader(MetaCore):
server = server or acc_details.get('server', '') server = server or acc_details.get('server', '')
return await asyncio.to_thread(self._login, login, password=password, server=server, timeout=timeout) return await asyncio.to_thread(self._login, login, password=password, server=server, timeout=timeout)
async def initialize(self, path: str = "", login: int = 0, password: str = "", server: str = "", async def initialize(self, path: str = None, login: int = 0, password: str = "", server: str = "",
timeout: int | None = None, portable=False) -> bool: 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.
@@ -97,16 +99,23 @@ class MetaTrader(MetaCore):
Returns: Returns:
bool: True if successful, False otherwise. bool: True if successful, False otherwise.
""" """
path = path or self.config.path async with asyncio.Lock() as _:
args = (str(path),) if path else () path = self.config.path if path is None else path
acc = self.config.account_info() args = (str(path),) if path else ()
kwargs = {key: value for key, value in (('login', login or acc.get('login')), acc = self.config.account_info()
('password', password or acc.get('password')), kwargs = {key: value for key, value in (('login', login or acc.get('login')),
('server', server or acc.get('server')), ('password', password or acc.get('password')),
('timeout', timeout or 60000), ('server', server or acc.get('server')),
('portable', portable)) if key is not None} ('timeout', timeout or 60000),
res = await asyncio.to_thread(self._initialize, *args, **kwargs) ('portable', portable)) if key is not None}
return res 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
async def shutdown(self) -> None: async def shutdown(self) -> None:
""" """
@@ -228,7 +237,8 @@ class MetaTrader(MetaCore):
res = await self._handler(api) res = await self._handler(api)
return res return res
async def order_calc_margin(self, action: OrderType, symbol: str, volume: float, price: float) -> float | None: async def order_calc_margin(self, action: Literal[OrderType.BUY, OrderType.SELL],
symbol: str, volume: float, price: float) -> float | None:
api = {'func': self._order_calc_margin, 'args': (action, symbol, volume, price), api = {'func': self._order_calc_margin, 'args': (action, symbol, volume, price),
'error_msg': 'Error in calculating margin.'} 'error_msg': 'Error in calculating margin.'}
res = await self._handler(api) res = await self._handler(api)
+2 -5
View File
@@ -327,13 +327,9 @@ class SymbolInfo(Base):
exchange: str exchange: str
formula: str formula: str
isin: str isin: str
name: str
page: str page: str
path: str path: str
name: str = ''
def __init__(self, **kwargs):
assert 'name' in kwargs, "Symbol Object Must be initialized with a name"
super().__init__(**kwargs)
def __repr__(self): def __repr__(self):
return '%(class)s(name=%(name)s)' % {'class': self.__class__.__name__, 'name': self.name} return '%(class)s(name=%(name)s)' % {'class': self.__class__.__name__, 'name': self.name}
@@ -346,6 +342,7 @@ class SymbolInfo(Base):
def __hash__(self): def __hash__(self):
return hash(self.name) return hash(self.name)
# return hash(id(self))
class BookInfo(Base): class BookInfo(Base):
+23 -24
View File
@@ -54,34 +54,35 @@ class TaskQueue:
item = (priority, item) item = (priority, item)
self.priority_tasks.add(item) if item.must_complete else ... self.priority_tasks.add(item) if item.must_complete else ...
self.queue.put_nowait(item) self.queue.put_nowait(item)
except asyncio.QueueFull: except asyncio.QueueFull:
logger.error(f"Queue is full") logger.error("Queue is full")
async def worker(self): async def worker(self):
while True: while True:
if isinstance(self.queue, asyncio.PriorityQueue): try:
_, item = await self.queue.get() if isinstance(self.queue, asyncio.PriorityQueue):
_, item = self.queue.get_nowait()
else: else:
item = await self.queue.get() item = self.queue.get_nowait()
if not self.stop or item.must_complete:
await item.run()
self.queue.task_done() if not self.stop or item.must_complete:
self.priority_tasks.discard(item) await item.run()
if self.stop and len(self.priority_tasks) == 0: self.queue.task_done()
print('All priority tasks completed') self.priority_tasks.discard(item)
self.cancel() if self.stop and len(self.priority_tasks) == 0:
break logger.info('All priority tasks completed')
self.cancel()
break
except Exception as err:
logger.error(f"Error {err} occurred in worker")
def sigint_handle(self, sig, frame): def sigint_handle(self, sig, frame):
print('SIGINT received, cleaning up...') logger.info('SIGINT received, cleaning up...')
if self.on_exit == 'complete_priority' and self.priority_tasks: if self.on_exit == 'complete_priority' and self.priority_tasks:
print(f'Completing {len(self.priority_tasks)} priority tasks...') logger.info(f'Completing {len(self.priority_tasks)} priority tasks...')
self.stop = True self.stop = True
else: else:
self.cancel() self.cancel()
@@ -100,26 +101,24 @@ class TaskQueue:
await asyncio.wait_for(task, timeout = timeout or self.timeout) await asyncio.wait_for(task, timeout = timeout or self.timeout)
except TimeoutError: except TimeoutError:
print(f"Timed out after {loop.time() - start} seconds. {self.queue.qsize()} tasks remaining") logger.warning(f"Timed out after {loop.time() - start} seconds. {self.queue.qsize()} tasks remaining")
if self.on_exit == 'complete_priority' and self.priority_tasks: if self.on_exit == 'complete_priority' and self.priority_tasks:
print(f'Completing {len(self.priority_tasks)} priority tasks...') logger.info(f'Completing {len(self.priority_tasks)} priority tasks...')
self.stop = True self.stop = True
await self.queue.join() await self.queue.join()
else: else:
self.cancel() self.cancel()
except asyncio.CancelledError: except asyncio.CancelledError:
print('Tasks cancelled') logger.debug('All tasks cancelled')
finally: finally:
print(f'Exiting queue after {(loop.time() - start)} seconds.' logger.info(f'Exiting queue after {(loop.time() - start)} seconds.'
f'{self.queue.qsize()} tasks remaining, {len(self.priority_tasks)} are priority tasks') f'{self.queue.qsize()} tasks remaining, {len(self.priority_tasks)} are priority tasks')
self.cancel() self.cancel()
def cancel(self): def cancel(self):
cancelled = [task.cancel() for task in self.tasks if not task.done()] [task.cancel() for task in self.tasks if not task.done()]
print(f'Cancelled {len(cancelled)} worker tasks') if cancelled else ...
self.tasks.clear() self.tasks.clear()
+17 -4
View File
@@ -1,4 +1,17 @@
from .strategies import * from .account import Account
from .traders import * from .backtest_runner import BackTestRunner
from .symbols import * from .bot_factory import Bot
from .candle_patterns import * from .candle import Candle, Candles
from .executor import Executor
from .history import History
from .order import Order
from .positions import Positions
from .ram import RAM
from .result import Result
from .symbol import Symbol
from .ticks import Tick, Ticks
from .trader import Trader
from .strategy import Strategy
from .sessions import Sessions, Session
from .trade_records import TradeRecords
from .terminal import Terminal
@@ -1,13 +1,14 @@
from logging import getLogger from logging import getLogger
from typing import Self from typing import Self
from .core.models import AccountInfo from ..core.base import _Base
from .core.exceptions import LoginError from ..core.models import AccountInfo
from ..core.exceptions import LoginError
logger = getLogger(__name__) logger = getLogger(__name__)
class Account(AccountInfo): class Account(_Base, AccountInfo):
"""A class for managing a trading account. A singleton class. """A class for managing a trading account. A singleton class.
A subclass of AccountInfo. All AccountInfo attributes are available in this class. A subclass of AccountInfo. All AccountInfo attributes are available in this class.
@@ -23,15 +24,9 @@ class Account(AccountInfo):
def __new__(cls, *args, **kwargs): def __new__(cls, *args, **kwargs):
if not hasattr(cls, '_instance'): if not hasattr(cls, '_instance'):
cls._instance = super().__new__(cls) cls._instance = super().__new__(cls)
cls._instance.exclude = cls._instance.exclude | {'_instance'} cls._instance.connected = False
return cls._instance return cls._instance
async def refresh(self):
"""Refreshes the account instance with the latest account details from the MetaTrader 5 terminal"""
account_info = await self.mt5.account_info()
acc = account_info._asdict()
self.set_attributes(**acc)
async def __aenter__(self) -> Self: async def __aenter__(self) -> Self:
"""Connect to a trading account and return the account instance. """Connect to a trading account and return the account instance.
Async context manager for the Account class. Async context manager for the Account class.
@@ -42,6 +37,7 @@ class Account(AccountInfo):
Raises: Raises:
LoginError: If login fails LoginError: If login fails
""" """
await self.mt5.initialize()
self.connected = await self.mt5.login() self.connected = await self.mt5.login()
if not self: if not self:
raise LoginError('Login failed') raise LoginError('Login failed')
@@ -50,3 +46,10 @@ class Account(AccountInfo):
async def __aexit__(self, exc_type, exc_val, exc_tb): async def __aexit__(self, exc_type, exc_val, exc_tb):
await self.mt5.shutdown() await self.mt5.shutdown()
self.connected = False self.connected = False
async def refresh(self):
"""Refreshes the account instance with the latest account details from the MetaTrader 5 terminal"""
account_info = await self.mt5.account_info()
acc = account_info._asdict()
self.connected = True
self.set_attributes(**acc)
@@ -2,27 +2,28 @@ import asyncio
import signal import signal
from logging import getLogger from logging import getLogger
from .event_manager import EventManager from ..core.event_manager import EventManager
from .meta_tester import MetaTester from ..contrib.backtesting.backtest_engine import BackTestEngine
from .backtest_engine import BackTestEngine from .strategy import Strategy
from .strategy_tester import StrategyTester from ..core.meta_backtester import MetaBackTester
logger = getLogger(__name__) logger = getLogger(__name__)
class BackTester: class BackTestRunner:
def __init__(self, *, strategies: list[StrategyTester] = None, backtest_engine: BackTestEngine = None): def __init__(self, *, strategies: list[Strategy] = None, backtest_engine: BackTestEngine = None):
self.strategies = strategies or [] self.strategies = strategies or []
self.event_manager = EventManager() self.event_manager = EventManager()
self.mt5 = MetaTester(backtest_engine=backtest_engine) self.mt5 = MetaBackTester(backtest_engine=backtest_engine)
signal.signal(signal.SIGINT, self.event_manager.sigint_handler) signal.signal(signal.SIGINT, self.event_manager.sigint_handler)
async def run(self): async def run(self):
try: try:
await self.mt5.initialize() await self.mt5.initialize()
await self.mt5.login()
strategies = [strategy for strategy in self.strategies if await strategy.symbol.init()] strategies = [strategy for strategy in self.strategies if await strategy.symbol.init()]
self.event_manager.num_main_tasks = len(strategies) self.event_manager.num_main_tasks = len(strategies)
tasks = [*[asyncio.create_task(strategy.test()) for strategy in strategies], tasks = [*[asyncio.create_task(strategy.run_strategy()) for strategy in strategies],
asyncio.create_task(self.event_manager.event_monitor())] asyncio.create_task(self.event_manager.event_monitor())]
self.event_manager.add_tasks(*tasks) self.event_manager.add_tasks(*tasks)
await asyncio.gather(*tasks, return_exceptions=True) if strategies else ... await asyncio.gather(*tasks, return_exceptions=True) if strategies else ...
@@ -4,8 +4,8 @@ from typing import Type, Iterable, Callable, Coroutine
import logging import logging
from .executor import Executor from .executor import Executor
from .core.config import Config from ..core.config import Config
from .core.meta_trader import MetaTrader from ..core.meta_trader import MetaTrader
from .symbol import Symbol as Symbol from .symbol import Symbol as Symbol
from .strategy import Strategy as Strategy from .strategy import Strategy as Strategy
@@ -51,37 +51,38 @@ class Bot:
SystemExit if sign in was not successful SystemExit if sign in was not successful
""" """
try: try:
await self.mt.initialize()
login = await self.mt.login() login = await self.mt.login()
if not login: if not login:
logger.warning(f"Unable to sign in to MetaTrder 5 Terminal") logger.critical(f"Unable to sign in to MetaTrder 5 Terminal")
raise SystemExit raise SystemExit
logger.info("Login Successful") logger.info("Login Successful")
await self.init_strategies() await self.init_strategies()
self.add_coroutine(self.config.task_queue.start) self.add_coroutine(coroutine=self.config.task_queue.run)
except Exception as err: except Exception as err:
logger.error(f"{err}. Bot initialization failed") logger.error(f"{err}. Bot initialization failed")
raise SystemExit raise SystemExit
def add_function(self, func: Callable[..., ...], **kwargs: dict): def add_function(self, *, function: Callable[..., ...], **kwargs: dict):
"""Add a function to the executor. """Add a function to the executor.
Args: Args:
func (Callable): A function to be executed function (Callable): A function to be executed
**kwargs (dict): Keyword arguments for the function **kwargs (dict): Keyword arguments for the function
""" """
self.executor.add_function(func, kwargs) self.executor.add_function(function=function, kwargs=kwargs)
def add_coroutine(self, coro: Coroutine[..., ...], **kwargs): def add_coroutine(self, *, coroutine: Callable[..., ...] | Coroutine, **kwargs):
"""Add a coroutine to the executor. """Add a coroutine to the executor.
Args: Args:
coro (Coroutine): A coroutine to be executed coroutine (Coroutine): A coroutine to be executed
**kwargs (dict): keyword arguments for the coroutine **kwargs (dict): keyword arguments for the coroutine
Returns: Returns:
""" """
self.executor.add_coroutine(coro, kwargs) self.executor.add_coroutine(coroutine=coroutine, kwargs=kwargs)
def execute(self): def execute(self):
"""Execute the bot.""" """Execute the bot."""
@@ -92,7 +93,7 @@ class Bot:
await self.initialize() await self.initialize()
await self.executor.execute() await self.executor.execute()
def add_strategy(self, strategy: Strategy): def add_strategy(self, *, strategy: Strategy):
"""Add a strategy to the executor. An added strategy will only run if it's symbol was successfully initialized. """Add a strategy to the executor. An added strategy will only run if it's symbol was successfully initialized.
Args: Args:
@@ -101,15 +102,15 @@ class Bot:
Notes: Notes:
Make sure the symbol has been added to the market Make sure the symbol has been added to the market
""" """
self.executor.add_worker(strategy) self.executor.add_strategy(strategy=strategy)
def add_strategies(self, strategies: Iterable[Strategy]): def add_strategies(self, *, strategies: Iterable[Strategy]):
"""Add multiple strategies at the same time """Add multiple strategies at the same time
Args: Args:
strategies: A list of strategies strategies: A list of strategies
""" """
[self.add_strategy(strategy) for strategy in strategies] [self.add_strategy(strategy=strategy) for strategy in strategies]
def add_strategy_all(self, *, strategy: Type[Strategy], params: dict | None = None, def add_strategy_all(self, *, strategy: Type[Strategy], params: dict | None = None,
symbols: list[Symbol] = None, **kwargs): symbols: list[Symbol] = None, **kwargs):
@@ -122,21 +123,21 @@ class Bot:
**kwargs: Additional keyword arguments for the strategy **kwargs: Additional keyword arguments for the strategy
""" """
[ [
self.add_strategy(strategy(symbol=symbol, params=params, **kwargs)) self.add_strategy(strategy=strategy(symbol=symbol, params=params, **kwargs))
for symbol in symbols for symbol in symbols
] ]
@staticmethod @staticmethod
async def init_strategy(strategy: Strategy) -> tuple[bool, Strategy]: async def init_strategy(*, strategy: Strategy) -> tuple[bool, Strategy]:
"""Initialize a single strategy. This method is called internally by the bot.""" """Initialize a single strategy. This method is called internally by the bot."""
res = await strategy.symbol.init() res = await strategy.symbol.init()
return res, strategy return res, strategy
async def init_strategies(self): async def init_strategies(self):
"""Initialize the symbols for the current trading session. This method is called internally by the bot.""" """Initialize the symbols for the current trading session. This method is called internally by the bot."""
tasks = [self.init_strategy(strategy) for strategy in self.executor.workers] tasks = [self.init_strategy(strategy=strategy) for strategy in self.executor.strategy_runners]
for task in asyncio.as_completed(tasks): for task in asyncio.as_completed(tasks):
res = await task res = await task
if not res[0]: if not res[0]:
logger.warning(f"Failed to initialize symbol {res[1].symbol}") logger.warning(f"Failed to initialize symbol {res[1].symbol}")
self.executor.workers.remove(res[1]) self.executor.strategy_runners.remove(res[1])
@@ -1,6 +1,6 @@
"""Candle and Candles classes for handling bars from the MetaTrader 5 terminal.""" """Candle and Candles classes for handling bars from the MetaTrader 5 terminal."""
import time
from typing import Type, TypeVar, Generic, Iterable from typing import Type, Self, Iterable
from logging import getLogger from logging import getLogger
from pandas import DataFrame, Series from pandas import DataFrame, Series
@@ -8,7 +8,7 @@ import pandas as pd
import pandas_ta as ta import pandas_ta as ta
import mplfinance as mplt import mplfinance as mplt
from .core.constants import TimeFrame from ..core.constants import TimeFrame
logger = getLogger(__name__) logger = getLogger(__name__)
@@ -47,8 +47,11 @@ class Candle:
""" """
if not all(i in kwargs for i in ['open', 'high', 'low', 'close']): if not all(i in kwargs for i in ['open', 'high', 'low', 'close']):
raise ValueError("Candle must be instantiated with open, high, low and close prices") raise ValueError("Candle must be instantiated with open, high, low and close prices")
self.time = kwargs.pop('time', 0) self.time = kwargs.pop('time', time.monotonic_ns())
self.Index = kwargs.pop('Index', 0) self.Index = kwargs.pop('Index', 0)
self.real_volume = kwargs.pop('real_volume', 0)
self.spread = kwargs.pop('spread', 0)
self.tick_volume = kwargs.pop('tick_volume', 0)
self.set_attributes(**kwargs) self.set_attributes(**kwargs)
def __repr__(self): def __repr__(self):
@@ -56,27 +59,32 @@ class Candle:
% {"class": self.__class__.__name__, "open": self.open, "high": self.high, % {"class": self.__class__.__name__, "open": self.open, "high": self.high,
"low": self.low, "close": self.close, "time": self.time, 'Index': self.Index}) "low": self.low, "close": self.close, "time": self.time, 'Index': self.Index})
def __str__(self): def __eq__(self, other: Self):
return str(self.dict()) return self.time == other.time
def __eq__(self, other: "Candle"): def __lt__(self, other: Self):
eq = self.open == other.open and self.high == other.high and self.low == other.low and self.close == other.close return self.time < other.time
return eq
def __hash__(self): def __hash__(self):
return hash(self.time) return hash(self.time)
def __lt__(self, other: "Candle"):
return self.time < other.time
def __gt__(self, other: "Candle"):
return self.time > other.time
def __getitem__(self, item): def __getitem__(self, item):
return self.__dict__[item] return self.__dict__[item]
def __setitem__(self, key, value):
self.__dict__[key] = value
def __iter__(self):
return iter(self.__dict__.items())
def keys(self):
return self.__dict__.keys()
def values(self):
return self.__dict__.values()
def set_attributes(self, **kwargs): def set_attributes(self, **kwargs):
"""Set keyword arguments as instance attributes """Set keyword arguments as instance attributes and values.
Keyword Args: Keyword Args:
**kwargs: Instance attributes and values as keyword arguments **kwargs: Instance attributes and values as keyword arguments
@@ -99,7 +107,7 @@ class Candle:
""" """
return self.open > self.close return self.open > self.close
def dict(self, exclude: set = None, include: set = None) -> dict: def dict(self, *, exclude: set = None, include: set = None) -> dict:
""" """
Returns a dictionary of the instance attributes. Returns a dictionary of the instance attributes.
@@ -112,14 +120,10 @@ class Candle:
exclude = exclude or set() exclude = exclude or set()
include = include or set() include = include or set()
keys = include or set(self.__dict__.keys()).difference(exclude) keys = include or set(self.__dict__.keys()).difference(exclude)
return {k: v for k, v in self.__dict__.items() if k in keys} return {k: v for k, v in self if k in keys}
_Candle = TypeVar("_Candle", bound=Candle) class Candles:
_Candles = TypeVar("_Candles", bound="Candles")
class Candles(Generic[_Candle]):
"""An iterable container class of Candle objects in chronological order. """An iterable container class of Candle objects in chronological order.
Attributes: Attributes:
@@ -155,7 +159,7 @@ class Candles(Generic[_Candle]):
timeframe: TimeFrame timeframe: TimeFrame
_data: DataFrame _data: DataFrame
def __init__(self, *, data: DataFrame | _Candles | Iterable, flip=False, candle_class: Type[_Candle] = None): def __init__(self, *, data: DataFrame | Self | Iterable, flip=False, candle_class: Candle = None):
"""A container class of Candle objects in chronological order. """A container class of Candle objects in chronological order.
Args: Args:
@@ -183,10 +187,10 @@ class Candles(Generic[_Candle]):
def __len__(self): def __len__(self):
return len(self._data.index) return len(self._data.index)
def __contains__(self, item: _Candle): def __contains__(self, item: Self):
return item.time == self[item.Index].time return item.time == self[item.Index].time
def __getitem__(self, index) -> _Candle | _Candles | Series: def __getitem__(self, index) -> Self | Self | Series:
if isinstance(index, slice): if isinstance(index, slice):
cls = self.__class__ cls = self.__class__
data = self._data.iloc[index] data = self._data.iloc[index]
@@ -212,6 +216,7 @@ class Candles(Generic[_Candle]):
def __getattr__(self, item): def __getattr__(self, item):
if item in self._data.columns: if item in self._data.columns:
return self._data[item] return self._data[item]
if item == 'Index': if item == 'Index':
return Series(self._data.index) return Series(self._data.index)
raise AttributeError(f"Attribute {item} not defined on class {self.__class__.__name__}") raise AttributeError(f"Attribute {item} not defined on class {self.__class__.__name__}")
@@ -224,6 +229,10 @@ class Candles(Generic[_Candle]):
tf = self.time[1] - self.time[0] tf = self.time[1] - self.time[0]
return TimeFrame.get(abs(tf)) return TimeFrame.get(abs(tf))
@property
def columns(self) -> DataFrame:
return self._data.columns
@property @property
def ta(self): def ta(self):
"""Access to the pandas_ta library for performing technical analysis on the underlying data attribute. """Access to the pandas_ta library for performing technical analysis on the underlying data attribute.
@@ -247,7 +256,7 @@ class Candles(Generic[_Candle]):
"""The original data passed to the class as a pandas DataFrame""" """The original data passed to the class as a pandas DataFrame"""
return self._data return self._data
def rename(self, inplace=True, **kwargs) -> _Candles: def rename(self, inplace=True, **kwargs) -> Self:
"""Rename columns of the candles class. """Rename columns of the candles class.
Keyword Args: Keyword Args:
@@ -274,13 +283,13 @@ class Candles(Generic[_Candle]):
data.index = pd.to_datetime(data['time'], unit='s') data.index = pd.to_datetime(data['time'], unit='s')
return mplt.make_addplot(data[columns], **kwargs) return mplt.make_addplot(data[columns], **kwargs)
def visualize(self, *, count: int = 50, type='candle', savefig: str | dict = None, addplot: dict = None, def visualize(self, *, count: int = 50, _type='candle', savefig: str | dict = None, addplot: dict = None,
style: str = 'charles', ylabel: str = 'Price', title: str = 'Chart', **kwargs): style: str = 'charles', ylabel: str = 'Price', title: str = 'Chart', **kwargs):
"""Visualize the candles using the mplfinance library. """Visualize the candles using the mplfinance library.
Args: Args:
count (int): The number of candles to visualize, counting from behind, i.e the most recent candles. count (int): The number of candles to visualize, counting from behind, i.e the most recent candles.
Defaults to 50. Defaults to 50.
type: Type of chart, defaults to candle _type: Type of chart, defaults to candle
savefig (str|dict): The path to save the figure or a dictionary of parameters to pass to the savefig method. savefig (str|dict): The path to save the figure or a dictionary of parameters to pass to the savefig method.
addplot: Additional plots to add to the chart. Defaults to None. They should match the dimension of the addplot: Additional plots to add to the chart. Defaults to None. They should match the dimension of the
original data which is specified via the count parameter. original data which is specified via the count parameter.
@@ -290,7 +299,7 @@ class Candles(Generic[_Candle]):
kwargs: valid kwargs for the plot function. kwargs: valid kwargs for the plot function.
""" """
kwargs |= {key: arg for key, arg in (('savefig', savefig), ('addplot', addplot), ('style', style), kwargs |= {key: arg for key, arg in (('savefig', savefig), ('addplot', addplot), ('style', style),
('ylabel', ylabel), ('title', title), ('type', type)) if arg} ('ylabel', ylabel), ('title', title), ('type', _type)) if arg}
data = self._data[-count:] data = self._data[-count:]
data.index = pd.to_datetime(data['time'], unit='s') data.index = pd.to_datetime(data['time'], unit='s')
mplt.plot(data, **kwargs) mplt.plot(data, **kwargs)
@@ -13,49 +13,50 @@ class Executor:
Attributes: Attributes:
executor (ThreadPoolExecutor): The executor object. executor (ThreadPoolExecutor): The executor object.
workers (list): List of strategies. strategy_runners (list): List of strategies.
coroutines (dict[Coroutine, dict]): A dictionary of coroutines to run in the executor coroutines (dict[Coroutine, dict]): A dictionary of coroutines to run in the executor
functions (dict[Callable, dict]): A dictionary of functions to run in the executor functions (dict[Callable, dict]): A dictionary of functions to run in the executor
loop (asyncio.AbstractEventLoop): The event loop
""" """
loop: asyncio.AbstractEventLoop
def __init__(self): def __init__(self):
self.executor = ThreadPoolExecutor self.executor = ThreadPoolExecutor
self.workers: list[Strategy] = [] self.strategy_runners: list[Strategy] = []
self.coroutines: dict[Coroutine | Callable: dict] = {} self.coroutines: dict[Coroutine | Callable: dict] = {}
self.functions: dict[Callable: dict] = {} self.functions: dict[Callable: dict] = {}
def add_function(self, func: Callable, kwargs: dict): def add_function(self, *, function: Callable, kwargs: dict):
self.functions[func] = kwargs self.functions[function] = kwargs
def add_coroutine(self, coro: Coroutine, kwargs: dict): def add_coroutine(self, *, coroutine: Coroutine, kwargs: dict):
self.coroutines[coro] = kwargs self.coroutines[coroutine] = kwargs
def add_workers(self, strategies: tuple[Strategy]): def add_strategies(self, *, strategies: tuple[Strategy]):
"""Add multiple strategies at once """Add multiple strategies at once
Args: Args:
strategies (Sequence[Strategy]): A sequence of strategies. strategies (Sequence[Strategy]): A sequence of strategies.
""" """
self.workers.extend(strategies) self.strategy_runners.extend(strategies)
def add_worker(self, strategy: Strategy): def add_strategy(self, *, strategy: Strategy):
"""Add a strategy instance to the list of workers """Add a strategy instance to the list of workers
Args: Args:
strategy (Strategy): A strategy object strategy (Strategy): A strategy object
""" """
self.workers.append(strategy) self.strategy_runners.append(strategy)
@staticmethod def run_strategy(self, strategy: Strategy):
def trade(strategy: Strategy):
"""Wraps the coroutine trade method of each strategy with 'asyncio.run'. """Wraps the coroutine trade method of each strategy with 'asyncio.run'.
Args: Args:
strategy (Strategy): A strategy object strategy (Strategy): A strategy object
""" """
asyncio.run(strategy.trade()) self.loop.run_until_complete(strategy.run_strategy())
def run(self, func, kwargs: dict): def run_coroutine(self, func, kwargs: dict):
""" """
Run a coroutine function Run a coroutine function
@@ -64,11 +65,11 @@ class Executor:
kwargs: A dictionary of keyword arguments for the function kwargs: A dictionary of keyword arguments for the function
""" """
try: try:
asyncio.run(func(**kwargs)) self.loop.run_until_complete(func(**kwargs))
except Exception as err: except Exception as err:
logger.error(f'Error: {err}. Unable to run function') logger.error(f'Error: {err}. Unable to run function')
async def execute(self, workers: int = 5): async def execute(self, *, workers: int = 5):
"""Run the strategies with a threadpool executor. """Run the strategies with a threadpool executor.
Args: Args:
@@ -77,10 +78,11 @@ class Executor:
Notes: Notes:
No matter the number specified, the executor will always use a minimum of 5 workers. No matter the number specified, the executor will always use a minimum of 5 workers.
""" """
workers_ = sum([len(self.workers), len(self.functions), len(self.coroutines)]) workers_ = sum([len(self.strategy_runners), len(self.functions), len(self.coroutines)])
workers = max(workers, workers_) workers = max(workers, workers_)
loop = asyncio.get_running_loop() self.loop = asyncio.get_running_loop()
with self.executor(max_workers=workers) as executor: with self.executor(max_workers=workers) as executor:
[loop.run_in_executor(executor, self.trade, worker) for worker in self.workers] [self.loop.run_in_executor(executor, self.run_strategy, worker) for worker in self.strategy_runners]
[loop.run_in_executor(executor, self.run, coro, kwargs) for coro, kwargs in self.coroutines.items()] [self.loop.run_in_executor(executor, self.run_coroutine, coroutine, kwargs) for coroutine,
[loop.run_in_executor(executor, func, kwargs) for func, kwargs in self.functions.items()] kwargs in self.coroutines.items()]
[self.loop.run_in_executor(executor, function, kwargs) for function, kwargs in self.functions.items()]
@@ -6,11 +6,11 @@ import pytz
from pandas import DataFrame from pandas import DataFrame
import pandas as pd import pandas as pd
from .core.config import Config from ..core.config import Config
from .core.meta_trader import MetaTrader, CopyTicks, OrderType from ..core.meta_trader import MetaTrader, CopyTicks, OrderType
from .core.models import TradeDeal, TradeOrder from ..core.models import TradeDeal, TradeOrder
from .contrib.backtester.meta_tester import MetaTester from ..core.meta_backtester import MetaBackTester
from .utils import backoff_decorator from .._utils import backoff_decorator
logger = getLogger(__name__) logger = getLogger(__name__)
@@ -27,7 +27,7 @@ class History:
mt5 (MetaTrader): MetaTrader instance mt5 (MetaTrader): MetaTrader instance
config (Config): Config instance config (Config): Config instance
""" """
mt5: MetaTrader | MetaTester mt5: MetaTrader | MetaBackTester
config: Config config: Config
def __init__(self, *, date_from: datetime | int, date_to: datetime | int, group: str = '', use_utc: bool = True): def __init__(self, *, date_from: datetime | int, date_to: datetime | int, group: str = '', use_utc: bool = True):
@@ -42,11 +42,11 @@ class History:
group (str): Filter for selecting history by symbols. Defaults to an empty string group (str): Filter for selecting history by symbols. Defaults to an empty string
""" """
self.config = Config() self.config = Config()
self.mt5 = MetaTrader() if self.config.mode == 'live' else MetaTester() self.mt5 = MetaTrader() if self.config.mode != 'backtest' else MetaBackTester()
date_from = date_from if isinstance(date_from, datetime) else datetime.fromtimestamp(date_from) date_from = date_from if isinstance(date_from, datetime) else datetime.fromtimestamp(date_from)
date_to = date_to if isinstance(date_to, datetime) else datetime.fromtimestamp(date_to) date_to = date_to if isinstance(date_to, datetime) else datetime.fromtimestamp(date_to)
self.date_from = date_from.astimezone(pytz.utc) if use_utc else date_from self.date_from = date_from.astimezone(pytz.UTC) if use_utc else date_from
self.date_to = date_to.astimezone(pytz.utc) if use_utc else date_to self.date_to = date_to.astimezone(pytz.UTC) if use_utc else date_to
self.group = group self.group = group
self.deals: tuple[TradeDeal, ...] = () self.deals: tuple[TradeDeal, ...] = ()
self.orders: tuple[TradeOrder, ...] = () self.orders: tuple[TradeOrder, ...] = ()
@@ -111,12 +111,12 @@ class History:
logger.warning(f'Failed to get orders') logger.warning(f'Failed to get orders')
return tuple() return tuple()
def get_orders_by_ticket(self, ticket: int) -> tuple[TradeOrder, ...]: def get_orders_by_ticket(self, *, ticket: int) -> tuple[TradeOrder, ...]:
"""filter orders by ticket""" """filter orders by ticket"""
return tuple(sorted((order for order in self.orders if order.ticket == ticket), key=lambda x: x.time_done_msc)) return tuple(sorted((order for order in self.orders if order.ticket == ticket), key=lambda x: x.time_done_msc))
def get_orders_by_position(self, position: int) -> tuple[TradeOrder, ...]: def get_orders_by_position(self, *, position: int) -> tuple[TradeOrder, ...]:
""" filter orders by position""" """ filter orders by position"""
return tuple(sorted((order for order in self.orders if order.position_id == position), return tuple(sorted((order for order in self.orders if order.position_id == position),
key=lambda x: x.time_done_msc)) key=lambda x: x.time_done_msc))
@@ -1,15 +1,15 @@
from logging import getLogger from logging import getLogger
from .core.models import TradeRequest, TradeOrder from ..core.models import TradeRequest, TradeOrder, OrderCheckResult, OrderSendResult
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 .utils import backoff_decorator, error_handler from ..core.base import _Base
from MetaTrader5 import OrderCheckResult, OrderSendResult from .._utils import backoff_decorator, error_handler
logger = getLogger(__name__) logger = getLogger(__name__)
class Order(TradeRequest): class Order(_Base, TradeRequest):
"""Trade order related functions and properties. Subclass of TradeRequest.""" """Trade order related functions and properties. Subclass of TradeRequest."""
def __init__(self, **kwargs): def __init__(self, **kwargs):
@@ -25,7 +25,7 @@ class Order(TradeRequest):
type_time (OrderTime.DAY): Order time type_time (OrderTime.DAY): Order time
type_filling (OrderFilling.FOK): Order filling type_filling (OrderFilling.FOK): Order filling
""" """
kwargs = {'action': TradeAction.DEAL, OrderTime.DAY: self.type_time, 'type_filling': OrderFilling.FOK, **kwargs} kwargs = {'action': TradeAction.DEAL, 'type_time': OrderTime.DAY, 'type_filling': OrderFilling.FOK, **kwargs}
super().__init__(**kwargs) super().__init__(**kwargs)
async def orders_total(self): async def orders_total(self):
@@ -81,7 +81,7 @@ class Order(TradeRequest):
res = await self.mt5.order_check(req) res = await self.mt5.order_check(req)
if res is None: if res is None:
raise OrderError(f'Order check failed for {self.symbol}') raise OrderError(f'Order check failed for {self.symbol}')
return res return OrderCheckResult(**res._asdict())
@backoff_decorator @backoff_decorator
async def send(self) -> OrderSendResult: async def send(self) -> OrderSendResult:
@@ -96,7 +96,7 @@ class Order(TradeRequest):
res = await self.mt5.order_send(self.dict) res = await self.mt5.order_send(self.dict)
if res is None: if res is None:
raise OrderError(f'Failed to send order {self.symbol}') raise OrderError(f'Failed to send order {self.symbol}')
return res return OrderSendResult(**res._asdict())
async def calc_margin(self) -> float | None: async def calc_margin(self) -> float | None:
"""Return the required margin in the account currency to perform a specified trading operation. """Return the required margin in the account currency to perform a specified trading operation.
@@ -2,13 +2,14 @@
import asyncio import asyncio
from logging import getLogger from logging import getLogger
from .core.meta_trader import MetaTrader 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 .contrib.backtester.meta_tester import MetaTester from .._utils import backoff_decorator
from ..core.meta_backtester import MetaBackTester
from .order import Order from .order import Order
from .utils import backoff_decorator
logger = getLogger(__name__) logger = getLogger(__name__)
@@ -19,13 +20,13 @@ class Positions:
Attributes: Attributes:
mt5 (MetaTrader): MetaTrader instance. mt5 (MetaTrader): MetaTrader instance.
""" """
mt5: MetaTrader | MetaTester mt5: MetaTrader | MetaBackTester
positions: tuple[TradePosition, ...] positions: tuple[TradePosition, ...]
def __init__(self): def __init__(self):
"""Get Open Positions""" """Get Open Positions"""
self.config = Config() self.config = Config()
self.mt5 = MetaTrader() if self.config.mode == 'live' else MetaTester() self.mt5 = MetaTrader() if self.config.mode != 'backtest' else MetaBackTester()
self.positions = () self.positions = ()
@backoff_decorator @backoff_decorator
@@ -68,7 +69,7 @@ class Positions:
return tuple(TradePosition(**pos._asdict()) for pos in (positions or ())) return tuple(TradePosition(**pos._asdict()) for pos in (positions or ()))
@staticmethod @staticmethod
async def close(ticket: int, symbol: str, price: float, volume: float, order_type: OrderType) -> OrderSendResult: async def close(*, ticket: int, symbol: str, price: float, volume: float, order_type: OrderType) -> OrderSendResult:
"""Close an open position for the trading account using the ticket and other parameters. """Close an open position for the trading account using the ticket and other parameters.
Args: Args:
@@ -83,14 +84,14 @@ class Positions:
return await order.send() return await order.send()
@staticmethod @staticmethod
async def close_by(pos: TradePosition) -> OrderSendResult: async def close_by(*, position: TradePosition) -> OrderSendResult:
"""Close an open position for the trading account.""" """Close an open position for the trading account."""
order = Order(position=pos.ticket, symbol=pos.symbol, volume=pos.volume, type=pos.type.opposite, order = Order(position=position.ticket, symbol=position.symbol, volume=position.volume,
price=pos.price_current, action=TradeAction.DEAL) type=position.type.opposite, price=position.price_current, action=TradeAction.DEAL)
return await order.send() return await order.send()
@staticmethod @staticmethod
async def close_position(position: TradePosition): async def close_position(*, position: TradePosition):
"""Close an open position for the trading account. Using a position object.""" """Close an open position for the trading account. Using a position object."""
order = Order(position=position.ticket, symbol=position.symbol, volume=position.volume, order = Order(position=position.ticket, symbol=position.symbol, volume=position.volume,
type=position.type.opposite, price=position.price_current, action=TradeAction.DEAL) type=position.type.opposite, price=position.price_current, action=TradeAction.DEAL)
@@ -103,6 +104,6 @@ class Positions:
int: Return number of positions closed. int: Return number of positions closed.
""" """
positions = self.positions or await self.get_positions() positions = self.positions or await self.get_positions()
results = await asyncio.gather(*(self.close_position(position) for position in positions), results = await asyncio.gather(*(self.close_position(position=position) for position in positions),
return_exceptions=True) return_exceptions=True)
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)])
+19 -15
View File
@@ -7,23 +7,27 @@ class RAM:
account: Account account: Account
risk_to_reward: float risk_to_reward: float
risk: float risk: float
min_amount: float = 0 min_amount: float
max_amount: float = 0 max_amount: float
loss_limit: int = 3 loss_limit: int
open_limit: int = 5 open_limit: int
def __init__(self, *, risk_to_reward: float = 1, risk: float = 0.01, **kwargs): def __init__(self, **kwargs):
"""Initialize Risk Assessment and Management with the provided keyword arguments. """Initialize Risk Assessment and Management with the provided keyword arguments.
Keyword Args: Keyword Args:
risk_to_reward (float): Risk to reward ratio. Defaults to 1 risk_to_reward (float): Risk to reward ratio. Defaults to 1
risk (float): Percentage of account balance to risk per trade 0.01 # 1% risk (float): Percentage of capital to risk per trade 0.01 # 1%
kwargs: extra keyword arguments are set as object attributes kwargs: extra keyword arguments are set as object attributes
""" """
self.risk_to_reward = risk_to_reward
self.risk = risk
self.account = Account() self.account = Account()
[setattr(self, key, value) for key, value in kwargs.items()] self.positions = Positions()
self.risk_to_reward = kwargs.get('risk_to_reward', 1)
self.risk = kwargs.get('risk', 0.01)
self.min_amount = kwargs.get('min_amount', 0)
self.max_amount = kwargs.get('max_amount', 0)
self.loss_limit = kwargs.get('loss_limit', 1)
self.open_limit = kwargs.get('open_limit', 1)
async def get_amount(self) -> float: async def get_amount(self) -> float:
"""Calculate the amount to risk per trade as a percentage of margin_free. """Calculate the amount to risk per trade as a percentage of margin_free.
@@ -38,14 +42,14 @@ class RAM:
return amount return amount
async def check_losing_positions(self) -> bool: async def check_losing_positions(self) -> bool:
"""Check if the number of losing positions is greater than or equal the loss limit """Check if the number of losing positions is greater than the loss limit
Returns: Returns:
bool: True if the number of losing positions is less than the loss limit bool: True if the number of losing positions is less than or equal the loss limit
""" """
positions = await Positions().get_positions() positions = await self.positions.get_positions()
loosing = [position for position in positions if position.profit < 0] loosing = [position for position in positions if position.profit < 0]
return len(loosing) < self.loss_limit return len(loosing) <= self.loss_limit
async def check_open_positions(self) -> bool: async def check_open_positions(self) -> bool:
"""Check if the number of open positions is greater than or equal the loss limit. """Check if the number of open positions is greater than or equal the loss limit.
@@ -53,5 +57,5 @@ class RAM:
Returns: Returns:
bool: True if the number of open positions is less than the open limit bool: True if the number of open positions is less than the open limit
""" """
positions = await Positions().get_positions() positions = await self.positions.get_positions()
return len(positions) < self.open_limit return len(positions) <= self.open_limit
@@ -4,10 +4,8 @@ from logging import getLogger
from typing import Iterable, Literal from typing import Iterable, Literal
from asyncio import Lock from asyncio import Lock
from _typeshed import SupportsWrite, SupportsRead from ..core.config import Config
from ..core.models import OrderSendResult
from .core.config import Config
from .core.models import OrderSendResult
logger = getLogger(__name__) logger = getLogger(__name__)
@@ -22,7 +20,7 @@ class Result:
""" """
config: Config config: Config
def __init__(self, result: OrderSendResult, parameters: dict = None, name: str = ''): def __init__(self, *, result: OrderSendResult, parameters: dict = None, name: str = ''):
""" """
Prepare result data Prepare result data
Args: Args:
@@ -33,7 +31,7 @@ class Result:
self.config = Config() self.config = Config()
self.parameters = parameters or {} self.parameters = parameters or {}
self.result = result self.result = result
self.name = name or parameters.get('name', 'Trades') self.name = name or self.parameters.get('name', 'Trades')
self.lock = Lock() self.lock = Lock()
def get_data(self) -> dict: def get_data(self) -> dict:
@@ -48,8 +46,10 @@ class Result:
trade_record_mode = trade_record_mode or self.config.trade_record_mode trade_record_mode = trade_record_mode or self.config.trade_record_mode
if trade_record_mode == 'csv': if trade_record_mode == 'csv':
await self.to_csv() await self.to_csv()
else: elif trade_record_mode == 'json':
await self.to_json() await self.to_json()
else:
logger.error(f"Invalid trade record mode: {trade_record_mode}")
async def to_csv(self): async def to_csv(self):
"""Record trade results and associated parameters as a csv file """Record trade results and associated parameters as a csv file
@@ -61,13 +61,13 @@ class Result:
file.touch(exist_ok=True) if not file.exists() else ... file.touch(exist_ok=True) if not file.exists() else ...
read_file = file.open('r', newline='') read_file = file.open('r', newline='')
reader: Iterable[dict] = csv.DictReader(read_file) reader: Iterable[dict] = csv.DictReader(read_file)
read_file.close()
rows: list[dict] = [] rows: list[dict] = []
headers = set() headers = set()
[(rows.append(row), headers.update(row.keys())) for row in reader] [(rows.append(row), headers.update(row.keys())) for row in reader]
rows.append(data) rows.append(data)
headers.update(data.keys()) headers.update(data.keys())
with file.open('w', newline='') as write_file: # type: SupportsWrite[str] read_file.close()
with file.open('w', newline='') as write_file:
writer = csv.DictWriter(write_file, fieldnames=headers, restval=None, extrasaction='ignore') writer = csv.DictWriter(write_file, fieldnames=headers, restval=None, extrasaction='ignore')
writer.writeheader() writer.writeheader()
writer.writerows(rows) writer.writerows(rows)
@@ -93,12 +93,16 @@ class Result:
try: try:
file = self.config.records_dir / f"{self.name}.json" file = self.config.records_dir / f"{self.name}.json"
data = self.get_data() data = self.get_data()
file.touch(exist_ok=True) if not file.exists() else ... if not file.exists():
with file.open('r') as fh: # type: SupportsRead[str] file.touch()
with file.open('w') as fh:
json.dump([], fh, indent=2)
with file.open('r') as fh:
rows = json.load(fh) rows = json.load(fh)
rows.append(data) rows.append(data)
with file.open('w') as fh: # type: SupportsWrite[str] with file.open('w') as fh:
json.dump(rows, fh, indent=2, skipkeys=True, default=self.serialize) json.dump(rows, fh, indent=2, skipkeys=True, default=self.serialize)
except Exception as err: except Exception as err:
@@ -1,19 +1,24 @@
import asyncio import asyncio
from datetime import time, timedelta, datetime from datetime import time, timedelta, datetime
from typing import Literal, Callable, Iterable from typing import Literal, Callable, Iterable, NamedTuple
from logging import getLogger from logging import getLogger
import pytz import pytz
from . import TradePosition from ..core.models import OrderSendResult, TradePosition
from .core.models import OrderSendResult from ..core.config import Config
from ..core.event_manager import EventManager
from .positions import Positions from .positions import Positions
from .core.config import Config
from.contrib.backtester.event_manager import EventManager
logger = getLogger(__name__) logger = getLogger(__name__)
class Duration(NamedTuple):
hours: int
minutes: int
seconds: int
def delta(obj: time) -> timedelta: def delta(obj: time) -> timedelta:
"""Get the timedelta of a datetime.time object. """Get the timedelta of a datetime.time object.
@@ -73,19 +78,9 @@ class Session:
self.config = Config() self.config = Config()
def __contains__(self, item: time): def __contains__(self, item: time):
if self.start > self.end: span = (delta(self.end) - delta(self.start)).seconds
end = timedelta(days=1, hours=self.start.hour, minutes=self.start.minute, seconds=self.start.second, item_span = (delta(self.end) - delta(item)).seconds
microseconds=self.start.microsecond) return item_span <= span
start = delta(self.start)
if item < self.start and item < self.end:
item = timedelta(days=1, hours=item.hour, minutes=item.minute, seconds=item.second,
microseconds=item.microsecond)
else:
item = delta(item)
else:
start = delta(self.start)
end = delta(self.end)
return start <= item < end
def __str__(self): def __str__(self):
return f'{self.start}<-->{self.end}' return f'{self.start}<-->{self.end}'
@@ -94,7 +89,7 @@ class Session:
return f'{self.start}<-->{self.end}' return f'{self.start}<-->{self.end}'
def __len__(self): def __len__(self):
return (delta(self.start) - delta(self.end)).seconds return int((delta(self.end) - delta(self.start)).seconds)
def in_session(self) -> bool: def in_session(self) -> bool:
"""Check if the current time is within the session.""" """Check if the current time is within the session."""
@@ -110,9 +105,15 @@ class Session:
"""Call the action specified in on_end or custom_end.""" """Call the action specified in on_end or custom_end."""
await self.action(action=self.on_end) await self.action(action=self.on_end)
def duration(self) -> Duration:
"""Get the duration of the session in seconds."""
hours, seconds = divmod(len(self), 3600)
minutes, seconds = divmod(seconds, 60)
return Duration(hours=hours, minutes=minutes, seconds=seconds)
async def close_positions(self, *, positions: tuple[TradePosition, ...]): async def close_positions(self, *, positions: tuple[TradePosition, ...]):
results = asyncio.gather(*(self.positions_manager.close_position(pos) for pos in positions), results = asyncio.gather(*(self.positions_manager.close_position(position=position) for position in positions),
return_exceptions=True) return_exceptions=True)
closed = pending = 0 closed = pending = 0
for result in results: for result in results:
@@ -169,7 +170,7 @@ class Session:
"""Get the seconds until the session starts from the current time in seconds.""" """Get the seconds until the session starts from the current time in seconds."""
if self.config.mode == 'backtest': if self.config.mode == 'backtest':
now = datetime.fromtimestamp(self.config.backtest_engine.cursor.time, tz=pytz.UTC).time() now = datetime.fromtimestamp(self.config.backtest_engine.cursor.time, tz=pytz.UTC).time()
secs = delta(self.start) - delta(now) secs = (delta(self.start) - delta(now)).seconds
else: else:
secs = (delta(self.start) - delta(datetime.now(tz=pytz.UTC).time())).seconds secs = (delta(self.start) - delta(datetime.now(tz=pytz.UTC).time())).seconds
return secs return secs
@@ -189,6 +190,7 @@ class Sessions:
check: Check if the current session has started and if not, wait until it starts. check: Check if the current session has started and if not, wait until it starts.
""" """
sessions: list[Session] sessions: list[Session]
current_session: Session | None
def __init__(self, *, sessions: Iterable[Session]): def __init__(self, *, sessions: Iterable[Session]):
self.sessions = list(sessions) self.sessions = list(sessions)
@@ -196,43 +198,47 @@ class Sessions:
self.current_session = None self.current_session = None
self.config = Config() self.config = Config()
def find(self, obj: time) -> Session | None: def find(self, *, moment: time = None) -> Session | None:
"""Find a session that contains a datetime.time object. """Find a session that contains a datetime.time object, if not found return None.
Args: Keyword Args:
obj (datetime.time): A datetime.time object. moment (datetime.time | None): A datetime.time object. if not provided, the current time is used.
Returns: Returns:
Session | None: A Session object or None if not found. Session | None: A Session object or None if not found.
""" """
moment = moment or datetime.now(tz=pytz.UTC).time() if self.config.mode == 'live' else (
datetime.fromtimestamp(self.config.backtest_engine.cursor.time, tz=pytz.UTC).time())
for session in self.sessions: for session in self.sessions:
if obj in session: if moment in session:
return session return session
return None return None
def find_next(self, obj: time) -> Session: def find_next(self, *, moment: time = None) -> Session:
"""Find the next session that contains a datetime.time object. """Find the next session that contains a datetime.time object.
Args: Args:
obj (datetime.time): A datetime.time object. moment (datetime.time | None): A datetime.time object, if not provided, the current time is used.
Returns: Returns:
Session: A Session object. Session: A Session object.
""" """
moment = moment or datetime.now(tz=pytz.UTC).time() if self.config.mode == 'live' else (
datetime.fromtimestamp(self.config.backtest_engine.cursor.time, tz=pytz.UTC).time())
for session in self.sessions: for session in self.sessions:
if delta(obj) < delta(session.start): if delta(moment) < delta(session.start):
return session return session
return self.sessions[0] return self.sessions[0]
def __contains__(self, item: time): def __contains__(self, moment: time):
return True if self.find(item) is not None else False return True if self.find(moment=moment) is not None else False
async def __aenter__(self): async def __aenter__(self):
await self.check() await self.check()
return self return self
async def __aexit__(self, exc_type, exc_val, exc_tb): async def __aexit__(self, exc_type, exc_val, exc_tb):
await self.current_session.close() await self.current_session.close() if self.current_session is not None else ...
async def check(self): async def check(self):
"""Check if the current session has started and if not, wait until it starts.""" """Check if the current session has started and if not, wait until it starts."""
@@ -242,9 +248,9 @@ class Sessions:
if self.config.mode == 'backtest': if self.config.mode == 'backtest':
now = datetime.fromtimestamp(self.config.backtest_engine.cursor.time, tz=pytz.UTC).time() now = datetime.fromtimestamp(self.config.backtest_engine.cursor.time, tz=pytz.UTC).time()
else: else:
now = datetime.now().time() now = datetime.now(tz=pytz.UTC).time()
next_session = self.find(now) next_session = self.find(moment=now)
if next_session and self.current_session is None: if next_session and self.current_session is None:
self.current_session = next_session self.current_session = next_session
@@ -255,11 +261,12 @@ class Sessions:
await self.current_session.close() await self.current_session.close()
self.current_session = next_session self.current_session = next_session
await self.current_session.begin() await self.current_session.begin()
return
if next_session is None and self.current_session is not None: if next_session is None and self.current_session is not None:
await self.current_session.close() await self.current_session.close()
next_session = self.find_next(now) next_session = self.find_next(moment=now)
secs = next_session.until() + 10 secs = next_session.until() + 10
logger.info(f'sleeping for {secs} seconds until next {next_session} session') logger.info(f'sleeping for {secs} seconds until next {next_session} session')
sleep_func = asyncio.sleep if self.config.mode == 'live' else backtest_sleep sleep_func = asyncio.sleep if self.config.mode == 'live' else backtest_sleep
-3
View File
@@ -1,3 +0,0 @@
from .finger_trap import FingerTrap
from .tracker import Tracker
# from .finger_trap_back_test import FingerTrapTest, FingerTrapSingleTest
@@ -1,30 +0,0 @@
from .finger_trap import FingerTrap
from ...contrib.backtester.strategy_tester import StrategyTester
class FingerTrapTest(StrategyTester, FingerTrap):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.set_up()
async def test(self):
print(f"Backtesting {self.symbol}")
while True:
await self.event_manager.acquire()
try:
await self.event_manager.wait()
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)
finally:
self.event_manager.release()
+159
View File
@@ -0,0 +1,159 @@
"""The base class for creating strategies."""
import asyncio
from time import time
from typing import TypeVar
from abc import ABC, abstractmethod
from datetime import time as dtime
from logging import getLogger
from ..core.meta_trader import MetaTrader
from ..core import Config
from ..core.meta_backtester import MetaBackTester
from ..core.event_manager import EventManager
from .sessions import Sessions, Session
from .symbol import Symbol as _Symbol
Symbol = TypeVar("Symbol", bound=_Symbol)
logger = getLogger(__name__)
class Strategy(ABC):
"""The base class for creating strategies.
Attributes:
name (str): The name of the strategy.
symbol (Symbol): The Financial Instrument as a Symbol Object
parameters (Dict): A dictionary of parameters for the strategy.
sessions (Sessions): The sessions to use for the strategy.
running (bool): A flag to indicate if the strategy is running.
Notes:
Define the name of a strategy as a class attribute. If not provided, the class name will be used as the name.
"""
name: str
symbol: Symbol
sessions: Sessions
mt5: MetaTrader | MetaBackTester
config: Config
running: bool
parameters = {}
event_manager = EventManager
current_session = Session
def __init__(self, *, symbol: Symbol, params: dict = None, sessions: Sessions = None, name=''):
"""Initiate the parameters dict and add name and symbol fields.
Use class name as strategy name if name is not provided
Args:
symbol (Symbol): The Financial instrument
params (Dict): Trading strategy parameters
"""
self.parameters = self.parameters | (params or {})
self.symbol = symbol
self.name = name or self.__class__.__name__
self.parameters["symbol"] = symbol.name
self.parameters["name"] = self.name
self.running = True
self.sessions = sessions or Sessions(sessions=[Session(start=0, end=dtime(hour=23, minute=59, second=59))])
self.config = Config()
self.mt5 = MetaTrader() if self.config.mode == 'live' else MetaBackTester()
self.event_manager = EventManager()
def __repr__(self):
return f"{self.name}({self.symbol!r})"
def __getattr__(self, item):
if item in self.parameters:
return self.parameters[item]
raise AttributeError(f'{item} not an attribute of {self.name}')
def __setattr__(self, key, value):
if key in self.parameters:
self.parameters[key] = value
super().__setattr__(key, value)
async def __aenter__(self):
await self.sessions.check()
self.current_session = self.sessions.current_session
async def __aexit__(self, exc_type, exc_val, exc_tb):
try:
await self.current_session.close() if self.current_session else ...
self.running = False
except Exception as err:
logger.error(f"Error: {err}")
@staticmethod
async def live_sleep(*, secs: float):
"""Sleep for the needed amount of seconds in between requests to the terminal.
computes the accurate amount of time needed to sleep ensuring that the next request is made at the start of
a new bar and making cooperative multitasking possible.
Args:
secs (float): The time in seconds. Usually the timeframe you are trading on.
"""
mod = time() % secs
secs = secs - mod if mod != 0 else mod
await asyncio.sleep(secs + 0.2)
async def sleep(self, *, secs: float):
"""Sleep for the needed amount of seconds in between requests to the terminal.
computes the accurate amount of time needed to sleep ensuring that the next request is made at the start of
a new bar and making cooperative multitasking possible.
Args:
secs (float): The time in seconds. Usually the timeframe you are trading on.
"""
if self.config.mode == 'live':
await self.live_sleep(secs=secs)
elif self.config.mode == 'backtest':
await self.backtest_sleep(secs=secs)
async def backtest_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:
await self.event_manager.wait()
async def run_strategy(self):
"""Run the strategy."""
if self.config.mode == 'live':
await self.live_strategy()
elif self.config.mode == 'backtest':
await self.backtest_strategy()
async def live_strategy(self):
"""Run the strategy."""
while self.running:
async with self as _:
await self.sessions.check()
await self.trade()
async def backtest_strategy(self):
"""Backtest the strategy."""
async with self as _:
while self.running:
async with self.event_manager.condition:
await self.sessions.check()
await self.event_manager.wait()
await self.trade()
@abstractmethod
async def trade(self):
"""Place trades using this method. This is the main method of the strategy.
It will be called by the strategy runner.
"""
raise NotImplementedError("Implement this method in your subclass")
async def test(self):
await self.trade()
@@ -3,19 +3,19 @@ import asyncio
from datetime import datetime from datetime import datetime
from logging import getLogger from logging import getLogger
from .core.constants import TimeFrame, CopyTicks from ..core.constants import TimeFrame, CopyTicks
from .core.models import SymbolInfo, BookInfo from ..core.base import _Base
from ..core.models import SymbolInfo, BookInfo
from .._utils import round_off, backoff_decorator
from .ticks import Tick from .ticks import Tick
from .account import Account from .account import Account
from .candle import Candles from .candle import Candles
from .ticks import Ticks from .ticks import Ticks
from .utils import round_off, backoff_decorator
logger = getLogger(__name__) logger = getLogger(__name__)
class Symbol(SymbolInfo): class Symbol(_Base, SymbolInfo):
"""Main class for handling a financial instrument. A subclass of SymbolInfo it has attributes and methods """Main class for handling a financial instrument. A subclass of SymbolInfo it has attributes and methods
for working with a financial instrument. for working with a financial instrument.
@@ -36,18 +36,10 @@ class Symbol(SymbolInfo):
Args: Args:
name (str): Name of the financial instrument name (str): Name of the financial instrument
""" """
assert 'name' in kwargs, "Symbol Object Must be initialized with a name"
super().__init__(**kwargs) super().__init__(**kwargs)
self.account = Account() self.account = Account()
@property
def pip(self):
"""Returns the pip value of the symbol. This is ten times the point value for forex symbols.
Returns:
float: The pip value of the symbol.
"""
return self.point * 10
@backoff_decorator @backoff_decorator
async def info_tick(self, *, name: str = "") -> Tick: async def info_tick(self, *, name: str = "") -> Tick:
"""Get the current price tick of a financial instrument. """Get the current price tick of a financial instrument.
@@ -152,7 +144,7 @@ class Symbol(SymbolInfo):
""" """
return await self.mt5.market_book_release(self.name) return await self.mt5.market_book_release(self.name)
def check_volume(self, volume) -> tuple[bool, float]: def check_volume(self, *, volume) -> tuple[bool, float]:
"""Check if the volume is within the limits of the symbol. If not, return the nearest limit. """Check if the volume is within the limits of the symbol. If not, return the nearest limit.
Args: Args:
@@ -167,7 +159,7 @@ class Symbol(SymbolInfo):
else: else:
return check, self.volume_min if volume <= self.volume_min else self.volume_max return check, self.volume_min if volume <= self.volume_min else self.volume_max
def round_off_volume(self, volume: float, round_down: bool = False) -> float: def round_off_volume(self, *, volume: float, round_down: bool = False) -> float:
"""Round off the volume to the nearest volume step. """Round off the volume to the nearest volume step.
Args: Args:
@@ -179,7 +171,7 @@ class Symbol(SymbolInfo):
""" """
return round_off(value=volume, step=self.volume_step, round_down=round_down) return round_off(value=volume, step=self.volume_step, round_down=round_down)
async def check_amount(self, amount: float) -> float: async def check_amount(self, *, amount: float) -> float:
if self.currency_profit != self.account.currency: if self.currency_profit != self.account.currency:
amount = await self.convert_currency(amount=amount, base=self.currency_profit, quote=self.account.currency) amount = await self.convert_currency(amount=amount, base=self.currency_profit, quote=self.account.currency)
return amount return amount
@@ -212,27 +204,19 @@ class Symbol(SymbolInfo):
Returns: Returns:
float: Amount in terms of the quote currency float: Amount in terms of the quote currency
Raises:
ValueError: If conversion is impossible
""" """
try: try:
pair = f'{base}{quote}' pair = f'{base}{quote}'
if self.account.has_symbol(pair): tick = await self.info_tick(name=pair)
tick = await self.info_tick(name=pair) if tick is not None:
if tick is not None: return amount / tick.ask
return amount / tick.ask
pair = f'{quote}{base}' pair = f'{quote}{base}'
if self.account.has_symbol(pair): tick = await self.info_tick(name=pair)
tick = await self.info_tick(name=pair) if tick is not None:
if tick is not None: return amount * tick.bid
return amount * tick.bid
except Exception as err: except Exception as err:
logger.warning(f'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}')
raise ValueError(f'Currency Conversion Failed: {err}')
else:
logger.warning(f'Currency conversion failed: Unable to convert {amount} in {quote} to {base}')
@backoff_decorator @backoff_decorator
async def copy_rates_from(self, *, timeframe: TimeFrame, async def copy_rates_from(self, *, timeframe: TimeFrame,
@@ -2,20 +2,22 @@
from typing import NamedTuple from typing import NamedTuple
from logging import getLogger from logging import getLogger
from .core.models import TerminalInfo from ..core.models import TerminalInfo
from ..core.base import _Base
logger = getLogger(__name__) logger = getLogger(__name__)
Version = NamedTuple("Version", (('version', str), ('build', int), ('release_date', str)))
class Terminal(TerminalInfo):
class Terminal(_Base, TerminalInfo):
"""Terminal Class. Get information about the MetaTrader 5 terminal. The class is a subclass of the TerminalInfo """Terminal Class. Get information about the MetaTrader 5 terminal. The class is a subclass of the TerminalInfo
class. It inherits all the attributes and methods of the TerminalInfo class and adds some useful methods. class. It inherits all the attributes and methods of the TerminalInfo class and adds some useful methods.
Notes: Notes:
Other attributes are defined in the TerminalInfo Class Other attributes are defined in the TerminalInfo Class
""" """
version: Version | None = None
Version = NamedTuple("Version", (('version', str), ('build', int), ('release_date', str)))
async def initialize(self) -> bool: async def initialize(self) -> bool:
"""Establish a connection with the MetaTrader 5 terminal. There are three call options. Call without parameters. """Establish a connection with the MetaTrader 5 terminal. There are three call options. Call without parameters.
@@ -26,30 +28,29 @@ class Terminal(TerminalInfo):
Returns: Returns:
bool: True if successful else False bool: True if successful else False
""" """
self.connected = await self.mt5.initialize(**self.config.account_info()) self.connected = await self.mt5.initialize()
if not self.connected: if not self.connected:
err = await self.mt5.last_error() err = await self.mt5.last_error()
logger.critical(f'Failed to initialize Terminal. Error Code: {err}') logger.warning(f'Failed to initialize Terminal. Error Code: {err}')
raise SystemExit info = await self.info()
return self.connected await self.get_version()
return bool(self.connected and info and self.version)
async def version(self): async def get_version(self) -> Version | None:
"""Get the MetaTrader 5 terminal version. This method returns the terminal version, build and release date as """Get the MetaTrader 5 terminal version. This method returns the terminal version, build and release date as
a tuple of three values a tuple of three values
Returns: Returns:
Version: version of tuple as Version object Version: version of tuple as Version object
Raises:
ValueError: If the terminal version cannot be obtained
""" """
res = await self.mt5.version() res = await self.mt5.version()
if res is None: if res is None:
raise ValueError('Failed to get terminal version') logger.error('Failed to get terminal version')
return self.Version(*res) return None
self.version = Version(*res)
return self.version
async def info(self): async def info(self) -> TerminalInfo | None:
"""Get the connected MetaTrader 5 client terminal status and settings. gets terminal info in the form of a """Get the connected MetaTrader 5 client terminal status and settings. gets terminal info in the form of a
named tuple structure (namedtuple). Return None in case of an error. The info on the error can be named tuple structure (namedtuple). Return None in case of an error. The info on the error can be
obtained using last_error(). obtained using last_error().
@@ -58,7 +59,9 @@ class Terminal(TerminalInfo):
Terminal: Terminal status and settings as a terminal object. Terminal: Terminal status and settings as a terminal object.
""" """
info = await self.mt5.terminal_info() info = await self.mt5.terminal_info()
self.set_attributes(**info._asdict()) if info:
self.set_attributes(**info._asdict())
return info
async def symbols_total(self) -> int: async def symbols_total(self) -> int:
"""Get the number of all financial instruments in the MetaTrader 5 terminal. """Get the number of all financial instruments in the MetaTrader 5 terminal.
+36 -14
View File
@@ -1,15 +1,14 @@
"""Module for working with price ticks.""" """Module for working with price ticks."""
from typing import TypeVar, Iterable from typing import Iterable, Self
import time
from pandas import DataFrame, Series from pandas import DataFrame, Series
import pandas_ta as ta import pandas_ta as ta
import mplfinance as mplt import mplfinance as mplt
import pandas as pd import pandas as pd
from .core.constants import TickFlag from ..core.constants import TickFlag
Self = TypeVar('Self', bound='Ticks')
class Tick: class Tick:
@@ -37,11 +36,13 @@ class Tick:
Index: int Index: int
def __init__(self, **kwargs): def __init__(self, **kwargs):
"""Initialize the Tick class. Set attributes from keyword arguments. bid, ask, last, time and volume must be """Initialize the Tick class. Set attributes from keyword arguments. bid, ask, last and volume must be
present""" present"""
if not all(key in kwargs for key in ['bid', 'ask', 'last', 'volume', 'time']): if not all(key in kwargs for key in ['bid', 'ask', 'last', 'volume']):
raise ValueError("bid, ask, last and volume, time must be present in the keyword arguments") raise ValueError("bid, ask, last and volume, time must be present in the keyword arguments")
self.Index = kwargs.pop('Index', 0) self.Index = kwargs.pop('Index', 0)
self.time = kwargs.pop('time', time.monotonic())
self.time_msc = int(self.time * 1000)
self.set_attributes(**kwargs) self.set_attributes(**kwargs)
def __repr__(self): def __repr__(self):
@@ -49,6 +50,30 @@ class Tick:
% {"class": self.__class__.__name__, "time": self.time, "bid": self.bid, % {"class": self.__class__.__name__, "time": self.time, "bid": self.bid,
"ask": self.ask, "last": self.last, "volume": self.volume, 'Index': self.Index}) "ask": self.ask, "last": self.last, "volume": self.volume, 'Index': self.Index})
def __eq__(self, other: Self):
return self.time == other.time
def __lt__(self, other: Self):
return self.time < other.time
def __hash__(self):
return hash(self.time)
def __getitem__(self, item):
return self.__dict__[item]
def __setitem__(self, key, value):
self.__dict__[key] = value
def __iter__(self):
return iter(self.__dict__.items())
def keys(self):
return self.__dict__.keys()
def values(self):
return self.__dict__.values()
def dict(self, exclude: set = None, include: set = None) -> dict: def dict(self, exclude: set = None, include: set = None) -> dict:
""" """
Returns a dictionary of the instance attributes. Returns a dictionary of the instance attributes.
@@ -70,9 +95,6 @@ class Tick:
setattr(self, key, value) setattr(self, key, value)
_Ticks = TypeVar('_Ticks', bound='Ticks')
class Ticks: class Ticks:
"""Container class for price ticks. Arrange in chronological order. Supports iteration, slicing and assignment""" """Container class for price ticks. Arrange in chronological order. Supports iteration, slicing and assignment"""
time: Series time: Series
@@ -85,7 +107,7 @@ class Ticks:
volume_real: Series volume_real: Series
Index: Series Index: Series
def __init__(self, *, data: DataFrame | Iterable, flip=False): def __init__(self, *, data: DataFrame | Iterable | Self, flip=False):
"""Initialize the Ticks class. Creates a DataFrame of price ticks from the data argument. """Initialize the Ticks class. Creates a DataFrame of price ticks from the data argument.
Args: Args:
@@ -162,7 +184,7 @@ class Ticks:
"""DataFrame of price ticks arranged in chronological order.""" """DataFrame of price ticks arranged in chronological order."""
return self._data return self._data
def rename(self, inplace=True, **kwargs) -> _Ticks | None: def rename(self, inplace=True, **kwargs) -> Self | None:
"""Rename columns of the candle class. """Rename columns of the candle class.
Keyword Args: Keyword Args:
@@ -190,13 +212,13 @@ class Ticks:
data.index = pd.to_datetime(data['time'], unit='s') data.index = pd.to_datetime(data['time'], unit='s')
return mplt.make_addplot(data[columns], **kwargs) return mplt.make_addplot(data[columns], **kwargs)
def visualize(self, *, count: int = 50, type='candle', savefig: str | dict = None, addplot: dict = None, def visualize(self, *, count: int = 50, _type='candle', savefig: str | dict = None, addplot: dict = None,
style: str = 'charles', ylabel: str = 'Price', title: str = 'Chart', **kwargs): style: str = 'charles', ylabel: str = 'Price', title: str = 'Chart', **kwargs):
"""Visualize the candles using the mplfinance library. """Visualize the candles using the mplfinance library.
Args: Args:
count (int): The number of candles to visualize, counting from behind, i.e the most recent candles. count (int): The number of candles to visualize, counting from behind, i.e the most recent candles.
Defaults to 50. Defaults to 50.
type: Type of chart, defaults to candle _type: Type of chart, defaults to candle
savefig (str|dict): The path to save the figure or a dictionary of parameters to pass to the savefig method. savefig (str|dict): The path to save the figure or a dictionary of parameters to pass to the savefig method.
addplot: Additional plots to add to the chart. Defaults to None. They should match the dimension of the addplot: Additional plots to add to the chart. Defaults to None. They should match the dimension of the
original data which is specified via the count parameter. original data which is specified via the count parameter.
@@ -206,7 +228,7 @@ class Ticks:
kwargs: valid kwargs for the plot function. kwargs: valid kwargs for the plot function.
""" """
kwargs |= {key: arg for key, arg in (('savefig', savefig), ('addplot', addplot), ('style', style), kwargs |= {key: arg for key, arg in (('savefig', savefig), ('addplot', addplot), ('style', style),
('ylabel', ylabel), ('title', title), ('type', type)) if arg} ('ylabel', ylabel), ('title', title), ('type', _type)) if arg}
data = self._data[-count:] data = self._data[-count:]
data.index = pd.to_datetime(data['time'], unit='s') data.index = pd.to_datetime(data['time'], unit='s')
mplt.plot(data, **kwargs) mplt.plot(data, **kwargs)
@@ -7,9 +7,9 @@ import csv
import logging import logging
from typing import Iterable from typing import Iterable
from .core.config import Config from ..core.config import Config
from .core.meta_trader import MetaTrader from ..core.meta_trader import MetaTrader
from .contrib.backtester.meta_tester import MetaTester from ..core.meta_backtester import MetaBackTester
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -23,7 +23,7 @@ class TradeRecords:
from the config from the config
""" """
config: Config config: Config
mt5: MetaTrader | MetaTester mt5: MetaTrader | MetaBackTester
def __init__(self, *, records_dir: Path | str = ''): def __init__(self, *, records_dir: Path | str = ''):
"""Initialize the Records class. The main method of this class is update_records which you should call to update """Initialize the Records class. The main method of this class is update_records which you should call to update
@@ -33,7 +33,7 @@ class TradeRecords:
records_dir (Path): Absolute path to directory containing record of placed trades. records_dir (Path): Absolute path to directory containing record of placed trades.
""" """
self.config = Config() self.config = Config()
self.mt5 = MetaTrader() if self.config.mode == 'live' else MetaTester() self.mt5 = MetaTrader() if self.config.mode == 'live' else MetaBackTester()
self.records_dir = records_dir or self.config.records_dir self.records_dir = records_dir or self.config.records_dir
async def get_csv_records(self): async def get_csv_records(self):
@@ -63,16 +63,15 @@ class TradeRecords:
file: Trade record file in csv format file: Trade record file in csv format
""" """
try: try:
fr = open(file, mode='r', newline='') with open(file, mode='r', newline='') as fr:
reader: Iterable[dict] | csv.DictReader = csv.DictReader(fr) reader: Iterable[dict] | csv.DictReader = csv.DictReader(fr)
rows = [row for row in reader] rows = [row for row in reader]
rows = await self.update_rows(rows=rows) rows = await self.update_rows(rows=rows)
fr.close()
fw = open(file, mode='w', newline='') with open(file, mode='w', newline='') as fw: # type: SupportsWrite[str]
writer = csv.DictWriter(fw, fieldnames=reader.fieldnames, extrasaction='ignore', restval=None) writer = csv.DictWriter(fw, fieldnames=reader.fieldnames, extrasaction='ignore', restval=None)
writer.writeheader() writer.writeheader()
writer.writerows(rows) writer.writerows(rows)
fw.close()
except Exception as err: except Exception as err:
logger.error(f'Error: {err}. Unable to read and update csv trade records') logger.error(f'Error: {err}. Unable to read and update csv trade records')
@@ -82,14 +81,13 @@ class TradeRecords:
file: Trade record file in csv format file: Trade record file in csv format
""" """
try: try:
fh = open(file, mode='r') with open(file, mode='r') as fh: # type: SupportsRead[str | bytes]
data = json.load(fh) data = json.load(fh)
rows = [row for row in data] rows = [row for row in data]
rows = await self.update_rows(rows=rows) rows = await self.update_rows(rows=rows)
fh.close()
fh = open(file, mode='w') with open(file, mode='w') as fh: # type: SupportsWrite[str]
json.dump(rows, fh, indent=2) json.dump(rows, fh, indent=2)
fh.close()
except Exception as err: except Exception as err:
logger.error(f'Error: {err}. Unable to read and update json trade records') logger.error(f'Error: {err}. Unable to read and update json trade records')
@@ -109,7 +107,7 @@ class TradeRecords:
return row return row
deals = [deal for deal in deals if (deal.order != deal.position_id and deal.position_id == order deals = [deal for deal in deals if (deal.order != deal.position_id and deal.position_id == order
and deal.entry == 1)] and deal.entry == 1)]
deals.sort(key=lambda x: x.time_msc) deals.sort(key=lambda deal: deal.time_msc)
deal = deals[-1] deal = deals[-1]
row.update(actual_profit=deal.profit, win=deal.profit > 0, closed=True) row.update(actual_profit=deal.profit, win=deal.profit > 0, closed=True)
return row return row
@@ -3,33 +3,40 @@ from abc import ABC, abstractmethod
from datetime import datetime from datetime import datetime
from typing import TypeVar from typing import TypeVar
from logging import getLogger from logging import getLogger
from zoneinfo import ZoneInfo import pytz
from ..core.models import OrderType, OrderSendResult, OrderCheckResult
from ..core.config import Config
from ..core.task_queue import QueueItem
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 .ticks import Tick from .ticks import Tick
from .ram import RAM from .ram import RAM
from .core.models import OrderType, OrderSendResult
from .core.config import Config
from .result import Result
from .core.task_queue import QueueItem
logger = getLogger(__name__) logger = getLogger(__name__)
Symbol = TypeVar("Symbol", bound=_Symbol) Symbol = TypeVar("Symbol", bound=_Symbol)
class Trader(ABC): class Trader(ABC):
"""Base class for creating a Trader object. Handles the creation of an order and the placing of trades. """Base class for creating and managing orders.
Handles the creation and placing of an order.
It is an abstract class and must be subclassed to implement the place_trade method.
It has a set of methods that can be used to set the order limits and stop levels for the order.
Attributes: Attributes:
symbol (Symbol): The financial instrument. symbol (Symbol): The financial instrument.
ram (RAM): RAM instance ram (RAM): RAM instance
order (Order): Trade order order (Order): Trade order
parameters (dict): Parameters of the trading strategy used to place the trade
Class Attributes: Class Attributes:
config (Config): Config instance. config (Config): Config instance.
""" """
config: Config config: Config
ram: RAM
parameters: dict
def __init__(self, *, symbol: Symbol, ram: RAM = None): def __init__(self, *, symbol: Symbol, ram: RAM = None):
"""Initializes the order object and RAM instance """Initializes the order object and RAM instance
@@ -82,47 +89,55 @@ class Trader(ABC):
self.symbol.digits) self.symbol.digits)
self.order.price = tick.bid self.order.price = tick.bid
async def check_order(self) -> bool: async def check_order(self) -> OrderCheckResult | None:
"""Check order before sending it to the broker. """Check order before sending it to the broker.
Returns: Returns:
bool: True if order can go through else false bool: True if order can go through else false
""" """
check = await self.order.check() check = await self.order.check()
if check.retcode != 0:
logger.warning(f"Invalid order for {self.symbol} due to {check.comment}")
return False
return True
async def send_order(self) -> OrderSendResult: if check is None:
logger.warning(f"{self.order.mt5.error}: Order check failed")
return check
if check.retcode != 0:
logger.warning(f"Invalid order for due to {check.comment}")
else:
logger.info(f"Order check passed for {self.symbol}")
return check
async def send_order(self) -> OrderSendResult | None:
"""Send the order to the broker.""" """Send the order to the broker."""
result = await self.order.send() result = await self.order.send()
if result.retcode != 10009: if result is None:
logger.warning(f"Unable to place order for {self.symbol} due to {result.comment}") logger.warning(f"{self.order.mt5.error}: Failed to place order.")
return result return result
logger.info(f"Placed Trade for {self.symbol}")
if result.retcode != 10009:
logger.warning(f"Unable to place order for due to {result.comment}")
return result
logger.info("Order placed successfully")
return result return result
async def record_trade(self, result: OrderSendResult, parameters: dict = None, name: str = '', exclude: set = None): async def record_trade(self, *, result: OrderSendResult, parameters: dict = None, name: str = ''):
"""Record the trade in csv or json. """Record the trade in csv or json.
Args: Args:
result (OrderSendResult): Result of the order send result (OrderSendResult): Result of the order send
parameters: parameters of the trading strategy used to place the trade parameters: parameters of the trading strategy used to place the trade
name: Name of the trading strategy name: Name of the trading strategy
exclude: Exclude these fields from the recorded trade
""" """
if result.retcode != 10009 or not self.config.record_trades: if result.retcode != 10009 or not self.config.record_trades:
return return
params = parameters or self.parameters params = parameters or {}
params = {k: v for k, v in params.items() if k not in (exclude or set())}
profit = result.profit or await self.order.calc_profit() profit = result.profit or await self.order.calc_profit()
params["expected_profit"] = profit params["expected_profit"] = profit
date = datetime.utcnow() date = datetime.now(tz=pytz.UTC)
date = date.replace(tzinfo=ZoneInfo("UTC"))
params["date"] = str(date.date()) params["date"] = str(date.date())
params["time"] = str(date.time()) params["time"] = str(date.time())
res = Result(result=result, parameters=params, name=name) res = Result(result=result, parameters=params, name=name)
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
async def place_trade(self, *args, **kwargs): async def place_trade(self, *args, **kwargs):
-117
View File
@@ -1,117 +0,0 @@
"""This module contains the Records class, which is used to read and update trade records from csv files."""
import asyncio
from pathlib import Path
import csv
import logging
from typing import Iterable
from .contrib.backtester.meta_tester import MetaTester
from .core.config import Config
from .core.meta_trader import MetaTrader
logger = logging.getLogger(__name__)
class Records:
"""This utility class read trade records from csv files, and update them based on their closing positions.
Attributes:
config: Config object
records_dir(Path): Absolute path to directory containing record of placed trades, If not given takes the default
from the config
"""
config: Config
mt5: MetaTrader
def __init__(self, records_dir: Path | str = ''):
"""Initialize the Records class. The main method of this class is update_records which you should call to update
all the records specified in the records_dir.
Keyword Args:
records_dir (Path): Absolute path to directory containing record of placed trades.
"""
self.config = Config()
self.mt5 = MetaTrader() if self.config.mode == 'live' else MetaTester()
self.records_dir = records_dir or self.config.records_dir
async def get_records(self):
"""Get trade records from records_dir folder
Yields:
files: Trade record files
"""
for file in self.records_dir.iterdir():
if file.is_file() and file.name.endswith('.csv'):
yield file
async def read_update(self, file: Path):
"""Read and update trade records
Args:
file: Trade record file
"""
try:
fr = open(file, mode='r', newline='')
reader: Iterable[dict] | csv.DictReader = csv.DictReader(fr)
rows = [row for row in reader]
rows = await self.update_rows(rows)
fr.close()
fw = open(file, mode='w', newline='')
writer = csv.DictWriter(fw, fieldnames=reader.fieldnames, extrasaction='ignore', restval=None)
writer.writeheader()
writer.writerows(rows)
fw.close()
except Exception as err:
logger.error(f'Error: {err}. Unable to read and update trade records')
async def update_row(self, row: dict) -> dict:
"""Update a single row of entered trade in the csv file with the actual profit.
Args:
row: A dictionary from the dictionary writer object of the csv file.
Returns:
dict: A dictionary with the actual profit and win status.
"""
try:
order = int(row['order'])
deals = await self.mt5.history_deals_get(position=order)
if not deals or len(deals) <= 1:
return row
deals = [deal for deal in deals if (deal.order != deal.position_id and deal.position_id == order
and deal.entry == 1)]
deals.sort(key=lambda x: x.time_msc)
deal = deals[-1]
row.update(actual_profit=deal.profit, win=deal.profit > 0, closed=True)
return row
except Exception as err:
logging.error(f'Error: {err}. Unable to update trade record')
return row
async def update_rows(self, rows: list[dict]) -> list[dict]:
"""Update the rows of entered trades in the csv file with the actual profit.
Args:
rows: A list of dictionaries from the dictionary writer object of the csv file.
Returns:
list[dict]: A list of dictionaries with the actual profit and win status.
"""
closed, unclosed = [], []
for row in rows:
if (row.get('closed', 'FALSE')).title() == 'True':
closed.append(row)
else:
unclosed.append(row)
unclosed = await asyncio.gather(*[self.update_row(row) for row in unclosed])
return closed + list(unclosed)
async def update_records(self):
"""Update trade records in the records_dir folder."""
records = [self.read_update(record) async for record in self.get_records()]
await asyncio.gather(*records)
async def update_record(self, file: Path | str):
"""Update a single trade record file."""
await self.read_update(file)
-83
View File
@@ -1,83 +0,0 @@
"""The base class for creating strategies."""
import asyncio
from time import time
from typing import TypeVar
from abc import ABC, abstractmethod
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 .sessions import Sessions, Session
Symbol = TypeVar("Symbol", bound=_Symbol)
class Strategy(ABC):
"""The base class for creating strategies.
Attributes:
name (str): The name of the strategy.
symbol (Symbol): The Financial Instrument as a Symbol Object
parameters (Dict): A dictionary of parameters for the strategy.
sessions (Sessions): The sessions to use for the strategy.
Notes:
Define the name of a strategy as a class attribute. If not provided, the class name will be used as the name.
"""
name: str
symbol: Symbol
sessions: Sessions
mt5: MetaTrader
config: Config
parameters = {}
def __init__(self, *, symbol: Symbol, params: dict = None, sessions: Sessions = None, name=''):
"""Initiate the parameters dict and add name and symbol fields.
Use class name as strategy name if name is not provided
Args:
symbol (Symbol): The Financial instrument
params (Dict): Trading strategy parameters
"""
self.parameters = self.parameters | (params or {})
self.symbol = symbol
self.name = name or self.__class__.__name__
self.parameters["symbol"] = symbol.name
self.parameters["name"] = self.name
self.sessions = sessions or Sessions(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()
def __repr__(self):
return f"{self.name}({self.symbol!r})"
def __getattr__(self, item):
if item in self.parameters:
return self.parameters[item]
raise AttributeError(f'{item} not an attribute of {self.name}')
def __setattr__(self, key, value):
if key in self.__dict__.get('parameters', {}):
self.parameters[key] = value
super().__setattr__(key, value)
@staticmethod
async def sleep(secs: float):
"""Sleep for the needed amount of seconds in between requests to the terminal.
computes the accurate amount of time needed to sleep ensuring that the next request is made at the start of
a new bar and making cooperative multitasking possible.
Args:
secs (float): The time in seconds. Usually the timeframe you are trading on.
"""
mod = time() % secs
secs = secs - mod if mod != 0 else mod
await asyncio.sleep(secs + 0.2)
@abstractmethod
async def trade(self):
"""Place trades using this method. This is the main method of the strategy.
It will be called by the strategy runner.
"""
-14
View File
@@ -1,14 +0,0 @@
import pytest
from aiomql import Config
import MetaTrader5
@pytest.fixture(scope='session', autouse=True)
def config():
config = Config(filename='test.json')
return config
@pytest.fixture(scope='session', autouse=True)
def metatrader5():
return MetaTrader5
+93
View File
@@ -0,0 +1,93 @@
import asyncio
import json
import shutil
from logging import getLogger
from pathlib import Path
import pytest
from aiomql.core import Config
from aiomql.core.meta_trader import MetaTrader
logger = getLogger(__name__)
async def cleanup():
try:
shutil.rmtree(Path('tests/configs'), ignore_errors=True)
Path.unlink(Path('tests/test.json'), missing_ok=True)
shutil.rmtree(Path('tests/trade_records'), ignore_errors=True)
await close_all_positions()
await MetaTrader().shutdown()
except Exception as err:
logger.error(f"Failed to complete cleanup: {err}")
async def close_all_positions():
try:
mt = MetaTrader()
positions = await mt.positions_get()
tasks = []
for position in positions:
order_type = mt.ORDER_TYPE_BUY if position.type == mt.ORDER_TYPE_SELL else mt.ORDER_TYPE_SELL
req = {'action': mt.TRADE_ACTION_DEAL, 'symbol': position.symbol, 'volume': position.volume,
'type': order_type, 'position': position.ticket, 'price': position.price_current}
tasks.append(mt.order_send(req))
await asyncio.gather(*tasks)
except Exception as err:
logger.error(f"Failed to close all positions: {err}")
@pytest.fixture(scope='session', autouse=True)
async def config(request):
Path('tests/configs').mkdir(exist_ok=True)
with open('aiomql.json', 'r') as fh, open('tests/configs/test2.json', 'w') as fh1, open('tests/test.json', 'w') as fh2:
data = json.load(fh)
json.dump(data, fh1, indent=2)
json.dump(data, fh2, indent=2)
config = Config(filename='test.json', root='tests')
yield config
await cleanup()
@pytest.fixture(scope='session')
async def mt():
mt = MetaTrader()
await mt.initialize()
await mt.login()
yield mt
await mt.shutdown()
@pytest.fixture(scope='function')
async def sell_order(mt):
sym = 'BTCUSD'
sym_info = await mt.symbol_info(sym)
return {'action': mt.TRADE_ACTION_DEAL, 'symbol': sym, 'volume': sym_info.volume_min,
'type': mt.ORDER_TYPE_SELL, 'price': sym_info.bid}
@pytest.fixture(scope='function')
async def buy_order(mt):
sym = 'BTCUSD'
sym_info = await mt.symbol_info(sym)
dsl = (sym_info.trade_stops_level + sym_info.spread) * 2 * sym_info.point
sl = sym_info.ask - dsl
tp = sym_info.ask + dsl
return {'action': mt.TRADE_ACTION_DEAL, 'symbol': sym, 'volume': sym_info.volume_min,
'type': mt.ORDER_TYPE_BUY, 'price': sym_info.ask, 'sl': sl, 'tp': tp}
@pytest.fixture(scope='class')
async def make_orders(mt):
sym = 'BTCUSD'
sym_info = await mt.symbol_info(sym)
dsl = (sym_info.trade_stops_level + sym_info.spread) * 2 * sym_info.point
sl = sym_info.ask - dsl
tp = sym_info.ask + dsl
req = {'action': mt.TRADE_ACTION_DEAL, 'symbol': sym, 'volume': sym_info.volume_min,
'type': mt.ORDER_TYPE_BUY, 'price': sym_info.ask, 'sl': sl, 'tp': tp}
await mt.order_send(req)
req['type'] = mt.ORDER_TYPE_SELL
req['price'] = sym_info.bid
req['sl'] = sym_info.bid + dsl
req['tp'] = sym_info.bid - dsl
await mt.order_send(req)
+4
View File
@@ -0,0 +1,4 @@
[pytest]
asyncio_default_fixture_loop_scope = session
addopts = --rootdir=tests --capture=tee-sys --last-failed
asyncio_mode = auto
-10
View File
@@ -1,10 +0,0 @@
{
"login": 31288540,
"password": "nwa0#anaEze",
"server": "Deriv-Demo",
"demo": 5463204,
"fin": 24251812,
"deriv-demo": 5463204,
"deriv-real": 31288540,
"mode": "backtest"
}
+20
View File
@@ -0,0 +1,20 @@
import pytest
from aiomql.lib.account import Account
class TestAccount:
@classmethod
def setup_class(cls):
cls.account = Account()
@pytest.fixture(scope='class', autouse=True)
async def refresh(self):
await self.account.refresh()
async def test_connected(self):
assert self.account.connected is True
async def test_account_info(self):
acc_info = await self.account.mt5.account_info()
assert acc_info.login == self.account.login
assert acc_info.server == self.account.server
+60
View File
@@ -0,0 +1,60 @@
import pytest
from aiomql.core.base import Base
class ChildClass(Base):
attr: int
attr2: str
cls_attr: int = 10
class TestBaseClass:
@pytest.fixture
def child(self):
return ChildClass(attr=1, attr2="test")
def test_repr(self, child):
repr_str = repr(child)
assert repr_str.startswith("ChildClass(")
assert "attr=1" in repr_str
assert "attr2=test" in repr_str
def test_set_attributes(self, child):
child.set_attributes(attr3=3.14, attr2='str')
assert child.attr2 == 'str'
assert getattr(child, 'attr3', None) is None
def test_annotations(self, child):
annotations = child.annotations
assert isinstance(annotations, dict)
def test_get_dict(self, child):
child.set_attributes(attr2='test')
result = child.get_dict()
assert result["attr"] == 1
assert result["attr2"] == "test"
def test_get_dict_with_exclude(self, child):
child.set_attributes(attr2='test')
result = child.get_dict(exclude={"attr"})
assert "attr" not in result
assert result["attr2"] == "test"
def test_get_dict_with_include(self, child):
child.set_attributes(attr3=3.14)
result = child.get_dict(include={"attr"})
assert result["attr"] == 1
assert "attr2" not in result
def test_class_vars(self, child):
class_vars = child.class_vars
assert isinstance(class_vars, dict)
assert 'cls_attr' in class_vars
assert 'attr' not in class_vars
def test_dict_property(self, child):
child.set_attributes(attr2="test")
dict_prop = child.dict
assert dict_prop["attr"] == 1
assert dict_prop["attr2"] == "test"
assert dict_prop["cls_attr"] == 10
+106
View File
@@ -0,0 +1,106 @@
from datetime import datetime
import pytest
import pytz
import pandas as pd
from aiomql.lib.candle import Candle, Candles
from aiomql.core.meta_trader import MetaTrader
from aiomql.core.constants import TimeFrame
class TestCandle:
@classmethod
def setup_class(cls):
cls.bullish_candle = Candle(open=1.3421, high=1.3462, low=1.3405, close=1.3452, time=0, Index=0)
cls.bearish_candle = Candle(open=1.3452, high=1.3405, low=1.3462, close=1.3421, time=1, Index=1)
def test_repr(self):
repr_str = repr(self.bearish_candle)
assert repr_str.startswith("Candle(")
assert "open=" in repr_str
assert "high=" in repr_str
assert "low=" in repr_str
assert "close=" in repr_str
def test_set_attributes(self):
self.bearish_candle.set_attributes(ema=10)
assert self.bearish_candle.ema == 10
def test_compare(self):
assert self.bearish_candle > self.bullish_candle
assert self.bullish_candle != self.bearish_candle
assert self.bullish_candle < self.bearish_candle
def test_dict(self):
self.bearish_candle.set_attributes(ema=10)
result = self.bearish_candle.dict(exclude={'time'})
result2 = self.bearish_candle.dict(include={'close', 'high'})
assert result['open'] == 1.3452
assert result['ema'] == 10
assert 'time' not in result
assert set(result2.keys()) == {'close', 'high'}
def test_dictionary_properties(self):
self.bearish_candle['ema'] = 4
assert self.bearish_candle['ema'] == 4
def test_candle_type(self):
assert self.bearish_candle.is_bearish()
assert self.bullish_candle.is_bullish()
class TestCandles:
@pytest.fixture(scope='class')
async def candles(self):
mt = MetaTrader()
start = datetime(day=5, month=10, year=2023)
rates = await mt.copy_rates_from("BTCUSD", mt.TIMEFRAME_H1, start, 200)
return Candles(data=rates)
def test_get_series(self, candles):
series = candles['open']
assert isinstance(series, pd.Series)
assert len(series) == 200
def test_get_candle(self, candles):
candle = candles[10]
assert isinstance(candle, Candle)
assert candle in candles
assert candle.Index == 10
def test_slice(self, candles):
sliced = candles[10:15]
assert len(sliced) == 5
assert isinstance(sliced, Candles)
def test_setitem(self, candles):
new_series = candles.open
new_series = new_series * 2
candles['double_open'] = new_series
assert 'double_open' in candles.data.columns
def test_getattr(self, candles):
open_series = candles.open
assert isinstance(open_series, pd.Series)
assert open_series.equals(candles.data['open'])
def test_iter(self, candles):
l_5 = candles[-5:]
assert all(isinstance(candle, Candle) for candle in l_5)
def test_timeframe(self, candles):
tf = candles.timeframe
assert tf == TimeFrame.H1
def test_ta_and_rename(self, candles):
ema = candles.ta.ema(close='open', length=10, append=True)
assert 'EMA_10' in candles.data.columns
candles.rename(inplace=True, EMA_10='ema')
assert 'ema' in candles.data.columns
def test_ta_lib(self, candles):
fas = candles.ta_lib.above(candles.open, candles.close)
assert isinstance(fas, pd.Series)
candles['fas'] = fas
assert 'fas' in candles.data.columns
+29
View File
@@ -0,0 +1,29 @@
from aiomql.core.config import Config
from aiomql.contrib.backtesting import BackTestEngine
class TestConfig:
def test_singleton(self, config):
config2 = Config(filename='test.json')
assert config is config2
def test_set_attributes(self, config):
config.set_attributes(timeout=5000, record_trades=False)
assert config.timeout == 5000
assert config.record_trades is False
def test_backtest_engine(self, config):
engine = BackTestEngine()
config.backtest_engine = engine
assert config.backtest_engine is engine
def test_account_info(self, config):
account_info = config.account_info()
assert isinstance(account_info, dict)
assert 'login' in account_info
assert 'password' in account_info
assert 'server' in account_info
def test_load_config(self, config):
config.load_config(file='tests/configs/test2.json')
assert config.filename == 'test2.json'
+49
View File
@@ -0,0 +1,49 @@
from datetime import datetime
import pytest
from aiomql.lib.history import History
class TestHistory:
@pytest.fixture(scope='class', autouse=True)
async def init(self, make_orders):
await self.history.init()
@classmethod
def setup_class(cls):
now = datetime.now()
cls.start = now.replace(hour=0)
cls.end = now.replace(hour=23)
history = History(date_from=cls.start, date_to=cls.end)
cls.history = history
async def test_init(self):
assert self.history.total_deals > 0
assert self.history.total_orders > 0
async def test_get_deals(self):
deals = await self.history.get_deals()
assert len(deals) > 0
async def test_get_deals_by_ticket(self):
ticket = self.history.deals[0].order
deals = self.history.get_deals_by_ticket(ticket=ticket)
assert len(deals) > 0
async def test_get_deals_by_position(self):
position = self.history.deals[0].position_id
deals = self.history.get_deals_by_position(position=position)
assert len(deals) > 0
async def test_get_orders(self):
orders = await self.history.get_orders()
assert len(orders) > 0
async def test_get_orders_by_ticket(self):
ticket = self.history.orders[0].ticket
orders = self.history.get_orders_by_ticket(ticket=ticket)
assert len(orders) > 0
async def test_get_orders_by_position(self):
position = self.history.orders[0].position_id
orders = self.history.get_orders_by_position(position=position)
assert len(orders) > 0
@@ -1,199 +1,207 @@
from datetime import datetime, timedelta from datetime import datetime, timedelta
import pytest
import pytz import pytz
import pytest_asyncio import MetaTrader5
from aiomql import MetaTrader, TimeFrame, OrderType, CopyTicks
from . import metatrader5 from aiomql import MetaTrader
class TestMetaTrader: class TestMetaTrader:
@classmethod @classmethod
def setup_class(cls, metatrader5): def setup_class(cls):
tz = pytz.timezone('Etc/UTC')
cls.mt = MetaTrader() cls.mt = MetaTrader()
cls.mt5 = metatrader5 cls.mt5 = MetaTrader5
cls.symbol = "Volatility 100 Index" cls.symbol = "BTCUSD"
now = datetime.now(tz=tz) now = datetime.now(tz=pytz.UTC)
cls.start = now - timedelta(hours=24) cls.start = now - timedelta(hours=10)
cls.end = now + timedelta(hours=2) cls.end = now + timedelta(hours=1)
cls.tf = cls.mt.TIMEFRAME_H1 cls.tf = cls.mt.TIMEFRAME_H1
@pytest.mark.asyncio @classmethod
def teardown_class(cls):
cls.mt._shutdown()
async def test_initialize(self): async def test_initialize(self):
res = await self.mt.initialize() res = await self.mt.initialize()
assert res == True assert res == True
@pytest.mark.asyncio
async def test_login(self): async def test_login(self):
res = await self.mt.login() res = await self.mt.login()
assert res == True assert res == True
@pytest.mark.asyncio
async def test_last_error(self): async def test_last_error(self):
res = await self.mt.last_error() res = await self.mt.last_error()
assert isinstance(res, tuple) assert isinstance(res, tuple)
assert res[0] == 1 assert res[0] == 1
assert res[1] == 'Successful' assert res[1] == 'Success'
@pytest.mark.asyncio
async def test_version(self): async def test_version(self):
res = await self.mt.version() res = await self.mt.version()
res2 = self.mt5.version() res2 = self.mt5.version()
assert res is not None assert res is not None
assert res == res2 assert res == res2
@pytest.mark.asyncio
async def test_account_info(self): async def test_account_info(self):
res = await self.mt.account_info() res = await self.mt.account_info()
res2 = self.mt5.account_info() res2 = self.mt5.account_info()
assert res is not None assert res is not None
assert res == res2 assert res == res2
@pytest.mark.asyncio
async def test_terminal_info(self): async def test_terminal_info(self):
res = await self.mt.terminal_info() res = await self.mt.terminal_info()
res2 = await self.mt5.terminal_info() res2 = self.mt5.terminal_info()
assert res is not None assert res is not None
assert res == res2 assert res == res2
@pytest.mark.asyncio
async def test_symbols_total(self): async def test_symbols_total(self):
res = await self.mt.symbols_total() res = await self.mt.symbols_total()
res2 = await self.mt5.symbols_total() res2 = self.mt5.symbols_total()
assert isinstance(res, int) assert isinstance(res, int)
assert res assert res == res2
@pytest.mark.asyncio
async def test_symbols_get(self): async def test_symbols_get(self):
res = await self.mt.symbols_get() res = await self.mt.symbols_get()
res2 = await self.mt5.symbols_get() res2 = self.mt5.symbols_get()
assert res is not None assert res is not None
assert res == res2 assert len(res) == len(res2)
@pytest.mark.asyncio
async def test_symbol_info(self): async def test_symbol_info(self):
res = await self.mt.symbol_info(self.symbol) res = await self.mt.symbol_info(self.symbol)
res2 = await self.mt5.symbol_info(self.symbol) res2 = self.mt5.symbol_info(self.symbol)
assert res is not None assert res is not None
assert res == res2 assert res == res2
@pytest.mark.asyncio
async def test_symbol_info_tick(self): async def test_symbol_info_tick(self):
res = await self.mt.symbol_info_tick(self.symbol) res = await self.mt.symbol_info_tick(self.symbol)
res2 = await self.mt5.symbol_info_tick(self.symbol) res2 = self.mt5.symbol_info_tick(self.symbol)
assert res is not None assert res is not None
assert res == res2 assert res == res2
@pytest.mark.asyncio
async def test_symbol_select(self): async def test_symbol_select(self):
res = await self.mt.symbol_select(self.symbol, True) res = await self.mt.symbol_select(self.symbol, True)
assert res == True assert res == True
@pytest.mark.asyncio
async def test_market_book_add(self): async def test_market_book_add(self):
res = await self.mt.market_book_add(self.symbol) res = await self.mt.market_book_add(self.symbol)
assert res == True assert res == True
@pytest.mark.asyncio
async def test_market_book_get(self): async def test_market_book_get(self):
res = await self.mt.market_book_get(self.symbol) res = await self.mt.market_book_get(self.symbol)
res2 = await self.mt5.market_book_get(self.symbol) res2 = self.mt5.market_book_get(self.symbol)
assert res is not None assert res is not None
assert res == res2 assert res == res2
@pytest.mark.asyncio
async def test_market_book_release(self): async def test_market_book_release(self):
res = await self.mt.market_book_release(self.symbol) res = await self.mt.market_book_release(self.symbol)
assert res == True assert res == True
@pytest.mark.asyncio
async def test_copy_rates_from(self): async def test_copy_rates_from(self):
res = await self.mt.copy_rates_from(self.symbol, self.tf, self.start, 10) res = await self.mt.copy_rates_from(self.symbol, self.tf, self.start, 10)
assert res is not None assert res is not None
assert res.shape[0] == 10 assert res.shape[0] == 10
@pytest.mark.asyncio
async def test_copy_rates_from_pos(self): async def test_copy_rates_from_pos(self):
res = await self.mt.copy_rates_from_pos(self.symbol, self.tf, 0, 10) res = await self.mt.copy_rates_from_pos(self.symbol, self.tf, 0, 10)
assert res is not None assert res is not None
assert res.shape[0] == 10 assert res.shape[0] == 10
@pytest.mark.asyncio
async def test_copy_rates_range(self): async def test_copy_rates_range(self):
res = await self.mt.copy_rates_range(self.symbol, TimeFrame.M1, datetime.now(), datetime.now()) res = await self.mt.copy_rates_range(self.symbol, self.tf, self.start, self.end)
assert res is not None assert res is not None
assert res.shape[0] == 10
@pytest.mark.asyncio
async def test_copy_ticks_from(self): async def test_copy_ticks_from(self):
res = await self.mt.copy_ticks_from(self.symbol, datetime.now(), 10, CopyTicks.ALL) res = await self.mt.copy_ticks_from(self.symbol, self.start, 10, self.mt.COPY_TICKS_ALL)
assert res is not None assert res is not None
assert res.shape[0] == 10
@pytest.mark.asyncio
async def test_copy_ticks_range(self): async def test_copy_ticks_range(self):
res = await self.mt.copy_ticks_range(self.symbol, datetime.now(), datetime.now(), CopyTicks.ALL) res = await self.mt.copy_ticks_range(self.symbol, self.start, self.end, self.mt.COPY_TICKS_ALL)
res2 = self.mt5.copy_ticks_range(self.symbol, self.start, self.end, self.mt5.COPY_TICKS_ALL)
assert res is not None assert res is not None
assert res.shape[0] == res2.shape[0]
@pytest.mark.asyncio
async def test_orders_total(self): async def test_orders_total(self):
res = await self.mt.orders_total() res = await self.mt.orders_total()
assert isinstance(res, int) assert isinstance(res, int)
@pytest.mark.asyncio
async def test_orders_get(self): async def test_orders_get(self):
res = await self.mt.orders_get() res = await self.mt.orders_get()
assert res is not None assert res is not None
assert isinstance(res, tuple)
@pytest.mark.asyncio assert len(res) == 0
async def test_order_calc_margin(self):
res = await self.mt.order_calc_margin(OrderType.BUY, self.symbol, 1.0, 1.0) async def test_order_calc_margin(self, sell_order):
price = sell_order['price']
volume = sell_order['volume']
type_ = sell_order['type']
res = await self.mt.order_calc_margin(type_, self.symbol, volume, price)
assert isinstance(res, float) assert isinstance(res, float)
@pytest.mark.asyncio async def test_order_calc_profit(self, buy_order):
async def test_order_calc_profit(self): volume = buy_order['volume']
res = await self.mt.order_calc_profit(OrderType.BUY, self.symbol, 1.0, 1.0, 1.1) price_open = buy_order['price']
price_close = buy_order['tp']
type_ = buy_order['type']
res = await self.mt.order_calc_profit(type_, self.symbol, volume, price_open, price_close)
assert isinstance(res, float) assert isinstance(res, float)
@pytest.mark.asyncio async def test_order_check(self, buy_order):
async def test_order_check(self): res = await self.mt.order_check(buy_order)
request = {"action": OrderType.BUY, "symbol": self.symbol, "volume": 1.0, "price": 1.0}
res = await self.mt.order_check(request)
assert res is not None assert res is not None
assert res.retcode == 0
@pytest.mark.asyncio
async def test_order_send(self): async def test_order_send(self, sell_order):
request = {"action": OrderType.BUY, "symbol": self.symbol, "volume": 1.0, "price": 1.0} res = await self.mt.order_send(sell_order)
res = await self.mt.order_send(request)
assert res is not None assert res is not None
assert res.retcode == 10009
@pytest.mark.asyncio
async def test_positions_total(self): async def test_positions_total(self):
res = await self.mt.positions_total() res = await self.mt.positions_total()
assert isinstance(res, int) assert isinstance(res, int)
assert res >= 0
@pytest.mark.asyncio
async def test_positions_get(self): async def test_positions_get(self):
res = await self.mt.positions_get() res = await self.mt.positions_get()
assert res is not None assert res is not None
assert isinstance(res, tuple)
@pytest.mark.asyncio assert len(res) >= 0
async def test_history_orders_total(self): async def test_history_orders_total(self):
res = await self.mt.history_orders_total(datetime.now(), datetime.now()) res = await self.mt.history_orders_total(self.start, self.end)
assert isinstance(res, int) assert isinstance(res, int)
assert res >= 0
@pytest.mark.asyncio
async def test_history_orders_get(self): async def test_history_orders_get(self):
res = await self.mt.history_orders_get(datetime.now(), datetime.now()) res = await self.mt.history_orders_get(self.start, self.end)
assert res is not None assert res is not None
assert isinstance(res, tuple)
@pytest.mark.asyncio assert len(res) >= 0
async def test_history_deals_total(self): async def test_history_deals_total(self):
res = await self.mt.history_deals_total(datetime.now(), datetime.now()) res = await self.mt.history_deals_total(self.start, self.end)
assert isinstance(res, int) assert isinstance(res, int)
assert res >= 0
@pytest.mark.asyncio
async def test_history_deals_get(self): async def test_history_deals_get(self):
res = await self.mt.history_deals_get(datetime.now(), datetime.now()) res = await self.mt.history_deals_get(self.start, self.end)
assert res is not None assert res is not None
assert isinstance(res, tuple)
assert len(res) >= 0
# sym = order_request['symbol']
# sym_info = await self.mt.symbol_info(sym)
# dsl = (sym_info.trade_stops_level + sym_info.spread) * 2 * sym_info.point
# order_request['volume'] = sym_info.volume_min
# order_request['price'] = sym_info.ask
# order_request['tp'] = round(sym_info.ask + dsl, sym_info.digits)
# order_request['sl'] = round(sym_info.ask - dsl, sym_info.digits)
# sym = order_request['symbol']
# sym_info = await self.mt.symbol_info(sym)
# dsl = (sym_info.trade_stops_level + sym_info.spread) * 2 * sym_info.point
# order_request['volume'] = sym_info.volume_min
# order_request['price'] = sym_info.ask
# order_request['tp'] = round(sym_info.ask + dsl, sym_info.digits)
# order_request['sl'] = round(sym_info.ask - dsl, sym_info.digits)
+27
View File
@@ -0,0 +1,27 @@
from aiomql.lib.order import Order
class TestOrder:
async def test_check(self, sell_order):
order = Order(**sell_order)
check = await order.check()
assert check.retcode == 0
async def test_send(self, buy_order):
order = Order(**buy_order)
send = await order.send()
assert send.retcode == 10009
async def test_margin(self, buy_order):
order = Order(**buy_order)
margin = await order.calc_margin()
assert margin is not None
assert margin > 0
assert isinstance(margin, float)
async def test_profit(self, buy_order):
order = Order(**buy_order)
profit = await order.calc_profit()
assert profit is not None
assert profit > 0
assert isinstance(profit, float)
+30
View File
@@ -0,0 +1,30 @@
import pytest
from aiomql.lib.positions import Positions
class TestPositions:
@pytest.fixture(scope='class', autouse=True)
async def init(self, make_orders):
await self.positions.get_positions()
@classmethod
def setup_class(cls):
cls.positions = Positions()
@pytest.mark.order(1)
async def test_get_positions(self):
await self.positions.get_positions()
assert len(self.positions.positions) >= 0
async def test_get_position_by_ticket(self):
ticket = self.positions.positions[0].ticket
position = await self.positions.get_position_by_ticket(ticket=ticket)
assert position is not None
assert position.ticket == ticket
async def test_get_position_by_symbol(self):
symbol = self.positions.positions[0].symbol
positions = await self.positions.get_position_by_symbol(symbol=symbol)
assert len(positions) >= 0
assert positions[0].symbol == symbol
+22
View File
@@ -0,0 +1,22 @@
from aiomql.lib.ram import RAM
class TestRAM:
@classmethod
def setup_class(cls):
cls.ram = RAM(min_amount=5, max_amount=10, loss_limit=3, open_limit=5)
async def test_get_amount(self):
res = await self.ram.get_amount()
assert self.ram.min_amount <= res <= self.ram.max_amount
async def test_checks(self, buy_order, sell_order, mt):
for i in range(self.ram.open_limit+1):
if i % 2 == 0:
await mt.order_send(buy_order)
else:
await mt.order_send(sell_order)
res1 = await self.ram.check_losing_positions()
res2 = await self.ram.check_open_positions()
assert res2 is False
assert isinstance(res1, bool)
+44
View File
@@ -0,0 +1,44 @@
import asyncio
import pytest
from aiomql.lib.result import Result
from aiomql.core.models import OrderSendResult
class TestResult:
@pytest.fixture(scope='class')
def parameters(self):
return {'name': 'test_trades', 'ema': 20, 'rsi': 14}
@pytest.fixture(scope='function')
async def order_results(self, mt, sell_order, buy_order, parameters):
res1 = await mt.order_send(sell_order)
res2 = await mt.order_send(buy_order)
res1 = Result(result=OrderSendResult(**res1._asdict()), parameters=parameters)
res2 = Result(result=OrderSendResult(**res2._asdict()), parameters=parameters)
return res1, res2
async def test_get_data(self, order_results):
res1, res2 = order_results
data1 = res1.get_data()
data2 = res2.get_data()
assert data1['actual_profit'] == data2['actual_profit'] == 0
assert data1['closed'] == data2['closed'] == False
assert data1['win'] == data2['win'] == False
async def test_csv(self, order_results):
res1, res2 = order_results
await asyncio.gather(res1.save(), res2.save())
assert res1.config.records_dir.exists()
record = res1.config.records_dir / f"{res1.name}.csv"
assert record.exists()
async def test_json(self, order_results):
res1, res2 = order_results
await asyncio.gather(res1.save(trade_record_mode='json'), res2.save(trade_record_mode='json'))
assert res1.config.records_dir.exists()
record = res1.config.records_dir / f"{res1.name}.json"
assert record.exists()
+68
View File
@@ -0,0 +1,68 @@
from datetime import datetime, time
import pytest
import pytz
from aiomql.lib.sessions import Session, Sessions, delta
class TestSessions:
@pytest.fixture(scope='class')
def make_sessions(self, make_session):
london, all_day, over_night = make_session
return Sessions(sessions=[london, all_day, over_night])
@pytest.fixture(scope='class')
def make_session(self):
end = time(hour=16, minute=59, second=59, microsecond=999_999, tzinfo=pytz.UTC)
london = Session(start=8, end=end, name='London', on_end='close_all')
start, end = time(hour=0, tzinfo=pytz.UTC), time(hour=23, minute=59, second=59, tzinfo=pytz.UTC)
all_day = Session(start=start, end=end, name='AllDay', on_end='close_all')
end = time(hour=6, minute=59, second=59, microsecond=999_999, tzinfo=pytz.UTC)
over_night = Session(start=18, end=end, name='OverNight', on_end='close_all')
return london, all_day, over_night
def test_session_attributes(self, make_session):
london, all_day, over_night = make_session
period = over_night.duration()
assert london.name == 'London'
assert london.start == time(hour=8, tzinfo=pytz.UTC)
assert london.end.hour == 16
assert period.hours == 12
assert period.minutes == period.seconds == 59
def test_session_intervals(self, make_session):
london, all_day, over_night = make_session
two_am = time(hour=2, tzinfo=pytz.UTC)
noon = time(hour=12, tzinfo=pytz.UTC)
now = datetime.now(pytz.UTC).time()
hours_till_london_starts = (delta(london.start) - delta(now)).seconds // 3600
assert hours_till_london_starts == london.until() // 3600
assert two_am in over_night
assert noon in london
assert two_am not in london
assert noon not in over_night
# all_day session is always open
assert all_day.in_session()
async def test_sessions(self, make_session):
london, all_day, over_night = make_session
sessions = Sessions(sessions=[london, over_night])
now = time(hour=21, tzinfo=pytz.UTC)
noon = time(hour=12, tzinfo=pytz.UTC)
mid_nite = time(hour=0, tzinfo=pytz.UTC)
next_sess = sessions.find_next(moment=now)
noon_sess = sessions.find(moment=noon)
no_sess = sessions.find(moment=time(hour=17, tzinfo=pytz.UTC))
mid_nite_sess = sessions.find(moment=mid_nite)
current_sess = sessions.find(moment=now)
assert current_sess.name == 'OverNight'
assert noon_sess.name == 'London'
assert no_sess is None
assert next_sess.name == 'London'
assert mid_nite_sess.name == 'OverNight'
current = datetime.now(pytz.UTC).time()
if current.hour not in (7, 17):
await sessions.check()
assert sessions.current_session is not None
assert sessions.current_session.name in ('London', 'OverNight')
+53
View File
@@ -0,0 +1,53 @@
from datetime import datetime, timedelta
import pytest
from aiomql.lib.symbol import Symbol
from aiomql.lib.candle import Candles
from aiomql.lib.ticks import Ticks
class TestSymbol:
@pytest.fixture(scope='class', autouse=True)
async def btc(self):
symbol = Symbol(name='BTCUSD')
select = getattr(symbol, 'select', False)
if select is False:
await symbol.init()
return symbol
async def test_symbol_attributes(self, btc):
assert btc.name == 'BTCUSD'
assert btc.select is True
assert btc.tick is not None
async def test_volume(self, btc):
volume = btc.volume_min - btc.volume_step
success, volume = btc.check_volume(volume=volume)
assert success is False
volume = btc.volume_min + btc.volume_step * 2
success, volume = btc.check_volume(volume=volume)
assert success is True
volume = btc.volume_min + btc.volume_step * 2.5
volume = btc.round_off_volume(volume=volume, round_down=True)
assert volume == btc.volume_min + btc.volume_step * 2
async def test_rates(self, btc):
start = datetime(year=2023, month=10, day=5)
end = start + timedelta(hours=9)
rates_from = await btc.copy_rates_from(timeframe=btc.mt5.TIMEFRAME_H1, date_from=start, count=10)
assert isinstance(rates_from, Candles)
assert len(rates_from) == 10
rates_from_pos = await btc.copy_rates_from_pos(timeframe=btc.mt5.TIMEFRAME_H1, count=10, start_position=0)
assert isinstance(rates_from_pos, Candles)
assert len(rates_from_pos) == 10
rates_range = await btc.copy_rates_range(timeframe=btc.mt5.TIMEFRAME_H1, date_from=start, date_to=end)
assert isinstance(rates_range, Candles)
assert len(rates_range) == 10
ticks_from = await btc.copy_ticks_from(date_from=start, count=10)
assert isinstance(ticks_from, Ticks)
assert len(ticks_from) == 10
end = start + timedelta(seconds=20)
ticks_from_pos = await btc.copy_ticks_range(date_from=start, date_to=end)
assert isinstance(ticks_from_pos, Ticks)
assert len(ticks_from_pos) >= 10
+17
View File
@@ -0,0 +1,17 @@
import pytest
from aiomql.lib.terminal import Terminal
class TestTerminal:
@pytest.fixture(scope='class', autouse=True)
async def init_terminal(self):
terminal = Terminal()
init = await terminal.initialize()
return init, terminal
async def test_terminal(self, init_terminal):
init, terminal = init_terminal
assert init is True
assert terminal.connected is True
assert terminal.version is not None
+26
View File
@@ -0,0 +1,26 @@
from datetime import datetime
from aiomql.lib.ticks import Ticks, Tick
from pandas import Series
class TestTicks:
async def test_tick(self, mt):
btc_tick = await mt.symbol_info_tick("BTCUSD")
btc_tick = Tick(**btc_tick._asdict())
tick_dict = btc_tick.dict(include={'ask', 'bid', 'time', 'volume'})
assert isinstance(btc_tick, Tick)
assert isinstance(tick_dict, dict)
assert 'ask' in tick_dict
assert 'bid' in tick_dict
assert 'volume_real' not in tick_dict
async def test_ticks(self, mt):
start = datetime(year=2023, month=10, day=5)
ticks = await mt.copy_ticks_from("BTCUSD", start, 10, mt.COPY_TICKS_ALL)
ticks = Ticks(data=ticks)
assert isinstance(ticks, Ticks)
assert len(ticks) == 10
assert isinstance(ticks[0], Tick)
bids = ticks['bid']
assert len(bids) == 10
assert isinstance(bids, Series)