This commit is contained in:
Ichinga Samuel
2024-09-11 06:16:27 +01:00
parent a10faf5e60
commit e8fd4ab4dc
19 changed files with 201 additions and 216 deletions
+6
View File
@@ -75,12 +75,18 @@ class Account(AccountInfo):
async def _login(self, *, acc: dict, tries=3, **kwargs) -> bool:
res = False
if tries == 0:
return False
init_args = {**acc} | {'path': self.config.path} | {**kwargs}
ini = await self.mt5.initialize(**init_args)
if ini:
res = await self.mt5.login(**acc)
if not res:
await self.mt5.shutdown()
if ini and res:
return True
else:
+3 -42
View File
@@ -1,43 +1,4 @@
from functools import wraps
from dataclasses import dataclass, fields, field
from typing import ClassVar
def fun(a, b=6, *c, **d):
print(f"{a=}, {b=}, {c=}, {d=}")
def dd(func):
@wraps(func)
def wrapper(*args, **kwargs):
print(func.__name__)
return func(*args, **kwargs)
return wrapper
class C:
def __new__(cls, *args, **kwargs):
if not hasattr(cls, '_instance'):
cls._instance = super().__new__(cls)
cls._instance.tasks = []
# cls.__init__(a)
[setattr(cls._instance, k, v) for k, v in kwargs.items()]
return cls._instance
def __init__(self, *args, **kwargs):
print('receiving args')
@dataclass
class D:
b: int = 0
c: str = ''
_fields: list[ClassVar[str]] = field(default_factory=list)
@dd
def setattrs(self, **kwargs):
[setattr(self, k, v) for k, v in kwargs.items() if k in self.fields]
@property
def fields(self):
return self._fields or [name for f in fields(self) if (name := f.name) != '_fields']
d = D()
d.setattrs(r=3)
fun(1, 3, 4, 5, six=6, seven=7)
+12 -4
View File
@@ -1,7 +1,7 @@
import asyncio
from asyncio import Condition, Task
from typing import Self
from datetime import datetime
from ...core import Config
@@ -25,11 +25,15 @@ class EventManager:
def __init__(self, *, num_tasks: int = 0):
self.num_main_tasks = num_tasks or self.num_main_tasks
def add_task(self, *task: Task):
self.tasks.extend(task)
def add_task(self, *tasks: Task):
self.tasks.extend(tasks)
def sigint_handler(self, sig, frame):
print(self.config.test_data)
print('KeyboardInterrupt')
print(self.config.test_data.orders)
for task in self.tasks:
task.cancel()
# self.mt.cancel()
async def acquire(self):
await self.condition.acquire()
@@ -45,6 +49,10 @@ class EventManager:
await self.config.test_data.tracker()
self.config.test_data.next()
self.condition.notify_all()
if (timestamp := self.config.test_data.cursor.time % int(60 * 60 * 24)) == 0:
print(f"Time: {datetime.fromtimestamp(timestamp)}")
else:
print(f"Time: {datetime.fromtimestamp(timestamp)}")
await asyncio.sleep(0)
async def wait(self):
+32 -32
View File
@@ -4,7 +4,6 @@ from pathlib import Path
import lzma
from datetime import datetime
from logging import getLogger
import asyncio
from typing import Sequence, ClassVar
import pytz
@@ -13,7 +12,7 @@ from pandas import DataFrame
from ...core.meta_trader import MetaTrader
from ...core.config import Config
from ...core.constants import TimeFrame,
from ...core.constants import TimeFrame, CopyTicks
from ...core.task_queue import TaskQueue, QueueItem
from ...utils import backoff_decorator
@@ -76,31 +75,6 @@ class GetData:
self.mt5 = MetaTrader()
self.task_queue = TaskQueue()
async def get_data(self):
""""""
qitems = [QueueItem(self.get_symbol_rates), QueueItem(self.get_symbol_ticks),
QueueItem(self.get_symbol_prices), QueueItem(self.get_symbol_info),
QueueItem(self.get_account_info), QueueItem(self.get_symbols_info), QueueItem(self.get_version)]
[self.task_queue.add(item=item, priority=0) for item in qitems]
await self.task_queue.run()
# terminal, version = await asyncio.gather(self.get_terminal_info(), self.get_version())
# self.data.setattrs(account=account, symbols=symbols, prices=prices, ticks=ticks, rates=rates,
# span=self.span, range=self.range, terminal=terminal, version=version, name=self.name)
def pickle_data(self):
""""""
fh = open(f'{self.config.test_data_dir}/{self.name}', 'wb')
pickle.dump(self.data, fh)
fh.close()
async def compress_data(self):
""""""
bdata = pickle.dumps(self.data)
name = self.name + 'xz'
with lzma.open(f'{self.config.test_data_dir}/{name}', 'w') as fh:
fh.write(bdata)
@classmethod
def dump_data(cls, data: Data, name: str | Path, compress: bool = False):
""""""
@@ -136,6 +110,32 @@ class GetData:
logger.error(f"Error: {err}")
return None
async def get_data(self):
""""""
q_items = [QueueItem(self.get_symbol_rates, must_complete=True),
QueueItem(self.get_symbol_ticks, must_complete=True),
QueueItem(self.get_symbol_prices, must_complete=True),
QueueItem(self.get_symbol_info, must_complete=True),
QueueItem(self.get_account_info, must_complete=True),
QueueItem(self.get_symbols_info, must_complete=True),
QueueItem(self.get_version, must_complete=True),
QueueItem(self.get_terminal_info, must_complete=True)]
[self.task_queue.add(item=item, priority=0) for item in q_items]
await self.task_queue.run()
def pickle_data(self):
""""""
fh = open(f'{self.config.test_data_dir}/{self.name}', 'wb')
pickle.dump(self.data, fh)
fh.close()
async def compress_data(self):
""""""
bdata = pickle.dumps(self.data)
name = self.name + 'xz'
with lzma.open(f'{self.config.test_data_dir}/{name}', 'w') as fh:
fh.write(bdata)
async def get_terminal_info(self):
""""""
terminal = await self.mt5.terminal_info()
@@ -149,19 +149,19 @@ class GetData:
async def get_symbols_info(self):
""""""
[self.task_queue.add(QueueItem(self.get_symbol_info, symbol, must_complete=True), priority=4) for symbol in self.symbols]
[self.task_queue.add(item=QueueItem(self.get_symbol_info, symbol)) for symbol in self.symbols]
async def get_symbols_ticks(self):
""""""
[self.task_queue.add(QueueItem(self.get_symbol_ticks, symbol)) for symbol in self.symbols]
[self.task_queue.add(item=QueueItem(self.get_symbol_ticks, symbol)) for symbol in self.symbols]
async def get_symbols_prices(self):
""""""
[self.task_queue.add(QueueItem(self.get_symbol_prices, symbol)) for symbol in self.symbols]
[self.task_queue.add(item=QueueItem(self.get_symbol_prices, symbol)) for symbol in self.symbols]
async def get_symbols_rates(self):
""""""
[self.task_queue.add(QueueItem(self.get_symbol_rates, symbol, timeframe, must_complete=True), priority=4)
[self.task_queue.add(item=QueueItem(self.get_symbol_rates, symbol, timeframe), priority=4)
for symbol in self.symbols for timeframe in self.timeframes]
@backoff_decorator(max_retries=5)
@@ -203,4 +203,4 @@ class GetData:
res = pd.DataFrame(res)
res.drop_duplicates(subset=['time'], keep='last', inplace=True)
res.set_index('time', inplace=True, drop=False)
self.data.rates[symbol].setdefault(timeframe, res)
self.data.rates.setdefault(symbol, {})[timeframe.name] = res
+6 -6
View File
@@ -62,13 +62,13 @@ class MetaTester(MetaTrader):
async def shutdown(self) -> None:
await super().shutdown() if self.config.use_terminal_for_backtesting else ...
self.test_data.save()
name = self.test_data.data.name
if self.config.compress_test_data:
name += '.xz'
name = self.config.test_data_dir/name
GetData.dump_data(data=self.test_data.data, name=name, compress=self.config.compress_test_data)
# self.test_data.save()
# name = self.test_data.data.name
# if self.config.compress_test_data:
# name += '.xz'
# name = self.config.test_data_dir/name
# GetData.dump_data(data=self.test_data.data, name=name, compress=self.config.compress_test_data)
@error_handler(msg='test data not available', exe=AttributeError)
async def terminal_info(self) -> TerminalInfo:
@@ -1,4 +1,5 @@
import asyncio
import signal
from .event_manager import EventManager
from .get_data import GetData
@@ -15,6 +16,7 @@ class StrategyTester:
self.test_data = test_data or self.get_test_data(name=test_data_file)
self.config.test_data = self.test_data
self.event_manager = EventManager(num_tasks=len(self.strategies))
signal.signal(signal.SIGINT, self.event_manager.sigint_handler)
def get_test_data(self, name: str) -> TestData | None:
name = f"{self.config.test_data_dir_name}/{name or self.config.test_data_file}"
@@ -27,9 +29,14 @@ class StrategyTester:
await self.mt5.login(**acc)
async def run(self):
await self.start()
tasks = [*[asyncio.create_task(strategy.test()) for strategy in self.strategies],
asyncio.create_task(self.event_manager.event_monitor())]
self.event_manager.add_task(*tasks)
await asyncio.gather(*tasks)
await self.mt5.shutdown()
try:
await self.start()
tasks = [*[asyncio.create_task(strategy.test()) for strategy in self.strategies],
asyncio.create_task(self.event_manager.event_monitor())]
self.event_manager.add_task(*tasks)
await asyncio.gather(*tasks, return_exceptions=True)
# self.mt = asyncio.create_task(asyncio.gather(*tasks))
except Exception as err:
print(f"Error {err} occurred in StrategyTester")
finally:
await self.mt5.shutdown()
@@ -16,5 +16,6 @@ class TestStrategy:
mod = time % secs
secs = secs - mod if mod != 0 else mod
time = self.config.test_data.cursor.time + secs
print(f"Sleeping for {secs} seconds")
while time > self.config.test_data.cursor.time:
await self.event_manager.wait()
+11 -6
View File
@@ -8,7 +8,7 @@ logger = getLogger(__name__)
class QueueItem:
def __init__(self, task_item: Callable | Coroutine, must_complete=False, *args, **kwargs):
def __init__(self, task_item: Callable | Coroutine, *args, must_complete: bool = False, **kwargs):
self.task_item = task_item
self.args = args
self.kwargs = kwargs
@@ -30,7 +30,7 @@ class QueueItem:
self.task_item(*self.args, **self.kwargs)
except Exception as err:
logger.error(f"Error {err} occurred in {self.func.__name__} with args {self.args} and kwargs {self.kwargs}")
logger.error(f"Error {err} occurred in {self.task_item.__name__} with args {self.args} and kwargs {self.kwargs}")
class TaskQueue:
@@ -44,7 +44,7 @@ class TaskQueue:
self.timeout = timeout
self.stop = False
self.on_exit = on_exit
signal(SIGINT, self.sigint_handle)
# signal(SIGINT, self.sigint_handle)
def add(self, *, item: QueueItem, priority=3):
try:
@@ -71,19 +71,24 @@ class TaskQueue:
self.queue.task_done()
self.priority_tasks.discard(item)
if self.stop and len(self.priority_tasks) == 0:
print('All priority tasks completed')
self.cancel()
break
def sigint_handle(self, sig, frame):
print('SIGINT received, cleaning up...')
if self.on_exit == 'complete_priority':
if self.on_exit == 'complete_priority' and self.priority_tasks:
print(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
async def run(self, timeout: int = 0):
signal(SIGINT, self.sigint_handle)
loop = asyncio.get_running_loop()
start = loop.time()
@@ -116,4 +121,4 @@ class TaskQueue:
def cancel(self):
cancelled = [task.cancel() for task in self.tasks if not task.done()]
print(f'Cancelled {len(cancelled)} worker tasks') if cancelled else ...
self.tasks.clear()
self.tasks.clear()
+20 -23
View File
@@ -1,4 +1,3 @@
import asyncio
from datetime import datetime
from logging import getLogger
@@ -9,6 +8,10 @@ from .core.config import Config
from .core.meta_trader import MetaTrader, CopyTicks, OrderType
from .core.models import TradeDeal, TradeOrder
from .contrib.backtester.meta_tester import MetaTester
from .utils import backoff_decorator
logger = getLogger(__name__)
@@ -26,7 +29,7 @@ class History:
mt5 (MetaTrader): MetaTrader instance
config (Config): Config instance
"""
mt5: MetaTrader
mt5: MetaTrader | MetaTester
config: Config
def __init__(self, *, date_from: datetime | int = None, date_to: datetime | int = None,
@@ -44,7 +47,7 @@ class History:
position (int): Filter for selecting history deals by position
"""
self.config = Config()
self.mt5 = MetaTrader()
self.mt5 = MetaTrader() if self.config.mode == 'live' else MetaTester()
self.date_from = date_from
self.date_to = date_to
self.group = group
@@ -67,30 +70,24 @@ class History:
self.total_deals = len(self.deals)
self.total_orders = len(self.orders)
async def get_deals(self, *, date_from: datetime | int = None, date_to: datetime | int = None, group: str = '',
retries: int = 3) -> tuple[TradeDeal, ...]:
@backoff_decorator
async def get_deals(self, *, date_from: datetime | int = None, date_to: datetime | int = None, group: str = '')\
-> tuple[TradeDeal, ...]:
"""Get deals from trading history using the parameters set in the constructor.
Returns:
tuple[TradeDeal]: A list of trade deals
"""
if retries < 1:
logger.warning(f'Failed to get deals: {self.mt5.error}')
return tuple()
date_from, date_to, group = date_from or self.date_from, date_to or self.date_to, group or self.group
deals = await self.mt5.history_deals_get(date_from=date_from, date_to=date_to, group=group)
if deals is not None:
return tuple(TradeDeal(**deal._asdict()) for deal in deals)
if self.mt5.error.is_connection_error():
await asyncio.sleep(retries)
return await self.get_deals(date_from=date_from, date_to=date_to, group=group, retries=retries-1)
logger.warning(f'Failed to get deals: {self.mt5.error}')
logger.warning(f'Failed to get deals')
return tuple()
@backoff_decorator
async def get_deals_ticket(self, *, ticket: int = None) -> tuple[TradeDeal, ...]:
"""Call specifying the order ticket. Return all deals having the specified order ticket in the DEAL_ORDER
property.
@@ -106,6 +103,7 @@ class History:
deals = await self.mt5.history_deals_get(ticket=ticket)
return tuple(sorted([TradeDeal(**deal._asdict()) for deal in deals or []], key=lambda x: x.time_msc))
@backoff_decorator
async def get_deals_position(self, *, position: int = None) -> tuple[TradeDeal, ...]:
"""
Get all deals with the specified position ticket in the DEAL_POSITION_ID property
@@ -135,6 +133,7 @@ class History:
total_deals = await self.mt5.history_deals_total(date_from, date_to)
return total_deals
@backoff_decorator
async def get_orders(self, *, date_from: datetime | int = None, date_to: datetime | int = None, group: str = '',
retries: int = 3) -> tuple[TradeOrder, ...]:
"""Get orders from trading history using the parameters set in the constructor or the method arguments.
@@ -142,30 +141,27 @@ class History:
Returns:
list[TradeOrder]: A list of trade orders
"""
if retries < 1:
logger.warning(f'Failed to get orders: {self.mt5.error}')
return tuple()
date_from, date_to, group = date_from or self.date_from, date_to or self.date_to, group or self.group
orders = await self.mt5.history_orders_get(date_from=date_from, date_to=date_to, group=group)
if orders is not None:
return tuple(TradeOrder(**order._asdict()) for order in orders)
if self.mt5.error.is_connection_error():
await asyncio.sleep(retries)
return await self.get_orders(date_from=date_from, date_to=date_to, group=group, retries=retries - 1)
logger.warning(f'Failed to get orders: {self.mt5.error}')
logger.warning(f'Failed to get orders')
return tuple()
@backoff_decorator
async def get_order_ticket(self, ticket: int | None = None) -> TradeOrder | None:
ticket = ticket or self.ticket
assert isinstance(ticket, int), 'ticket not provided'
orders = await self.mt5.history_orders_get(ticket=ticket)
if orders and (order := orders[0]).ticket == ticket:
return TradeOrder(**order._asdict())
return None
@backoff_decorator
async def get_orders_position(self, position: int = None) -> tuple[TradeOrder, ...]:
"""
Call specifying the position ticket. Return all orders with a position ticket specified in the
@@ -182,6 +178,7 @@ class History:
orders = await self.mt5.history_orders_get(position=position)
return tuple(sorted([TradeOrder(**order._asdict()) for order in orders or []], key=lambda x: x.time_done_msc))
@backoff_decorator
async def orders_total(self, date_from: int | datetime = None, date_to: int | datetime = None) -> int:
"""Get total number of orders within the specified period in the constructor.
+15 -18
View File
@@ -1,9 +1,10 @@
import asyncio
from logging import getLogger
from .core.models import TradeRequest, OrderSendResult, OrderCheckResult, TradeOrder
from .core.constants import TradeAction, OrderTime, OrderFilling
from .core.exceptions import OrderError
from .utils import backoff_decorator
logger = getLogger(__name__)
@@ -38,26 +39,25 @@ class Order(TradeRequest):
"""
return await self.mt5.orders_total()
async def get_order(self, *, ticket: int, retries: int = 3) -> TradeOrder | None:
@backoff_decorator
async def get_order(self, *, ticket: int) -> TradeOrder | None:
"""
Get the order by ticket number.
Args:
ticket (int): Order ticket number
retries (int): Number of retries
Returns:
"""
if retries < 1:
return None
orders = await self.mt5.orders_get(ticket=ticket)
if orders and (order := orders[0]).ticket == ticket:
return TradeOrder(**order._asdict())
if self.mt5.error.is_connection_error():
await asyncio.sleep(retries)
return await self.get_order(ticket=ticket, retries=retries-1)
return None
async def get_orders(self, *, ticket: int = 0, symbol: str = '', group: str = '', retries=3)\
-> tuple[TradeOrder, ...]:
@backoff_decorator
async def get_orders(self, *, ticket: int = 0, symbol: str = '', group: str = '') -> tuple[TradeOrder, ...]:
"""Get the list of active orders for the current symbol.
Keyword Args:
ticket (int): Order ticket number
@@ -66,16 +66,13 @@ class Order(TradeRequest):
Returns:
tuple[TradeOrder]: A Tuple of active trade orders as TradeOrder objects
"""
if retries < 1:
return tuple()
symbol = getattr(self, 'symbol', symbol)
orders = await self.mt5.orders_get(symbol=symbol, ticket=ticket, group=group)
if orders is not None:
orders = (TradeOrder(**order._asdict()) for order in orders)
return tuple(orders)
if self.mt5.error.is_connection_error():
await asyncio.sleep(retries)
return await self.get_orders(ticket=ticket, symbol=symbol, group=group, retries=retries-1)
return tuple()
async def check(self, **kwargs) -> OrderCheckResult:
@@ -90,7 +87,7 @@ class Order(TradeRequest):
req = self.dict | kwargs
res = await self.mt5.order_check(req)
if res is None:
raise OrderError(f'Failed to check order due to {self.mt5.error.description}')
raise OrderError(f'Order check failed for {self.symbol}')
return OrderCheckResult(**res._asdict())
async def send(self) -> OrderSendResult:
@@ -104,7 +101,7 @@ class Order(TradeRequest):
"""
res = await self.mt5.order_send(self.dict)
if res is None:
raise OrderError(f'Failed to send order {self.symbol} due to {self.mt5.error.description}')
raise OrderError(f'Failed to send order {self.symbol}')
res = OrderSendResult(**res._asdict())
try:
profit = await self.calc_profit()
@@ -126,7 +123,7 @@ class Order(TradeRequest):
"""
res = await self.mt5.order_calc_margin(self.type, self.symbol, self.volume, self.price)
if res is None:
raise OrderError(f'Failed to calculate margin for {self.symbol} due to {self.mt5.error.description}')
raise OrderError(f'Failed to calculate margin for {self.symbol}')
return res
async def calc_profit(self, **kwargs) -> float | None:
+17 -11
View File
@@ -2,8 +2,15 @@
import asyncio
from logging import getLogger
from .core import MetaTrader, TradePosition, TradeAction, OrderType
from .core.meta_trader import MetaTrader
from .core.models import TradePosition, TradeAction
from .core.constants import OrderType
from .core.config import Config
from .contrib.backtester.meta_tester import MetaTester
from .order import Order
from .utils import backoff_decorator
logger = getLogger(__name__)
@@ -18,7 +25,7 @@ class Positions:
ticket (int): Position ticket.
mt5 (MetaTrader): MetaTrader instance.
"""
mt5: MetaTrader
mt5: MetaTrader | MetaTester
def __init__(self, *, symbol: str = "", group: str = "", ticket: int = 0):
"""Get Open Positions.
@@ -30,7 +37,8 @@ class Positions:
ticket (int): Position ticket
"""
self.mt5 = MetaTrader()
self.config = Config()
self.mt5 = MetaTrader() if self.config.mode == 'live' else MetaTester()
self.symbol = symbol
self.group = group
self.ticket = ticket
@@ -43,7 +51,8 @@ class Positions:
"""
return await self.mt5.positions_total()
async def positions_get(self, symbol: str = '', group: str = '', ticket: int = 0, retries=3) -> list[TradePosition]:
@backoff_decorator
async def positions_get(self, symbol: str = '', group: str = '', ticket: int = 0) -> list[TradePosition]:
"""Get open positions with the ability to filter by symbol or ticket.
Keyword Args:
@@ -55,17 +64,12 @@ class Positions:
Returns:
list[TradePosition]: A list of open trade positions
"""
if retries < 1:
logger.warning(f'Failed to get positions for {symbol or self.symbol}. {self.mt5.error}')
return []
positions = await self.mt5.positions_get(group=group or self.group, symbol=symbol or self.symbol,
ticket=ticket or self.ticket)
if positions is not None:
return [TradePosition(**pos._asdict()) for pos in positions]
if self.mt5.error.is_connection_error():
await asyncio.sleep(retries)
return await self.positions_get(symbol, group, ticket, retries - 1)
logger.warning(f'Failed to get positions for {symbol or self.symbol}. {self.mt5.error}')
logger.warning(f'Failed to get positions for {symbol or self.symbol}')
return []
async def position_get(self, *, ticket: int) -> TradePosition | None:
@@ -78,8 +82,10 @@ class Positions:
"""
positions = await self.positions_get(ticket=ticket)
position = positions[0] if positions else None
if position is None or position.ticket != ticket:
return None
return position
async def close(self, *, ticket: int, symbol: str, price: float, volume: float, order_type: OrderType):
+4 -2
View File
@@ -6,7 +6,9 @@ import csv
import logging
from typing import Iterable
from .core import Config, MetaTrader
from .contrib.backtester.meta_tester import MetaTester
from .core.config import Config
from .core.meta_trader import MetaTrader
logger = logging.getLogger(__name__)
@@ -30,7 +32,7 @@ class Records:
records_dir (Path): Absolute path to directory containing record of placed trades.
"""
self.config = Config()
self.mt5 = MetaTrader()
self.mt5 = MetaTrader() if self.config.mode == 'live' else MetaTester()
self.records_dir = records_dir or self.config.records_dir
async def get_records(self):
+11 -1
View File
@@ -2,8 +2,9 @@ import csv
import json
from logging import getLogger
from typing import Iterable, Literal
from asyncio import Lock
from .core import Config
from .core.config import Config
from .core.models import OrderSendResult
logger = getLogger(__name__)
@@ -31,6 +32,7 @@ class Result:
self.parameters = parameters or {}
self.result = result
self.name = name or parameters.get('name', 'Trades')
self.lock = Lock()
def get_data(self) -> dict:
res = self.result.get_dict(exclude={'retcode', 'comment', 'retcode_external', 'request_id', 'request'})
@@ -50,6 +52,7 @@ class Result:
async def to_csv(self):
"""Record trade results and associated parameters as a csv file
"""
await self.lock.acquire()
try:
data = self.get_data()
file = self.config.records_dir / f"{self.name}.csv"
@@ -67,6 +70,9 @@ class Result:
except Exception as err:
logger.error(f'Unable to save to csv: {err}')
finally:
self.lock.release()
@staticmethod
def serialize(value) -> str:
"""Serialize the trade records and strategy parameters
@@ -79,6 +85,7 @@ class Result:
async def to_json(self):
"""Save trades and strategy parameters in a json file
"""
await self.lock.acquire()
try:
file = self.config.records_dir / f"{self.name}.json"
data = self.get_data()
@@ -92,3 +99,6 @@ class Result:
json.dump(rows, fh, indent=2, skipkeys=True, default=self.serialize)
except Exception as err:
logger.error(f"Unable to save as json file: {err}")
finally:
self.lock.release()
+13 -1
View File
@@ -5,6 +5,8 @@ from typing import Literal, Callable
from logging import getLogger
from .positions import Positions
from .core.config import Config
from.contrib.backtester.event_manager import EventManager
logger = getLogger(__name__)
@@ -18,6 +20,15 @@ def delta(obj: time) -> timedelta:
return timedelta(hours=obj.hour, minutes=obj.minute, seconds=obj.second, microseconds=obj.microsecond)
async def backtest_sleep(secs):
"""A custom function to call when the session starts."""
em = EventManager()
async with em.condition:
while em.config.test_data.cursor.time < (em.config.test_data.cursor.time + secs):
await em.condition.wait()
class Session:
"""A session is a time period between two datetime.time objects specified in utc.
@@ -202,6 +213,7 @@ class Sessions:
current_session = self.find_next(now)
secs = current_session.until() + 10
logger.info(f'sleeping for {secs} seconds until next {current_session} session')
await sleep(secs)
sleep_func = sleep if Config().mode == 'live' else backtest_sleep
await sleep_func(secs)
self.current_session = current_session
await self.current_session.begin()
+3 -2
View File
@@ -8,6 +8,7 @@ from datetime import time as dtime
from .core.meta_trader import MetaTrader
from .symbol import Symbol as _Symbol
from .core import Config
from .contrib.backtester.meta_tester import MetaTester
from .sessions import Sessions, Session
Symbol = TypeVar("Symbol", bound=_Symbol)
@@ -47,7 +48,7 @@ class Strategy(ABC):
self.parameters["name"] = self.name
self.sessions = sessions or Sessions(Session(start=0, end=dtime(hour=23, minute=59, second=59)))
self.config = Config()
self.mt5 = MetaTrader()
self.mt5 = MetaTrader() if self.config.mode == 'live' else MetaTester()
def __repr__(self):
return f"{self.name}({self.symbol!r})"
@@ -79,4 +80,4 @@ class Strategy(ABC):
async def trade(self):
"""Place trades using this method. This is the main method of the strategy.
It will be called by the strategy runner.
"""
"""
+26 -56
View File
@@ -9,7 +9,7 @@ from .ticks import Tick
from .account import Account
from .candle import Candles
from .ticks import Ticks
from .utils import round_off
from .utils import round_off, backoff_decorator
logger = getLogger(__name__)
@@ -48,6 +48,7 @@ class Symbol(SymbolInfo):
"""
return self.point * 10
@backoff_decorator
async def info_tick(self, *, name: str = "") -> Tick:
"""Get the current price tick of a financial instrument.
@@ -65,7 +66,7 @@ class Symbol(SymbolInfo):
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}. {self.mt5.error}')
raise ValueError(f'Could not get tick for {name or self.name}.')
async def symbol_select(self, *, enable: bool = True) -> bool:
"""Select a symbol in the MarketWatch window or remove a symbol from the window.
@@ -81,7 +82,8 @@ class Symbol(SymbolInfo):
self.select = await self.mt5.symbol_select(self.name, enable)
return self.select
async def info(self, retries=3) -> SymbolInfo:
@backoff_decorator
async def info(self) -> SymbolInfo:
"""Get data on the specified financial instrument and update the symbol object properties
Returns:
@@ -90,18 +92,13 @@ class Symbol(SymbolInfo):
Raises:
ValueError: If request was unsuccessful and None was returned
"""
if retries < 1:
raise ValueError(f'Could not get info for {self.name}. {self.mt5.error}')
info = await self.mt5.symbol_info(self.name)
if info:
info = info._asdict()
info['swap_rollover3days'] = info.get('swap_rollover3days', 0) % 7
self.set_attributes(**info)
return SymbolInfo(**info)
if self.mt5.error.is_connection_error():
await asyncio.sleep(retries)
return await self.info(retries=retries - 1)
raise ValueError(f'Could not get info for {self.name}. {self.mt5.error}')
raise ValueError(f'Could not get info for {self.name}')
async def init(self) -> bool:
"""Initialized the symbol by pulling properties from the terminal
@@ -131,7 +128,8 @@ class Symbol(SymbolInfo):
"""
return await self.mt5.market_book_add(self.name)
async def book_get(self, retries=3) -> tuple[BookInfo, ...]:
@backoff_decorator
async def book_get(self) -> tuple[BookInfo, ...]:
"""Returns a tuple of BookInfo featuring Market Depth entries for the specified symbol.
Returns:
@@ -140,16 +138,13 @@ class Symbol(SymbolInfo):
Raises:
ValueError: If request was unsuccessful and None was returned
"""
if retries < 1:
raise ValueError(f'Could not get book info for {self.name}. {self.mt5.error}')
infos = await self.mt5.market_book_get(self.name)
if infos is not None:
book_infos = (BookInfo(**info._asdict()) for info in infos)
return tuple(book_infos)
if self.mt5.error.is_connection_error():
await asyncio.sleep(retries)
return await self.book_get(retries=retries - 1)
raise ValueError(f'Could not get book info for {self.name}. {self.mt5.error}')
raise ValueError(f'Could not get book info for {self.name}')
async def book_release(self) -> bool:
"""Cancels subscription of the MetaTrader 5 terminal to the Market Depth change events for a specified symbol.
@@ -241,8 +236,9 @@ class Symbol(SymbolInfo):
else:
logger.warning(f'Currency conversion failed: Unable to convert {amount} in {quote} to {base}')
@backoff_decorator
async def copy_rates_from(self, *, timeframe: TimeFrame,
date_from: datetime | int, count: int = 500, retries=3) -> Candles:
date_from: datetime | int, count: int = 500) -> Candles:
"""
Get bars from the MetaTrader 5 terminal starting from the specified date.
@@ -260,19 +256,14 @@ class Symbol(SymbolInfo):
Raises:
ValueError: If request was unsuccessful and None was returned
"""
if retries < 1:
raise ValueError(f'Could not get rates for {self.name}. {self.mt5.error}')
rates = await self.mt5.copy_rates_from(self.name, timeframe, date_from, count)
if rates is not None:
return Candles(data=rates)
if self.mt5.error.is_connection_error():
await asyncio.sleep(retries)
return await self.copy_rates_from(timeframe=timeframe, date_from=date_from,
count=count, retries=retries - 1)
raise ValueError(f'Could not get rates for {self.name}. {self.mt5.error}')
raise ValueError(f'Could not get rates for {self.name}.')
@backoff_decorator
async def copy_rates_from_pos(self, *, timeframe: TimeFrame, count: int = 500,
start_position: int = 0, retries=3) -> Candles:
start_position: int = 0) -> Candles:
"""Get bars from the MetaTrader 5 terminal starting from the specified index.
Args:
@@ -289,19 +280,14 @@ class Symbol(SymbolInfo):
Raises:
ValueError: If request was unsuccessful and None was returned
"""
if retries < 1:
raise ValueError(f'Could not get rates for {self.name}. {self.mt5.error}')
rates = await self.mt5.copy_rates_from_pos(self.name, timeframe, start_position, count)
if rates is not None:
return Candles(data=rates)
if self.mt5.error.is_connection_error():
await asyncio.sleep(retries)
return await self.copy_rates_from_pos(timeframe=timeframe, count=count,
start_position=start_position, retries=retries - 1)
raise ValueError(f'Could not get rates for {self.name}. {self.mt5.error}')
raise ValueError(f'Could not get rates for {self.name}.')
@backoff_decorator
async def copy_rates_range(self, *, timeframe: TimeFrame, date_from: datetime | int,
date_to: datetime | int, retries=3) -> Candles:
date_to: datetime | int) -> Candles:
"""Get bars in the specified date range from the MetaTrader 5 terminal.
Args:
@@ -321,21 +307,15 @@ class Symbol(SymbolInfo):
Raises:
ValueError: If request was unsuccessful and None was returned
"""
if retries < 1:
raise ValueError(f'Could not get rates for {self.name}. {self.mt5.error}')
rates = await self.mt5.copy_rates_range(symbol=self.name, timeframe=timeframe, date_from=date_from,
date_to=date_to)
if rates is not None:
return Candles(data=rates)
if self.mt5.error.is_connection_error():
await asyncio.sleep(retries)
return await self.copy_rates_range(timeframe=timeframe, date_from=date_from,
date_to=date_to, retries=retries - 1)
raise ValueError(f'Could not get rates for {self.name}. {self.mt5.error}')
raise ValueError(f'Could not get rates for {self.name}.')
@backoff_decorator
async def copy_ticks_from(self, *, date_from: datetime | int, count: int = 100,
flags: CopyTicks = CopyTicks.ALL, retries=3) -> Ticks:
flags: CopyTicks = CopyTicks.ALL) -> Ticks:
"""
Get ticks from the MetaTrader 5 terminal starting from the specified date.
@@ -352,19 +332,14 @@ class Symbol(SymbolInfo):
Raises:
ValueError: If request was unsuccessful and None was returned
"""
if retries < 1:
raise ValueError(f'Could not get ticks for {self.name}. {self.mt5.error}')
ticks = await self.mt5.copy_ticks_from(self.name, date_from, count, flags)
if ticks is not None:
return Ticks(data=ticks)
if self.mt5.error.is_connection_error():
await asyncio.sleep(retries)
return await self.copy_ticks_from(date_from=date_from, count=count, flags=flags, retries=retries - 1)
raise ValueError(f'Could not get ticks for {self.name}. {self.mt5.error}')
raise ValueError(f'Could not get ticks for {self.name}.')
@backoff_decorator
async def copy_ticks_range(self, *, date_from: datetime | int, date_to: datetime | int,
flags: CopyTicks = CopyTicks.ALL, retries=3) -> Ticks:
flags: CopyTicks = CopyTicks.ALL) -> Ticks:
"""Get ticks for the specified date range from the MetaTrader 5 terminal.
Args:
@@ -383,12 +358,7 @@ class Symbol(SymbolInfo):
Raises:
ValueError: If request was unsuccessful and None was returned.
"""
if retries < 1:
raise ValueError(f'Could not get ticks for {self.name}. {self.mt5.error}')
ticks = await self.mt5.copy_ticks_range(self.name, date_from, date_to, flags)
if ticks is not None:
return Ticks(data=ticks)
if self.mt5.error.is_connection_error():
await asyncio.sleep(retries)
return await self.copy_ticks_range(date_from=date_from, date_to=date_to, flags=flags, retries=retries - 1)
raise ValueError(f'Could not get ticks for {self.name}. {self.mt5.error}')
raise ValueError(f'Could not get ticks for {self.name}.')
+1 -2
View File
@@ -1,7 +1,7 @@
"""Terminal related functions and properties"""
from typing import NamedTuple
from logging import getLogger
from .core.models import TerminalInfo
logger = getLogger(__name__)
@@ -59,7 +59,6 @@ class Terminal(TerminalInfo):
"""
info = await self.mt5.terminal_info()
self.set_attributes(**info._asdict())
return self
async def symbols_total(self) -> int:
"""Get the number of all financial instruments in the MetaTrader 5 terminal.
+5 -3
View File
@@ -7,7 +7,9 @@ import csv
import logging
from typing import Iterable
from .core import Config, MetaTrader
from .core.config import Config
from .core.meta_trader import MetaTrader
from .contrib.backtester.meta_tester import MetaTester
logger = logging.getLogger(__name__)
@@ -21,7 +23,7 @@ class TradeRecords:
from the config
"""
config: Config
mt5: MetaTrader
mt5: MetaTrader | MetaTester
def __init__(self, *, records_dir: Path | str = ''):
"""Initialize the Records class. The main method of this class is update_records which you should call to update
@@ -31,7 +33,7 @@ class TradeRecords:
records_dir (Path): Absolute path to directory containing record of placed trades.
"""
self.config = Config()
self.mt5 = MetaTrader()
self.mt5 = MetaTrader() if self.config.mode == 'live' else MetaTester()
self.records_dir = records_dir or self.config.records_dir
async def get_csv_records(self):
+2 -1
View File
@@ -12,6 +12,7 @@ from .ram import RAM
from .core.models import OrderType, OrderSendResult
from .core.config import Config
from .result import Result
from .core.task_queue import QueueItem
logger = getLogger(__name__)
Symbol = TypeVar("Symbol", bound=_Symbol)
@@ -121,7 +122,7 @@ class Trader(ABC):
params["date"] = str(date.date())
params["time"] = str(date.time())
res = Result(result=result, parameters=params, name=name)
self.config.task_queue.add_task(res.save)
self.config.task_queue.add(item=QueueItem(res.save, must_complete=True))
@abstractmethod
async def place_trade(self, *args, **kwargs):