diff --git a/celery-stubs/contrib/abortable.pyi b/celery-stubs/contrib/abortable.pyi index da61dbc..8735ae2 100644 --- a/celery-stubs/contrib/abortable.pyi +++ b/celery-stubs/contrib/abortable.pyi @@ -1,4 +1,4 @@ -from typing import Any, Generic, ParamSpec, TypeVar +from typing import Any, Generic, ParamSpec, TypeVar, overload from celery import Task from celery.result import AsyncResult @@ -20,4 +20,7 @@ class AbortableTask(Task[_P, _R_co], Generic[_P, _R_co]): @override def AsyncResult(self, task_id: str) -> AbortableAsyncResult[_R_co]: ... # type: ignore[override] # pyright: ignore[reportIncompatibleMethodOverride] + @overload def is_aborted(self, *, task_id: str, **kwargs: Any) -> bool: ... + @overload + def is_aborted(self, **kwargs: Any) -> bool: ... diff --git a/tests/test_celery.py b/tests/test_celery.py index 59a7791..4577047 100644 --- a/tests/test_celery.py +++ b/tests/test_celery.py @@ -17,6 +17,7 @@ if TYPE_CHECKING: from collections.abc import Iterator + from celery.contrib.abortable import AbortableTask from celery.contrib.django.task import DjangoTask app = celery.Celery() @@ -364,3 +365,7 @@ def test_celery_top_level_exports() -> None: def test_djangotask(task: DjangoTask[[int, int], Any]) -> None: task.delay_on_commit(1, 2) task.apply_async_on_commit((1, 2), countdown=10.0) + + +def test_abortabletask(task: AbortableTask[[], None]) -> None: + task.is_aborted()