This commit is contained in:
Ichinga Samuel
2024-09-09 17:58:23 +01:00
parent bda2dbf308
commit 8928bcaab8
2 changed files with 132 additions and 66 deletions
+39 -43
View File
@@ -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)
+93 -23
View File
@@ -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()