Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@ classifiers = [
]
dependencies = [
"redis>=8.0.0,<9",
"taskiq>=0.12.0",
"taskiq>=0.13.0",
]

[project.urls]
Expand All @@ -42,8 +42,8 @@ dev = [
]
lint = [
"black>=25.11.0",
"mypy>=2.0.0",
"ruff>=0.14.7",
"mypy>=2.3.1",
"ruff>=0.16.9",
]
test = [
"fakeredis>=2.32.1",
Expand Down Expand Up @@ -112,6 +112,7 @@ lint.select = [
"RUF", # Specific to Ruff checks
"FA102", # Future annotations
"UP", # Pyupgrade
"LOG", # Logging
]
lint.ignore = [
"D105", # Missing docstring in magic method
Expand All @@ -123,6 +124,7 @@ lint.ignore = [
"ANN401", # typing.Any are disallowed in `**kwargs
"PLR0913", # Too many arguments for function call
"D106", # Missing docstring in public nested class
"PLR0917", # Too many positional arguments
]
exclude = [".venv/"]
line-length = 88
Expand Down
10 changes: 4 additions & 6 deletions taskiq_redis/list_schedule_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,6 @@
from redis.asyncio import BlockingConnectionPool, Redis
from taskiq import ScheduledTask, ScheduleSource
from taskiq.abc.serializer import TaskiqSerializer
from taskiq.compat import model_dump, model_validate
from taskiq.serializers import PickleSerializer
from typing_extensions import Self

Expand Down Expand Up @@ -64,7 +63,7 @@ async def startup(self) -> None:
logger.info("Migrating schedules from previous source")
await self._previous_schedule_source.startup()
schedules = await self._previous_schedule_source.get_schedules()
logger.info(f"Found {len(schedules)}")
logger.info("Found %d", len(schedules))
for schedule in schedules:
await self.add_schedule(schedule)
if self._delete_schedules_after_migration:
Expand Down Expand Up @@ -143,8 +142,7 @@ async def delete_schedule(self, schedule_id: str) -> None:
raw_schedule = await redis.getdel(self._get_data_key(schedule_id))
if raw_schedule is not None:
logger.debug("Deleting schedule %s", schedule_id)
schedule = model_validate(
ScheduledTask,
schedule = ScheduledTask.model_validate(
self._serializer.loadb(raw_schedule), # type: ignore[arg-type]
)
# We need to remove the schedule from the cron or time list.
Expand All @@ -162,7 +160,7 @@ async def add_schedule(self, schedule: "ScheduledTask") -> None:
# At first we set data key which contains the schedule data.
await redis.set(
f"{self._prefix}:data:{schedule.schedule_id}",
self._serializer.dumpb(model_dump(schedule)),
self._serializer.dumpb(schedule.model_dump(mode="json")),
)
# Then we add the schedule to the cron or time list.
# This is an optimization, so we can get all the schedules
Expand Down Expand Up @@ -229,7 +227,7 @@ async def get_schedules(self) -> list["ScheduledTask"]:
buffer = buffer[self._buffer_size :]

return [
model_validate(ScheduledTask, self._serializer.loadb(schedule)) # type: ignore[arg-type]
ScheduledTask.model_validate(self._serializer.loadb(schedule)) # type: ignore[arg-type]
for schedule in schedules
if schedule
]
Expand Down
33 changes: 14 additions & 19 deletions taskiq_redis/redis_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,6 @@
from redis.asyncio.connection import Connection
from taskiq import AsyncResultBackend
from taskiq.abc.serializer import TaskiqSerializer
from taskiq.compat import model_dump, model_validate
from taskiq.depends.progress_tracker import TaskProgress
from taskiq.result import TaskiqResult
from taskiq.serializers import PickleSerializer
Expand Down Expand Up @@ -107,7 +106,7 @@ async def set_result(
:param result: TaskiqResult instance.
"""
name = self._task_name(task_id)
value = self.serializer.dumpb(model_dump(result))
value = self.serializer.dumpb(result.model_dump(mode="json"))
async with Redis(connection_pool=self.redis_pool) as redis:
if self.result_ex_time:
await redis.set(name=name, value=value, ex=self.result_ex_time)
Expand Down Expand Up @@ -154,8 +153,7 @@ async def get_result(
if result_value is None:
raise ResultIsMissingError

taskiq_result = model_validate(
TaskiqResult[_ReturnType],
taskiq_result = TaskiqResult[_ReturnType].model_validate(
self.serializer.loadb(result_value), # type: ignore[arg-type]
)

Expand All @@ -179,7 +177,7 @@ async def set_progress(
:param result: task's TaskProgress instance.
"""
name = self._task_name(task_id) + PROGRESS_KEY_SUFFIX
value = self.serializer.dumpb(model_dump(progress))
value = self.serializer.dumpb(progress.model_dump(mode="json"))
async with Redis(connection_pool=self.redis_pool) as redis:
if self.result_ex_time:
await redis.set(name=name, value=value, ex=self.result_ex_time)
Expand All @@ -206,8 +204,7 @@ async def get_progress(
if result_value is None:
return None

return model_validate(
TaskProgress[_ReturnType],
return TaskProgress[_ReturnType].model_validate(
self.serializer.loadb(result_value), # type: ignore[arg-type]
)

Expand Down Expand Up @@ -286,7 +283,7 @@ async def set_result(
:param result: TaskiqResult instance.
"""
name = self._task_name(task_id)
value = self.serializer.dumpb(model_dump(result))
value = self.serializer.dumpb(result.model_dump(mode="json"))
async with self.redis as redis:
if self.result_ex_time:
await redis.set(name=name, value=value, ex=self.result_ex_time)
Expand Down Expand Up @@ -331,8 +328,9 @@ async def get_result(
if result_value is None:
raise ResultIsMissingError

taskiq_result: TaskiqResult[_ReturnType] = model_validate(
TaskiqResult[_ReturnType],
taskiq_result: TaskiqResult[_ReturnType] = TaskiqResult[
_ReturnType
].model_validate(
self.serializer.loadb(result_value), # type: ignore[arg-type]
)

Expand All @@ -356,7 +354,7 @@ async def set_progress(
:param result: task's TaskProgress instance.
"""
name = self._task_name(task_id) + PROGRESS_KEY_SUFFIX
value = self.serializer.dumpb(model_dump(progress))
value = self.serializer.dumpb(progress.model_dump(mode="json"))
async with self.redis as redis:
if self.result_ex_time:
await redis.set(name=name, value=value, ex=self.result_ex_time)
Expand All @@ -382,8 +380,7 @@ async def get_progress(
if result_value is None:
return None

return model_validate(
TaskProgress[_ReturnType],
return TaskProgress[_ReturnType].model_validate(
self.serializer.loadb(result_value), # type: ignore[arg-type]
)

Expand Down Expand Up @@ -470,7 +467,7 @@ async def set_result(
:param result: TaskiqResult instance.
"""
name = self._task_name(task_id)
value = self.serializer.dumpb(model_dump(result))
value = self.serializer.dumpb(result.model_dump(mode="json"))
async with self._acquire_master_conn() as redis:
if self.result_ex_time:
await redis.set(name=name, value=value, ex=self.result_ex_time)
Expand Down Expand Up @@ -517,8 +514,7 @@ async def get_result(
if result_value is None:
raise ResultIsMissingError

taskiq_result = model_validate(
TaskiqResult[_ReturnType],
taskiq_result = TaskiqResult[_ReturnType].model_validate(
self.serializer.loadb(result_value), # type: ignore[arg-type]
)

Expand All @@ -542,7 +538,7 @@ async def set_progress(
:param result: task's TaskProgress instance.
"""
name = self._task_name(task_id) + PROGRESS_KEY_SUFFIX
value = self.serializer.dumpb(model_dump(progress))
value = self.serializer.dumpb(progress.model_dump(mode="json"))
async with self._acquire_master_conn() as redis:
if self.result_ex_time:
await redis.set(name=name, value=value, ex=self.result_ex_time)
Expand All @@ -569,8 +565,7 @@ async def get_progress(
if result_value is None:
return None

return model_validate(
TaskProgress[_ReturnType],
return TaskProgress[_ReturnType].model_validate(
self.serializer.loadb(result_value), # type: ignore[arg-type]
)

Expand Down
17 changes: 7 additions & 10 deletions taskiq_redis/schedule_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,6 @@
)
from taskiq import ScheduleSource
from taskiq.abc.serializer import TaskiqSerializer
from taskiq.compat import model_dump, model_validate
from taskiq.scheduler.scheduled_task import ScheduledTask
from taskiq.serializers import PickleSerializer

Expand Down Expand Up @@ -50,8 +49,7 @@ def __init__(
**connection_kwargs: Any,
) -> None:
warnings.warn(
"RedisScheduleSource is deprecated. "
"Please switch to ListRedisScheduleSource",
"RedisScheduleSource is deprecated. Please switch to ListRedisScheduleSource", # noqa: E501
DeprecationWarning,
stacklevel=2,
)
Expand Down Expand Up @@ -81,7 +79,7 @@ async def add_schedule(self, schedule: ScheduledTask) -> None:
async with Redis(connection_pool=self.connection_pool) as redis:
await redis.set(
f"{self.prefix}:{schedule.schedule_id}",
self.serializer.dumpb(model_dump(schedule)),
self.serializer.dumpb(schedule.model_dump(mode="json")),
)

async def get_schedules(self) -> list[ScheduledTask]:
Expand All @@ -103,7 +101,7 @@ async def get_schedules(self) -> list[ScheduledTask]:
if buffer:
schedules.extend(await redis.mget(buffer))
return [
model_validate(ScheduledTask, self.serializer.loadb(schedule)) # type: ignore[arg-type]
ScheduledTask.model_validate(self.serializer.loadb(schedule)) # type: ignore[arg-type]
for schedule in schedules
if schedule
]
Expand Down Expand Up @@ -163,7 +161,7 @@ async def add_schedule(self, schedule: ScheduledTask) -> None:
"""
await self.redis.set(
f"{self.prefix}:{schedule.schedule_id}",
self.serializer.dumpb(model_dump(schedule)),
self.serializer.dumpb(schedule.model_dump(mode="json")),
)

async def get_schedules(self) -> list[ScheduledTask]:
Expand All @@ -177,8 +175,7 @@ async def get_schedules(self) -> list[ScheduledTask]:
schedules = []
async for key in self.redis.scan_iter(f"{self.prefix}:*"):
raw_schedule = await self.redis.get(key)
parsed_schedule = model_validate(
ScheduledTask,
parsed_schedule = ScheduledTask.model_validate(
self.serializer.loadb(raw_schedule), # type: ignore[arg-type]
)
schedules.append(parsed_schedule)
Expand Down Expand Up @@ -255,7 +252,7 @@ async def add_schedule(self, schedule: ScheduledTask) -> None:
async with self._acquire_master_conn() as redis:
await redis.set(
f"{self.prefix}:{schedule.schedule_id}",
self.serializer.dumpb(model_dump(schedule)),
self.serializer.dumpb(schedule.model_dump(mode="json")),
)

async def get_schedules(self) -> list[ScheduledTask]:
Expand All @@ -277,7 +274,7 @@ async def get_schedules(self) -> list[ScheduledTask]:
if buffer:
schedules.extend(await redis.mget(buffer))
return [
model_validate(ScheduledTask, self.serializer.loadb(schedule)) # type: ignore[arg-type]
ScheduledTask.model_validate(self.serializer.loadb(schedule)) # type: ignore[arg-type]
for schedule in schedules
if schedule
]
Expand Down
Loading
Loading