Skip to content
Open
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
119 changes: 80 additions & 39 deletions taskiq/receiver/params_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,62 @@
logger = getLogger(__name__)


def _parse_arg(
param_name: str,
annot: Any,
argnum: int,
message: TaskiqMessage,
) -> None:
"""
Parse a positional argument by its annotation, in place.

:param param_name: name of the parameter.
:param annot: type annotation of the parameter.
:param argnum: index of the argument in message.args.
:param message: incoming message.
"""
value = message.args[argnum]
if value is None:
return
logger.debug("Trying to parse %s as %s", param_name, annot)
try:
# trying to parse found value as in type annotation.
message.args[argnum] = parse_obj_as(annot, value)
except (ValueError, RuntimeError) as exc:
logger.warning(
"Can't parse argument %d for task %s. Reason: %s",
argnum,
message.task_name,
exc,
exc_info=True,
)


def _parse_kwarg(param_name: str, annot: Any, message: TaskiqMessage) -> None:
"""
Parse a keyword argument by its annotation, in place.

:param param_name: name of the parameter.
:param annot: type annotation of the parameter.
:param message: incoming message.
"""
value = message.kwargs.get(param_name)
if value is None:
return
logger.debug("Trying to parse %s as %s", param_name, annot)
try:
# trying to parse found value as in type annotation.
message.kwargs[param_name] = parse_obj_as(annot, value)
except (ValueError, RuntimeError) as exc:
logger.warning(
"Can't parse argument %s for task %s. Reason: %s",
param_name,
message.task_name,
exc,
exc_info=True,
)


