diff --git a/src/aiomql/contrib/backtester/get_data.py b/src/aiomql/contrib/backtester/get_data.py index fedaefe..9af30fd 100644 --- a/src/aiomql/contrib/backtester/get_data.py +++ b/src/aiomql/contrib/backtester/get_data.py @@ -13,7 +13,8 @@ from pandas import DataFrame from ...core.meta_trader import MetaTrader from ...core.config import Config -from ...core.constants import TimeFrame, CopyTicks +from ...core.constants import TimeFrame, +from ...core.task_queue import TaskQueue, QueueItem from ...utils import backoff_decorator @@ -73,16 +74,19 @@ class GetData: self.span = range(start := int(self.start.timestamp()), diff + start) self.data = Data(name=name) self.mt5 = MetaTrader() + self.task_queue = TaskQueue() async def get_data(self): """""" - rates, ticks, prices, symbols, account = await asyncio.gather(self.get_symbols_rates(), self.get_symbols_ticks(), - self.get_symbols_prices(), self.get_symbols_info(), - self.get_account_info()) - terminal, version = await asyncio.gather(self.get_terminal_info(), self.get_version()) + 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) + # 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): """""" @@ -98,7 +102,7 @@ class GetData: fh.write(bdata) @classmethod - def dump_data(cls, data: Data, name: str | Path, compress: bool = False) -> None: + def dump_data(cls, data: Data, name: str | Path, compress: bool = False): """""" try: fo = open(name, 'wb') @@ -114,7 +118,7 @@ class GetData: @classmethod - def load_data(cls, *, name: str | Path, compressed=False) -> Data | None: + def load_data(cls, *, name: str | Path, compressed=False): """""" try: fo = open(name, 'rb') @@ -132,79 +136,71 @@ class GetData: logger.error(f"Error: {err}") return None - async def get_terminal_info(self) -> dict[str, [str | int | bool | float]]: + async def get_terminal_info(self): """""" terminal = await self.mt5.terminal_info() - return terminal._asdict() + terminal = terminal._asdict() + self.data.setattrs(terminal=terminal) - async def get_version(self) -> tuple[int, int, str]: + async def get_version(self): """""" version = await self.mt5.version() - return version + self.data.setattrs(version=version) - async def get_symbols_info(self) -> dict[str, dict]: + async def get_symbols_info(self): """""" - tasks = [self.get_symbol_info(symbol) for symbol in self.symbols] - res = await asyncio.gather(*tasks) - return {symbol: info for symbol, info in res} + [self.task_queue.add(QueueItem(self.get_symbol_info, symbol, must_complete=True), priority=4) for symbol in self.symbols] - async def get_symbols_ticks(self) -> dict[str, DataFrame]: + async def get_symbols_ticks(self): """""" - tasks = [self.get_symbol_ticks(symbol) for symbol in self.symbols] - res = await asyncio.gather(*tasks) - return {symbol: tick for symbol, tick in res} + [self.task_queue.add(QueueItem(self.get_symbol_ticks, symbol)) for symbol in self.symbols] - async def get_symbols_prices(self) -> dict[str, DataFrame]: + async def get_symbols_prices(self): """""" - tasks = [self.get_symbol_prices(symbol) for symbol in self.symbols] - res = await asyncio.gather(*tasks) - return {symbol: prices for symbol, prices in res} + [self.task_queue.add(QueueItem(self.get_symbol_prices, symbol)) for symbol in self.symbols] - async def get_symbols_rates(self) -> dict[str: dict[TimeFrame: DataFrame]]: + async def get_symbols_rates(self): """""" - tasks = [self.get_symbol_rates(symbol, timeframe) for symbol in self.symbols for timeframe in self.timeframes] - res = await asyncio.gather(*tasks) - data = {} - for symbol, timeframe, rates in res: - data.setdefault(symbol, {}).setdefault(timeframe.name, rates) - return data - + [self.task_queue.add(QueueItem(self.get_symbol_rates, symbol, timeframe, must_complete=True), priority=4) + for symbol in self.symbols for timeframe in self.timeframes] + @backoff_decorator(max_retries=5) - async def get_account_info(self) -> dict: + async def get_account_info(self): """""" res = await self.mt5.account_info() - return res._asdict() + res = res._asdict() + self.data.setattrs(account=res) @backoff_decorator(max_retries=5) - async def get_symbol_info(self, symbol: str) -> tuple[str, dict]: + async def get_symbol_info(self, symbol: str): """""" res = await self.mt5.symbol_info(symbol) - return symbol, res._asdict() + self.data.symbols[symbol] = res._asdict() @backoff_decorator(max_retries=5) - async def get_symbol_ticks(self, symbol: str) -> tuple[str, DataFrame]: + async def get_symbol_ticks(self, symbol: str): """""" res = await self.mt5.copy_ticks_range(symbol, self.start, self.end, CopyTicks.ALL) res = pd.DataFrame(res) res.drop_duplicates(subset=['time'], keep='last', inplace=True) res.set_index('time', inplace=True, drop=False) - return symbol, res + self.data.ticks[symbol] = res @backoff_decorator(max_retries=5) - async def get_symbol_prices(self, symbol: str) -> tuple[str, DataFrame]: + async def get_symbol_prices(self, symbol: str): """""" res = await self.mt5.copy_ticks_range(symbol, self.start, self.end, CopyTicks.ALL) res = pd.DataFrame(res) res.drop_duplicates(subset=['time'], keep='last', inplace=True) res.set_index('time', inplace=True, drop=False) res = res.reindex(self.span, method='nearest') - return symbol, res + self.data.prices[symbol] = res @backoff_decorator(max_retries=5) - async def get_symbol_rates(self, symbol: str, timeframe: TimeFrame) -> tuple[str, TimeFrame, DataFrame]: + async def get_symbol_rates(self, symbol: str, timeframe: TimeFrame): """""" res = await self.mt5.copy_rates_range(symbol, timeframe, self.start, self.end) res = pd.DataFrame(res) res.drop_duplicates(subset=['time'], keep='last', inplace=True) res.set_index('time', inplace=True, drop=False) - return symbol, timeframe, res + self.data.rates[symbol].setdefault(timeframe, res) diff --git a/src/aiomql/core/task_queue.py b/src/aiomql/core/task_queue.py index 06043d8..e03c6c5 100644 --- a/src/aiomql/core/task_queue.py +++ b/src/aiomql/core/task_queue.py @@ -1,49 +1,119 @@ import asyncio -from typing import Coroutine, Callable, Awaitable +from typing import Coroutine, Callable, Literal +from signal import signal, SIGINT from logging import getLogger + logger = getLogger(__name__) class QueueItem: - def __init__(self, task: Callable | Awaitable | Coroutine, *args, **kwargs): - self.task = task + def __init__(self, task_item: Callable | Coroutine, must_complete=False, *args, **kwargs): + self.task_item = task_item self.args = args self.kwargs = kwargs + self.must_complete = must_complete + self.time = asyncio.get_event_loop().time() + + def __hash__(self): + return id(self) + + def __lt__(self, other): + return self.time < other.time async def run(self): try: - if asyncio.iscoroutinefunction(self.task): - return await self.task(*self.args, **self.kwargs) + if asyncio.iscoroutinefunction(self.task_item): + await self.task_item(*self.args, **self.kwargs) + else: - return self.task(*self.args, **self.kwargs) + self.task_item(*self.args, **self.kwargs) + except Exception as err: - logger.error(f"Error in running {getattr(self.task, '__name__', str(self.task))}" - f" with {str(self.args)}, {self.kwargs}: {err}") + logger.error(f"Error {err} occurred in {self.func.__name__} with args {self.args} and kwargs {self.kwargs}") class TaskQueue: - def __init__(self): - self.queue = asyncio.Queue() + def __init__(self, size: int = 0, workers: int = 50, timeout: int = None, queue: asyncio.Queue = None, + on_exit: Literal['cancel', 'complete_priority'] = 'complete_priority'): + + self.queue = queue or asyncio.PriorityQueue(maxsize=size) + self.workers = workers + self.tasks = [] + self.priority_tasks = set() # tasks that must complete + self.timeout = timeout + self.stop = False + self.on_exit = on_exit + signal(SIGINT, self.sigint_handle) - def add(self, item: QueueItem): + def add(self, *, item: QueueItem, priority=3): try: - self.queue.put_nowait(item) + if not self.stop: + if isinstance(self.queue, asyncio.PriorityQueue): + self.priority_tasks.add(item) if item.must_complete else ... + item = (priority, item) + self.queue.put_nowait(item) + except asyncio.QueueFull: - return + logger.error(f"Queue is full, could not add {item}") async def worker(self): while True: - try: - item: QueueItem = self.queue.get_nowait() + if isinstance(self.queue, asyncio.PriorityQueue): + _, item = await self.queue.get() + + else: + item = await self.queue.get() + + if not self.stop or item.must_complete: await item.run() - self.queue.task_done() - except asyncio.QueueEmpty: - return - def add_task(self, item: Callable | Awaitable | Coroutine, *args, **kwargs): - self.add(QueueItem(item, *args, **kwargs)) - asyncio.create_task(self.worker()) + self.queue.task_done() + self.priority_tasks.discard(item) - async def start(self): - await self.queue.join() + def sigint_handle(self, sig, frame): + print('SIGINT received, cleaning up...') + + if self.on_exit == 'complete_priority': + 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): + loop = asyncio.get_running_loop() + start = loop.time() + + 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) + + except TimeoutError: + print(f"Timed out after {loop.time() - start} seconds. {self.queue.qsize()} tasks remaining") + + if self.on_exit == 'complete_priority' and self.priority_tasks: + print(f'Completing {len(self.priority_tasks)} priority tasks...') + self.stop = True + await self.queue.join() + + else: + self.cancel() + + except asyncio.CancelledError: + print('Tasks cancelled') + + finally: + print(f'Exiting queue after {(loop.time() - start)} seconds.' + f'{self.queue.qsize()} tasks remaining, {len(self.priority_tasks)} are priority tasks') + + self.cancel() + + 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() \ No newline at end of file