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
3 changes: 0 additions & 3 deletions .github/workflows/test.yml
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,6 @@ jobs:
strategy:
matrix:
py_version: ["3.10", "3.11", "3.12", "3.13", "3.14"]
pydantic_ver: ["<2", ">=2.5,<3"]
os: [ubuntu-latest, windows-latest, macos-latest]
runs-on: "${{ matrix.os }}"
steps:
Expand All @@ -57,8 +56,6 @@ jobs:
version: "latest"
- name: Install deps
run: uv sync --all-extras
- name: Setup pydantic version
run: uv pip install "pydantic ${{ matrix.pydantic_ver }}"
- name: Run pytest check
run: uv run pytest -vv -n auto --cov="taskiq" .
- name: Generate report
Expand Down
3 changes: 1 addition & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -27,9 +27,8 @@ classifiers = [
dependencies = [
"aiohttp>=3",
"anyio>=4",
"packaging>=19",
"pycron>=3.0.0",
"pydantic>=1.0,<=3.0",
"pydantic>=2.5,<3.0",
"taskiq_dependencies>=1.3.1,<2",
"typing-extensions>=3.10.0.0; python_version < '3.11'"
]
Expand Down
75 changes: 6 additions & 69 deletions taskiq/compat.py
Original file line number Diff line number Diff line change
@@ -1,79 +1,16 @@
from collections.abc import Hashable
from functools import lru_cache
from importlib.metadata import version
from typing import Any, TypeVar

import pydantic
from packaging.version import Version, parse

PYDANTIC_VER = parse(version("pydantic"))
T = TypeVar("T", bound=Hashable)

Model = TypeVar("Model", bound="pydantic.BaseModel")
IS_PYDANTIC2 = Version("2.0") <= PYDANTIC_VER

if IS_PYDANTIC2:
T = TypeVar("T", bound=Hashable)
@lru_cache
def create_type_adapter(annot: type[T]) -> pydantic.TypeAdapter[T]:
return pydantic.TypeAdapter(annot)

@lru_cache
def create_type_adapter(annot: type[T]) -> pydantic.TypeAdapter[T]:
return pydantic.TypeAdapter(annot)

def parse_obj_as(annot: type[T], obj: Any) -> T:
return create_type_adapter(annot).validate_python(obj)

def model_validate(
model_class: type[Model],
message: dict[str, Any],
) -> Model:
return model_class.model_validate(message)

def model_dump(instance: Model) -> dict[str, Any]:
return instance.model_dump(mode="json")

def model_validate_json(
model_class: type[Model],
message: str | bytes | bytearray,
) -> Model:
return model_class.model_validate_json(message)

def model_dump_json(instance: Model) -> str:
return instance.model_dump_json()

def model_copy(
instance: Model,
update: dict[str, Any] | None = None,
deep: bool = False,
) -> Model:
return instance.model_copy(update=update, deep=deep)

validate_call = pydantic.validate_call

else:
parse_obj_as = pydantic.parse_obj_as # type: ignore

def model_validate(
model_class: type[Model],
message: dict[str, Any],
) -> Model:
return model_class.parse_obj(message)

def model_dump(instance: Model) -> dict[str, Any]:
return instance.dict()

def model_validate_json(
model_class: type[Model],
message: str | bytes | bytearray,
) -> Model:
return model_class.parse_raw(message) # type: ignore[arg-type]

def model_dump_json(instance: Model) -> str:
return instance.json()

def model_copy(
instance: Model,
update: dict[str, Any] | None = None,
deep: bool = False,
) -> Model:
return instance.copy(update=update, deep=deep)

validate_call = pydantic.validate_arguments # type: ignore
def parse_obj_as(annot: type[T], obj: Any) -> T:
return create_type_adapter(annot).validate_python(obj)
16 changes: 3 additions & 13 deletions taskiq/depends/progress_tracker.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
import enum
from typing import Generic, TypeVar

from pydantic import BaseModel, ConfigDict
from taskiq_dependencies import Depends

from taskiq.compat import IS_PYDANTIC2
from taskiq.context import Context

_ProgressType = TypeVar("_ProgressType")
Expand All @@ -18,18 +18,8 @@ class TaskState(str, enum.Enum):
RETRY = "RETRY"


if IS_PYDANTIC2:
from pydantic import BaseModel, ConfigDict

class _TaskProgressConfig(BaseModel):
model_config = ConfigDict(arbitrary_types_allowed=True)

else:
from pydantic.generics import GenericModel

class _TaskProgressConfig(GenericModel): # type: ignore[no-redef]
class Config:
arbitrary_types_allowed = True
class _TaskProgressConfig(BaseModel):
model_config = ConfigDict(arbitrary_types_allowed=True)


class TaskProgress(_TaskProgressConfig, Generic[_ProgressType]):
Expand Down
5 changes: 2 additions & 3 deletions taskiq/formatters/json_formatter.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
from taskiq.abc.formatter import TaskiqFormatter
from taskiq.compat import model_dump_json, model_validate_json
from taskiq.message import BrokerMessage, TaskiqMessage


Expand All @@ -16,7 +15,7 @@ def dumps(self, message: TaskiqMessage) -> BrokerMessage:
return BrokerMessage(
task_id=message.task_id,
task_name=message.task_name,
message=model_dump_json(message).encode(),
message=message.model_dump_json().encode(),
labels=message.labels,
)

Expand All @@ -27,4 +26,4 @@ def loads(self, message: bytes) -> TaskiqMessage:
:param message: broker's message.
:return: parsed taskiq message.
"""
return model_validate_json(TaskiqMessage, message)
return TaskiqMessage.model_validate_json(message)
5 changes: 2 additions & 3 deletions taskiq/formatters/proxy_formatter.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
from typing import TYPE_CHECKING

from taskiq.abc.formatter import TaskiqFormatter
from taskiq.compat import model_dump, model_validate
from taskiq.message import BrokerMessage, TaskiqMessage

if TYPE_CHECKING:
Expand All @@ -24,7 +23,7 @@ def dumps(self, message: TaskiqMessage) -> BrokerMessage:
return BrokerMessage(
task_id=message.task_id,
task_name=message.task_name,
message=self.broker.serializer.dumpb(model_dump(message)),
message=self.broker.serializer.dumpb(message.model_dump(mode="json")),
labels=message.labels,
)

Expand All @@ -35,4 +34,4 @@ def loads(self, message: bytes) -> TaskiqMessage:
:param message: broker's message.
:return: parsed taskiq message.
"""
return model_validate(TaskiqMessage, self.broker.serializer.loadb(message))
return TaskiqMessage.model_validate(self.broker.serializer.loadb(message))
3 changes: 1 addition & 2 deletions taskiq/kicker.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,6 @@
from pydantic import BaseModel

from taskiq.abc.middleware import TaskiqMiddleware
from taskiq.compat import model_dump
from taskiq.exceptions import SendTaskError
from taskiq.labels import prepare_label
from taskiq.message import TaskiqMessage
Expand Down Expand Up @@ -294,7 +293,7 @@ def _prepare_arg(cls, arg: Any) -> Any:
:return: Formatted argument.
"""
if isinstance(arg, BaseModel):
arg = model_dump(arg)
arg = arg.model_dump(mode="json")
if is_dataclass(arg):
if isinstance(arg, type):
raise ValueError(
Expand Down
16 changes: 2 additions & 14 deletions taskiq/middlewares/opentelemetry_middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,9 @@
from collections.abc import Generator
from contextlib import AbstractContextManager
from datetime import datetime, timezone
from importlib.metadata import version
from typing import Any, TypeVar

import psutil
from packaging.version import Version, parse

try:
import opentelemetry # noqa: F401
Expand Down Expand Up @@ -34,16 +32,6 @@
# Taskiq Context key
CTX_KEY = "__otel_task_span"

# unlike pydantic v2, v1 includes CTX_KEY by default
# excluding it here
PYDANTIC_VER = parse(version("pydantic"))
IS_PYDANTIC1 = Version("2.0") > PYDANTIC_VER
if IS_PYDANTIC1:
if TaskiqMessage.__exclude_fields__: # type: ignore[attr-defined]
TaskiqMessage.__exclude_fields__.update(CTX_KEY) # type: ignore
else:
TaskiqMessage.__exclude_fields__ = {CTX_KEY} # type: ignore

# Taskiq Context attributes
TASKIQ_CONTEXT_ATTRIBUTES = [
"_retries",
Expand Down Expand Up @@ -115,8 +103,8 @@ def attach_context(

if ctx_dict is None:
ctx_dict = {}
# use object.__setattr__ directly
# to skip pydantic v1 setattr
# use object.__setattr__ directly since CTX_KEY is not a declared model field,
# and pydantic forbids setting undeclared attributes
object.__setattr__(message, CTX_KEY, ctx_dict)

ctx_dict[(message.task_id, is_publish)] = (span, activation, token)
Expand Down
7 changes: 3 additions & 4 deletions taskiq/middlewares/taskiq_admin_middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,6 @@
import aiohttp

from taskiq.abc.middleware import TaskiqMiddleware
from taskiq.compat import model_dump
from taskiq.message import TaskiqMessage
from taskiq.result import TaskiqResult

Expand Down Expand Up @@ -126,7 +125,7 @@ async def post_send(self, message: TaskiqMessage) -> None:

:param message: kicked message.
"""
dict_message: dict[str, Any] = model_dump(message)
dict_message: dict[str, Any] = message.model_dump(mode="json")
await self._spawn_request(
f"/api/tasks/{message.task_id}/queued",
{
Expand All @@ -149,7 +148,7 @@ async def pre_execute(self, message: TaskiqMessage) -> TaskiqMessage:
:param message: incoming parsed taskiq message.
:return: modified message.
"""
dict_message: dict[str, Any] = model_dump(message)
dict_message: dict[str, Any] = message.model_dump(mode="json")
await self._spawn_request(
f"/api/tasks/{message.task_id}/started",
{
Expand Down Expand Up @@ -177,7 +176,7 @@ async def post_execute(
:param message: incoming message.
:param result: result of execution for current task.
"""
dict_result: dict[str, Any] = model_dump(result)
dict_result: dict[str, Any] = result.model_dump(mode="json")
await self._spawn_request(
f"/api/tasks/{message.task_id}/executed",
{
Expand Down
Loading
Loading