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