diff --git a/check.py b/check.py new file mode 100644 index 0000000..de6c57d --- /dev/null +++ b/check.py @@ -0,0 +1,62 @@ +import asyncio +from datetime import datetime, UTC + +from aiomql.contrib.backtesting.get_data import BackTestData, GetData +from aiomql.contrib.backtesting.backtest_engine import BackTestEngine +from aiomql.core.constants import TimeFrame +from aiomql.core.meta_trader import MetaTrader + + +async def get_data(): + mt = MetaTrader() + await mt.login() + start = datetime(2024, 2, 1, tzinfo=UTC) + end = datetime(2024, 2, 3, tzinfo=UTC) + symbols = ['BTCUSD', "ETHUSD", "SOLUSD"] + timeframes = [TimeFrame.H1, TimeFrame.H2, TimeFrame.M5] + g_data = GetData(start=start, end=end, symbols=symbols, timeframes=timeframes, + name='test_data') + s = datetime.now().timestamp() + await g_data.get_data(workers=500) + e = datetime.now().timestamp() + print(e - s, 'for getting data') + + s = datetime.now().timestamp() + g_data.save_data() + e = datetime.now().timestamp() + print(e - s, 'for saving data') + + +async def test_data(): + start = datetime.now().timestamp() + td = GetData.load_data(name='backtesting/test_data.pkl') + end = datetime.now().timestamp() + print(end-start, 'seconds') + print(td.name) + print(td.version) + print(len(td.ticks.keys())) + print(len(td.symbols.keys())) + print(td.rates.keys()) + +async def back_test_engine(): + td = GetData.load_data(name='backtesting/test_data.pkl') + bt = BackTestEngine(name='test_data_2', start=datetime(2024, 2, 1), + end=datetime(2024, 2, 7)) + # bt.next() + bt.next() + # print(bt.cursor) + # now = bt.cursor.time + # r = bt.cursor.index + # bt.fast_forward(steps=20) + # print(bt.cursor.time == now + 20) + # print(bt.cursor.index == r + 20) + # print(bt.cursor, r) + # print(bt.cursor) + print(bt.range, bt.span) + bt.go_to(time=datetime(2024, 2, 13)) + print(bt.cursor) + print(datetime.fromtimestamp(bt.cursor.time)) + print(bt.range, bt.span) + + +asyncio.run(get_data()) diff --git a/docs/core/constants.md b/docs/core/constants.md index 7bdde03..98e516a 100644 --- a/docs/core/constants.md +++ b/docs/core/constants.md @@ -186,9 +186,10 @@ The number of seconds in a TIMEFRAME ### Example + ```python t = TimeFrame.H1 -print(t.time) # 3600 +print(t.seconds) # 3600 ``` diff --git a/docs/account.md b/docs/lib/account.md similarity index 100% rename from docs/account.md rename to docs/lib/account.md diff --git a/docs/bot_builder.md b/docs/lib/bot_builder.md similarity index 100% rename from docs/bot_builder.md rename to docs/lib/bot_builder.md diff --git a/docs/candle.md b/docs/lib/candle.md similarity index 100% rename from docs/candle.md rename to docs/lib/candle.md diff --git a/docs/executor.md b/docs/lib/executor.md similarity index 100% rename from docs/executor.md rename to docs/lib/executor.md diff --git a/docs/history.md b/docs/lib/history.md similarity index 100% rename from docs/history.md rename to docs/lib/history.md diff --git a/docs/order.md b/docs/lib/order.md similarity index 100% rename from docs/order.md rename to docs/lib/order.md diff --git a/docs/positions.md b/docs/lib/positions.md similarity index 100% rename from docs/positions.md rename to docs/lib/positions.md diff --git a/docs/ram.md b/docs/lib/ram.md similarity index 100% rename from docs/ram.md rename to docs/lib/ram.md diff --git a/docs/records.md b/docs/lib/records.md similarity index 100% rename from docs/records.md rename to docs/lib/records.md diff --git a/docs/result.md b/docs/lib/result.md similarity index 100% rename from docs/result.md rename to docs/lib/result.md diff --git a/docs/sessions.md b/docs/lib/sessions.md similarity index 100% rename from docs/sessions.md rename to docs/lib/sessions.md diff --git a/docs/strategy.md b/docs/lib/strategy.md similarity index 100% rename from docs/strategy.md rename to docs/lib/strategy.md diff --git a/docs/symbol.md b/docs/lib/symbol.md similarity index 100% rename from docs/symbol.md rename to docs/lib/symbol.md diff --git a/docs/terminal.md b/docs/lib/terminal.md similarity index 100% rename from docs/terminal.md rename to docs/lib/terminal.md diff --git a/docs/ticks.md b/docs/lib/ticks.md similarity index 100% rename from docs/ticks.md rename to docs/lib/ticks.md diff --git a/docs/trade_records.md b/docs/lib/trade_records.md similarity index 100% rename from docs/trade_records.md rename to docs/lib/trade_records.md diff --git a/docs/trader.md b/docs/lib/trader.md similarity index 100% rename from docs/trader.md rename to docs/lib/trader.md diff --git a/docs/utils.md b/docs/lib/utils.md similarity index 100% rename from docs/utils.md rename to docs/lib/utils.md diff --git a/google_pylintrc b/google_pylintrc new file mode 100644 index 0000000..5c771e6 --- /dev/null +++ b/google_pylintrc @@ -0,0 +1,399 @@ +# This Pylint rcfile contains a best-effort configuration to uphold the +# best-practices and style described in the Google Python style guide: +# https://google.github.io/styleguide/pyguide.html +# +# Its canonical open-source location is: +# https://google.github.io/styleguide/pylintrc + +[MAIN] + +# Files or directories to be skipped. They should be base names, not paths. +ignore=third_party + +# Files or directories matching the regex patterns are skipped. The regex +# matches against base names, not paths. +ignore-patterns= + +# Pickle collected data for later comparisons. +persistent=no + +# List of plugins (as comma separated values of python modules names) to load, +# usually to register additional checkers. +load-plugins= + +# Use multiple processes to speed up Pylint. +jobs=4 + +# Allow loading of arbitrary C extensions. Extensions are imported into the +# active Python interpreter and may run arbitrary code. +unsafe-load-any-extension=no + + +[MESSAGES CONTROL] + +# Only show warnings with the listed confidence levels. Leave empty to show +# all. Valid levels: HIGH, INFERENCE, INFERENCE_FAILURE, UNDEFINED +confidence= + +# Enable the message, report, category or checker with the given id(s). You can +# either give multiple identifier separated by comma (,) or put this option +# multiple time (only on the command line, not in the configuration file where +# it should appear only once). See also the "--disable" option for examples. +#enable= + +# Disable the message, report, category or checker with the given id(s). You +# can either give multiple identifiers separated by comma (,) or put this +# option multiple times (only on the command line, not in the configuration +# file where it should appear only once).You can also use "--disable=all" to +# disable everything first and then reenable specific checks. For example, if +# you want to run only the similarities checker, you can use "--disable=all +# --enable=similarities". If you want to run only the classes checker, but have +# no Warning level messages displayed, use"--disable=all --enable=classes +# --disable=W" +disable=R, + abstract-method, + apply-builtin, + arguments-differ, + attribute-defined-outside-init, + backtick, + bad-option-value, + basestring-builtin, + buffer-builtin, + c-extension-no-member, + consider-using-enumerate, + cmp-builtin, + cmp-method, + coerce-builtin, + coerce-method, + delslice-method, + div-method, + eq-without-hash, + execfile-builtin, + file-builtin, + filter-builtin-not-iterating, + fixme, + getslice-method, + global-statement, + hex-method, + idiv-method, + implicit-str-concat, + import-error, + import-self, + import-star-module-level, + input-builtin, + intern-builtin, + invalid-str-codec, + locally-disabled, + long-builtin, + long-suffix, + map-builtin-not-iterating, + misplaced-comparison-constant, + missing-function-docstring, + metaclass-assignment, + next-method-called, + next-method-defined, + no-absolute-import, + no-init, # added + no-member, + no-name-in-module, + no-self-use, + nonzero-method, + oct-method, + old-division, + old-ne-operator, + old-octal-literal, + old-raise-syntax, + parameter-unpacking, + print-statement, + raising-string, + range-builtin-not-iterating, + raw_input-builtin, + rdiv-method, + reduce-builtin, + relative-import, + reload-builtin, + round-builtin, + setslice-method, + signature-differs, + standarderror-builtin, + suppressed-message, + sys-max-int, + trailing-newlines, + unichr-builtin, + unicode-builtin, + unnecessary-pass, + unpacking-in-except, + useless-else-on-loop, + useless-suppression, + using-cmp-argument, + wrong-import-order, + xrange-builtin, + zip-builtin-not-iterating, + + +[REPORTS] + +# Set the output format. Available formats are text, parseable, colorized, msvs +# (visual studio) and html. You can also give a reporter class, eg +# mypackage.mymodule.MyReporterClass. +output-format=text + +# Tells whether to display a full report or only the messages +reports=no + +# Python expression which should return a note less than 10 (10 is the highest +# note). You have access to the variables errors warning, statement which +# respectively contain the number of errors / warnings messages and the total +# number of statements analyzed. This is used by the global evaluation report +# (RP0004). +evaluation=10.0 - ((float(5 * error + warning + refactor + convention) / statement) * 10) + +# Template used to display messages. This is a python new-style format string +# used to format the message information. See doc for all details +#msg-template= + + +[BASIC] + +# Good variable names which should always be accepted, separated by a comma +good-names=main,_ + +# Bad variable names which should always be refused, separated by a comma +bad-names= + +# Colon-delimited sets of names that determine each other's naming style when +# the name regexes allow several styles. +name-group= + +# Include a hint for the correct naming format with invalid-name +include-naming-hint=no + +# List of decorators that produce properties, such as abc.abstractproperty. Add +# to this list to register other decorators that produce valid properties. +property-classes=abc.abstractproperty,cached_property.cached_property,cached_property.threaded_cached_property,cached_property.cached_property_with_ttl,cached_property.threaded_cached_property_with_ttl + +# Regular expression matching correct function names +function-rgx=^(?:(?PsetUp|tearDown|setUpModule|tearDownModule)|(?P_?[A-Z][a-zA-Z0-9]*)|(?P_?[a-z][a-z0-9_]*))$ + +# Regular expression matching correct variable names +variable-rgx=^[a-z][a-z0-9_]*$ + +# Regular expression matching correct constant names +const-rgx=^(_?[A-Z][A-Z0-9_]*|__[a-z0-9_]+__|_?[a-z][a-z0-9_]*)$ + +# Regular expression matching correct attribute names +attr-rgx=^_{0,2}[a-z][a-z0-9_]*$ + +# Regular expression matching correct argument names +argument-rgx=^[a-z][a-z0-9_]*$ + +# Regular expression matching correct class attribute names +class-attribute-rgx=^(_?[A-Z][A-Z0-9_]*|__[a-z0-9_]+__|_?[a-z][a-z0-9_]*)$ + +# Regular expression matching correct inline iteration names +inlinevar-rgx=^[a-z][a-z0-9_]*$ + +# Regular expression matching correct class names +class-rgx=^_?[A-Z][a-zA-Z0-9]*$ + +# Regular expression matching correct module names +module-rgx=^(_?[a-z][a-z0-9_]*|__init__)$ + +# Regular expression matching correct method names +method-rgx=(?x)^(?:(?P_[a-z0-9_]+__|runTest|setUp|tearDown|setUpTestCase|tearDownTestCase|setupSelf|tearDownClass|setUpClass|(test|assert)_*[A-Z0-9][a-zA-Z0-9_]*|next)|(?P_{0,2}[A-Z][a-zA-Z0-9_]*)|(?P_{0,2}[a-z][a-z0-9_]*))$ + +# Regular expression which should only match function or class names that do +# not require a docstring. +no-docstring-rgx=(__.*__|main|test.*|.*test|.*Test)$ + +# Minimum line length for functions/classes that require docstrings, shorter +# ones are exempt. +docstring-min-length=12 + + +[TYPECHECK] + +# List of decorators that produce context managers, such as +# contextlib.contextmanager. Add to this list to register other decorators that +# produce valid context managers. +contextmanager-decorators=contextlib.contextmanager,contextlib2.contextmanager + +# List of module names for which member attributes should not be checked +# (useful for modules/projects where namespaces are manipulated during runtime +# and thus existing member attributes cannot be deduced by static analysis. It +# supports qualified module names, as well as Unix pattern matching. +ignored-modules= + +# List of class names for which member attributes should not be checked (useful +# for classes with dynamically set attributes). This supports the use of +# qualified names. +ignored-classes=optparse.Values,thread._local,_thread._local + +# List of members which are set dynamically and missed by pylint inference +# system, and so shouldn't trigger E1101 when accessed. Python regular +# expressions are accepted. +generated-members= + + +[FORMAT] + +# Maximum number of characters on a single line. +max-line-length=80 + +# TODO(https://github.com/pylint-dev/pylint/issues/3352): Direct pylint to exempt +# lines made too long by directives to pytype. + +# Regexp for a line that is allowed to be longer than the limit. +ignore-long-lines=(?x)( + ^\s*(\#\ )??$| + ^\s*(from\s+\S+\s+)?import\s+.+$) + +# Allow the body of an if to be on the same line as the test if there is no +# else. +single-line-if-stmt=yes + +# Maximum number of lines in a module +max-module-lines=99999 + +# String used as indentation unit. The internal Google style guide mandates 2 +# spaces. Google's externaly-published style guide says 4, consistent with +# PEP 8. Here, we use 2 spaces, for conformity with many open-sourced Google +# projects (like TensorFlow). +indent-string=' ' + +# Number of spaces of indent required inside a hanging or continued line. +indent-after-paren=4 + +# Expected format of line ending, e.g. empty (any line ending), LF or CRLF. +expected-line-ending-format= + + +[MISCELLANEOUS] + +# List of note tags to take in consideration, separated by a comma. +notes=TODO + + +[STRING] + +# This flag controls whether inconsistent-quotes generates a warning when the +# character used as a quote delimiter is used inconsistently within a module. +check-quote-consistency=yes + + +[VARIABLES] + +# Tells whether we should check for unused import in __init__ files. +init-import=no + +# A regular expression matching the name of dummy variables (i.e. expectedly +# not used). +dummy-variables-rgx=^\*{0,2}(_$|unused_|dummy_) + +# List of additional names supposed to be defined in builtins. Remember that +# you should avoid to define new builtins when possible. +additional-builtins= + +# List of strings which can identify a callback function by name. A callback +# name must start or end with one of those strings. +callbacks=cb_,_cb + +# List of qualified module names which can have objects that can redefine +# builtins. +redefining-builtins-modules=six,six.moves,past.builtins,future.builtins,functools + + +[LOGGING] + +# Logging modules to check that the string format arguments are in logging +# function parameter format +logging-modules=logging,absl.logging,tensorflow.io.logging + + +[SIMILARITIES] + +# Minimum lines number of a similarity. +min-similarity-lines=4 + +# Ignore comments when computing similarities. +ignore-comments=yes + +# Ignore docstrings when computing similarities. +ignore-docstrings=yes + +# Ignore imports when computing similarities. +ignore-imports=no + + +[SPELLING] + +# Spelling dictionary name. Available dictionaries: none. To make it working +# install python-enchant package. +spelling-dict= + +# List of comma separated words that should not be checked. +spelling-ignore-words= + +# A path to a file that contains private dictionary; one word per line. +spelling-private-dict-file= + +# Tells whether to store unknown words to indicated private dictionary in +# --spelling-private-dict-file option instead of raising a message. +spelling-store-unknown-words=no + + +[IMPORTS] + +# Deprecated modules which should not be used, separated by a comma +deprecated-modules=regsub, + TERMIOS, + Bastion, + rexec, + sets + +# Create a graph of every (i.e. internal and external) dependencies in the +# given file (report RP0402 must not be disabled) +import-graph= + +# Create a graph of external dependencies in the given file (report RP0402 must +# not be disabled) +ext-import-graph= + +# Create a graph of internal dependencies in the given file (report RP0402 must +# not be disabled) +int-import-graph= + +# Force import order to recognize a module as part of the standard +# compatibility libraries. +known-standard-library= + +# Force import order to recognize a module as part of a third party library. +known-third-party=enchant, absl + +# Analyse import fallback blocks. This can be used to support both Python 2 and +# 3 compatible code, which means that the block might have code that exists +# only in one or another interpreter, leading to false positives when analysed. +analyse-fallback-blocks=no + + +[CLASSES] + +# List of method names used to declare (i.e. assign) instance attributes. +defining-attr-methods=__init__, + __new__, + setUp + +# List of member names, which should be excluded from the protected access +# warning. +exclude-protected=_asdict, + _fields, + _replace, + _source, + _make + +# List of valid names for the first argument in a class method. +valid-classmethod-first-arg=cls, + class_ + +# List of valid names for the first argument in a metaclass class method. +valid-metaclass-classmethod-first-arg=mcs diff --git a/sample_backtest.py b/sample_backtest.py new file mode 100644 index 0000000..122f158 --- /dev/null +++ b/sample_backtest.py @@ -0,0 +1,28 @@ +import asyncio +import logging +from datetime import datetime, UTC + +from aiomql.lib.backtest_runner import BackTestRunner +from aiomql.core import Config +from aiomql.contrib.strategies import FingerTrap, Chaos +from aiomql.contrib.symbols import ForexSymbol +from aiomql.contrib.backtesting import BackTestEngine + + +async def back_tester(): + config = Config() + config.mode = 'backtest' + logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s') + syms = ['Volatility 75 Index', 'Volatility 100 Index', 'Volatility 50 Index'] + symbols = [ForexSymbol(name=sym) for sym in syms] + stgs = [Chaos(symbol=symbol) for symbol in symbols] + start = datetime(2021, 1, 1, tzinfo=UTC) + end = datetime(2021, 1, 7, tzinfo=UTC) + back_test_engine = BackTestEngine(start=start, end=end) + await back_test_engine.setup_account(balance=100) + back_test_runner = BackTestRunner(strategies=stgs, backtest_engine=back_test_engine) + await back_test_runner.run() + + +asyncio.run(back_tester()) +print("Bot executed successfully") diff --git a/sample_bot.py b/sample_bot.py new file mode 100644 index 0000000..f7bad54 --- /dev/null +++ b/sample_bot.py @@ -0,0 +1,22 @@ +import logging + +from aiomql.lib.bot import Bot +from aiomql.core import Config +from aiomql.contrib.strategies import FingerTrap, Chaos +from aiomql.contrib.symbols import ForexSymbol + + +def bot(): + logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s') + syms = ['Volatility 75 Index', 'Volatility 100 Index', 'Volatility 50 Index'] + symbols = [ForexSymbol(name=sym) for sym in syms] + stgs = [Chaos(symbol=symbol) for symbol in symbols] + config = Config() + bot = Bot() + bot.executor.timeout = 60 + bot.add_strategies(strategies=stgs) + bot.execute() + + +bot() +print("Bot executed successfully") diff --git a/src/aiomql/_utils.py b/src/aiomql/_utils.py index 4ba1538..71f03f9 100644 --- a/src/aiomql/_utils.py +++ b/src/aiomql/_utils.py @@ -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()) -def backoff_decorator(func=None, *, max_retries: int = 3, retries: int = 0, error='') -> callable: +def backoff_decorator(func=None, *, max_retries: int = 2, retries: int = 0, error='') -> callable: if func is None: return partial(backoff_decorator, max_retries=max_retries, retries=retries, error=error) @@ -41,16 +41,16 @@ def backoff_decorator(func=None, *, max_retries: int = 3, retries: int = 0, erro retries = 0 return res except Exception as err: - logger.error(f'Error in {func.__name__}: {err}') - await asyncio.sleep(retries + random.randint(1, max_retries)) + logger.error('Error in %s: %s', func.__name__, err) + await asyncio.sleep(2**retries + random.uniform(0, 1)) await wrapper(*args, **kwargs) return wrapper -def error_handler(func=None, *, msg='', exe = Exception, response=None): +def error_handler(func=None, *, msg='', exe = Exception, response=None, log_error_msg=True): if func is None: - return partial(error_handler, msg=msg, exe=exe, response=response) + return partial(error_handler, msg=msg, exe=exe, response=response, log_error_msg=log_error_msg) @wraps(func) async def wrapper(*args, **kwargs): @@ -58,14 +58,15 @@ def error_handler(func=None, *, msg='', exe = Exception, response=None): res = await func(*args, **kwargs) return res except exe as err: - logger.error(f'Error in {func.__name__}: {msg or err}') + if log_error_msg: + logger.error(f'Error in {func.__name__}: {msg or err}') return response return wrapper -def error_handler_sync(func=None, *, msg='', exe=Exception, response=None): +def error_handler_sync(func=None, *, msg='', exe=Exception, response=None, log_error_msg=True): if func is None: - return partial(error_handler, msg=msg, exe=exe, response=response) + return partial(error_handler, msg=msg, exe=exe, response=response, log_error_msg=log_error_msg) @wraps(func) def wrapper(*args, **kwargs): @@ -73,17 +74,18 @@ def error_handler_sync(func=None, *, msg='', exe=Exception, response=None): res = func(*args, **kwargs) return res except exe as err: - logger.error(f'Error in {func.__name__}: {msg or err}') + if log_error_msg: + logger.error(f'Error in {func.__name__}: {msg or err}') return response return wrapper -def round_down(value: int, base: int) -> int: - return value if value % base == 0 else value - (value % base) +def round_down(value: int | float, base: int) -> int: + return int(value) if value % base == 0 else int(value - (value % base)) -def round_up(value: int, base: int) -> int: - return value if value % base == 0 else value + base - (value % base) +def round_up(value: int | float, base: int) -> int: + return int(value) if value % base == 0 else int(value + base - (value % base)) # noinspection PyShadowingNames diff --git a/src/aiomql/contrib/backtesting/__init__.py b/src/aiomql/contrib/backtesting/__init__.py index c57b924..3f124fa 100644 --- a/src/aiomql/contrib/backtesting/__init__.py +++ b/src/aiomql/contrib/backtesting/__init__.py @@ -1,4 +1,4 @@ -from .get_data import GetData, TestData +from .get_data import GetData, BackTestData from .backtest_engine import BackTestEngine from .backtest_account import BackTestAccount from .trades_manager import PositionsManager, OrdersManager, DealsManager diff --git a/src/aiomql/contrib/backtesting/backtest_engine.py b/src/aiomql/contrib/backtesting/backtest_engine.py index b088041..829b8a6 100644 --- a/src/aiomql/contrib/backtesting/backtest_engine.py +++ b/src/aiomql/contrib/backtesting/backtest_engine.py @@ -1,13 +1,13 @@ import asyncio -from datetime import datetime +from datetime import datetime, UTC from typing import Literal from itertools import zip_longest import random from functools import cached_property from logging import getLogger +from math import ceil import pandas as pd -import pytz import numpy as np from pandas import DataFrame from MetaTrader5 import (Tick, SymbolInfo, AccountInfo, TradeOrder, TradePosition, TradeDeal, @@ -19,7 +19,7 @@ from ...core.constants import (TimeFrame, OrderType, TradeAction, AccountStopOut 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 BackTestData, GetData, Cursor from .backtest_account import BackTestAccount from .trades_manager import PositionsManager, OrdersManager, DealsManager @@ -30,6 +30,7 @@ class BackTestEngine: mt5: MetaTrader span: range range: range + speed: int cursor: Cursor iter: zip_longest rates: dict[str, dict[int, DataFrame]] @@ -39,45 +40,71 @@ class BackTestEngine: deals: DealsManager positions: PositionsManager _account: BackTestAccount + stop_testing: bool + use_terminal: bool + restart: bool + stop_time: int | None - - def __init__(self, *, data: TestData = None, speed: int = 1, start: float | datetime = 0, - end: float | datetime = 0, restart: bool = False, name: str = ''): - self._data = data or TestData() + def __init__(self, *, data: BackTestData = None, speed: int = 1, start: float | datetime = 0, + end: float | datetime = 0, restart: bool = True, use_terminal: bool = None, name: str = '', + stop_time: float | datetime = None): + self._data = data or BackTestData() 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.prepare_data() - 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.setup_test_range(start=start, end=end, speed=speed, restart=restart) + self.setup_data(restart=restart) + start, end = (self.span[0], self.span[-1]) if len(self.span) >= 2 else ((now := datetime.now(UTC).timestamp()), now) + self.name = name or self._data.name or f"backtest_data_{datetime.now(tz=UTC):%d_%m_%y}" + self.stop_testing = False + self.use_terminal = self.config.use_terminal_for_backtesting if use_terminal is None else use_terminal + if stop_time is not None: + val = stop_time.astimezone(tz=UTC) if isinstance(start, datetime) else datetime.fromtimestamp(start, tz=UTC) + stop_time = int(val.timestamp()) + self.stop_time = stop_time def __next__(self) -> Cursor: try: index, time = next(self.iter) + if self.stop_time and time >= self.stop_time: + raise StopIteration self.cursor = Cursor(index=index, time=time) return self.cursor except StopIteration: - logger.warning('End of time') + logger.critical("End of the test range") + self.stop_testing = True def __repr__(self): return f"{self.__class__.__name__}()" - def set_up(self, *, start: float | datetime = 0, end: float | datetime = 0, speed: int = 1, restart: bool = False): - span_start = (int(start.timestamp()) if isinstance(start, datetime) else int(start)) or self._data.span.start - span_end = (int(end.timestamp()) if isinstance(end, datetime) else int(end)) or self._data.span.stop + def setup_test_range(self, *, start: float | datetime = None, end: float | datetime = None, speed: int = 1, + restart: bool = True): + if self._data.span and self._data.range: + start = start or self._data.span[0] + end = end or self._data.span[-1] + 1 + start = start.astimezone(tz=UTC) if isinstance(start, datetime) else datetime.fromtimestamp(start, tz=UTC) + end = end.astimezone(tz=UTC) if isinstance(end, datetime) else datetime.fromtimestamp(end, tz=UTC) + span_start = int(start.timestamp()) + span_end = int(end.timestamp()) + self.speed = speed self.span = range(span_start, span_end, speed) self.range = range(0, span_end - span_start, speed) self.iter = zip_longest(self.range, self.span) - if restart and self._data.cursor is not None: + if restart is False and self._data.cursor is not None: self.cursor = self._data.cursor self.go_to(time=self.cursor.time) else: - self.cursor: Cursor = self._data.cursor or Cursor(index=self.range.start, time=self.span.start) + self.cursor = Cursor(index=self.range.start, time=self.span.start) + + def setup_data(self, *, restart: bool = True): + if restart is True: + self.orders = OrdersManager() + self.positions = PositionsManager() + self.deals = DealsManager() + self._account = BackTestAccount() + return - def prepare_data(self): orders = {} for ticket, order in self._data.orders.items(): orders[ticket] = TradeOrder((order.get(k) for k in TradeOrder.__match_args__)) @@ -94,7 +121,7 @@ class BackTestEngine: deals[ticket] = TradeDeal((deal.get(k) for k in TradeDeal.__match_args__)) self.deals = DealsManager(data=deals) - self._account: BackTestAccount = BackTestAccount(**self._data.account) + self._account = BackTestAccount(**self._data.account) def next(self) -> Cursor: return next(self) @@ -103,23 +130,20 @@ class BackTestEngine: def data(self): return self._data - def reset(self): + def reset(self, clear_data: bool = False): self.iter = zip_longest(self.range, self.span) self.cursor = Cursor(index=self.range.start, time=self.span.start) + if clear_data: + self.setup_data(restart=True) def go_to(self, *, time: datetime | float): - time = int(time.timestamp()) if isinstance(time, datetime) else int(time) + time = time.astimezone(tz=UTC) if isinstance(time, datetime) else datetime.fromtimestamp(time, tz=UTC) + time = int(time.timestamp()) steps = time - self.cursor.time - - if steps > 0: + if 0 <= steps < (len(self.range) - 1): self.fast_forward(steps=steps) return - - range_start = time - self.span.start - span = range(time, self.span.stop, self.span.step) - range_ = range(range_start, self.range.stop, self.range.step) - self.iter = zip_longest(range_, span) - self.cursor = Cursor(index=range_.start, time=span.start) + raise ValueError("Can't go back in time or beyond the limits of the range") def fast_forward(self, *, steps: int): for _ in range(steps): @@ -133,7 +157,9 @@ class BackTestEngine: pos_tasks = [self.check_position(ticket=ticket) for ticket in self.positions._open_positions] await asyncio.gather(*pos_tasks) profit = sum(pos.profit for pos in self.positions.open_positions) + profit = round(profit, self._account.currency_digits) self.update_account(profit=profit) + self.check_account() def wrap_up(self): try: @@ -143,23 +169,29 @@ class BackTestEngine: self._data.open_positions = self.positions._open_positions self._data.margins = self.positions.margins self._data.account = self._account.asdict() + self._data.cursor = self.cursor + self._data.span = self.span + self._data.range = self.range + self._data.account = self._account.asdict() name = self._data.name or self.name + self._data.name = name path = self.config.backtest_dir/f"{name}.pkl" GetData.pickle_data(data=self._data, name=path) except Exception as err: - print(err) + logger.error("Error in wrap_up: %s", err) @async_cache async def get_price_tick(self, *, symbol: str, time: int) -> Tick | None: try: - if self.config.use_terminal_for_backtesting: + if self.use_terminal: + time = datetime.fromtimestamp(time, tz=UTC) tick = await self.mt5.copy_ticks_from(symbol, time, 1, CopyTicks.ALL) return Tick(tick[-1]) if tick else None - tick = self.prices[symbol].loc[self.cursor.time] + tick = self.prices[symbol].loc[time] return Tick(tick) except Exception as exe: - logger.error(f"Error Getting Price Tick: {exe}") + logger.error("Error Getting Price Tick: %s", exe) @error_handler async def check_order(self, *, ticket: int): @@ -178,17 +210,39 @@ class BackTestEngine: if not (tp and sl): return + deal = {'ticket': random.randint(100_000_000, 999_999_999), 'order': ticket, + 'symbol': symbol, 'commission': 0, 'swap': 0, 'position_id': ticket, + 'fee': 0, 'time': self.cursor.time, 'time_msc': self.cursor.time * 1000, + 'price': tick.bid, 'type': DealType(order_type), + 'reason': DealReason.EXPERT, 'entry': DealEntry.OUT, 'profit': 0} + match order_type: case OrderType.BUY: - if tp >= tick.bid or sl <= tick.bid: - self.close_position(ticket=ticket) + if tick.bid >= tp or tick.bid <= sl: + res = self.close_position(ticket=ticket) + if res: + pos = self.positions.get(ticket) + deal.update({'profit': pos.profit, 'volume': pos.volume}) + self.deals[deal['ticket']] = TradeDeal((deal.get(k, 0) for k in TradeDeal.__match_args__)) + case OrderType.SELL: - if tp <= tick.ask or sl >= tick.ask: - self.close_position(ticket=ticket) + if tick.ask <= tp or tick.ask >= sl: + res = self.close_position(ticket=ticket) + if res: + pos = self.positions.get(ticket) + deal.update({'profit': pos.profit, 'volume': pos.volume}) + self.deals[deal['ticket']] = TradeDeal((deal.get(k, 0) for k in TradeDeal.__match_args__)) case _: ... + def check_account(self): + account = self._account + level = account.margin_level if account.margin_so_mode == AccountStopOutMode.PERCENT else account.margin_so_call + if level < account.margin_so_call and level != 0: + logger.critical("Account has burned out!!! Please top up to continue trading") + self.stop_testing = True + @error_handler async def check_position(self, *, ticket: int): """ @@ -208,7 +262,7 @@ class BackTestEngine: time_update=self.cursor.time) await self.check_order(ticket=ticket) - @error_handler_sync(response=False) + # @error_handler_sync(response=False) def close_position(self, *, ticket: int) -> bool: """ Close an open position for the trading account using the position ticket. @@ -219,13 +273,18 @@ class BackTestEngine: Returns: bool: True if the position is closed successfully, False otherwise """ - position = self.positions[ticket] - margin = self.positions.get_margin(ticket=ticket) - self.positions.delete_margin(ticket=ticket) - self.positions.close(ticket=ticket) - self.orders.update(ticket=ticket, time_done=self.cursor.time, time_done_msc=self.cursor.time*1000) - self.update_account(gain=position.profit, margin=-margin) - return True + try: + position = self.positions[ticket] + margin = self.positions.get_margin(ticket=ticket) + self.positions.delete_margin(ticket=ticket) + self.positions.close(ticket=ticket) + self.orders.update(ticket=ticket, time_done=self.cursor.time, time_done_msc=self.cursor.time*1000) + gain = round(position.profit, self._account.currency_digits) + self.update_account(gain=gain, margin=-margin) + return True + except Exception as exe: + logger.error("Error Closing Position %d: %s", ticket, exe) + return False @error_handler(response=False) def modify_stops(self, *, ticket: int, sl: int, tp: int) -> bool: @@ -258,10 +317,10 @@ class BackTestEngine: level = self._account.equity / self._account.margin * 100 self._account.margin_level = level if mode == AccountStopOutMode.PERCENT else self._account.margin_free - def deposit(self, amount: float): + def deposit(self, *, amount: float): self.update_account(gain=amount) - def withdraw(self, amount: float): + def withdraw(self, *, amount: float): assert amount <= self._account.balance, 'Insufficient funds' self.update_account(gain=-amount) @@ -272,7 +331,7 @@ class BackTestEngine: 'balance': self._account.balance, **{k: v for k, v in kwargs.items() if k in self._account.__match_args__}} - if self.config.use_terminal_for_backtesting: + if self.use_terminal: acc_info = await self.mt5.account_info() default = {**acc_info._asdict(), **default} @@ -282,12 +341,12 @@ class BackTestEngine: @cached_property def prices(self) -> dict[str, DataFrame]: prices = {} - for symbol in self._data.prices.keys(): - res = self._data.prices[symbol] + for symbol in self._data.ticks.keys(): + res = self._data.ticks[symbol] res = pd.DataFrame(res) res.drop_duplicates(subset=['time'], keep='last', inplace=True) res.set_index('time', inplace=True, drop=False) - res.reindex(self.span) # fill in missing values with NaN + res = res.reindex(self.span, copy=True, method='ffill') # fill in missing values with NaN prices[symbol] = res return prices @@ -297,8 +356,6 @@ class BackTestEngine: for symbol in self._data.ticks.keys(): res = self._data.ticks[symbol] res = pd.DataFrame(res) - res.drop_duplicates(subset=['time'], keep='last', inplace=True) - res.set_index('time', inplace=True, drop=False) ticks[symbol] = res return ticks @@ -309,9 +366,7 @@ class BackTestEngine: for timeframe in self._data.rates[symbol].keys(): res = self._data.rates[symbol][timeframe] res = pd.DataFrame(res) - res.drop_duplicates(subset=['time'], keep='last', inplace=True) - res.set_index('time', inplace=True, drop=False) - rates[symbol][timeframe] = res + rates.setdefault(symbol, {})[timeframe] = res return rates @cached_property @@ -322,11 +377,11 @@ class BackTestEngine: return symbols @error_handler - async def order_send(self, *, request: dict) -> OrderSendResult: + async def order_send(self, *, request: dict, use_terminal=False) -> OrderSendResult: osr = {'retcode': 10013, 'comment': 'Invalid request', 'request': TradeRequest(request.get(k, (0 if k != 'comment' else '')) for k in TradeRequest.__match_args__)} - current_tick = await self.get_price_tick(request.get('symbol'), self.cursor.time) + current_tick = await self.get_price_tick(symbol=request.get('symbol'), time=self.cursor.time) if current_tick is None: osr['comment'] = 'Market is closed' osr['retcode'] = 10018 @@ -344,9 +399,10 @@ class BackTestEngine: # closing an order by an opposite order using a position ticket and Deal action if action == TradeAction.DEAL and current_position and order_type.opposite == current_position.type: - res = self.close_position(current_position.ticket) + res = self.close_position(ticket=current_position.ticket) if res: price_current = current_tick.ask if order_type == OrderType.BUY else current_tick.bid + # self.orders.update(ticket=ticket, time_done=self.cursor.time, time_done_msc=self.cursor.time * 1000) trade_order.update({'position_id': current_position.ticket, 'ticket': order_ticket, 'time_setup': current_tick.time, 'time_setup_msc': current_tick.time_msc, 'time_done': current_tick.time, 'time_done_msc': current_tick.time_msc, @@ -357,7 +413,7 @@ class BackTestEngine: # TODO: calculate commission and swap if possible or necessary deal = {'ticket': deal_ticket, 'position_id': current_position.ticket, 'order': order_ticket, - 'symbol': symbol, 'time': current_tick.time, + 'symbol': symbol, 'time': current_tick.time, 'profit': current_position.profit, 'time_msc': current_tick.time_msc, 'volume': current_position.volume, 'price': price_current, 'type': DealType(order_type), 'reason': DealReason.EXPERT, 'entry': DealEntry.OUT, 'comment': '', 'external_id': ''} @@ -371,7 +427,7 @@ class BackTestEngine: return OrderSendResult((osr.get(k, 0) for k in OrderSendResult.__match_args__)) if action == TradeAction.SLTP and current_position: - check = await self.order_check(position_id) + check = await self.order_check(request=request, use_terminal=use_terminal) if check.retcode != 0: osr = {'retcode': check.retcode, 'comment': check.comment, 'request': check.request} return OrderSendResult((osr.get(k, 0) for k in OrderSendResult.__match_args__)) @@ -382,7 +438,7 @@ class BackTestEngine: return OrderSendResult((osr.get(k, 0) for k in OrderSendResult.__match_args__)) if action == TradeAction.DEAL and order_type in (OrderType.BUY, OrderType.SELL): - check = await self.order_check(request=request) + check = await self.order_check(request=request, use_terminal=use_terminal) if check.retcode != 0: osr = {'retcode': check.retcode, 'comment': check.comment, 'request': check.request} return OrderSendResult((osr.get(k, 0) for k in OrderSendResult.__match_args__)) @@ -397,7 +453,7 @@ class BackTestEngine: deal = {'ticket': deal_ticket, 'order': order_ticket, 'symbol': symbol, 'commission': 0, 'swap': 0, 'position_id': order_ticket, 'fee': 0, 'time': current_tick.time, 'time_msc': current_tick.time_msc, 'volume': volume, 'price': price, 'type': DealType(order_type), 'reason': DealReason.EXPERT, - 'entry': DealEntry.IN} + 'entry': DealEntry.IN, 'profit': 0} # ToDo: set time_expiration based on order_type_time trade_order.update({'ticket': order_ticket, 'symbol': symbol, 'volume': volume, 'price': price, @@ -411,15 +467,17 @@ class BackTestEngine: self.deals[deal_ticket] = deal self.positions[order.ticket] = pos self.orders[order.ticket] = order - osr.update({'order': order_ticket, 'price': price, 'volume': volume, 'bid': current_tick.bid, + osr.update({'comment': 'Request completed', 'retcode': 10009, 'order': order_ticket, 'price': price, + 'volume': volume, 'bid': current_tick.bid, 'ask': current_tick.ask, 'deal': deal_ticket}) - margin = await self.order_calc_margin(action=action, symbol=symbol, volume=volume, price=price) + margin = await self.order_calc_margin(action=action, symbol=symbol, volume=volume, price=price, + use_terminal=use_terminal) self.positions.set_margin(ticket=order_ticket, margin=margin) self.update_account(margin=margin) return OrderSendResult((osr.get(k, 0) for k in OrderSendResult.__match_args__)) @error_handler - async def order_check(self, *, request: dict) -> OrderCheckResult: + async def order_check(self, *, request: dict, use_terminal: bool = False) -> OrderCheckResult: ocr = {'retcode': 10013, 'balance': 0, 'profit': 0, 'margin': 0, 'equity': 0, 'margin_free': 0, 'margin_level': 0, 'comment': 'Invalid request', 'request': TradeRequest(request.get(k, (0 if k != 'comment' else '')) for k in @@ -430,7 +488,8 @@ class BackTestEngine: # check margin and confirm order can go through for a deal action and buy or sell order type if action == TradeAction.DEAL and order_type in (OrderType.BUY, OrderType.SELL) and position_id is None: - margin = await self.order_calc_margin(action, symbol, volume, price) + margin = await self.order_calc_margin(action=action, symbol=symbol, volume=volume, + price=price, use_terminal=use_terminal) if margin is None: return OrderCheckResult((ocr.get(k, 0) for k in OrderCheckResult.__match_args__)) @@ -439,7 +498,6 @@ class BackTestEngine: level = self._account.equity / used_margin * 100 if used_margin else float('inf') margin_level = level if self._account.margin_so_mode == AccountStopOutMode.PERCENT else free_margin - ocr.update({'margin_level': margin_level, 'margin': margin, 'margin_free': free_margin}) # check if the account has enough money @@ -472,7 +530,7 @@ class BackTestEngine: ocr['retcode'] = 0 return OrderCheckResult((ocr.get(k, 0) for k in OrderCheckResult.__match_args__)) - if self.mt5.config.use_terminal_for_backtesting: + if use_terminal or self.use_terminal: ocr_t = await self.mt5.order_check(request) if ocr_t.retcode in (10013, 10014): return ocr_t @@ -490,28 +548,28 @@ class BackTestEngine: @error_handler async def get_terminal_info(self) -> TerminalInfo: - if self.config.use_terminal_for_backtesting: + if self.use_terminal: res = await self.mt5.terminal_info() return res return TerminalInfo(self._data.terminal) @error_handler async def get_version(self) -> tuple[int, int, str]: - if self.config.use_terminal_for_backtesting: + if self.use_terminal: res = await self.mt5.version() return res return self._data.version @error_handler async def get_symbols_total(self) -> int: - if self.config.use_terminal_for_backtesting: + if self.use_terminal: syms = await self.mt5.symbols_total() return syms return len(self.symbols) @error_handler - async def get_symbols(self, group: str = '') -> tuple[SymbolInfo, ...]: - if self.config.use_terminal_for_backtesting: + async def get_symbols(self, *, group: str = '') -> tuple[SymbolInfo, ...]: + if self.use_terminal: syms = await self.mt5.symbols_get(group=group) return syms return tuple(list(self.symbols.values())) @@ -522,93 +580,109 @@ class BackTestEngine: @error_handler async def get_symbol_info_tick(self, *, symbol: str) -> Tick | None: - tick = await self.get_price_tick(symbol, self.cursor.time) + tick = await self.get_price_tick(symbol=symbol, time=self.cursor.time) return tick @error_handler async def get_symbol_info(self, *, symbol: str) -> SymbolInfo: - if self.config.use_terminal_for_backtesting: + if self.use_terminal: info = await self.mt5.symbol_info(symbol) else: info = self.symbols[symbol] - - tick = await self.get_symbol_info_tick(symbol) + tick = await self.get_symbol_info_tick(symbol=symbol) info = info._asdict() | {'bid': tick.bid, 'bidhigh': tick.bid, 'bidlow': tick.bid, 'ask': tick.ask, 'askhigh': tick.ask, 'asklow': tick.bid, 'last': tick.last, 'volume_real': tick.volume_real} return SymbolInfo((info.get(key) for key in SymbolInfo.__match_args__)) @error_handler - async def get_rates_from(self, symbol: str, timeframe: TimeFrame, date_from: datetime | float, count: int) -> np.ndarray: - if self.config.use_terminal_for_backtesting: + async def get_rates_from(self, *, symbol: str, timeframe: TimeFrame, date_from: + datetime | float, count: int) -> np.ndarray: + date_from = date_from.astimezone(tz=UTC) if isinstance(date_from, datetime) else datetime.fromtimestamp( + date_from, tz=UTC) + if self.use_terminal: rates = await self.mt5.copy_rates_from(symbol, timeframe, date_from, count) return rates rates = self.rates[symbol][timeframe] - start = int(datetime.timestamp(date_from)) if isinstance(date_from, datetime) else int(date_from) - start = round_down(start, timeframe.time) + start = int(date_from.timestamp()) + start = round_down(start, timeframe.seconds) rates = rates[rates.time <= start].iloc[-count:] return np.fromiter((tuple(i) for i in rates.iloc), dtype=self.get_dtype(df=rates)) - @error_handler - async def get_rates_from_pos(self, symbol: str, timeframe: TimeFrame, start_pos: int, count: int) -> np.ndarray: - if self.config.use_terminal_for_backtesting: - now = datetime.now(tz=pytz.UTC) + # @error_handler + async def get_rates_from_pos(self, *, symbol: str, timeframe: TimeFrame, start_pos: int, count: int) -> np.ndarray: + if self.use_terminal: + # Todo: Optimize this!!! + now = datetime.now(tz=UTC).timestamp() b_now = self.cursor.time - diff = (now.timestamp() - b_now) // timeframe.time + diff = ceil((now - b_now) / timeframe.seconds) start_pos = int(diff + start_pos) + # print('hehk', start_pos, count) res = await self.mt5.copy_rates_from_pos(symbol, timeframe, start_pos, count) return res rates = self.rates[symbol][timeframe] - end = abs(self.cursor.index - start_pos) - start = abs(end - count) - rates = rates.iloc[start:end] + + # the current time rounded up to a multiple of the timeframe in seconds and then subtracted by the start_pos + # multiplied by the timeframe in seconds gives the time of the last candlestick in the range when using + # copy_rates_from_pos + end = int(round_down(self.cursor.time, timeframe.seconds)) - start_pos * timeframe.seconds + start = end - count * timeframe.seconds + rates = rates[(rates.time > start) & (rates.time <= end)] return np.fromiter((tuple(i) for i in rates.iloc), dtype=self.get_dtype(df=rates)) @error_handler - async def get_rates_range(self, symbol: str, timeframe: TimeFrame, date_from: datetime | float, + async def get_rates_range(self, *, symbol: str, timeframe: TimeFrame, date_from: datetime | float, date_to: datetime | float) -> np.ndarray: - if self.config.use_terminal_for_backtesting: + date_from = date_from.astimezone(tz=UTC) if isinstance(date_from, datetime) else datetime.fromtimestamp( + date_from, tz=UTC) + date_to = date_to.astimezone(tz=UTC) if isinstance(date_to, datetime) else datetime.fromtimestamp( + date_to, tz=UTC) + if self.use_terminal: rates = await self.mt5.copy_rates_range(symbol, timeframe, date_from, date_to) return rates rates = self.rates[symbol][timeframe] - start = int(datetime.timestamp(date_from)) if isinstance(date_from, datetime) else int(date_from) - start = round_down(start, timeframe.time) - end = int(datetime.timestamp(date_to)) if isinstance(date_to, datetime) else int(date_to) - end = round_up(end, timeframe.time) + start = round_up(int(date_from.timestamp()), timeframe.seconds) + end = round_up(int(date_to.timestamp()), timeframe.seconds) rates = rates[(rates.time >= start) & (rates.time <= end)] return np.fromiter((tuple(i) for i in rates.iloc), dtype=self.get_dtype(df=rates)) @error_handler async def get_ticks_from(self, *, symbol: str, date_from: datetime | float, count: int, - flags: CopyTicks) -> np.ndarray: - if self.config.use_terminal_for_backtesting: + flags: CopyTicks = CopyTicks.ALL) -> np.ndarray: + date_from = date_from.astimezone(tz=UTC) if isinstance(date_from, datetime) else datetime.fromtimestamp( + date_from, tz=UTC) + if self.use_terminal: ticks = await self.mt5.copy_ticks_from(symbol, date_from, count, flags) return ticks ticks = self.ticks[symbol] - start = int(datetime.timestamp(date_from)) if isinstance(date_from, datetime) else int(date_from) - ticks = ticks[ticks.time <= start].iloc[-count:] - return np.fromiter((tuple(i) for i in ticks.iloc), dtype=self.get_dtype(df=ticks)) + start = int(date_from.timestamp()) + rates = ticks[ticks.time <= start].iloc[-count:] + return np.fromiter((tuple(i) for i in rates.iloc), dtype=self.get_dtype(df=rates)) @error_handler async def get_ticks_range(self, *, symbol: str, date_from: datetime | float, - date_to: datetime | float, flags: CopyTicks) -> np.ndarray: - if self.config.use_terminal_for_backtesting: + date_to: datetime | float, flags: CopyTicks = CopyTicks.ALL) -> np.ndarray: + date_from = date_from.astimezone(tz=UTC) if isinstance(date_from, datetime) else datetime.fromtimestamp( + date_from, tz=UTC) + date_to = date_to.astimezone(tz=UTC) if isinstance(date_to, datetime) else datetime.fromtimestamp( + date_to, tz=UTC) + if self.use_terminal: ticks = await self.mt5.copy_ticks_range(symbol, date_from, date_to, flags) return ticks ticks = self.ticks[symbol] - start = int(datetime.timestamp(date_from)) if isinstance(date_from, datetime) else int(date_from) - end = int(datetime.timestamp(date_to)) if isinstance(date_to, datetime) else int(date_to) - ticks = ticks[(ticks.time >= start) & (ticks.time <= end)] - return np.fromiter((tuple(i) for i in ticks.iloc), dtype=self.get_dtype(df=ticks)) + start = int(date_from.timestamp()) + end = int(date_to.timestamp()) + rates = ticks[(ticks.time >= start) & (ticks.time <= end)] + return np.fromiter((tuple(i) for i in rates.iloc), dtype=self.get_dtype(df=ticks)) @error_handler async def order_calc_margin(self, *, action: Literal[OrderType.BUY, OrderType.SELL], symbol: str, volume: float, - price: float): - if self.mt5.config.use_terminal_for_backtesting: + price: float, use_terminal: bool = False): + if use_terminal or self.use_terminal: return await self.mt5.order_calc_margin(action, symbol, volume, price) sym = self.symbols[symbol] @@ -617,9 +691,9 @@ class BackTestEngine: @error_handler async def order_calc_profit(self, *, action: Literal[OrderType.BUY, OrderType.SELL], symbol: str, volume: float, - price_open: float, price_close: float): + price_open: float, price_close: float, use_terminal = False): - if self.mt5.config.use_terminal_for_backtesting: + if use_terminal or self.use_terminal: return await self.mt5.order_calc_profit(action, symbol, volume, price_open, price_close) sym = self.symbols[symbol] @@ -638,7 +712,7 @@ class BackTestEngine: return 0 @error_handler_sync - def get_orders(self, symbol: str = '', group: str = '', ticket: int = None) -> tuple[TradeOrder, ...]: + def get_orders(self, *, symbol: str = '', group: str = '', ticket: int = None) -> tuple[TradeOrder, ...]: """ Get pending orders from the terminal history. This has to do with pending orders, which this backtester doesn't support yet. diff --git a/src/aiomql/contrib/backtesting/get_data.py b/src/aiomql/contrib/backtesting/get_data.py index 856adcc..be8a473 100644 --- a/src/aiomql/contrib/backtesting/get_data.py +++ b/src/aiomql/contrib/backtesting/get_data.py @@ -1,12 +1,11 @@ from dataclasses import dataclass, field, fields import pickle from pathlib import Path -from datetime import datetime +from datetime import datetime, UTC from logging import getLogger from typing import Sequence, NamedTuple import MetaTrader5 -import pytz from numpy import ndarray from ...core.meta_trader import MetaTrader @@ -24,13 +23,12 @@ class Cursor(NamedTuple): @dataclass -class TestData: +class BackTestData: name: str = '' terminal: dict[str, [str | int | bool | float]] = field(default_factory=dict) version: tuple[int, int, str] = (0, 0, '') account: dict = field(default_factory=dict) symbols: dict[str, dict] = field(default_factory=dict) - prices: dict[str, ndarray] = field(default_factory=dict) ticks: dict[str, ndarray] = field(default_factory=dict) rates: dict[str, dict[int, ndarray]] = field(default_factory=dict) span: range = range(0) @@ -41,6 +39,7 @@ class TestData: open_positions: set[int, ...] = field(default_factory=lambda: set()) cursor: Cursor = None margins: dict[int, float] = field(default_factory=lambda: {}) + fully_loaded: bool = True def __str__(self): return f"{self.name}" @@ -57,27 +56,27 @@ class TestData: class GetData: - data: TestData + data: BackTestData def __init__(self, *, start: datetime, end: datetime, symbols: Sequence[str], - timeframes: Sequence[TimeFrame], name: str = '', tz: str = 'Etc/UTC'): + timeframes: Sequence[TimeFrame], name: str = ''): """""" self.config = Config() - self.tz = pytz.timezone(tz) - self.start = start.replace(tzinfo=self.tz) - self.end = end.replace(tzinfo=self.tz) + self.start = start.astimezone(tz=UTC) + self.end = end.astimezone(tz=UTC) self.symbols = set(symbols) self.timeframes = set(timeframes) self.name = name or f"{start:%d-%m-%y}_{end:%d-%m-%y}" - diff = int((self.end - self.start).total_seconds()) - self.range = range(diff) - self.span = range(start := int(self.start.timestamp()), diff + start) - self.data = TestData(name=name, span=self.span, range=self.range) + span_start = int(self.start.timestamp()) + span_end = int(self.end.timestamp()) + self.range = range(0, span_end - span_start) + self.span = range(span_start, span_end) + self.data = BackTestData(name=self.name, span=self.span, range=self.range) self.mt5 = MetaTrader() - self.task_queue = TaskQueue(workers=250) + self.task_queue = TaskQueue(workers=500, mode='finite', on_exit='cancel') @classmethod - def pickle_data(cls, *, data: TestData, name: str | Path): + def pickle_data(cls, *, data: BackTestData, name: str | Path): """""" try: with open(name, 'wb') as fo: @@ -85,9 +84,8 @@ class GetData: except Exception as err: logger.error(f"Error in dump_data: {err}") - @classmethod - def load_data(cls, *, name: str | Path) -> TestData: + def load_data(cls, *, name: str | Path) -> BackTestData: """""" try: with open(name, 'rb') as fo: @@ -97,7 +95,8 @@ class GetData: logger.error(f"Error: {err}") def save_data(self, *, name: str | Path = ''): - name = name or self.name + name = name or (self.name + '.pkl' if not self.name.endswith('.pkl') else self.name) + name = Path(self.config.backtest_dir) / name if not isinstance(name, Path) else name with open(name, 'wb') as fo: pickle.dump(self.data, fo, protocol=pickle.HIGHEST_PROTOCOL) @@ -108,7 +107,6 @@ class GetData: q_items = [QueueItem(self.get_symbols_rates), QueueItem(self.get_symbols_ticks), - QueueItem(self.get_symbols_prices), QueueItem(self.get_symbols_info), ] @@ -125,21 +123,34 @@ class GetData: await self.task_queue.run() + if self.data.fully_loaded is False: + logger.warning("Data not fully loaded") + self.data = BackTestData(name=self.name, span=self.span, range=self.range, fully_loaded=False) + async def get_terminal_info(self): """""" terminal = await self.mt5.terminal_info() + if terminal is None: + self.data.fully_loaded = False + self.task_queue.stop_queue() terminal = terminal._asdict() self.data.set_attrs(terminal=terminal) async def get_version(self): """""" version = await self.mt5.version() + if version is None: + self.data.fully_loaded = False + self.task_queue.stop_queue() self.data.set_attrs(version=version) @backoff_decorator async def get_account_info(self): """""" res = await self.mt5.account_info() + if res is None: + self.data.fully_loaded = False + self.task_queue.stop_queue() res = res._asdict() self.data.set_attrs(account=res) @@ -153,11 +164,6 @@ class GetData: [self.task_queue.add(item=QueueItem(self.get_symbol_ticks, symbol=symbol)) for symbol in self.symbols if self.data.ticks.get(symbol) is None] - async def get_symbols_prices(self): - """""" - [self.task_queue.add(item=QueueItem(self.get_symbol_prices, symbol=symbol)) - for symbol in self.symbols if self.data.prices.get(symbol) is None] - async def get_symbols_rates(self): """""" [self.task_queue.add(item=QueueItem(self.get_symbol_rates, symbol=symbol, timeframe=timeframe), priority=4) @@ -168,22 +174,25 @@ class GetData: async def get_symbol_info(self, *, symbol: str): """""" res = await self.mt5.symbol_info(symbol) + if res is None: + self.data.fully_loaded = False + self.task_queue.stop_queue() self.data.symbols[symbol] = res._asdict() @backoff_decorator async def get_symbol_ticks(self, *, symbol: str): """""" res = await self.mt5.copy_ticks_range(symbol, self.start, self.end, MetaTrader5.COPY_TICKS_ALL) + if res is None: + self.data.fully_loaded = False + self.task_queue.stop_queue() self.data.ticks[symbol] = res - @backoff_decorator - async def get_symbol_prices(self, *, symbol: str): - """""" - res = await self.mt5.copy_ticks_range(symbol, self.start, self.end, MetaTrader5.COPY_TICKS_ALL) - self.data.prices[symbol] = res - @backoff_decorator async def get_symbol_rates(self, *, symbol: str, timeframe: TimeFrame): """""" res = await self.mt5.copy_rates_range(symbol, timeframe, self.start, self.end) - self.data.rates.setdefault(symbol, {})[timeframe] = res + if res is None: + self.data.fully_loaded = False + self.task_queue.stop_queue() + self.data.rates.setdefault(symbol, {})[int(timeframe)] = res diff --git a/src/aiomql/contrib/backtesting/trades_manager.py b/src/aiomql/contrib/backtesting/trades_manager.py index a9b210d..00b3a7d 100644 --- a/src/aiomql/contrib/backtesting/trades_manager.py +++ b/src/aiomql/contrib/backtesting/trades_manager.py @@ -78,7 +78,7 @@ class PositionsManager(TradeManager): return item.ticket in self._open_positions def __getitem__(self, item): - if item in self: + if item in self._open_positions: return super().__getitem__(item) raise KeyError('Position not found') @@ -90,6 +90,11 @@ class PositionsManager(TradeManager): self._open_positions.discard(key) del self._data[key] + @property + def margin(self): + """Returns the total margin of all open positions""" + return sum(self.margins.values()) + def close(self, *, ticket: int) -> bool: is_open = ticket in self._open_positions self._open_positions.discard(ticket) @@ -104,7 +109,7 @@ class PositionsManager(TradeManager): def set_margin(self, *, ticket: int, margin: float): self.margins[ticket] = margin - def positions_get(self, *, ticket: int = None, symbol: str = None, group: None) -> tuple[TradePosition, ...]: + def positions_get(self, *, ticket: int = None, symbol: str = None, group: None = None) -> tuple[TradePosition, ...]: if ticket: return tuple(position for position in self.open_positions if position.ticket == ticket) @@ -120,7 +125,7 @@ class PositionsManager(TradeManager): return tuple() def positions_total(self) -> int: - return len(self) + return len(self._open_positions) @property def open_positions(self) -> tuple[TradePosition, ...]: @@ -147,7 +152,7 @@ class OrdersManager(TradeManager): return tuple(order for order in self.values() if order.ticket == ticket) if position: - return tuple(order for order in self.values() if order.position == position) + return tuple(order for order in self.values() if order.position_id == position) return () @@ -174,7 +179,7 @@ class DealsManager(TradeManager): return tuple(deal for deal in self.values() if deal.ticket == ticket) if position and (ticket is None): - return tuple(deal for deal in self.values() if deal.position == position) + return tuple(deal for deal in self.values() if deal.position_id == position) return () diff --git a/src/aiomql/contrib/strategies/__init__.py b/src/aiomql/contrib/strategies/__init__.py index 625296d..08d0006 100644 --- a/src/aiomql/contrib/strategies/__init__.py +++ b/src/aiomql/contrib/strategies/__init__.py @@ -1,2 +1,2 @@ from .finger_trap import FingerTrap -from .tracker import Tracker +from .chaos import Chaos diff --git a/src/aiomql/contrib/strategies/chaos.py b/src/aiomql/contrib/strategies/chaos.py new file mode 100644 index 0000000..f1dbf31 --- /dev/null +++ b/src/aiomql/contrib/strategies/chaos.py @@ -0,0 +1,57 @@ +import random +from logging import getLogger + +from ...lib.strategy import Strategy +from ..symbols import ForexSymbol +from ..utils import Tracker +from ...core.constants import TimeFrame, OrderType +from ..traders.scalp_trader import ScalpTrader + +logger = getLogger(__name__) + + +class Chaos(Strategy): + """A chaotic strategy that buys and sells randomly.""" + ltf: TimeFrame + htf: TimeFrame + lcc: int + hcc: int + fast_ema: int + slow_ema: int + tracker: Tracker + parameters = {"fast_ema": 8, "slow_ema": 20, "ltf": TimeFrame.M1, "htf": TimeFrame.M2, "lcc": 100, "hcc": 100} + + def __init__(self, *, symbol: ForexSymbol, params: dict = None, sessions=None, name='Chaos'): + super().__init__(symbol=symbol, params=params, sessions=sessions, name=name) + self.tracker = Tracker(snooze=self.ltf.seconds) + self.trader = ScalpTrader(symbol=self.symbol) + + async def check_trend(self): + try: + candles = await self.symbol.copy_rates_from_pos(timeframe=self.htf, count=self.hcc) + if ((current := candles[-1]) and current.time < self.tracker.trend_time + and current.close == self.tracker.last_trend_price): + self.tracker.update(new=False, order_type=None, snooze=5) + return + self.tracker.update(new=True, trend_time=current.time, last_trend_price=current.close) + candles.ta.ema(length=self.slow_ema, append=True, fillna=0) + candles.ta.ema(length=self.fast_ema, append=True, fillna=0) + candles.rename(inplace=True, **{f"EMA_{self.fast_ema}": "fast", f"EMA_{self.slow_ema}": "slow"}) + order_type = random.choice([OrderType.BUY, OrderType.SELL]) + if order_type == OrderType.BUY: + self.tracker.update(trend="bullish", snooze=self.htf.seconds, order_type=OrderType.BUY) + else: + self.tracker.update(trend="bearish", snooze=self.htf.seconds, order_type=OrderType.SELL) + except Exception as err: + logger.error(f"{err}. Failed to check trend") + self.tracker.update(trend="ranging", snooze=self.ltf.seconds, order_type=None) + + async def trade(self): + print(f"Trading {self.symbol.name} with {self.__class__.__name__}") + await self.check_trend() + if self.tracker.order_type is not None: + await self.trader.place_trade(order_type=self.tracker.order_type, parameters=self.parameters) + self.tracker.update(order_type=None) + await self.sleep(secs=self.tracker.snooze) + else: + await self.sleep(secs=self.tracker.snooze) diff --git a/src/aiomql/contrib/strategies/finger_trap.py b/src/aiomql/contrib/strategies/finger_trap.py index dd92e2f..72eeaaa 100644 --- a/src/aiomql/contrib/strategies/finger_trap.py +++ b/src/aiomql/contrib/strategies/finger_trap.py @@ -1,4 +1,3 @@ -import asyncio import logging from ...lib.symbol import Symbol @@ -9,7 +8,7 @@ from ...core.constants import TimeFrame, OrderType from ...lib.sessions import Sessions from ..candle_patterns import find_bearish_fractal, find_bullish_fractal from ..traders import SimpleTrader -from .tracker import Tracker +from ..utils.tracker import Tracker logger = logging.getLogger(__name__) @@ -32,15 +31,15 @@ class FingerTrap(Strategy): name: str = 'FingerTrap'): super().__init__(symbol=symbol, params=params, sessions=sessions, name=name) self.trader = trader or SimpleTrader(symbol=self.symbol) - self.tracker: Tracker = Tracker(snooze=self.ttf.time) + self.tracker: Tracker = Tracker(snooze=self.ttf.seconds) async def check_trend(self): try: candles: Candles = await self.symbol.copy_rates_from_pos(timeframe=self.ttf, count=self.tcc) - if not ((current := candles[-1].time) >= self.tracker.trend_time): + if (current := candles[-1]) and current.time < self.tracker.trend_time: self.tracker.update(new=False, order_type=None) return - self.tracker.update(new=True, trend_time=current) + self.tracker.update(new=True, trend_time=current.time, last_trend_price=current.close) candles.ta.ema(length=self.slow_ema, append=True, fillna=0) candles.ta.ema(length=self.fast_ema, append=True, fillna=0) candles.rename(inplace=True, **{f"EMA_{self.fast_ema}": "fast", f"EMA_{self.slow_ema}": "slow"}) @@ -55,19 +54,19 @@ class FingerTrap(Strategy): elif fbs.iloc[-1] and cbf.iloc[-1] and current.is_bearish(): self.tracker.update(trend="bearish") else: - self.tracker.update(trend="ranging", snooze=self.ttf.time, order_type=None) + self.tracker.update(trend="ranging", snooze=self.ttf.seconds, order_type=None) self.tracker.update(trend="bullish") # remove this line except Exception as err: logger.error(f"{err} for {self.symbol} in {self.__class__.__name__}.check_trend") - self.tracker.update(snooze=self.ttf.time, order_type=None) + self.tracker.update(snooze=self.ttf.seconds, order_type=None) async def confirm_trend(self): try: candles = await self.symbol.copy_rates_from_pos(timeframe=self.etf, count=self.ecc) - if not ((current := candles[-1].time) >= self.tracker.entry_time): + if (current := candles[-1]) and current.time < self.tracker.entry_time: self.tracker.update(new=False, order_type=None) return - self.tracker.update(new=True, entry_time=current) + self.tracker.update(new=True, trend_time=current.time, last_entry_price=current.close) candles.ta.ema(length=self.entry_ema, append=True) candles.rename(**{f"EMA_{self.entry_ema}": "ema"}) candles['cae'] = candles.ta_lib.cross(candles.close, candles.ema) @@ -75,15 +74,15 @@ class FingerTrap(Strategy): current = candles[-1] if self.tracker.bullish and True or current.cae: # change True to current.cae sl = find_bullish_fractal(candles).low - self.tracker.update(snooze=self.ttf.time, order_type=OrderType.BUY, sl=sl) + self.tracker.update(snooze=self.ttf.seconds, order_type=OrderType.BUY, sl=sl) elif self.tracker.bearish and current.cbe: sl = find_bearish_fractal(candles).high - self.tracker.update(snooze=self.ttf.time, order_type=OrderType.SELL, sl=sl) + self.tracker.update(snooze=self.ttf.seconds, order_type=OrderType.SELL, sl=sl) else: - self.tracker.update(snooze=self.etf.time, order_type=None) + self.tracker.update(snooze=self.etf.seconds, order_type=None) except Exception as err: logger.error(f"{err} for {self.symbol} in {self.__class__.__name__}.confirm_trend") - self.tracker.update(snooze=self.etf.time, order_type=None) + self.tracker.update(snooze=self.etf.seconds, order_type=None) async def watch_market(self): await self.check_trend() @@ -92,22 +91,17 @@ class FingerTrap(Strategy): async def trade(self): logger.info(f"Trading {self.symbol}") - async with self.sessions as sess: - await self.sleep(self.ttf.time) - while True: - await sess.check() - try: - await self.watch_market() - if not self.tracker.new: - await asyncio.sleep(2) - continue - if self.tracker.order_type is None: - await self.sleep(self.tracker.snooze) - continue - 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: - logger.error(f"{err} For {self.symbol} in {self.__class__.__name__}.trade") - await self.sleep(self.ttf.time) + try: + await self.watch_market() + if self.tracker.new is False: + await self.sleep(secs=5) + return + if self.tracker.order_type is None: + await self.sleep(secs=self.tracker.snooze) + return + await self.trader.place_trade(order_type=self.tracker.order_type, parameters=self.parameters, + sl=self.tracker.sl) + await self.sleep(secs=self.tracker.snooze) + except Exception as err: + logger.error(f"{err} For {self.symbol} in {self.__class__.__name__}.trade") + await self.sleep(secs=self.ttf.seconds) diff --git a/src/aiomql/contrib/symbols/forex_symbol.py b/src/aiomql/contrib/symbols/forex_symbol.py index fdaf2bd..e23fa2f 100644 --- a/src/aiomql/contrib/symbols/forex_symbol.py +++ b/src/aiomql/contrib/symbols/forex_symbol.py @@ -1,5 +1,4 @@ from ...lib.symbol import Symbol -from ...core.exceptions import VolumeError class ForexSymbol(Symbol): @@ -16,7 +15,7 @@ class ForexSymbol(Symbol): """ return self.point * 10 - def compute_points(self, *, amount: float, volume) -> float: + def compute_points(self, *, amount: float, volume: float) -> float: """Compute the number of points required for a trade. Given the amount and the volume of the trade. Args: amount (float): Amount to trade @@ -25,69 +24,17 @@ class ForexSymbol(Symbol): points = amount / (volume * self.point * self.trade_contract_size) return points - async def compute_volume_points(self, *, amount: float, points: float, use_limits=False, round_down: bool = False, - adjust: float = False) -> tuple[float, float]: - """Compute the volume and points required for a trade. Given the amount and the number of points. + async def compute_volume_points(self, *, amount: float, points: float, round_down: bool = False) -> float: + """Compute the volume required for a trade. Given the amount and the number of points. Args: amount (float): Amount to trade points (float): Number of points round_down: round down the computed volume to the nearest step default True - adjust: Adjust the points if the computed volume is outside the range of permitted volumes - use_limits: Adjust the computed volume to the nearest permitted volume if the computed volume is outside """ - amount = await self.check_amount(amount) volume = amount / (self.point * points * self.trade_contract_size) - volume = self.round_off_volume(volume, round_down=round_down) - if (chk_vol := self.check_volume(volume))[0]: - if adjust: - points = self.compute_points(amount=amount, volume=volume) - return volume, points - if use_limits: - vol = chk_vol[1] - if adjust: - points = self.compute_points(amount=amount, volume=vol) - return vol, points - raise VolumeError(f"Incorrect Volume. Computed Volume outside the range of permitted volumes") + return self.round_off_volume(volume=volume, round_down=round_down) - async def compute_volume_sl(self, *, amount: float, price: float, sl: float, use_limits=False, adjust: bool = False, - round_down: bool = False) -> tuple[float, float]: - amount = await self.check_amount(amount) - volume = amount / ((price - sl) * self.trade_contract_size) - volume = self.round_off_volume(volume, round_down=round_down) - sign = volume / abs(volume) if volume else 1 - if (chk_vol := self.check_volume(abs(volume)))[0]: - if adjust: - sl = price - (amount / (volume * self.trade_contract_size)) - return abs(volume), sl - return abs(volume), sl - if use_limits: - vol = chk_vol[1] * sign - if adjust: - sl = price - (amount / (vol * self.trade_contract_size)) - return abs(vol), sl - raise VolumeError(f"Incorrect Volume. Computed Volume outside the range of permitted volumes") - - async def compute_volume(self, *, amount: float, points, use_limits=False, round_down=True) -> float: - """Compute volume given an amount to risk and target points. Round the computed volume to the nearest step. - - Args: - amount (float): Amount to risk. Given in terms of the account currency. - points (float): Target points. - use_limits (bool): If True, the computed volume checked against the maximum and minimum volume. - round_down: round down the computed volume to the nearest step default True - - Returns: - float: volume - - Raises: - VolumeError: If the computed volume is less than the minimum volume or greater than the maximum volume. - """ - amount = await self.check_amount(amount) - volume = amount / (self.point * points * self.trade_contract_size) - volume = self.round_off_volume(volume, round_down=round_down) - if self.check_volume(volume)[0]: - return volume - if use_limits: - return self.check_volume(volume)[1] - raise VolumeError(f"Incorrect Volume. Computed Volume outside the range of permitted volumes") + async def compute_volume_sl(self, *, amount: float, price: float, sl: float, round_down: bool = False) -> float: + volume = amount / (abs(price - sl) * self.trade_contract_size) + return self.round_off_volume(volume=volume, round_down=round_down) diff --git a/src/aiomql/contrib/traders/scalp_trader.py b/src/aiomql/contrib/traders/scalp_trader.py new file mode 100644 index 0000000..fcc0220 --- /dev/null +++ b/src/aiomql/contrib/traders/scalp_trader.py @@ -0,0 +1,29 @@ +from logging import getLogger + +from ...core.models import OrderType +from ...lib.trader import Trader + +logger = getLogger(__name__) + + +class ScalpTrader(Trader): + async def place_trade(self, *, order_type: OrderType, volume: float = None, parameters: dict = None): + """Places a trade based on the order_type and a given stop_loss + + Args: + order_type (OrderType): The order_type + volume (float): The volume to trade + parameters (dict): Parameters associated with the trade + """ + try: + self.parameters |= parameters or {} + volume = volume or self.symbol.volume_min + await self.create_order_no_stops(order_type=order_type, volume=volume) + if not await self.check_order(): + return + self.order.comment = self.parameters.get('name', self.__class__.__name__) + res = await self.send_order() + if res is not None: + await self.record_trade(result=res, parameters=self.parameters) + except Exception as err: + logger.error(f"{err} in {self.__class__.__name__}.place_trade for {self.symbol.name}") diff --git a/src/aiomql/contrib/traders/simple_trader.py b/src/aiomql/contrib/traders/simple_trader.py index 2691918..f7f13a4 100644 --- a/src/aiomql/contrib/traders/simple_trader.py +++ b/src/aiomql/contrib/traders/simple_trader.py @@ -1,47 +1,26 @@ from logging import getLogger -from ...lib.ram import RAM from ...core.models import OrderType from ...lib.trader import Trader -from ..symbols import ForexSymbol logger = getLogger(__name__) class SimpleTrader(Trader): - """A simple trader class""" - def __init__(self, *, symbol: ForexSymbol, ram: RAM = None): - """Initializes the order object and RAM instance + async def place_trade(self, *, order_type: OrderType, sl: float, parameters: dict = None): + """Places a trade based on the order_type and a given stop_loss Args: - symbol (Symbol): Financial instrument - ram (RAM): Risk Assessment and Management instance + order_type (OrderType): The order_type + sl (float): The stop_loss + parameters (dict): Parameters associated with the trade """ - ram = ram or RAM(risk_to_reward=2) - super().__init__(symbol=symbol, ram=ram) - - async def create_order(self, *, order_type: OrderType, sl: float): - amount = await self.ram.get_amount() - await self.symbol.info() - tick = await self.symbol.info_tick() - min_points = self.symbol.trade_stops_level + (self.symbol.spread * 1.5) - points = (tick.ask - sl) / self.symbol.point if order_type == OrderType.BUY else\ - (abs(tick.bid - sl) / self.symbol.point) - points = max(points, min_points) - self.order.type = order_type - volume, points = await self.symbol.compute_volume_points(amount=amount, points=points) - self.order.volume = volume - self.order.comment = self.parameters.get('name', self.__class__.__name__) - tick = await self.symbol.info_tick() - self.set_trade_stop_levels(points=points, tick=tick) - - async def place_trade(self, order_type: OrderType, sl: float, parameters: dict = None): - """Places a trade based on the order_type.""" try: self.parameters |= parameters or {} - await self.create_order(order_type=order_type, sl=sl) + await self.create_order_with_sl(order_type=order_type, sl=sl) if not await self.check_order(): return + self.order.comment = self.parameters.get('name', self.__class__.__name__) await self.send_order() except Exception as err: logger.error(f"{err} in {self.__class__.__name__}.place_trade for {self.symbol.name}") diff --git a/src/aiomql/contrib/utils/__init__.py b/src/aiomql/contrib/utils/__init__.py new file mode 100644 index 0000000..94b9b56 --- /dev/null +++ b/src/aiomql/contrib/utils/__init__.py @@ -0,0 +1 @@ +from .tracker import Tracker diff --git a/src/aiomql/contrib/strategies/tracker.py b/src/aiomql/contrib/utils/tracker.py similarity index 87% rename from src/aiomql/contrib/strategies/tracker.py rename to src/aiomql/contrib/utils/tracker.py index 5531a66..83ff312 100644 --- a/src/aiomql/contrib/strategies/tracker.py +++ b/src/aiomql/contrib/utils/tracker.py @@ -1,7 +1,7 @@ from dataclasses import dataclass from typing import Literal -from ...core.constants import OrderType +from aiomql.core.constants import OrderType @dataclass @@ -14,6 +14,8 @@ class Tracker: snooze: float = 0 trend_time: float = 0 entry_time: float = 0 + last_trend_price: float = 0 + last_entry_price: float = 0 new: bool = True order_type: OrderType = None sl: float = 0 @@ -34,4 +36,4 @@ class Tracker: self.bullish = True case "bearish": self.ranging = self.bullish = False - self.bearish = True \ No newline at end of file + self.bearish = True diff --git a/src/aiomql/core/__init__.py b/src/aiomql/core/__init__.py index 7ba1d44..a62c5af 100644 --- a/src/aiomql/core/__init__.py +++ b/src/aiomql/core/__init__.py @@ -1,5 +1,6 @@ from .meta_trader import MetaTrader from .config import Config +from .meta_backtester import MetaBackTester from .models import * from .constants import * from .base import Base, _Base diff --git a/src/aiomql/core/base.py b/src/aiomql/core/base.py index 0db9001..bb70a65 100644 --- a/src/aiomql/core/base.py +++ b/src/aiomql/core/base.py @@ -89,7 +89,7 @@ class Base: """ exclude, include = exclude or set(), include or set() filter_ = include or set(self.dict.keys()).difference(exclude) - return {key: value for key, value in self.dict.items() if key in filter_} + return {key: value for key, value in self.dict.items() if key in filter_ and value is not None} @property @cache @@ -115,7 +115,7 @@ class Base: try: _filter = self.exclude.difference(self.include) return {key: value for key, value in (self.class_vars | self.__dict__).items() if - key not in _filter} + key not in _filter and value is not None} except Exception as err: logger.warning(err) diff --git a/src/aiomql/core/config.py b/src/aiomql/core/config.py index a61afb6..cec0e92 100644 --- a/src/aiomql/core/config.py +++ b/src/aiomql/core/config.py @@ -60,10 +60,12 @@ class Config: _instance: Self mode: Literal['backtest', 'live'] use_terminal_for_backtesting: bool + shutdown: bool + force_shutdown: bool _defaults = {"timeout": 60000, "record_trades": True, "trade_record_mode": "csv", "mode": "live", - 'filename': "aiomql.json", "records_dir_name": "trade_records", "backtest_dir_name": "backtester", + 'filename': "aiomql.json", "records_dir_name": "trade_records", "backtest_dir_name": "backtesting", "use_terminal_for_backtesting": True, 'path': '', 'login': 0, 'password': '', - 'server': '', 'records_dir': None} + 'server': '', 'records_dir': None, 'shutdown': False, 'force_shutdown': False} def __new__(cls, *args, **kwargs): if not hasattr(cls, "_instance"): @@ -166,7 +168,7 @@ class Config: self.records_dir = self.root / self.records_dir_name self.records_dir.mkdir(parents=True, exist_ok=True) - if self.mode == "backtest" and not hasattr(self, "backtest_dir"): + if not hasattr(self, "backtest_dir"): self.backtest_dir = self.root / self.backtest_dir_name self.backtest_dir.mkdir(parents=True, exist_ok=True) diff --git a/src/aiomql/core/constants.py b/src/aiomql/core/constants.py index 00c7e80..9280e86 100644 --- a/src/aiomql/core/constants.py +++ b/src/aiomql/core/constants.py @@ -169,10 +169,6 @@ class TimeFrame(Repr, IntEnum): get: get a timeframe object from a time value in seconds """ __enum_name__ = "TIMEFRAME" - - def __str__(self): - return self.name - M1 = mt5.TIMEFRAME_M1 M2 = mt5.TIMEFRAME_M2 M3 = mt5.TIMEFRAME_M3 @@ -195,7 +191,7 @@ class TimeFrame(Repr, IntEnum): MN1 = mt5.TIMEFRAME_MN1 @property - def time(self): + def seconds(self): """The number of seconds in a TIMEFRAME Returns: @@ -203,20 +199,21 @@ class TimeFrame(Repr, IntEnum): Examples: >>> t = TimeFrame.H1 - >>> print(t.time) + >>> print(t.seconds) 3600 """ - times = {1: 60, 2: 120, 3: 180, 4: 240, 5: 300, 6: 360, 10: 600, 15: 900, 20: 1200, 30: 1800, 16385: 3600, + seconds = {1: 60, 2: 120, 3: 180, 4: 240, 5: 300, 6: 360, 10: 600, 15: 900, 20: 1200, 30: 1800, 16385: 3600, 16386: 7200, 16387: 10800, 16388: 14400, 16390: 21600, 16392: 28800, 16396: 43200, 16408: 86400, 32769: 604800, 49153: 2592000} - return times[self] + return seconds[self] @classmethod - def get(cls, time: int) -> 'TimeFrame': - times = {60: 1, 120: 2, 180: 3, 240: 4, 300: 5, 360: 6, 600: 10, 900: 15, 1200: 20, 1800: 30, 3600: 16385, + def get_timeframe(cls, time: int) -> 'TimeFrame': + """Get a timeframe object from a time value in seconds""" + time_frames = {60: 1, 120: 2, 180: 3, 240: 4, 300: 5, 360: 6, 600: 10, 900: 15, 1200: 20, 1800: 30, 3600: 16385, 7200: 16386, 10800: 16387, 14400: 16388, 21600: 16390, 28800: 16392, 43200: 16396, 86400: 16408, 604800: 32769, 2592000: 49153} - return TimeFrame(times[time]) + return TimeFrame(time_frames[time]) @classmethod @property diff --git a/src/aiomql/core/event_manager.py b/src/aiomql/core/event_manager.py index a3c9b53..fbb9a52 100644 --- a/src/aiomql/core/event_manager.py +++ b/src/aiomql/core/event_manager.py @@ -25,13 +25,17 @@ class EventManager: def __init__(self, *, num_tasks: int = 0): self.num_main_tasks = num_tasks or self.num_main_tasks + @property + def backtest_engine(self): + return self.config.backtest_engine + def add_tasks(self, *tasks: Task): self.tasks.extend(tasks) def sigint_handler(self, sig, frame): for task in self.tasks: task.cancel() - self.config.backtest_engine.wrap_up() + self.backtest_engine.wrap_up() async def acquire(self): await self.condition.acquire() @@ -40,16 +44,21 @@ class EventManager: self.condition.notify_all() async def event_monitor(self): + self.backtest_engine.next() while True: async with self.condition: if self.task_tracker == self.num_main_tasks: # all main tasks have been completed in the current cycle self.task_tracker = 0 - await self.config.backtest_engine.tracker() - self.config.backtest_engine.next() + await self.backtest_engine.tracker() + self.backtest_engine.next() self.condition.notify_all() - if ((timestamp := self.config.backtest_engine.cursor.time) % int(60 * 60 * 24)) == 0: - print(f"if Time: {datetime.fromtimestamp(timestamp)}") - await asyncio.sleep(0) + # print(f"Time: {self.backtest_engine.cursor.time}") + if self.backtest_engine.cursor.time % 3600 == 0: + print(datetime.strftime(self.backtest_engine.cursor.time, "%Y-%m-%d %H:%M:%S")) + if self.backtest_engine.stop_testing: + break + await asyncio.sleep(0) + self.backtest_engine.wrap_up() async def wait(self): self.task_tracker += 1 diff --git a/src/aiomql/core/meta_backtester.py b/src/aiomql/core/meta_backtester.py index e2b13b5..8a11d84 100644 --- a/src/aiomql/core/meta_backtester.py +++ b/src/aiomql/core/meta_backtester.py @@ -28,7 +28,7 @@ class MetaBackTester(MetaTrader): @backtest_engine.setter def backtest_engine(self, value: BackTestEngine): - if BackTestEngine is not None: + if value is not None: self.config.backtest_engine = value async def last_error(self) -> tuple[int, str]: diff --git a/src/aiomql/core/meta_trader.py b/src/aiomql/core/meta_trader.py index 32d6b3d..1972727 100644 --- a/src/aiomql/core/meta_trader.py +++ b/src/aiomql/core/meta_trader.py @@ -2,6 +2,7 @@ import asyncio from datetime import datetime from logging import getLogger from typing import Literal +from pathlib import Path import numpy as np from MetaTrader5 import (BookInfo, SymbolInfo, AccountInfo, Tick, TerminalInfo, TradeOrder, TradeDeal, @@ -101,6 +102,7 @@ class MetaTrader(MetaCore): """ async with asyncio.Lock() as _: path = self.config.path if path is None else path + path = "" if Path(path).exists() is False else path args = (str(path),) if path else () acc = self.config.account_info() kwargs = {key: value for key, value in (('login', login or acc.get('login')), @@ -118,14 +120,9 @@ class MetaTrader(MetaCore): return res async def shutdown(self) -> None: + """Closes the connection to the MetaTrader terminal. """ - Closes the connection to the MetaTrader terminal. - - Returns: - None: None - """ - res = await asyncio.to_thread(self._shutdown) - return res + self._shutdown() async def last_error(self) -> tuple[int, str]: try: diff --git a/src/aiomql/core/task_queue.py b/src/aiomql/core/task_queue.py index bd02f85..b24bc4c 100644 --- a/src/aiomql/core/task_queue.py +++ b/src/aiomql/core/task_queue.py @@ -1,6 +1,6 @@ import asyncio +import time from typing import Coroutine, Callable, Literal -from signal import signal, SIGINT from logging import getLogger logger = getLogger(__name__) @@ -14,7 +14,7 @@ class QueueItem: self.args = args self.kwargs = kwargs self.must_complete = False - self.time = asyncio.get_event_loop().time() + self.time = int(time.monotonic_ns()) def __hash__(self): return id(self) @@ -29,15 +29,15 @@ class QueueItem: else: self.task_item(*self.args, **self.kwargs) - + except Exception as err: logger.error(f"Error {err} occurred in {self.task_item.__name__} with args {self.args} and kwargs {self.kwargs}") class TaskQueue: - def __init__(self, size: int = 0, workers: int = 50, timeout: int = None, queue: asyncio.Queue = None, - on_exit: Literal['cancel', 'complete_priority'] = 'complete_priority'): - + def __init__(self, size: int = 0, workers: int = 10, timeout: int = None, queue: asyncio.Queue = None, + on_exit: Literal['cancel', 'complete_priority'] = 'complete_priority', + mode: Literal['finite', 'infinite'] = 'infinite', worker_timeout: int = 60): self.queue = queue or asyncio.PriorityQueue(maxsize=size) self.workers = workers self.tasks = [] @@ -45,15 +45,18 @@ class TaskQueue: self.timeout = timeout self.stop = False self.on_exit = on_exit + self.mode = mode + self.worker_timeout = worker_timeout def add(self, *, item: QueueItem, priority=3, must_complete=False): try: - if not self.stop: - item.must_complete = must_complete - if isinstance(self.queue, asyncio.PriorityQueue): - item = (priority, item) - self.priority_tasks.add(item) if item.must_complete else ... - self.queue.put_nowait(item) + if self.stop: + return + item.must_complete = must_complete + self.priority_tasks.add(item) if item.must_complete else ... + if isinstance(self.queue, asyncio.PriorityQueue): + item = (priority, item) + self.queue.put_nowait(item) except asyncio.QueueFull: logger.error("Queue is full") @@ -66,59 +69,86 @@ class TaskQueue: else: item = self.queue.get_nowait() - if not self.stop or item.must_complete: + if self.stop is False or item.must_complete: await item.run() self.queue.task_done() self.priority_tasks.discard(item) - if self.stop and len(self.priority_tasks) == 0: - logger.info('All priority tasks completed') + + if self.stop and (self.on_exit == 'cancel' or len(self.priority_tasks) == 0): self.cancel() break + + except asyncio.QueueEmpty: + if self.stop: + break + + if self.mode == 'finite': + break + + sleep = QueueItem(asyncio.sleep, 1) + self.add(item=sleep) + await asyncio.sleep(self.worker_timeout) + except Exception as err: - logger.error(f"Error {err} occurred in worker") - - def sigint_handle(self, sig, frame): - logger.info('SIGINT received, cleaning up...') - - if self.on_exit == 'complete_priority' and self.priority_tasks: - logger.info(f'Completing {len(self.priority_tasks)} priority tasks...') - self.stop = True - else: - self.cancel() - - self.on_exit = 'cancel' # force cancel on exit if SIGINT is received again + logger.error("%s: Error occurred in worker", err) async def run(self, timeout: int = 0): - signal(SIGINT, self.sigint_handle) - loop = asyncio.get_running_loop() - start = loop.time() - + start = time.perf_counter() try: self.tasks.extend(asyncio.create_task(self.worker()) for _ in range(self.workers)) - task = asyncio.create_task(self.queue.join()) - self.tasks.append(task) - await asyncio.wait_for(task, timeout = timeout or self.timeout) + timeout = timeout or self.timeout + queue_task = asyncio.create_task(self.queue.join()) + + if timeout: + main_task = asyncio.create_task(asyncio.wait_for(queue_task, timeout=timeout)) + else: + main_task = queue_task + self.tasks.append(main_task) + await main_task except TimeoutError: - logger.warning(f"Timed out after {loop.time() - start} seconds. {self.queue.qsize()} tasks remaining") + logger.warning("Timed out after %d seconds, %d tasks remaining", + time.perf_counter() - start, self.queue.qsize()) + self.stop = True - if self.on_exit == 'complete_priority' and self.priority_tasks: - logger.info(f'Completing {len(self.priority_tasks)} priority tasks...') - self.stop = True - await self.queue.join() - else: - self.cancel() + except asyncio.CancelledError as _: + logger.warning("Main task cancelled") - except asyncio.CancelledError: - logger.debug('All tasks cancelled') + except Exception as err: + logger.warning("%s: An error occurred in %s.run", err, self.__class__.__name__) finally: - logger.info(f'Exiting queue after {(loop.time() - start)} seconds.' - f'{self.queue.qsize()} tasks remaining, {len(self.priority_tasks)} are priority tasks') + await self.clean_up() + def stop_queue(self): + self.stop = True + self.on_exit = 'cancel' self.cancel() + async def clean_up(self): + try: + if self.on_exit == 'complete_priority' and len(self.priority_tasks) > 0: + logger.warning(f'Completing {len(self.priority_tasks)} priority tasks...') + queue_task = asyncio.create_task(self.queue.join()) + self.tasks.append(queue_task) + await queue_task + self.cancel() + + except asyncio.CancelledError as _: + ... + + except Exception as err: + logger.error(f"%s: Error occurred in %s.clean_up", err, self.__class__.__name__) + + finally: + self.cancel() + def cancel(self): - [task.cancel() for task in self.tasks if not task.done()] + for task in self.tasks: + try: + if not task.done(): + task.cancel() + except asyncio.CancelledError as _: + ... self.tasks.clear() diff --git a/src/aiomql/lib/__init__.py b/src/aiomql/lib/__init__.py index 10aa7d3..cbf24f7 100644 --- a/src/aiomql/lib/__init__.py +++ b/src/aiomql/lib/__init__.py @@ -1,6 +1,6 @@ from .account import Account from .backtest_runner import BackTestRunner -from .bot_factory import Bot +from .bot import Bot from .candle import Candle, Candles from .executor import Executor from .history import History diff --git a/src/aiomql/lib/backtest_runner.py b/src/aiomql/lib/backtest_runner.py index b0fbecf..538f4c8 100644 --- a/src/aiomql/lib/backtest_runner.py +++ b/src/aiomql/lib/backtest_runner.py @@ -3,6 +3,7 @@ import signal from logging import getLogger from ..core.event_manager import EventManager +from ..core.task_queue import TaskQueue from ..contrib.backtesting.backtest_engine import BackTestEngine from .strategy import Strategy from ..core.meta_backtester import MetaBackTester @@ -11,21 +12,24 @@ logger = getLogger(__name__) class BackTestRunner: - def __init__(self, *, strategies: list[Strategy] = None, backtest_engine: BackTestEngine = None): + def __init__(self, *, strategies: list[Strategy], backtest_engine: BackTestEngine): + self.backtest_engine = backtest_engine self.strategies = strategies or [] self.event_manager = EventManager() - self.mt5 = MetaBackTester(backtest_engine=backtest_engine) + self.mt5 = MetaBackTester() signal.signal(signal.SIGINT, self.event_manager.sigint_handler) + self.mt5.config.task_queue.worker_timeout = 5 async def run(self): try: 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.initialize()] self.event_manager.num_main_tasks = len(strategies) tasks = [*[asyncio.create_task(strategy.run_strategy()) for strategy in strategies], asyncio.create_task(self.event_manager.event_monitor())] self.event_manager.add_tasks(*tasks) - await asyncio.gather(*tasks, return_exceptions=True) if strategies else ... + tasks.append(asyncio.create_task(self.mt5.config.task_queue.run())) + await asyncio.gather(*tasks, return_exceptions=True) if len(strategies) else ... except Exception as err: logger.error(f"Error {err} occurred in StrategyTester") diff --git a/src/aiomql/lib/bot_factory.py b/src/aiomql/lib/bot.py similarity index 93% rename from src/aiomql/lib/bot_factory.py rename to src/aiomql/lib/bot.py index b409218..a7da8c1 100644 --- a/src/aiomql/lib/bot_factory.py +++ b/src/aiomql/lib/bot.py @@ -38,7 +38,7 @@ class Bot: funcs (dict): A dictionary of functions to run with their respective keyword arguments as a dictionary num_workers (int): Number of workers to run the functions """ - num_workers = num_workers or len(funcs) * 2 + num_workers = num_workers or len(funcs) with ProcessPoolExecutor(max_workers=num_workers) as executor: for bot, kwargs in funcs.items(): executor.submit(bot, **kwargs) @@ -58,7 +58,8 @@ class Bot: raise SystemExit logger.info("Login Successful") await self.init_strategies() - self.add_coroutine(coroutine=self.config.task_queue.run) + self.add_coroutine(coroutine=self.config.task_queue.run, on_separate_thread=True) + self.add_coroutine(coroutine=self.executor.exit) except Exception as err: logger.error(f"{err}. Bot initialization failed") raise SystemExit @@ -72,17 +73,18 @@ class Bot: """ self.executor.add_function(function=function, kwargs=kwargs) - def add_coroutine(self, *, coroutine: Callable[..., ...] | Coroutine, **kwargs): + def add_coroutine(self, *, coroutine: Callable[..., ...] | Coroutine, on_separate_thread=False, **kwargs): """Add a coroutine to the executor. Args: coroutine (Coroutine): A coroutine to be executed + on_separate_thread (bool): Run the coroutine **kwargs (dict): keyword arguments for the coroutine Returns: """ - self.executor.add_coroutine(coroutine=coroutine, kwargs=kwargs) + self.executor.add_coroutine(coroutine=coroutine, kwargs=kwargs, on_separate_thread=on_separate_thread) def execute(self): """Execute the bot.""" @@ -130,7 +132,7 @@ class Bot: @staticmethod async def init_strategy(*, strategy: Strategy) -> tuple[bool, Strategy]: """Initialize a single strategy. This method is called internally by the bot.""" - res = await strategy.symbol.init() + res = await strategy.symbol.initialize() return res, strategy async def init_strategies(self): diff --git a/src/aiomql/lib/candle.py b/src/aiomql/lib/candle.py index 20a95c0..26fc8ff 100644 --- a/src/aiomql/lib/candle.py +++ b/src/aiomql/lib/candle.py @@ -227,7 +227,7 @@ class Candles: @property def timeframe(self): tf = self.time[1] - self.time[0] - return TimeFrame.get(abs(tf)) + return TimeFrame.get_timeframe(abs(tf)) @property def columns(self) -> DataFrame: diff --git a/src/aiomql/lib/executor.py b/src/aiomql/lib/executor.py index 2656bdf..c6de372 100644 --- a/src/aiomql/lib/executor.py +++ b/src/aiomql/lib/executor.py @@ -1,8 +1,10 @@ import asyncio from concurrent.futures import ThreadPoolExecutor from typing import Coroutine, Callable +import os from logging import getLogger +from ..core.config import Config from .strategy import Strategy logger = getLogger(__name__) @@ -14,23 +16,31 @@ class Executor: Attributes: executor (ThreadPoolExecutor): The executor object. strategy_runners (list): List of strategies. - coroutines (dict[Coroutine, dict]): A dictionary of coroutines to run in the executor + coroutines (list[Coroutine]): A list of coroutines 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 + executor: ThreadPoolExecutor + tasks: list[asyncio.Task] + config: Config def __init__(self): - self.executor = ThreadPoolExecutor self.strategy_runners: list[Strategy] = [] - self.coroutines: dict[Coroutine | Callable: dict] = {} + self.coroutines: list[Coroutine] = [] + self.coroutine_threads: list[Coroutine] = [] self.functions: dict[Callable: dict] = {} + self.tasks = [] + self.no_of_running_strategies = 0 + self.config = Config() + self.timeout = None # Timeout for the executor. For testing purposes only - def add_function(self, *, function: Callable, kwargs: dict): + def add_function(self, *, function: Callable, kwargs: dict = None): + kwargs = kwargs or {} self.functions[function] = kwargs - def add_coroutine(self, *, coroutine: Coroutine, kwargs: dict): - self.coroutines[coroutine] = kwargs + def add_coroutine(self, *, coroutine: Callable | Coroutine, kwargs: dict = None, on_separate_thread=False): + kwargs = kwargs or {} + coroutine = coroutine(**kwargs) + self.coroutines.append(coroutine) if on_separate_thread is False else self.coroutine_threads.append(coroutine) def add_strategies(self, *, strategies: tuple[Strategy]): """Add multiple strategies at once @@ -48,26 +58,69 @@ class Executor: """ self.strategy_runners.append(strategy) + async def create_strategy_task(self, strategy: Strategy): + task = asyncio.create_task(strategy.run_strategy()) + self.tasks.append(task) + self.no_of_running_strategies += 1 + await task + def run_strategy(self, strategy: Strategy): """Wraps the coroutine trade method of each strategy with 'asyncio.run'. Args: strategy (Strategy): A strategy object """ - self.loop.run_until_complete(strategy.run_strategy()) + asyncio.run(self.create_strategy_task(strategy)) - def run_coroutine(self, func, kwargs: dict): - """ - Run a coroutine function + async def create_coroutine_task(self, coroutine: Coroutine): + task = asyncio.create_task(coroutine) + self.tasks.append(task) + await task + + async def create_coroutines_task(self): + """""" + task = asyncio.create_task(asyncio.gather(*self.coroutines, return_exceptions=True)) + self.tasks.append(task) + await task + + def run_coroutine_tasks(self): + """Run all coroutines in the executor""" + asyncio.run(self.create_coroutines_task()) + + def run_coroutine_task(self, coroutine): + asyncio.run(self.create_coroutine_task(coroutine)) + + @staticmethod + def run_function(function: Callable, kwargs: dict): + """Run a function Args: - func: The coroutine. A variadic function. + function: The function to run kwargs: A dictionary of keyword arguments for the function """ try: - self.loop.run_until_complete(func(**kwargs)) + function(**kwargs) except Exception as err: - logger.error(f'Error: {err}. Unable to run function') + logger.error(f'Error: {err}. Unable to run function: {function.__name__}') + + async def exit(self): + """Shutdown the executor""" + try: + while self.config.shutdown is False and self.config.force_shutdown is False: + if self.timeout is not None and self.no_of_running_strategies == len(self.strategy_runners): + self.timeout -= 1 + if self.timeout == 0: + self.config.shutdown = True + continue + for strategy in self.strategy_runners: + strategy.running = False + for task in self.tasks: + task.cancel() + self.executor.shutdown(wait=False, cancel_futures=True) + if self.config.force_shutdown: + os._exit(1) + except Exception as err: + logger.error(f"Error: {err}. Unable to shutdown executor") async def execute(self, *, workers: int = 5): """Run the strategies with a threadpool executor. @@ -80,9 +133,9 @@ class Executor: """ workers_ = sum([len(self.strategy_runners), len(self.functions), len(self.coroutines)]) workers = max(workers, workers_) - self.loop = asyncio.get_running_loop() - with self.executor(max_workers=workers) as executor: - [self.loop.run_in_executor(executor, self.run_strategy, worker) for worker in self.strategy_runners] - [self.loop.run_in_executor(executor, self.run_coroutine, coroutine, kwargs) for coroutine, - kwargs in self.coroutines.items()] - [self.loop.run_in_executor(executor, function, kwargs) for function, kwargs in self.functions.items()] + with ThreadPoolExecutor(max_workers=workers) as executor: + self.executor = executor + self.executor.submit(self.run_coroutine_tasks) + [self.executor.submit(self.run_coroutine_task, coroutine) for coroutine in self.coroutine_threads] + [self.executor.submit(self.run_strategy, strategy) for strategy in self.strategy_runners] + [self.executor.submit(function, **kwargs) for function, kwargs in self.functions.items()] diff --git a/src/aiomql/lib/history.py b/src/aiomql/lib/history.py index 99a0824..f73fbfe 100644 --- a/src/aiomql/lib/history.py +++ b/src/aiomql/lib/history.py @@ -53,7 +53,7 @@ class History: self.total_deals: int = 0 self.total_orders: int = 0 - async def init(self): + async def initialize(self): """Get history deals and orders""" deals, orders = await asyncio.gather(self.get_deals(), self.get_orders(), return_exceptions=True) self.deals = deals if isinstance(deals, tuple) else () diff --git a/src/aiomql/lib/order.py b/src/aiomql/lib/order.py index e9a044b..53c22ff 100644 --- a/src/aiomql/lib/order.py +++ b/src/aiomql/lib/order.py @@ -77,7 +77,7 @@ class Order(_Base, TradeRequest): Raises: OrderError: If not successful """ - req = self.dict | kwargs + req = self.request | kwargs res = await self.mt5.order_check(req) if res is None: raise OrderError(f'Order check failed for {self.symbol}') @@ -93,24 +93,22 @@ class Order(_Base, TradeRequest): Raises: OrderError: If not successful """ - res = await self.mt5.order_send(self.dict) + res = await self.mt5.order_send(self.request) if res is None: raise OrderError(f'Failed to send order {self.symbol}') return OrderSendResult(**res._asdict()) + @error_handler(log_error_msg=False) async def calc_margin(self) -> float | None: """Return the required margin in the account currency to perform a specified trading operation. Returns: float: Returns float value if successful - - Raises: - OrderError: If not successful """ res = await self.mt5.order_calc_margin(self.type, self.symbol, self.volume, self.price) return res - @error_handler(response=0) + @error_handler(response=0, log_error_msg=False) async def calc_profit(self) -> float: """Return profit in the account currency for a specified trading operation. @@ -121,3 +119,20 @@ class Order(_Base, TradeRequest): action, symbol, volume, price_open, price_close = self.type, self.symbol, self.volume, self.price, self.tp res = await self.mt5.order_calc_profit(action, symbol, volume, price_open, price_close) return res + + @error_handler(response=0, log_error_msg=False) + async def calc_loss(self) -> float: + """Return profit in the account currency for a specified trading operation. + + Returns: + float: Returns float value if successful + None: If not successful + """ + action, symbol, volume, price_open, price_close = self.type, self.symbol, self.volume, self.price, self.sl + res = await self.mt5.order_calc_profit(action, symbol, volume, price_open, price_close) + return res + + @property + def request(self) -> dict: + """Return the order request as a dictionary.""" + return {key: value for key, value in self.dict.items() if key in self.mt5.TradeRequest.__match_args__} diff --git a/src/aiomql/lib/positions.py b/src/aiomql/lib/positions.py index 49c92f7..9e42450 100644 --- a/src/aiomql/lib/positions.py +++ b/src/aiomql/lib/positions.py @@ -40,7 +40,7 @@ class Positions: if positions is not None: self.positions = tuple(TradePosition(**pos._asdict()) for pos in positions) return self.positions - logger.warning('Failed to get open positions for') + logger.warning('Failed to get open positions') return () async def get_position_by_ticket(self, *, ticket: int) -> TradePosition | None: @@ -83,9 +83,11 @@ class Positions: type=order_type.opposite) return await order.send() - @staticmethod - async def close_by(*, position: TradePosition) -> OrderSendResult: - """Close an open position for the trading account.""" + async def close_position_by_ticket(self, *, ticket: int) -> OrderSendResult | None: + """Close an open position using the ticket.""" + position = await self.get_position_by_ticket(ticket=ticket) + if position is None: + return None order = Order(position=position.ticket, symbol=position.symbol, volume=position.volume, type=position.type.opposite, price=position.price_current, action=TradeAction.DEAL) return await order.send() diff --git a/src/aiomql/lib/ram.py b/src/aiomql/lib/ram.py index ffdbba7..fe9eaf7 100644 --- a/src/aiomql/lib/ram.py +++ b/src/aiomql/lib/ram.py @@ -7,6 +7,7 @@ class RAM: account: Account risk_to_reward: float risk: float + fixed_amount: float | None min_amount: float max_amount: float loss_limit: int @@ -16,18 +17,23 @@ class RAM: """Initialize Risk Assessment and Management with the provided keyword arguments. Keyword Args: - risk_to_reward (float): Risk to reward ratio. Defaults to 1 - risk (float): Percentage of capital to risk per trade 0.01 # 1% - kwargs: extra keyword arguments are set as object attributes + risk_to_reward (float): Risk to reward ratio. Defaults to 2 + risk (float): Percentage of capital to risk per trade. Defaults to 1% + min_amount (float): Minimum amount to risk per trade. Defaults to 0 + max_amount (float): Maximum amount to risk per trade. Defaults to 0 + loss_limit (int): Maximum number of losing positions. Defaults to 1 + open_limit (int): Maximum number of open positions. Defaults to 1 + fixed_amount (float): Fixed amount to risk per trade. Defaults to None """ self.account = Account() self.positions = Positions() - self.risk_to_reward = kwargs.get('risk_to_reward', 1) - self.risk = kwargs.get('risk', 0.01) + self.risk_to_reward = kwargs.get('risk_to_reward', 2) + self.risk = kwargs.get('risk', 1) 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) + self.loss_limit = kwargs.get('loss_limit', 3) + self.open_limit = kwargs.get('open_limit', 3) + self.fixed_amount = kwargs.get('fixed_amount', None) async def get_amount(self) -> float: """Calculate the amount to risk per trade as a percentage of margin_free. @@ -35,14 +41,16 @@ class RAM: Returns: float: Amount to risk per trade """ + if self.fixed_amount: + return self.fixed_amount await self.account.refresh() - amount = self.account.margin_free * self.risk + amount = self.account.margin_free * (self.risk/100) if self.min_amount and self.max_amount: return max(self.min_amount, min(self.max_amount, amount)) return amount async def check_losing_positions(self) -> bool: - """Check if the number of losing positions is greater than the loss limit + """Check if the number of losing positions is less than the loss limit Returns: bool: True if the number of losing positions is less than or equal the loss limit @@ -52,7 +60,7 @@ class RAM: return len(loosing) <= self.loss_limit 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 less than or equal the loss limit. Returns: bool: True if the number of open positions is less than the open limit diff --git a/src/aiomql/lib/sessions.py b/src/aiomql/lib/sessions.py index 42ee2d1..eb9e013 100644 --- a/src/aiomql/lib/sessions.py +++ b/src/aiomql/lib/sessions.py @@ -1,10 +1,8 @@ import asyncio -from datetime import time, timedelta, datetime +from datetime import time, timedelta, datetime, UTC from typing import Literal, Callable, Iterable, NamedTuple from logging import getLogger -import pytz - from ..core.models import OrderSendResult, TradePosition from ..core.config import Config from ..core.event_manager import EventManager @@ -29,7 +27,7 @@ def delta(obj: time) -> timedelta: async def backtest_sleep(secs): - """A custom function to call when the session starts.""" + """An async sleep function for use during backtesting.""" em = EventManager() config = Config() sleep = config.backtest_engine.cursor.time + secs @@ -67,8 +65,8 @@ class Session: custom_end (Callable): A custom function to call when the session ends. Default is None. name (str): A name for the session. Default is a combination of start and end. """ - self.start = start if isinstance(start, time) else time(hour=start, tzinfo=pytz.UTC) - self.end = end if isinstance(end, time) else time(hour=end, tzinfo=pytz.UTC) + self.start = start.replace(tzinfo=UTC) if isinstance(start, time) else time(hour=start, tzinfo=UTC) + self.end = end if isinstance(end, time) else time(hour=end, tzinfo=UTC) self.on_start = on_start self.on_end = on_end self.custom_start = custom_start @@ -93,8 +91,8 @@ class Session: def in_session(self) -> bool: """Check if the current time is within the session.""" - now = datetime.now(tz=pytz.UTC).time() if self.config.mode == 'live'\ - else datetime.fromtimestamp(self.config.backtest_engine.cursor.time, tz=pytz.UTC).time() + now = datetime.now(tz=UTC).time() if self.config.mode == 'live'\ + else datetime.fromtimestamp(self.config.backtest_engine.cursor.time, tz=UTC).time() return now in self async def begin(self): @@ -169,10 +167,10 @@ class Session: def until(self): """Get the seconds until the session starts from the current time in seconds.""" 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=UTC).time() secs = (delta(self.start) - delta(now)).seconds else: - secs = (delta(self.start) - delta(datetime.now(tz=pytz.UTC).time())).seconds + secs = (delta(self.start) - delta(datetime.now(tz=UTC).time())).seconds return secs @@ -207,8 +205,8 @@ class Sessions: Returns: 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()) + moment = moment or datetime.now(tz=UTC).time() if self.config.mode == 'live' else ( + datetime.fromtimestamp(self.config.backtest_engine.cursor.time, tz=UTC).time()) for session in self.sessions: if moment in session: return session @@ -223,8 +221,8 @@ class Sessions: Returns: 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()) + moment = moment or datetime.now(tz=UTC).time() if self.config.mode == 'live' else ( + datetime.fromtimestamp(self.config.backtest_engine.cursor.time, tz=UTC).time()) for session in self.sessions: if delta(moment) < delta(session.start): return session @@ -246,9 +244,9 @@ class Sessions: return 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=UTC).time() else: - now = datetime.now(tz=pytz.UTC).time() + now = datetime.now(tz=UTC).time() next_session = self.find(moment=now) diff --git a/src/aiomql/lib/strategy.py b/src/aiomql/lib/strategy.py index 0d854ce..8d079dd 100644 --- a/src/aiomql/lib/strategy.py +++ b/src/aiomql/lib/strategy.py @@ -48,7 +48,7 @@ class Strategy(ABC): symbol (Symbol): The Financial instrument params (Dict): Trading strategy parameters """ - self.parameters = self.parameters | (params or {}) + self.parameters = {**self.parameters} | (params or {}) self.symbol = symbol self.name = name or self.__class__.__name__ self.parameters["symbol"] = symbol.name @@ -56,7 +56,7 @@ class Strategy(ABC): 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.mt5 = MetaTrader() if self.config.mode != 'backtest' else MetaBackTester() self.event_manager = EventManager() def __repr__(self): @@ -74,6 +74,7 @@ class Strategy(ABC): async def __aenter__(self): await self.sessions.check() + self.running = True self.current_session = self.sessions.current_session async def __aexit__(self, exc_type, exc_val, exc_tb): @@ -94,7 +95,7 @@ class Strategy(ABC): """ mod = time() % secs secs = secs - mod if mod != 0 else mod - await asyncio.sleep(secs + 0.2) + await asyncio.sleep(secs + 0.1) async def sleep(self, *, secs: float): """Sleep for the needed amount of seconds in between requests to the terminal. @@ -104,26 +105,35 @@ class Strategy(ABC): 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': + if self.config.mode == 'backtest': await self.backtest_sleep(secs=secs) + else: + await self.live_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: + # print(f"Sleeping for {secs} seconds") + try: + _time = self.config.backtest_engine.cursor.time + mod = _time % secs + # print('mod', mod) + secs = secs - mod if mod != 0 else mod + # print('secs', secs) + if self.event_manager.num_main_tasks == 1: + self.config.backtest_engine.fast_forward(steps=int(secs)) await self.event_manager.wait() - else: - await self.event_manager.wait() + + elif self.event_manager.num_main_tasks > 1: + _time = self.config.backtest_engine.cursor.time + secs + # print(_time, self.config.backtest_engine.cursor.time) + while _time > self.config.backtest_engine.cursor.time: + print(f"Time in sleep {self.symbol}: {self.config.backtest_engine.cursor.time}") + await self.event_manager.wait() + + # await self.event_manager.wait() + else: + await self.event_manager.wait() + except Exception as err: + logger.error(f"Error: {err} in backtest_sleep") async def run_strategy(self): """Run the strategy.""" @@ -134,19 +144,22 @@ class Strategy(ABC): async def live_strategy(self): """Run the strategy.""" - while self.running: - async with self as _: + async with self as _: + while self.running: 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() + try: + async with self as _: + while self.running: + async with self.event_manager.condition: + await self.sessions.check() + await self.event_manager.wait() + await self.test() + except Exception as err: + logger.error(f"Error: {err} in backtest_strategy") @abstractmethod async def trade(self): diff --git a/src/aiomql/lib/symbol.py b/src/aiomql/lib/symbol.py index 82f488e..2530ac5 100644 --- a/src/aiomql/lib/symbol.py +++ b/src/aiomql/lib/symbol.py @@ -41,24 +41,22 @@ class Symbol(_Base, SymbolInfo): self.account = Account() @backoff_decorator - async def info_tick(self, *, name: str = "") -> Tick: + async def info_tick(self, *, name: str = "") -> Tick | None: """Get the current price tick of a financial instrument. Args: - name: if name is supplied get price tick of that financial instrument + name: if name is supplied get price tick of that financial instrument. Optional unnamed parameter. Returns: Tick: Return a Tick Object - - Raises: - ValueError: If request was unsuccessful and None was returned + None: If request was unsuccessful """ tick = await self.mt5.symbol_info_tick(name or self.name) if tick is not None: tick = Tick(**tick._asdict()) setattr(self, 'tick', tick) if not name else ... return tick - raise ValueError(f'Could not get tick for {name or self.name}.') + return None async def symbol_select(self, *, enable: bool = True) -> bool: """Select a symbol in the MarketWatch window or remove a symbol from the window. @@ -75,25 +73,23 @@ class Symbol(_Base, SymbolInfo): return self.select @backoff_decorator - async def info(self) -> SymbolInfo: + async def info(self) -> SymbolInfo | None: """Get data on the specified financial instrument and update the symbol object properties Returns: (SymbolInfo): SymbolInfo if successful - - Raises: - ValueError: If request was unsuccessful and None was returned + (None): If request was unsuccessful """ info = await self.mt5.symbol_info(self.name) - if info: + if info is not None: info = info._asdict() info['swap_rollover3days'] = info.get('swap_rollover3days', 0) % 7 self.set_attributes(**info) return SymbolInfo(**info) - raise ValueError(f'Could not get info for {self.name}') + return None - async def init(self) -> bool: - """Initialized the symbol by pulling properties from the terminal + async def initialize(self) -> bool: + """Initialize the symbol by pulling properties from the terminal Returns: bool: Returns True if symbol info was successful initialized @@ -101,12 +97,12 @@ class Symbol(_Base, SymbolInfo): try: res = await asyncio.gather(self.symbol_select(), self.info(), self.info_tick(), self.book_add(), return_exceptions=True) - if all(res): + if any(res): return True - logger.warning(f'Unable to initialized {self}') + logger.warning('Unable to initialize %s', self.name) return False except Exception as err: - logger.warning(f'{err}: Unable to initialized {self}') + logger.warning('%s: Unable to initialize %s', err, self.name) return False async def book_add(self) -> bool: @@ -116,7 +112,10 @@ class Symbol(_Base, SymbolInfo): Returns: bool: True if successful, otherwise – False. """ - return await self.mt5.market_book_add(self.name) + res = await self.mt5.market_book_add(self.name) + if res is False: + logger.debug("Could not add %s to market book", self.name) + return res @backoff_decorator async def book_get(self) -> tuple[BookInfo, ...]: @@ -171,50 +170,41 @@ class Symbol(_Base, SymbolInfo): """ return round_off(value=volume, step=self.volume_step, round_down=round_down) - async def check_amount(self, *, amount: float) -> float: + async def amount_in_quote_currency(self, *, amount: float) -> float: + """Convert the amount to the quote currency of the symbol.""" 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, from_currency=self.account.currency, + to_currency=self.currency_profit) return amount - async def compute_volume(self, *args, **kwargs) -> float: + async def compute_volume(self) -> float: """Computes the volume required for a trade usually based on the amount and any other keyword arguments. This is a dummy method that returns the minimum volume of the symbol. It is meant to be overridden by a subclass that implements the computation of volume. - Keyword Args: - use_limits (bool): round up or round down the computed volume to the nearest volume limit i.e. volume_min - or volume_max - Returns: float: Returns the volume of the trade """ return self.volume_min - async def convert_currency(self, *, amount: float, base: str, quote: str) -> float: - """Convert from one currency to the other. Alias for currency_conversion""" - return await self.currency_conversion(amount=amount, base=base, quote=quote) - - async def currency_conversion(self, *, amount: float, base: str, quote: str) -> float: - """Convert from one currency to the other. - + async def convert_currency(self, *, amount: float, from_currency: str, to_currency: str) -> float: + """Convert a given amount from one currency to the other. Args: - amount: amount to convert given in terms of the quote currency - base: The base currency of the pair - quote: The quote currency of the pair - - Returns: - float: Amount in terms of the quote currency + amount: Amount to convert + from_currency: Currency to convert from + to_currency: Currency to convert to """ + base, quote = to_currency, from_currency try: - pair = f'{base}{quote}' - tick = await self.info_tick(name=pair) - if tick is not None: - return amount / tick.ask - pair = f'{quote}{base}' tick = await self.info_tick(name=pair) if tick is not None: - return amount * tick.bid + return round(amount * tick.bid, 2) + + pair = f'{base}{quote}' + tick = await self.info_tick(name=pair) + if tick is not None: + return round(amount / tick.ask, 2) except Exception as err: logger.warning(f'{err}: Currency conversion failed: Unable to convert {amount} in {quote} to {base}') diff --git a/src/aiomql/lib/trader.py b/src/aiomql/lib/trader.py index a8b06bc..982e47f 100644 --- a/src/aiomql/lib/trader.py +++ b/src/aiomql/lib/trader.py @@ -1,9 +1,7 @@ -"""Trader class module. Handles the creation of an order and the placing of trades""" from abc import ABC, abstractmethod -from datetime import datetime +from datetime import datetime, UTC from typing import TypeVar from logging import getLogger -import pytz from ..core.models import OrderType, OrderSendResult, OrderCheckResult @@ -12,8 +10,8 @@ from ..core.task_queue import QueueItem from .result import Result from .order import Order from .symbol import Symbol as _Symbol -from .ticks import Tick from .ram import RAM +from .._utils import error_handler logger = getLogger(__name__) Symbol = TypeVar("Symbol", bound=_Symbol) @@ -23,16 +21,13 @@ class Trader(ABC): """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: symbol (Symbol): The financial instrument. ram (RAM): RAM instance order (Order): Trade order parameters (dict): Parameters of the trading strategy used to place the trade - - Class Attributes: - config (Config): Config instance. + config (Config): The Config instance. """ config: Config ram: RAM @@ -47,47 +42,120 @@ class Trader(ABC): """ self.config = Config() self.symbol = symbol - self.order = Order(symbol=symbol.name) self.ram = ram or RAM() + self.order = Order(symbol=symbol.name) self.parameters = {} - def set_order_limits(self, *, pips: float, tick: Tick): - """Sets the stop loss and take profit for the order. This method uses pips as defined for forex instruments. + def set_trade_stop_levels_pips(self, *, pips: float): + """Sets the stop loss and take profit for the order. + This method uses pips as defined for forex instruments. It is assumed that order_type and price are already + set before calling this method. Args: - pips: Target pips - tick: Tick object + pips (float): Target pips """ pips = pips * self.symbol.pip sl, tp = pips, pips * self.ram.risk_to_reward + price = self.order.price if self.order.type == OrderType.BUY: - self.order.sl, self.order.tp = round(tick.ask - sl, self.symbol.digits), round(tick.ask + tp, + self.order.sl, self.order.tp = round(price - sl, self.symbol.digits), round(price + tp, self.symbol.digits) - self.order.price = tick.ask elif self.order.type == OrderType.SELL: - self.order.sl, self.order.tp = round(tick.bid + sl, self.symbol.digits), round(tick.bid - tp, + self.order.sl, self.order.tp = round(price + sl, self.symbol.digits), round(price - tp, self.symbol.digits) - self.order.price = tick.bid - else: - raise ValueError(f"Invalid order type: {self.order.type}") - def set_trade_stop_levels(self, *, points: float, tick: Tick): - """Set the stop loss and take profit levels of the order based on the points and price tick. + def set_trade_stop_levels_points(self, *, points: float, risk_to_reward: float = None): + """Set the stop loss and take profit levels of the order based on the points and the risk to reward ratio. + It is assumed that order_type and price are already set before calling this method. Args: - points: Target points - tick: Tick object + points (float): Target points + risk_to_reward (float): Risk to reward ratio """ points = points * self.symbol.point - sl, tp = points, points * self.ram.risk_to_reward + sl, tp = points, points * (risk_to_reward or self.ram.risk_to_reward) + price, digits = self.order.price, self.symbol.digits + if self.order.type == OrderType.BUY: - self.order.sl, self.order.tp = round(tick.ask - sl, self.symbol.digits), round(tick.ask + tp, - self.symbol.digits) - self.order.price = tick.ask - else: - self.order.sl, self.order.tp = round(tick.bid + sl, self.symbol.digits), round(tick.bid - tp, - self.symbol.digits) - self.order.price = tick.bid + self.order.sl, self.order.tp = round(price - sl, self.symbol.digits), round(price + tp, digits) + + elif self.order.type == OrderType.SELL: + self.order.sl, self.order.tp = round(price + sl, self.symbol.digits), round(price - tp, digits) + + async def create_order_with_stops(self, *, order_type: OrderType, sl: float, tp: float, + amount_to_risk: float = None): + """Create an order with stop loss and take profit levels. Use the amount to risk per trade to + calculate the volume. + + Args: + order_type (OrderType): Order type + sl (float): Stop loss in price + tp (float): Take profit in price + amount_to_risk (float): Amount to risk per trade in terms of the account currency. Optional parameter, + default is the amount as computed by the RAM instance. + """ + amount = amount_to_risk or await self.ram.get_amount() + amount = await self.symbol.amount_in_quote_currency(amount=amount) + tick = await self.symbol.info_tick() + price = tick.ask if order_type == OrderType.BUY else tick.bid + volume = await self.symbol.compute_volume_sl(amount=amount, price=price, sl=sl) + self.order.set_attributes(sl=sl, tp=tp, volume=volume, price=price, type=order_type) + + async def create_order_with_sl(self, *, order_type: OrderType, sl: float, amount_to_risk: float = None, + risk_to_reward: float = None): + """ + Create an order with a given stop_loss level. Use the amount to risk per trade to calculate the volume. + + Args: + order_type (OrderType): Order type + sl (float): Stop loss in price + amount_to_risk (float): Amount to risk per trade in terms of the account currency. Optional parameter, + default is the amount as computed by the RAM instance. + risk_to_reward (float): Risk to reward ratio. Optional parameter, default is the risk to reward ratio as + defined in the RAM instance. + """ + amount = amount_to_risk or await self.ram.get_amount() + amount = await self.symbol.amount_in_quote_currency(amount=amount) + tick = await self.symbol.info_tick() + price = tick.ask if order_type == OrderType.BUY else tick.bid + dsl = abs(price - sl) + dtp = dsl * (risk_to_reward or self.ram.risk_to_reward) + tp = price + dtp if order_type == OrderType.BUY else price - dtp + volume = await self.symbol.compute_volume_sl(amount=amount, price=price, sl=sl) + self.order.set_attributes(sl=sl, tp=tp, volume=volume, price=price, type=order_type) + + async def create_order_with_points(self, *, order_type: OrderType, points: float, + amount_to_risk: float = None, risk_to_reward: float = None): + """Create an order with specific points to risk. Use the amount to risk per trade to calculate the volume. + + Args: + order_type (OrderType): Order type + points (float): Points to risk + amount_to_risk (float): Amount to risk per trade in terms of the account currency. Optional parameter, + default is the amount as computed by the RAM instance. + risk_to_reward (float): Risk to reward ratio. Optional parameter, default is the risk to reward ratio as + defined in the RAM instance. + """ + self.order.type = order_type + amount = amount_to_risk or await self.ram.get_amount() + amount = await self.symbol.amount_in_quote_currency(amount=amount) + tick = await self.symbol.info_tick() + self.order.price = tick.ask if order_type == OrderType.BUY else tick.bid + volume = await self.symbol.compute_volume_points(amount=amount, points=points) + self.order.volume = volume + self.set_trade_stop_levels_points(points=points, risk_to_reward=risk_to_reward) + + async def create_order_no_stops(self, *, order_type: OrderType, volume: float = None): + """Create an order without setting stop loss and take profit. Using minimum lot size. + + Args: + order_type (OrderType): Order type + volume (float): Volume to trade with. Optional parameter, default is the minimum lot size. + """ + tick = await self.symbol.info_tick() + self.order.volume = volume or self.symbol.volume_min + self.order.price = tick.ask if order_type == OrderType.BUY else tick.bid + self.order.type = order_type async def check_order(self) -> OrderCheckResult | None: """Check order before sending it to the broker. @@ -121,23 +189,24 @@ class Trader(ABC): logger.info("Order placed successfully") return result + @error_handler async def record_trade(self, *, result: OrderSendResult, parameters: dict = None, name: str = ''): """Record the trade in csv or json. Args: result (OrderSendResult): Result of the order send - parameters: parameters of the trading strategy used to place the trade - name: Name of the trading strategy + parameters (dict): parameters of the trading strategy used to place the trade + name (str): Name of the trading strategy """ - if result.retcode != 10009 or not self.config.record_trades: + if self.config.record_trades is False or result.retcode != 10009: return - params = parameters or {} - profit = result.profit or await self.order.calc_profit() + params = {**parameters} or {} + profit = await self.order.calc_profit() params["expected_profit"] = profit - date = datetime.now(tz=pytz.UTC) - params["date"] = str(date.date()) - params["time"] = str(date.time()) - res = Result(result=result, parameters=params, name=name) - self.config.task_queue.add(item=QueueItem(res.save), must_complete=True) + date = datetime.now(tz=UTC) if self.config.mode == 'live' else datetime.fromtimestamp(self.config.backtest_engine.cursor.time, tz=UTC) + params["date"] = date.strftime("%Y-%m-%d %H:%M:%S.%f") + if self.config.record_trades: + res = Result(result=result, parameters=params, name=name) + self.config.task_queue.add(item=QueueItem(res.save), must_complete=True) @abstractmethod async def place_trade(self, *args, **kwargs): diff --git a/tests/actions/__init__.py b/tests/actions/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/backtest/__init__.py b/tests/backtest/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/backtest/conftest.py b/tests/backtest/conftest.py new file mode 100644 index 0000000..7310a71 --- /dev/null +++ b/tests/backtest/conftest.py @@ -0,0 +1,96 @@ +from datetime import datetime, UTC +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_backtester import MetaBackTester +from aiomql.contrib import BackTestEngine +from aiomql.lib import Positions, History, Order + +logger = getLogger(__name__) + + +async def cleanup(): + try: + shutil.rmtree(Path('tests/backtest/configs'), ignore_errors=True) + Path.unlink(Path('tests/backtest/test.json'), missing_ok=True) + shutil.rmtree(Path('tests/backtest/trade_records'), ignore_errors=True) + shutil.rmtree(Path('tests/backtest/backtesting'), ignore_errors=True) + await close_all_positions() + await MetaBackTester().shutdown() + except Exception as err: + logger.error(f"Failed to complete cleanup: {err}") + + +async def close_all_positions(): + try: + mt = MetaBackTester() + 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='package', autouse=True) +async def config(request): + Path('tests/backtest/configs').mkdir(exist_ok=True) + with open('aiomql.json', 'r') as fh, open('tests/backtest/configs/test2.json', 'w') as fh1, open('tests/backtest/test.json', 'w') as fh2: + data = json.load(fh) + data['mode'] = 'backtest' + json.dump(data, fh1, indent=2) + json.dump(data, fh2, indent=2) + config = Config(filename='test.json', root='tests/backtest') + yield config + await cleanup() + + +@pytest.fixture(scope='package', autouse=True) +async def mt(): + mt = MetaBackTester() + await mt.initialize() + await mt.login() + yield mt + await mt.shutdown() + +@pytest.fixture(scope='package') +async def period(): + return {'start': datetime(2024, 2, 1, hour=8, tzinfo=UTC), + 'end': datetime(2024, 2, 7, hour=16, tzinfo=UTC)} + +@pytest.fixture(scope='package') +async def backtest_engine(period): + start = period['start'] + end = period['end'] + return BackTestEngine(start=start, end=end, name='backtest_data') + + +@pytest.fixture(scope='function') +def order_sell(sell_order): + return Order(**sell_order) + + +@pytest.fixture(scope='function') +def order_buy(buy_order): + return Order(**buy_order) + + +@pytest.fixture(scope='package') +def positions(): + return Positions() + + +@pytest.fixture(scope='package') +def history(period): + start = period['start'] + end = period['end'] + return History(date_from=start, date_to=end) diff --git a/tests/backtest/integration/test_backtesting.py b/tests/backtest/integration/test_backtesting.py new file mode 100644 index 0000000..7f0479e --- /dev/null +++ b/tests/backtest/integration/test_backtesting.py @@ -0,0 +1,140 @@ +from aiomql.contrib import BackTestEngine, ForexSymbol, GetData +from aiomql.core import MetaBackTester +from aiomql.lib import Order + + +async def make_buy_sell_orders(): + sym = ForexSymbol(name='BTCUSD') + sym_info = await sym.mt5.symbol_info(sym.name) + dsl = (sym_info.trade_stops_level + sym_info.spread) * 2 * sym_info.point + sl = sym_info.ask - dsl + tp = sym_info.ask + dsl + buy_req = {'action': sym.mt5.TRADE_ACTION_DEAL, 'symbol': sym.name, 'volume': sym_info.volume_min, + 'type': sym.mt5.ORDER_TYPE_BUY, 'price': sym_info.ask, 'sl': sl, 'tp': tp} + + sell_req = buy_req.copy() + sell_req['type'] = sym.mt5.ORDER_TYPE_SELL + sell_req['price'] = sym_info.bid + del sell_req['tp'] + del sell_req['sl'] + return {'buy': Order(**buy_req), 'sell': Order(**sell_req)} + + +def test_trade_mode(config, backtest_engine, history, positions, order_sell, order_buy, btc_usd): + assert config.mode == 'backtest' + assert isinstance(backtest_engine, BackTestEngine) + assert isinstance(history.mt5, MetaBackTester) + assert isinstance(positions.mt5, MetaBackTester) + assert isinstance(order_sell.mt5, MetaBackTester) + assert isinstance(order_buy.mt5, MetaBackTester) + assert isinstance(btc_usd.mt5, MetaBackTester) + + +async def test_order_send(backtest_engine, order_sell, order_buy): + await backtest_engine.setup_account(balance=100) + so = await backtest_engine.order_send(request=order_sell.request) + bo = await backtest_engine.order_send(request=order_buy.request) + assert so.retcode == 10009 + assert bo.retcode == 10009 + backtest_engine.reset(clear_data=True) + + +async def test_positions(backtest_engine, positions, order_sell, order_buy): + await backtest_engine.setup_account(balance=100) + so = await backtest_engine.order_send(request=order_sell.request) + bo = await backtest_engine.order_send(request=order_buy.request) + all_positions = await positions.get_positions() + assert len(all_positions) == 2 + await positions.close_position_by_ticket(ticket=so.order) + all_positions = await positions.get_positions() + assert len(all_positions) == 1 + await positions.close_position_by_ticket(ticket=bo.order) + all_positions = await positions.get_positions() + assert len(all_positions) == 0 + backtest_engine.reset(clear_data=True) + + +async def test_history(backtest_engine, history, order_sell, order_buy, positions): + await backtest_engine.setup_account(balance=100) + so = await backtest_engine.order_send(request=order_sell.request) + await backtest_engine.order_send(request=order_buy.request) + await history.initialize() + assert len(history.orders) == 2 + assert len(history.deals) == 2 + await positions.close_position_by_ticket(ticket=so.order) + deals = await history.get_deals() + assert len(deals) == 3 + backtest_engine.reset(clear_data=True) + + +async def test_margin(backtest_engine, order_sell, order_buy): + await backtest_engine.setup_account(balance=100) + so_margin = await backtest_engine.order_calc_margin(action=order_sell.action, volume=order_sell.volume, + symbol=order_sell.symbol, price=order_sell.price) + bo_margin = await backtest_engine.order_calc_margin(action=order_buy.action, volume=order_buy.volume, + symbol=order_buy.symbol, price=order_buy.price) + total_margin = so_margin + bo_margin + await backtest_engine.order_send(request=order_sell.request) + await backtest_engine.order_send(request=order_buy.request) + # noinspection PyTestUnpassedFixture + assert backtest_engine.positions.margin == total_margin == backtest_engine._account.margin + backtest_engine.reset(clear_data=True) + + +async def test_account(backtest_engine, positions): + await backtest_engine.setup_account(balance=100) + backtest_engine.fast_forward(steps=100) + balance = backtest_engine._account.balance + equity = backtest_engine._account.equity + orders = await make_buy_sell_orders() + buy_order = orders['buy'] + sell_order = orders['sell'] + so = await backtest_engine.order_send(request=sell_order.request) + bo = await backtest_engine.order_send(request=buy_order.request) + backtest_engine.fast_forward(steps=22000) + all_pos = await positions.get_positions() + for _ in range(1000): + backtest_engine.fast_forward(steps=1) + await backtest_engine.tracker() + all_pos = await positions.get_positions() + if len(all_pos) == 1: + break + + deal = backtest_engine.deals.history_deals_get(position=bo.order) + bo_profit = deal[-1].profit + assert len(all_pos) == 1 + assert backtest_engine.positions.margin == backtest_engine._account.margin == backtest_engine.positions.margins[so.order] + profit = sum([pos.profit for pos in all_pos]) + n_balance = backtest_engine._account.balance + n_equity = backtest_engine._account.equity + assert backtest_engine._account.profit == profit + assert n_balance == balance + bo_profit + assert n_equity == equity + bo_profit + profit + so_pos = await positions.get_position_by_ticket(ticket=so.order) + gain = so_pos.profit + await positions.close_position(position=so_pos) + assert backtest_engine._account.balance == n_balance + gain + backtest_engine.reset(clear_data=True) + + +async def test_wrapup(positions, buy_order, sell_order, backtest_engine, config): + await backtest_engine.setup_account(balance=100) + backtest_engine.fast_forward(steps=500) + bo = await backtest_engine.order_send(request=buy_order) + await backtest_engine.order_send(request=sell_order) + backtest_engine.fast_forward(steps=5000) + await backtest_engine.tracker() + await positions.close_position_by_ticket(ticket=bo.order) + last_balance = backtest_engine._account.balance + last_equity = backtest_engine._account.equity + last_profit = backtest_engine._account.profit + backtest_engine.wrap_up() + tdata = GetData.load_data(name=config.backtest_dir / f'{backtest_engine.name}.pkl') + new_bte = BackTestEngine(data=tdata, restart=False) + assert new_bte._account.balance == last_balance + assert new_bte._account.equity == last_equity + assert new_bte._account.profit == last_profit + assert new_bte.span == backtest_engine.span + assert new_bte.range == backtest_engine.range + assert new_bte.name == backtest_engine.name + assert new_bte.cursor.time == backtest_engine.cursor.time diff --git a/tests/backtest/unit/test_deals_manager.py b/tests/backtest/unit/test_deals_manager.py new file mode 100644 index 0000000..c828e54 --- /dev/null +++ b/tests/backtest/unit/test_deals_manager.py @@ -0,0 +1,24 @@ +# noinspection PyTestUnpassedFixture +async def test_deals_manager(backtest_engine, sell_order, buy_order, period, positions): + backtest_engine.reset(clear_data=True) + await backtest_engine.setup_account(balance=100) + backtest_engine.fast_forward(steps=100) + await backtest_engine.order_send(request=sell_order) + bo = await backtest_engine.order_send(request=buy_order) + start = period['start'] + end = period['end'] + all_deals = backtest_engine.deals.get_deals_range(date_from=start, date_to=end) + assert len(all_deals) == 2 + backtest_engine.fast_forward(steps=10_000) + start2 = backtest_engine.cursor.time + bo2 = await backtest_engine.order_send(request=buy_order) + backtest_engine.fast_forward(steps=50) + end2 = backtest_engine.cursor.time + deals = backtest_engine.deals.history_deals_get(date_from=start2, date_to=end2) + assert len(deals) == 1 + assert deals[0].order == bo2.order + await positions.close_position_by_ticket(ticket=bo.order) + deals = backtest_engine.deals.history_deals_get(position=bo.order) + assert len(deals) <= 2 + orders = backtest_engine.deals.get_deals_range(date_from=start, date_to=end) + assert len(orders) == backtest_engine.deals.history_deals_total(date_from=start, date_to=end) == len(backtest_engine.deals._data.keys()) diff --git a/tests/backtest/unit/test_order_manager.py b/tests/backtest/unit/test_order_manager.py new file mode 100644 index 0000000..cd5c408 --- /dev/null +++ b/tests/backtest/unit/test_order_manager.py @@ -0,0 +1,24 @@ +# noinspection PyTestUnpassedFixture +async def test_orders_manager(backtest_engine, sell_order, buy_order, period, positions): + backtest_engine.reset(clear_data=True) + await backtest_engine.setup_account(balance=100) + backtest_engine.fast_forward(steps=100) + await backtest_engine.order_send(request=sell_order) + bo = await backtest_engine.order_send(request=buy_order) + start = period['start'] + end = period['end'] + all_orders = backtest_engine.orders.get_orders_range(date_from=start, date_to=end) + assert len(all_orders) == 2 + backtest_engine.fast_forward(steps=10_000) + start2 = backtest_engine.cursor.time + bo2 = await backtest_engine.order_send(request=buy_order) + backtest_engine.fast_forward(steps=50) + end2 = backtest_engine.cursor.time + orders = backtest_engine.orders.history_orders_get(date_from=start2, date_to=end2) + assert len(orders) == 1 + assert orders[0].ticket == bo2.order + await positions.close_position_by_ticket(ticket=bo.order) + orders = backtest_engine.orders.history_orders_get(position=bo.order) + assert len(orders) <= 2 + orders = backtest_engine.orders.get_orders_range(date_from=start, date_to=end) + assert len(orders) == backtest_engine.orders.history_orders_total(date_from=start, date_to=end) == len(backtest_engine.orders._data.keys()) diff --git a/tests/backtest/unit/test_positions_manager.py b/tests/backtest/unit/test_positions_manager.py new file mode 100644 index 0000000..4891e38 --- /dev/null +++ b/tests/backtest/unit/test_positions_manager.py @@ -0,0 +1,17 @@ +# noinspection PyTestUnpassedFixture +async def test_positions_manager(backtest_engine, sell_order, buy_order): + backtest_engine.reset(clear_data=True) + await backtest_engine.setup_account(balance=100) + backtest_engine.fast_forward(steps=100) + so = await backtest_engine.order_send(request=sell_order) + bo = await backtest_engine.order_send(request=buy_order) + all_pos = backtest_engine.positions.positions_get() + assert len(all_pos) == 2 + so_positions = backtest_engine.positions.positions_get(ticket=so.order) + so_position = so_positions[0] + assert so_position.ticket == so.order + btc_positions = backtest_engine.positions.positions_get(symbol='BTCUSD') + assert len(btc_positions) == 2 + assert backtest_engine.positions.positions_total() == 2 + backtest_engine.positions.close(ticket=bo.order) + assert backtest_engine.positions.positions_total() == 1 diff --git a/tests/conftest.py b/tests/conftest.py index 23f5e0b..5a6026c 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,93 +1,43 @@ -import asyncio -import json -import shutil -from logging import getLogger -from pathlib import Path - +from aiomql.lib.symbol import Symbol 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} +def btc_usd(): + return Symbol(name='BTCUSD') @pytest.fixture(scope='function') -async def buy_order(mt): - sym = 'BTCUSD' - sym_info = await mt.symbol_info(sym) +async def buy_order(btc_usd): + sym = btc_usd + sym_info = await sym.mt5.symbol_info(sym.name) 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} + return {'action': sym.mt5.TRADE_ACTION_DEAL, 'symbol': sym.name, 'volume': sym_info.volume_min, + 'type': sym.mt5.ORDER_TYPE_BUY, 'price': sym_info.ask, 'sl': sl, 'tp': tp} + + +@pytest.fixture(scope='function') +async def sell_order(btc_usd): + sym = btc_usd + sym_info = await sym.mt5.symbol_info(sym.name) + return {'action': sym.mt5.TRADE_ACTION_DEAL, 'symbol': sym.name, 'volume': sym_info.volume_min, + 'type': sym.mt5.ORDER_TYPE_SELL, 'price': sym_info.bid} + @pytest.fixture(scope='class') -async def make_buy_sell_orders(mt): - sym = 'BTCUSD' - sym_info = await mt.symbol_info(sym) +async def make_buy_sell_orders(): + sym = Symbol(name='BTCUSD') + sym_info = await sym.mt5.symbol_info(sym.name) 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 = {'action': sym.mt5.TRADE_ACTION_DEAL, 'symbol': sym.name, 'volume': sym_info.volume_min, + 'type': sym.mt5.ORDER_TYPE_BUY, 'price': sym_info.ask, 'sl': sl, 'tp': tp} + await sym.mt5.order_send(req) + req['type'] = sym.mt5.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) + await sym.mt5.order_send(req) diff --git a/tests/live/__init__.py b/tests/live/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/live/conftest.py b/tests/live/conftest.py new file mode 100644 index 0000000..5fbc35e --- /dev/null +++ b/tests/live/conftest.py @@ -0,0 +1,59 @@ +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/live/configs'), ignore_errors=True) + Path.unlink(Path('tests/live/test.json'), missing_ok=True) + shutil.rmtree(Path('tests/live/trade_records'), ignore_errors=True) + shutil.rmtree(Path('tests/live/backtesting'), 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='package', autouse=True) +async def config(request): + Path('tests/live/configs').mkdir(exist_ok=True) + with open('aiomql.json', 'r') as fh, open('tests/live/configs/test2.json', 'w') as fh1, open('tests/live/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/live') + yield config + await cleanup() + + +@pytest.fixture(scope='package', autouse=True) +async def mt(): + mt = MetaTrader() + await mt.initialize() + await mt.login() + yield mt + await mt.shutdown() diff --git a/tests/live/integration/test_backtesting.py b/tests/live/integration/test_backtesting.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/integration/test_results_records.py b/tests/live/integration/test_results_records.py similarity index 100% rename from tests/integration/test_results_records.py rename to tests/live/integration/test_results_records.py diff --git a/tests/live/unit/__init__.py b/tests/live/unit/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/unit/test_account.py b/tests/live/unit/test_account.py similarity index 100% rename from tests/unit/test_account.py rename to tests/live/unit/test_account.py diff --git a/tests/live/unit/test_backtest_engine.py b/tests/live/unit/test_backtest_engine.py new file mode 100644 index 0000000..2c18980 --- /dev/null +++ b/tests/live/unit/test_backtest_engine.py @@ -0,0 +1,269 @@ +from datetime import datetime, UTC + +from aiomql import TimeFrame +from aiomql.contrib.backtesting import BackTestEngine +from aiomql.contrib.backtesting.get_data import GetData +from aiomql._utils import round_down +from aiomql.core.constants import OrderType, TradeAction + +import pytest + + +class TestBackTestEngine: + @classmethod + def setup_class(cls): + cls.start = datetime(2024, 2, 1) + cls.end = datetime(2024, 2, 7) + cls.g_data = GetData(start=cls.start, end=cls.end, symbols=['BTCUSD', 'SOLUSD'], + timeframes=[TimeFrame.H1, TimeFrame.H2], name='test_engine') + cls.bte = BackTestEngine(start=cls.start, end=cls.end) + + @pytest.fixture(scope='class') + async def bte2(self): + await self.g_data.get_data() + bte2 = BackTestEngine(start=self.start, end=self.end, data=self.g_data.data, use_terminal=False) + await bte2.setup_account(balance=100) + return bte2 + + @pytest.fixture(scope='class') + async def sell_order(self): + sym = await self.bte.get_symbol_info(symbol='BTCUSD') + request = {'type': OrderType.SELL, 'symbol': 'BTCUSD', 'volume': sym.volume_min, + 'price': sym.bid, 'action': TradeAction.DEAL} + return request + + @pytest.fixture(scope='class') + async def buy_order(self): + sym = await self.bte.get_symbol_info(symbol='BTCUSD') + dsl = (sym.trade_stops_level + sym.spread) * 2 * sym.point + sl = sym.ask - dsl + tp = sym.ask + dsl + request = {'type': OrderType.BUY, 'symbol': 'BTCUSD', 'volume': sym.volume_min, + 'price': sym.ask, 'action': TradeAction.DEAL, 'sl': sl, 'tp': tp} + return request + + def modify_stops(self, order): + ... + + def test_span_and_range(self): + assert self.bte.range == range(0, int((self.end - self.start).total_seconds()), self.bte.speed) + assert self.bte.span == range(int(self.start.timestamp()), int(self.end.timestamp()), self.bte.speed) + assert len(self.bte.span) == len(self.bte.range) + + def test_cursor(self): + self.bte.next() + r, t = self.bte.cursor + self.bte.fast_forward(steps=100) + assert self.bte.cursor.time == t + 100 + assert self.bte.cursor.index == r + 100 + go_to = datetime(2024, 2, 3, tzinfo=UTC) + self.bte.go_to(time=go_to) + assert self.bte.cursor.time == int(datetime.timestamp(go_to)) + self.bte.reset() + assert self.bte.cursor.time == int(self.start.timestamp()) + + def test_speed(self): + self.bte.setup_test_range(start=self.start, end=self.end, speed=3600) + assert self.bte.speed == 3600 + self.bte.next() + now = datetime.fromtimestamp(self.bte.cursor.time, tz=UTC) + index = self.bte.cursor.index + self.bte.next() + assert self.bte.cursor.index == index + 3600 + assert self.bte.cursor.time == int(now.timestamp()) + 3600 + self.bte.setup_test_range(start=self.start, end=self.end) + assert self.bte.speed == 1 + + async def test_account(self): + await self.bte.setup_account(balance=100) + acc = self.bte.get_account_info() + self.bte.use_terminal_for_backtesting = False + self.bte.use_terminal_for_backtesting = True + assert acc.balance == 100 + assert acc.equity == 100 + assert acc.margin == 0 + assert acc.margin_free == 100 + assert acc.margin_level == 0 + self.bte.deposit(amount=50) + acc = self.bte.get_account_info() + assert acc.balance == 150 + assert acc.equity == 150 + assert acc.margin == 0 + assert acc.margin_free == 150 + assert acc.margin_level == 0 + self.bte.withdraw(amount=80) + acc = self.bte.get_account_info() + assert acc.balance == 70 + assert acc.equity == 70 + assert acc.margin == 0 + assert acc.margin_free == 70 + assert acc.margin_level == 0 + self.bte.update_account(profit=-5) + acc = self.bte.get_account_info() + assert acc.equity == 65 + assert acc.balance == 70 + assert acc.profit == -5 + assert acc.margin == 0 + assert acc.margin_free == 65 + assert acc.margin_level == 0 + self.bte.update_account(margin=2.5) + acc = self.bte.get_account_info() + assert acc.balance == 70 + assert acc.equity == 65 + assert acc.margin == 2.5 + assert acc.margin_free == 62.5 + assert acc.margin_level == 2600 + + async def test_bte2_init(self, bte2): + assert bte2._data.fully_loaded is True + assert bte2.span == self.bte.span + assert bte2.range == self.bte.range + assert bte2.use_terminal is False + + async def test_get_rates_from(self): + start = datetime(2024, 2, 3, 12, 43, tzinfo=UTC) + rates = await self.bte.get_rates_from(symbol='BTCUSD', timeframe=TimeFrame.H1, date_from=start, count=24) + assert len(rates) == 24 + + async def test_get_rates_from_2(self, bte2): + start = datetime(2024, 2, 3, 12, 12, tzinfo=UTC) + rates = await bte2.get_rates_from(symbol='BTCUSD', timeframe=TimeFrame.H1, date_from=start, count=24) + assert len(rates) == 24 + + async def test_get_rates_from_pos(self): + now = datetime(2024, 2, 3, 11, 55, tzinfo=UTC) + self.bte.go_to(time=now) + tf = TimeFrame.H2 + start_pos = 2 + rates = await self.bte.get_rates_from_pos(symbol='BTCUSD', timeframe=tf, start_pos=start_pos, count=24) + assert int(rates[-1][0]) == round_down(int(now.replace(hour=7).timestamp()), tf.seconds) + assert len(rates) == 24 + + async def test_get_rates_from_pos2(self, bte2): + now = datetime(2024, 2, 4, 12, 15, tzinfo=UTC) + bte2.go_to(time=now) + tf = TimeFrame.H1 + start_pos = 2 + rates = await bte2.get_rates_from_pos(symbol='BTCUSD', timeframe=tf, start_pos=start_pos, count=24) + assert int(rates[-1][0]) == round_down(int(now.replace(hour=10).timestamp()), tf.seconds) + # assert int(rates[-1][0]) == round_up(int(now.timestamp()), tf.seconds) - start_pos * tf.seconds + assert len(rates) == 24 + + async def test_get_rates_range(self): + start = datetime(2024, 2, 3, 12, tzinfo=UTC) + end = datetime(2024, 2, 4, 18, tzinfo=UTC) + rates = await self.bte.get_rates_range(symbol="BTCUSD", timeframe=TimeFrame.H1, date_from=start, date_to=end) + assert len(rates) == 31 + assert int(rates[-1][0]) == int(end.timestamp()) + + async def test_get_rates_range2(self, bte2): + start = datetime(2024, 2, 3, 12, tzinfo=UTC) + end = datetime(2024, 2, 4, 18, tzinfo=UTC) + rates = await bte2.get_rates_range(symbol="BTCUSD", timeframe=TimeFrame.H1, date_from=start, date_to=end) + assert len(rates) == 31 + assert int(rates[-1][0]) == int(end.timestamp()) + + async def test_get_ticks_from(self): + start = datetime(2024, 2, 3, 12, tzinfo=UTC) + ticks = await self.bte.get_ticks_from(symbol='BTCUSD', date_from=start, count=24) + assert len(ticks) == 24 + + async def test_get_ticks_from2(self, bte2): + start = datetime(2024, 2, 3, 12, tzinfo=UTC) + ticks = await bte2.get_ticks_from(symbol='BTCUSD', date_from=start, count=24) + assert len(ticks) == 24 + + async def test_get_ticks_range(self): + start = datetime(2024, 2, 3, 12, tzinfo=UTC) + end = datetime(2024, 2, 3, 15, tzinfo=UTC) + ticks = await self.bte.get_ticks_range(symbol="BTCUSD", date_from=start, date_to=end) + approx_total = (end - start).total_seconds() // 2 # assuming 2 ticks per second at least + assert len(ticks) >= approx_total + + async def test_get_ticks_range2(self, bte2): + start = datetime(2024, 2, 3, 12, tzinfo=UTC) + end = datetime(2024, 2, 3, 15, tzinfo=UTC) + ticks = await bte2.get_ticks_range(symbol="BTCUSD", date_from=start, date_to=end) + approx_total = (end - start).total_seconds() // 2 # assuming 2 ticks per second at least + assert len(ticks) >= approx_total + + async def test_price_tick(self, bte2): + moment = datetime(2024, 2, 3, 12, 12, tzinfo=UTC) + self.bte.reset() + self.bte.go_to(time=moment) + tick = await self.bte.get_price_tick(symbol='BTCUSD', time=self.bte.cursor.time) + assert tick is not None + assert isinstance(tick.ask, float) + assert tick.ask > 0 + bte2.reset() + bte2.go_to(time=moment) + tick2 = await bte2.get_price_tick(symbol='BTCUSD', time=bte2.cursor.time) + assert tick.ask == tick2.ask + + async def test_get_symbol_info(self, bte2): + moment = datetime(2024, 2, 3, 12, 12, tzinfo=UTC) + self.bte.reset() + self.bte.go_to(time=moment) + bte2.reset() + bte2.go_to(time=moment) + sym = 'BTCUSD' + sym_info = await self.bte.get_symbol_info(symbol=sym) + assert sym_info is not None + assert sym_info.name == sym + sym_info2 = await bte2.get_symbol_info(symbol=sym) + assert sym_info.ask == sym_info2.ask + + async def test_order_profit(self, bte2): + moment = datetime(2024, 2, 3, 12, 12, tzinfo=UTC) + self.bte.reset() + self.bte.go_to(time=moment) + bte2.reset() + bte2.go_to(time=moment) + sym = 'BTCUSD' + sym_info = await self.bte.get_symbol_info(symbol=sym) + dsl = (sym_info.trade_stops_level + sym_info.spread) * 2 * sym_info.point + tp = sym_info.ask + dsl + + profit = await self.bte.order_calc_profit(action=OrderType.BUY, symbol=sym, + volume=sym_info.volume_min, price_open=sym_info.ask, + price_close=tp) + assert profit > 0 + sym_info2 = await bte2.get_symbol_info(symbol=sym) + dsl2 = (sym_info2.trade_stops_level + sym_info2.spread) * 2 * sym_info2.point + tp2 = sym_info2.ask + dsl2 + profit2 = await bte2.order_calc_profit(action=OrderType.BUY, symbol=sym, + volume=sym_info2.volume_min, price_open=sym_info2.ask, + price_close=tp2) + assert profit == profit2 + + async def test_order_margin(self, bte2): + moment = datetime(2024, 2, 3, 12, 12, tzinfo=UTC) + self.bte.reset() + self.bte.go_to(time=moment) + bte2.reset() + bte2.go_to(time=moment) + sym = 'BTCUSD' + sym_info = await self.bte.get_symbol_info(symbol=sym) + margin = await self.bte.order_calc_margin(action=OrderType.SELL, symbol=sym, + volume=sym_info.volume_min, price=sym_info.bid) + assert margin > 0 + sym_info2 = await self.bte.get_symbol_info(symbol=sym) + margin2 = await bte2.order_calc_margin(action=OrderType.SELL, symbol=sym, + volume=sym_info2.volume_min, price=sym_info2.bid) + assert margin2 > 0 + + async def test_order_check(self, buy_order, sell_order): + ocr = await self.bte.order_check(request=buy_order) + assert ocr is not None + assert ocr.retcode == 0 + ocr2 = await self.bte.order_check(request=sell_order) + assert ocr2 is not None + assert ocr2.retcode == 0 + + async def test_order_send(self, buy_order, sell_order): + ocr = await self.bte.order_send(request=buy_order) + assert ocr is not None + assert ocr.retcode == 10009 + ocr2 = await self.bte.order_send(request=sell_order) + assert ocr2 is not None + assert ocr2.retcode == 10009 diff --git a/tests/unit/test_base.py b/tests/live/unit/test_base.py similarity index 100% rename from tests/unit/test_base.py rename to tests/live/unit/test_base.py diff --git a/tests/live/unit/test_bot_and_executor.py b/tests/live/unit/test_bot_and_executor.py new file mode 100644 index 0000000..b33f428 --- /dev/null +++ b/tests/live/unit/test_bot_and_executor.py @@ -0,0 +1,51 @@ +import asyncio + +import pytest + +from aiomql.lib.bot import Bot + + +class TestBotFactoryAndExecutor: + @classmethod + def setup_class(cls): + cls.bot = Bot() + + @pytest.fixture(scope='class', autouse=True) + async def initialize(self): + self.bot.add_coroutine(coroutine=self.coro_one) + self.bot.add_coroutine(coroutine=self.coro_two) + self.bot.add_function(function=self.fun_one) + self.bot.add_coroutine(coroutine=self.coro_thread, on_separate_thread=True) + await self.bot.initialize() + + @staticmethod + def fun_one(): + print('function one') + + @staticmethod + async def coro_thread(): + while True: + print('coroutine thread') + await asyncio.sleep(1) + + @staticmethod + async def coro_one(): + while True: + print('coroutine one') + await asyncio.sleep(1) + + @staticmethod + async def coro_two(): + while True: + print('coroutine two') + await asyncio.sleep(1) + + def test_add_workers(self): + assert len(self.bot.executor.coroutines) == 2 + # exit function already added + assert len(self.bot.executor.functions) == 2 + # task_queue already added coroutine_thread + assert len(self.bot.executor.coroutine_threads) == 2 + + + # def diff --git a/tests/unit/test_candles.py b/tests/live/unit/test_candles.py similarity index 100% rename from tests/unit/test_candles.py rename to tests/live/unit/test_candles.py diff --git a/tests/unit/test_config.py b/tests/live/unit/test_config.py similarity index 93% rename from tests/unit/test_config.py rename to tests/live/unit/test_config.py index 6b16720..48d7e45 100644 --- a/tests/unit/test_config.py +++ b/tests/live/unit/test_config.py @@ -25,5 +25,5 @@ class TestConfig: assert 'server' in account_info def test_load_config(self, config): - config.load_config(file='tests/configs/test2.json') + config.load_config(file='tests/live/configs/test2.json') assert config.filename == 'test2.json' diff --git a/tests/live/unit/test_get_data.py b/tests/live/unit/test_get_data.py new file mode 100644 index 0000000..c763ec1 --- /dev/null +++ b/tests/live/unit/test_get_data.py @@ -0,0 +1,48 @@ +from pathlib import Path +from datetime import datetime, UTC + +import pytest + +from aiomql.contrib.backtesting.get_data import GetData +from aiomql.core.constants import TimeFrame + + +class TestGetData: + @classmethod + def setup_class(cls): + cls.start = datetime(2024, 2, 1, tzinfo=UTC) + cls.end = datetime(2024, 2, 2, tzinfo=UTC) + cls.symbols = ['BTCUSD', "ETHUSD"] + cls.timeframes = [TimeFrame.H1, TimeFrame.H2] + cls.g_data = GetData(start=cls.start, end=cls.end, symbols=cls.symbols, timeframes=cls.timeframes, + name='test_data') + + @pytest.fixture(scope='class', autouse=True) + async def get_data(self): + await self.g_data.get_data() + self.g_data.save_data() + + def test_init(self): + assert self.g_data.start == self.start + assert self.g_data.end == self.end + assert self.g_data.symbols == set(self.symbols) + assert self.g_data.timeframes == set(self.timeframes) + assert self.g_data.name == 'test_data' + assert self.g_data.range == range(int((self.end - self.start).total_seconds())) + assert self.g_data.span == range(int(self.start.timestamp()), int(self.end.timestamp())) + + async def test_get_data(self): + assert self.g_data.data.fully_loaded is True + assert len(self.g_data.data.ticks.keys()) == 2 + assert len(self.g_data.data.symbols.keys()) == 2 + + async def test_save_data(self): + file = Path(self.g_data.config.backtest_dir / 'test_data.pkl') + assert file.exists() + + async def test_load_data(self): + data = GetData.load_data(name='tests/live/backtesting/test_data.pkl') + assert data.name == 'test_data' + assert data.fully_loaded is True + assert len(data.ticks.keys()) == 2 + assert len(data.symbols.keys()) == 2 diff --git a/tests/unit/test_history.py b/tests/live/unit/test_history.py similarity index 97% rename from tests/unit/test_history.py rename to tests/live/unit/test_history.py index 40cf9b4..94b2da4 100644 --- a/tests/unit/test_history.py +++ b/tests/live/unit/test_history.py @@ -6,7 +6,7 @@ from aiomql.lib.history import History class TestHistory: @pytest.fixture(scope='class', autouse=True) async def init(self, make_buy_sell_orders): - await self.history.init() + await self.history.initialize() @classmethod def setup_class(cls): diff --git a/tests/unit/test_meta_trader.py b/tests/live/unit/test_meta_trader.py similarity index 100% rename from tests/unit/test_meta_trader.py rename to tests/live/unit/test_meta_trader.py diff --git a/tests/unit/test_order.py b/tests/live/unit/test_order.py similarity index 100% rename from tests/unit/test_order.py rename to tests/live/unit/test_order.py diff --git a/tests/unit/test_positions.py b/tests/live/unit/test_positions.py similarity index 100% rename from tests/unit/test_positions.py rename to tests/live/unit/test_positions.py diff --git a/tests/unit/test_ram.py b/tests/live/unit/test_ram.py similarity index 100% rename from tests/unit/test_ram.py rename to tests/live/unit/test_ram.py diff --git a/tests/unit/test_result.py b/tests/live/unit/test_result.py similarity index 100% rename from tests/unit/test_result.py rename to tests/live/unit/test_result.py diff --git a/tests/unit/test_sessions.py b/tests/live/unit/test_sessions.py similarity index 77% rename from tests/unit/test_sessions.py rename to tests/live/unit/test_sessions.py index c26536d..bd6c07e 100644 --- a/tests/unit/test_sessions.py +++ b/tests/live/unit/test_sessions.py @@ -1,7 +1,6 @@ -from datetime import datetime, time +from datetime import datetime, time, UTC import pytest -import pytz from aiomql.lib.sessions import Session, Sessions, delta @@ -14,11 +13,11 @@ class TestSessions: @pytest.fixture(scope='class') def make_session(self): - end = time(hour=16, minute=59, second=59, microsecond=999_999, tzinfo=pytz.UTC) + end = time(hour=16, minute=59, second=59, microsecond=999_999, tzinfo=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) + start, end = time(hour=0, tzinfo=UTC), time(hour=23, minute=59, second=59, tzinfo=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) + end = time(hour=6, minute=59, second=59, microsecond=999_999, tzinfo=UTC) over_night = Session(start=18, end=end, name='OverNight', on_end='close_all') return london, all_day, over_night @@ -26,16 +25,16 @@ class TestSessions: 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.start == time(hour=8, tzinfo=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() + two_am = time(hour=2, tzinfo=UTC) + noon = time(hour=12, tzinfo=UTC) + now = datetime.now(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 @@ -48,12 +47,12 @@ class TestSessions: 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) + now = time(hour=21, tzinfo=UTC) + noon = time(hour=12, tzinfo=UTC) + mid_nite = time(hour=0, tzinfo=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)) + no_sess = sessions.find(moment=time(hour=17, tzinfo=UTC)) mid_nite_sess = sessions.find(moment=mid_nite) current_sess = sessions.find(moment=now) assert current_sess.name == 'OverNight' @@ -61,7 +60,7 @@ class TestSessions: assert no_sess is None assert next_sess.name == 'London' assert mid_nite_sess.name == 'OverNight' - current = datetime.now(pytz.UTC).time() + current = datetime.now(UTC).time() if current.hour not in (7, 17): await sessions.check() assert sessions.current_session is not None diff --git a/tests/unit/test_symbol.py b/tests/live/unit/test_symbol.py similarity index 98% rename from tests/unit/test_symbol.py rename to tests/live/unit/test_symbol.py index 076dd78..e7069bf 100644 --- a/tests/unit/test_symbol.py +++ b/tests/live/unit/test_symbol.py @@ -13,7 +13,7 @@ class TestSymbol: symbol = Symbol(name='BTCUSD') select = getattr(symbol, 'select', False) if select is False: - await symbol.init() + await symbol.initialize() return symbol async def test_symbol_attributes(self, btc): diff --git a/tests/live/unit/test_task_queue.py b/tests/live/unit/test_task_queue.py new file mode 100644 index 0000000..ba83fc9 --- /dev/null +++ b/tests/live/unit/test_task_queue.py @@ -0,0 +1,40 @@ +import asyncio + +from aiomql.core.task_queue import TaskQueue, QueueItem + + +class TestTaskQueue: + @classmethod + def setup_class(cls): + cls.task_queue = TaskQueue(timeout=5, worker_timeout=1) + cls.data = {} + + async def task_one(self): + for i in range(10): + await asyncio.sleep(0.5) + self.data.setdefault('task_one', {})[i] = f"task_one_{i}" + + async def task_two(self): + for i in range(10): + await asyncio.sleep(0.5) + self.data.setdefault('task_two', {})[i] = f"task_two_{i}" + + async def task_three(self): + for i in range(10): + self.data.setdefault('task_three', {})[i] = f"task_three_{i}" + await asyncio.sleep(10) + + async def test_queue(self): + item_one = QueueItem(self.task_one) + self.task_queue.add(item=item_one, must_complete=False) + assert len(self.task_queue.priority_tasks) == 0 + assert self.task_queue.queue.qsize() == 1 + self.task_queue.add(item=QueueItem(self.task_two), must_complete=True) + assert len(self.task_queue.priority_tasks) == 1 + assert self.task_queue.queue.qsize() == 2 + self.task_queue.add(item=QueueItem(self.task_three), must_complete=False) + await self.task_queue.run() + assert len(self.data['task_one']) >= 2 + assert len(self.data['task_two']) == 10 + assert len(self.data['task_three']) == 1 + assert len(self.task_queue.priority_tasks) == 0 diff --git a/tests/unit/test_terminal.py b/tests/live/unit/test_terminal.py similarity index 100% rename from tests/unit/test_terminal.py rename to tests/live/unit/test_terminal.py diff --git a/tests/unit/test_ticks.py b/tests/live/unit/test_ticks.py similarity index 100% rename from tests/unit/test_ticks.py rename to tests/live/unit/test_ticks.py diff --git a/tests/live/unit/test_trader.py b/tests/live/unit/test_trader.py new file mode 100644 index 0000000..de80f2d --- /dev/null +++ b/tests/live/unit/test_trader.py @@ -0,0 +1,69 @@ +from math import floor +import pytest + +from aiomql.lib.ram import RAM +from aiomql.contrib.traders import SimpleTrader +from aiomql.contrib.symbols import ForexSymbol +from aiomql.core.constants import OrderType + + +class TestTrader: + @classmethod + def setup_class(cls): + ram = RAM(fixed_amount=10) + cls.trader = SimpleTrader(symbol=ForexSymbol(name='BTCUSD'), ram=ram) + cls.simple_trader2 = SimpleTrader(symbol=ForexSymbol(name='EURJPY'), ram=ram) + + @pytest.fixture(scope='class', autouse=True) + async def initialize(self): + await self.trader.symbol.initialize() + await self.simple_trader2.symbol.initialize() + + async def test_create_order_no_stops(self): + await self.trader.create_order_no_stops(order_type=OrderType.BUY) + assert self.trader.order.volume == self.trader.symbol.volume_min + res = await self.trader.order.send() + assert res is not None + assert res.retcode == 10009 + + async def test_create_order_with_sl(self): + sl = (self.trader.symbol.trade_stops_level * 2 + self.trader.symbol.spread) * self.trader.symbol.point + tick = await self.trader.symbol.info_tick() + sl = tick.bid + sl + await self.trader.create_order_with_sl(order_type=OrderType.SELL, sl=sl) + res = await self.trader.order.send() + profit = floor(await self.trader.order.calc_profit()) + loss = -floor(abs(await self.trader.order.calc_loss())) + assert profit == -loss*self.trader.ram.risk_to_reward + assert profit == self.trader.ram.fixed_amount * self.trader.ram.risk_to_reward + assert loss == -self.trader.ram.fixed_amount + assert res is not None + assert res.retcode == 10009 + + async def test_create_order_with_points(self): + points = (self.trader.symbol.trade_stops_level * 2 + self.trader.symbol.spread) + await self.trader.create_order_with_points(order_type=OrderType.BUY, points=points) + res = await self.trader.order.send() + profit = floor(await self.trader.order.calc_profit()) + loss = -floor(abs(await self.trader.order.calc_loss())) + assert profit == -loss * self.trader.ram.risk_to_reward + assert profit == self.trader.ram.fixed_amount * self.trader.ram.risk_to_reward + assert loss == -self.trader.ram.fixed_amount + assert res is not None + assert res.retcode == 10009 + + async def test_create_order_with_stops(self): + sl = (self.trader.symbol.trade_stops_level * 2 + self.trader.symbol.spread) * self.trader.symbol.point + tp = sl * self.trader.ram.risk_to_reward + tick = await self.trader.symbol.info_tick() + sl = tick.ask - sl + tp = tick.ask + tp + await self.trader.create_order_with_stops(order_type=OrderType.BUY, sl=sl, tp=tp) + res = await self.trader.order.send() + profit = floor(await self.trader.order.calc_profit()) + loss = -floor(abs(await self.trader.order.calc_loss())) + assert profit == -loss * self.trader.ram.risk_to_reward + assert profit == self.trader.ram.fixed_amount * self.trader.ram.risk_to_reward + assert loss == -self.trader.ram.fixed_amount + assert res is not None + assert res.retcode == 10009 diff --git a/trade_records/Chaos.csv b/trade_records/Chaos.csv new file mode 100644 index 0000000..cb7bceb --- /dev/null +++ b/trade_records/Chaos.csv @@ -0,0 +1,7 @@ +actual_profit,deal,ask,ltf,htf,price,order,slow_ema,symbol,expected_profit,lcc,date,bid,name,win,hcc,fast_ema,volume,closed +0,127446532,643587.63,TIMEFRAME_M1,TIMEFRAME_M2,643587.63,436128326,20,Volatility 75 Index,0,100,2021-01-01 00:00:00.000000,643457.63,Chaos,False,100,8,0.001,False +0,337655374,2268.02,TIMEFRAME_M1,TIMEFRAME_M2,2267.34,580347053,20,Volatility 100 Index,0,100,2021-01-01 00:00:00.000000,2267.34,Chaos,False,100,8,0.5,False +0,895556166,189.1722,TIMEFRAME_M1,TIMEFRAME_M2,189.1452,645695617,20,Volatility 50 Index,0,100,2021-01-01 00:00:00.000000,189.1452,Chaos,False,100,8,4.0,False +0,215199127,643587.63,TIMEFRAME_M1,TIMEFRAME_M2,643587.63,274191555,20,Volatility 75 Index,0,100,2021-01-01 00:00:00.000000,643457.63,Chaos,False,100,8,0.001,False +0,949649264,2268.02,TIMEFRAME_M1,TIMEFRAME_M2,2268.02,740685256,20,Volatility 100 Index,0,100,2021-01-01 00:00:00.000000,2267.34,Chaos,False,100,8,0.5,False +0,939601387,189.1722,TIMEFRAME_M1,TIMEFRAME_M2,189.1452,969153963,20,Volatility 50 Index,0,100,2021-01-01 00:00:00.000000,189.1452,Chaos,False,100,8,4.0,False