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