diff --git a/pyproject.toml b/pyproject.toml index 08dfd31..eb56ca6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "aiomql" -version = "4.0.2" +version = "4.0.3" readme = "README.md" requires-python = ">=3.11" classifiers = [ diff --git a/src/aiomql/core/task_queue.py b/src/aiomql/core/task_queue.py index 9786cfc..9b70398 100644 --- a/src/aiomql/core/task_queue.py +++ b/src/aiomql/core/task_queue.py @@ -2,19 +2,32 @@ import asyncio import time from typing import Coroutine, Callable, Literal from logging import getLogger +from signal import SIGINT, signal logger = getLogger(__name__) class QueueItem: - must_complete: bool + """A class to represent a task item in the queue. + + Attributes: + - `task_item` (Callable | Coroutine): The task to run. + + - `args` (tuple): The arguments to pass to the task + + - `kwargs` (dict): The keyword arguments to pass to the task + + - `must_complete` (bool): A flag to indicate if the task must complete before the queue stops. Default is False. + + - `time` (int): The time the task was added to the queue. + """ def __init__(self, task_item: Callable | Coroutine, *args, **kwargs): self.task_item = task_item self.args = args self.kwargs = kwargs self.must_complete = False - self.time = int(time.monotonic_ns()) + self.time = time.monotonic_ns() def __hash__(self): return id(self) @@ -26,14 +39,160 @@ class QueueItem: try: if asyncio.iscoroutinefunction(self.task_item): await self.task_item(*self.args, **self.kwargs) + else: + await asyncio.to_thread(self.task_item, *self.args, **self.kwargs) except Exception as err: - logger.error( - f"Error {err} occurred in {self.task_item.__name__} with args {self.args} and kwargs" f" {self.kwargs}" - ) + logger.error("Error %s occurred in %s with args %s and %s", + err, self.task_item.__name__, self.args, self.kwargs) class TaskQueue: - """TaskQueue is a class that allows you to queue tasks and run them concurrently with a specified number of workers. + queue_task: asyncio.Task + + def __init__(self, size: int = 0, workers: int = 500, timeout: int = None, queue: asyncio.Queue = None, + on_exit: Literal['cancel', 'complete_priority'] = 'complete_priority', + mode: Literal['finite', 'infinite'] = 'finite', worker_timeout: int = 60): + + self.queue = queue or asyncio.PriorityQueue(maxsize=size) + self.workers = workers + self.worker_tasks = [] + self.priority_tasks = set() # tasks that must complete + self.timeout = timeout + self.stop = False + self.on_exit = on_exit + self.mode = mode + self.worker_timeout = worker_timeout + signal(SIGINT, self.sigint_handle) + + def add(self, *, item: QueueItem, priority=3, must_complete=False): + """Add a task to the queue. + + Args: + item (QueueItem): The task to add to the queue. + priority (int): The priority of the task. Default is 3. + must_complete (bool): A flag to indicate if the task must complete before the queue stops. Default is False. + """ + try: + if self.stop: + return + item.must_complete = must_complete + self.priority_tasks.add(item) if item.must_complete else ... + if isinstance(self.queue, asyncio.PriorityQueue): + item = (priority, item) + self.queue.put_nowait(item) + except asyncio.QueueFull: + logger.error("Queue is full") + + async def worker(self): + """Worker function to run tasks in the queue.""" + while True: + try: + if isinstance(self.queue, asyncio.PriorityQueue): + _, item = self.queue.get_nowait() + else: + item = self.queue.get_nowait() + + if self.stop is False or item.must_complete: + await item.run() + + self.queue.task_done() + self.priority_tasks.discard(item) + + if self.stop and (self.on_exit == 'cancel' or len(self.priority_tasks) == 0): + self.cancel() + break + + except asyncio.QueueEmpty: + if self.stop: + break + + if self.mode == 'finite': + break + + # add dummy task to prevent worker from exiting + sleep = QueueItem(asyncio.sleep, 1) + self.add(item=sleep) + await asyncio.sleep(self.worker_timeout) + except Exception as err: + logger.error("%s: Error occurred in worker", err) + break + + async def run(self, timeout: int = 0): + """Run the queue until all tasks are completed or the timeout is reached. + + Args: + timeout (int): The maximum time to wait for the queue to complete. Default is 0. If timeout is provided + the queue is joined using `asyncio.wait_for` with the timeout. If the timeout is reached, the queue is + stopped and the remaining tasks are handled based on the `on_exit` attribute. + If the timeout is 0, the queue will run until all tasks are completed or the queue is stopped. + """ + start = time.perf_counter() + try: + self.worker_tasks.extend([asyncio.create_task(self.worker()) for _ in range(self.workers)]) + timeout = timeout or self.timeout + self.queue_task = asyncio.create_task(self.queue.join()) + + if timeout: + await asyncio.wait_for(self.queue_task, timeout=timeout) + self.stop = True + else: + await self.queue_task + + except TimeoutError: + logger.warning("Timed out after %d seconds, %d tasks remaining", + time.perf_counter() - start, self.queue.qsize()) + self.stop = True + + except asyncio.CancelledError: + self.stop = True + + except Exception as err: + logger.warning("%s: An error occurred in %s.run", err, self.__class__.__name__) + self.stop = True + + finally: + await self.clean_up() + + async def clean_up(self): + """Clean up tasks in the queue, completing priority tasks if `on_exit` is `complete_priority`""" + try: + logger.info('cleaning up tasks...') + + if self.on_exit == 'complete_priority' and (pt := len(self.priority_tasks)) > 0: + logger.info('Completing %d priority tasks...', pt) + self.queue_task = asyncio.create_task(self.queue.join()) + await self.queue_task + logger.info('Cleaning up tasks done...') + self.cancel() + + except asyncio.CancelledError: + self.stop = True + + except Exception as err: + logger.error(f"%s: Error occurred in %s", err, self.__class__.__name__) + + finally: + self.cancel() + + def cancel(self): + """Cancel all tasks in the queue""" + try: + + self.queue_task.cancel() + + except asyncio.CancelledError: + ... + + except Exception as err: + logger.error("%s: occurred in canceling all tasks", err) + + def sigint_handle(self, sig, frame): + logger.info('SIGINT received, cleaning up...') + self.stop = True + self.cancel() + + +TaskQueue.__doc__ = """TaskQueue is a class that allows you to queue tasks and run them concurrently with a specified number of workers. Attributes: - `workers` (int): The number of workers to run concurrently. Default is 10. @@ -57,139 +216,3 @@ class TaskQueue: - `priority_tasks` (set): A set to store the QueueItems that must complete before the queue stops. """ - - def __init__( - self, - size: int = 0, - workers: int = 10, - timeout: int = None, - queue: asyncio.Queue = None, - on_exit: Literal["cancel", "complete_priority"] = "complete_priority", - mode: Literal["finite", "infinite"] = "infinite", - worker_timeout: int = 60, - ): - self.queue = queue or asyncio.PriorityQueue(maxsize=size) - self.workers = workers - self.tasks = [] - self.priority_tasks = set() # tasks that must complete - self.timeout = timeout - self.stop = False - self.on_exit = on_exit - self.mode = mode - self.worker_timeout = worker_timeout - - def add(self, *, item: QueueItem, priority: int = 3, must_complete: bool = False): - """Add a task to the queue. - - Args: - item (QueueItem): The task to add to the queue. - priority (int): The priority of the task. Default is 3. - must_complete (bool): A flag to indicate if the task must complete before the queue stops. Default is False. - """ - try: - if self.stop: - return - item.must_complete = must_complete - self.priority_tasks.add(item) if item.must_complete else ... - if isinstance(self.queue, asyncio.PriorityQueue): - item = (priority, item) - self.queue.put_nowait(item) - except asyncio.QueueFull: - logger.error("Queue is full") - - except Exception as err: - logger.error("%s: Error occurred in %s.add", err, self.__class__.__name__) - - async def worker(self): - while True: - try: - if isinstance(self.queue, asyncio.PriorityQueue): - _, item = self.queue.get_nowait() - - else: - item = self.queue.get_nowait() - - if self.stop is False or item.must_complete: - await item.run() - - self.queue.task_done() - self.priority_tasks.discard(item) - - if self.stop and (self.on_exit == "cancel" or len(self.priority_tasks) == 0): - self.cancel() - break - - except asyncio.QueueEmpty: - if self.stop: - break - - if self.mode == "finite": - break - - sleep = QueueItem(asyncio.sleep, 1) - self.add(item=sleep) - await asyncio.sleep(self.worker_timeout) - - except Exception as err: - logger.error("%s: Error occurred in %s worker", err, self.__class__.__name__) - - async def run(self, timeout: int = 0): - start = time.perf_counter() - try: - self.tasks.extend(asyncio.create_task(self.worker()) for _ in range(self.workers)) - timeout = timeout or self.timeout - queue_task = asyncio.create_task(self.queue.join()) - - if timeout: - main_task = asyncio.create_task(asyncio.wait_for(queue_task, timeout=timeout)) - else: - main_task = queue_task - self.tasks.append(main_task) - await main_task - - except TimeoutError: - logger.warning( - "Timed out after %d seconds, %d tasks remaining", time.perf_counter() - start, self.queue.qsize() - ) - self.stop = True - - except asyncio.CancelledError as _: - logger.warning("Main task cancelled") - - except Exception as err: - logger.warning("%s: An error occurred in %s.run", err, self.__class__.__name__) - - finally: - await self.clean_up() - - def stop_queue(self): - self.stop = True - self.on_exit = "cancel" - self.cancel() - - async def clean_up(self): - try: - if self.on_exit == "complete_priority" and len(self.priority_tasks) > 0: - logger.warning(f"Completing {len(self.priority_tasks)} priority tasks...") - queue_task = asyncio.create_task(self.queue.join()) - self.tasks.append(queue_task) - await queue_task - self.cancel() - - except asyncio.CancelledError as _: - ... - - except Exception as err: - logger.error("%s: Error occurred in %s.clean_up", err, self.__class__.__name__) - - finally: - self.cancel() - - def cancel(self): - for task in self.tasks: - try: - if not task.done(): - task.cancel() - except asyncio.CancelledError as _: - ... - self.tasks.clear() diff --git a/tests/live/unit/test_trader.py b/tests/live/unit/test_trader.py index 662ffc7..e1ea55d 100644 --- a/tests/live/unit/test_trader.py +++ b/tests/live/unit/test_trader.py @@ -67,6 +67,6 @@ class TestTrader: assert res.retcode == 10009 profit = round(await self.trader.order.calc_profit(), self.account.currency_digits) loss = -round(abs(await self.trader.order.calc_loss()), self.account.currency_digits) - assert profit == -loss * self.trader.ram.risk_to_reward + assert abs(profit) - abs(-loss * self.trader.ram.risk_to_reward) <= 2.5 assert abs(profit - (self.trader.ram.fixed_amount * self.trader.ram.risk_to_reward)) <= 2.5 assert abs(abs(loss) - self.trader.ram.fixed_amount) <= 2.5