def parse_params(
signature: inspect.Signature | None,
type_hints: dict[str, Any],
Expand Down Expand Up @@ -55,47 +111,32 @@ def parse_params(
return
argnum = -1
# Iterate over function's params.
for param_name in signature.parameters:
for param_name, param in signature.parameters.items():
# If parameter doesn't have an annotation.
annot = type_hints.get(param_name)
if annot is None:
continue
# Increment argument numbers. This is
# for positional arguments.
argnum += 1
# Value from incoming message.
value = None
logger.debug("Trying to parse %s as %s", param_name, annot)
# Check if we have positional arguments in passed message.
if argnum < len(message.args):
# Get positional argument.
value = message.args[argnum]
if value is None:
if param.kind in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
):
# Every positional-capable parameter occupies a slot in
# message.args, even if it has no type annotation.
argnum += 1
if annot is None:
continue
if argnum < len(message.args):
# This parameter was passed positionally.
_parse_arg(param_name, annot, argnum, message)
else:
# The parameter was passed as a kwarg or not at all.
_parse_kwarg(param_name, annot, message)
elif param.kind == inspect.Parameter.VAR_POSITIONAL:
# All remaining positional arguments belong to *args.
if annot is None:
continue
try:
# trying to parse found value as in type annotation.
message.args[argnum] = parse_obj_as(annot, value)
except (ValueError, RuntimeError) as exc:
logger.warning(
"Can't parse argument %d for task %s. Reason: %s",
argnum,
message.task_name,
exc,
exc_info=True,
)
for i in range(argnum + 1, len(message.args)):
_parse_arg(param_name, annot, i, message)
else:
# We try to get this parameter from kwargs.
value = message.kwargs.get(param_name)
if value is None:
# KEYWORD_ONLY and VAR_KEYWORD parameters are matched by name.
if annot is None or param.kind == inspect.Parameter.VAR_KEYWORD:
continue
try:
# trying to parse found value as in type annotation.
message.kwargs[param_name] = parse_obj_as(annot, value)
except (ValueError, RuntimeError) as exc:
logger.warning(
"Can't parse argument %s for task %s. Reason: %s",
param_name,
message.task_name,
exc,
exc_info=True,
)
_parse_kwarg(param_name, annot, message)
152 changes: 152 additions & 0 deletions tests/receiver/test_params_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
import logging
from collections.abc import Callable
from dataclasses import dataclass
from datetime import datetime
from typing import Any, get_type_hints

import pytest
Expand Down Expand Up @@ -262,3 +263,154 @@ def func(a: TestObj) -> None:
_helper(func, msg)
assert "Can't parse argument a" in caplog.text
assert msg.kwargs == {"a": {"a": "10", "b": "f3"}}


def test_unannotated_param_value_left_untouched() -> None:
def func(request_id, count: int) -> None: # type: ignore # noqa: ANN001
pass

msg = TaskiqMessage(
task_id="test",
task_name="test",
labels={},
labels_types={},
args=["12345", 3],
kwargs={},
)
_helper(func, msg)
assert msg.args == ["12345", 3]
assert isinstance(msg.args[0], str)


def test_annotated_param_parsed_after_unannotated() -> None:
def func(request_id, count: int) -> None: # type: ignore # noqa: ANN001
pass

msg = TaskiqMessage(
task_id="test",
task_name="test",
labels={},
labels_types={},
args=["12345", "5"],
kwargs={},
)
_helper(func, msg)
assert msg.args == ["12345", 5]


def test_datetime_param_parsed_after_unannotated() -> None:
def func(ts, when: datetime) -> None: # type: ignore # noqa: ANN001
pass

msg = TaskiqMessage(
task_id="test",
task_name="test",
labels={},
labels_types={},
args=["2026-08-30T10:00:00", "2026-08-30T10:00:00"],
kwargs={},
)
_helper(func, msg)
assert msg.args[0] == "2026-08-30T10:00:00"
assert msg.args[1] == datetime(2026, 8, 30, 10, 0)


def test_invalid_annotated_value_warns_about_right_arg(
caplog: pytest.LogCaptureFixture,
) -> None:
def func(ctx, count: int) -> None: # type: ignore # noqa: ANN001
pass

msg = TaskiqMessage(
task_id="test",
task_name="test",
labels={},
labels_types={},
args=["hello", "not-an-int"],
kwargs={},
)
with caplog.at_level(logging.WARNING):
_helper(func, msg)
assert "Can't parse argument 1" in caplog.text
assert msg.args == ["hello", "not-an-int"]


def test_varargs_all_elements_parsed() -> None:
def func(*values: int) -> None:
pass

msg = TaskiqMessage(
task_id="test",
task_name="test",
labels={},
labels_types={},
args=["1", "2", "3"],
kwargs={},
)
_helper(func, msg)
assert msg.args == [1, 2, 3]


def test_varargs_with_leading_arg() -> None:
def func(a: int, *values: int) -> None:
pass

msg = TaskiqMessage(
task_id="test",
task_name="test",
labels={},
labels_types={},
args=["1", "2", "3"],
kwargs={},
)
_helper(func, msg)
assert msg.args == [1, 2, 3]


def test_unannotated_varargs_untouched() -> None:
def func(*values) -> None: # type: ignore # noqa: ANN002
pass

msg = TaskiqMessage(
task_id="test",
task_name="test",
labels={},
labels_types={},
args=["1", 2],
kwargs={},
)
_helper(func, msg)
assert msg.args == ["1", 2]


def test_kwonly_param_parsed() -> None:
def func(a: int, *, b: int) -> None:
pass

msg = TaskiqMessage(
task_id="test",
task_name="test",
labels={},
labels_types={},
args=["1"],
kwargs={"b": "2"},
)
_helper(func, msg)
assert msg.args == [1]
assert msg.kwargs == {"b": 2}


def test_positional_params_parsed_from_kwargs() -> None:
def func(a: int, b: int) -> None:
pass

msg = TaskiqMessage(
task_id="test",
task_name="test",
labels={},
labels_types={},
args=[],
kwargs={"a": "1", "b": "2"},
)
_helper(func, msg)
assert msg.kwargs == {"a": 1, "b": 2}