Skip to content

Commit f96551b

Browse files
committed
fix(tasks): optimisation
1 parent afbc9b8 commit f96551b

2 files changed

Lines changed: 133 additions & 0 deletions

File tree

‎taskiq/_broker_resources.py‎

Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,38 @@
1+
"""Internal validation for resources owned by broker lifecycle coordinators."""
2+
3+
from collections.abc import Iterable
4+
from typing import TYPE_CHECKING
5+
6+
if TYPE_CHECKING: # pragma: no cover
7+
from taskiq.abc.broker import AsyncBroker
8+
9+
10+
def validate_unshared_broker_resources(
11+
brokers: Iterable["AsyncBroker"],
12+
*,
13+
error_type: type[ValueError] = ValueError,
14+
) -> None:
15+
"""Reject middleware and result backends owned by multiple brokers."""
16+
middleware_owners: dict[int, str] = {}
17+
backend_owners: dict[int, str] = {}
18+
for broker in brokers:
19+
for middleware in broker.middlewares:
20+
previous_owner = middleware_owners.setdefault(
21+
id(middleware),
22+
broker.broker_name,
23+
)
24+
if previous_owner != broker.broker_name:
25+
raise error_type(
26+
"One middleware instance cannot be lifecycle-owned by "
27+
f"brokers {previous_owner!r} and {broker.broker_name!r}.",
28+
)
29+
30+
previous_owner = backend_owners.setdefault(
31+
id(broker.result_backend),
32+
broker.broker_name,
33+
)
34+
if previous_owner != broker.broker_name:
35+
raise error_type(
36+
"One result backend instance cannot be lifecycle-owned by "
37+
f"brokers {previous_owner!r} and {broker.broker_name!r}.",
38+
)
Lines changed: 95 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,95 @@
1+
from datetime import datetime, timezone
2+
from typing import Any, Literal
3+
4+
import pytest
5+
6+
from taskiq.abc.schedule_source import ScheduleSource
7+
from taskiq.kicker import AsyncKicker
8+
from taskiq.scheduler.scheduled_task import ScheduledTask
9+
from taskiq.scheduler.scheduler import TaskiqScheduler
10+
from tests.utils import RecordingBroker
11+
12+
ScheduleKind = Literal["cron", "interval", "time"]
13+
14+
15+
class RecordingScheduleSource(ScheduleSource):
16+
"""Store schedules created by a kicker."""
17+
18+
def __init__(self) -> None:
19+
self.schedules: list[ScheduledTask] = []
20+
21+
async def get_schedules(self) -> list[ScheduledTask]:
22+
return self.schedules
23+
24+
async def add_schedule(self, schedule: ScheduledTask) -> None:
25+
self.schedules.append(schedule)
26+
27+
28+
async def create_schedule(
29+
kind: ScheduleKind,
30+
kicker: AsyncKicker[Any, Any],
31+
source: RecordingScheduleSource,
32+
) -> ScheduledTask:
33+
"""Create one schedule through the selected public kicker method."""
34+
if kind == "cron":
35+
created = await kicker.schedule_by_cron(source, "* * * * *")
36+
elif kind == "interval":
37+
created = await kicker.schedule_by_interval(source, 10)
38+
else:
39+
created = await kicker.schedule_by_time(
40+
source,
41+
datetime(2030, 1, 1, tzinfo=timezone.utc),
42+
)
43+
return created.task
44+
45+
46+
@pytest.mark.parametrize("kind", ["cron", "interval", "time"])
47+
async def test_all_schedule_methods_preserve_custom_task_id(
48+
kind: ScheduleKind,
49+
) -> None:
50+
broker = RecordingBroker()
51+
source = RecordingScheduleSource()
52+
kicker: AsyncKicker[Any, Any] = AsyncKicker("demo.task", broker, {})
53+
54+
scheduled = await create_schedule(
55+
kind,
56+
kicker.with_task_id("custom-task-id"),
57+
source,
58+
)
59+
60+
assert scheduled.task_id == "custom-task-id"
61+
assert source.schedules == [scheduled]
62+
63+
64+
async def test_interval_without_custom_task_id_keeps_deferred_generation() -> None:
65+
broker = RecordingBroker()
66+
source = RecordingScheduleSource()
67+
kicker: AsyncKicker[Any, Any] = AsyncKicker("demo.task", broker, {})
68+
69+
scheduled = await create_schedule("interval", kicker, source)
70+
71+
assert scheduled.task_id is None
72+
73+
74+
@pytest.mark.parametrize(
75+
("custom_task_id", "expected_task_id"),
76+
[
77+
pytest.param("custom-task-id", "custom-task-id", id="custom"),
78+
pytest.param(None, "generated-task-id", id="generated-at-dispatch"),
79+
],
80+
)
81+
async def test_scheduler_dispatch_uses_stored_or_generated_interval_task_id(
82+
custom_task_id: str | None,
83+
expected_task_id: str,
84+
) -> None:
85+
broker = RecordingBroker()
86+
broker.with_id_generator(lambda: "generated-task-id")
87+
source = RecordingScheduleSource()
88+
kicker: AsyncKicker[Any, Any] = AsyncKicker("demo.task", broker, {})
89+
kicker.with_task_id(custom_task_id)
90+
scheduled = await create_schedule("interval", kicker, source)
91+
scheduler = TaskiqScheduler(broker, [source])
92+
93+
await scheduler.on_ready(source, scheduled)
94+
95+
assert broker.sent[0][0].task_id == expected_task_id

0 commit comments

Comments
 (0)