From 65fe923e32e5b799dd7c6267f29fb5cff5d57f9b Mon Sep 17 00:00:00 2001 From: Vaibhav Tiwari Date: Fri, 14 Aug 2026 14:26:43 -0400 Subject: [PATCH] feat: introduce API to fail a message in udf/transformer Signed-off-by: Vaibhav Tiwari --- packages/pynumaflow/pynumaflow/_constants.py | 1 + .../pynumaflow/batchmapper/__init__.py | 3 +- .../pynumaflow/batchmapper/_dtypes.py | 6 +- .../pynumaflow/pynumaflow/mapper/__init__.py | 2 + .../pynumaflow/pynumaflow/mapper/_dtypes.py | 6 +- .../pynumaflow/mapstreamer/__init__.py | 3 +- .../pynumaflow/mapstreamer/_dtypes.py | 6 +- .../pynumaflow/sourcetransformer/__init__.py | 2 + .../pynumaflow/sourcetransformer/_dtypes.py | 6 +- .../tests/batchmap/test_messages.py | 10 +++- .../pynumaflow/tests/map/test_async_mapper.py | 55 ++++++++++++++++++- .../pynumaflow/tests/map/test_messages.py | 10 +++- .../pynumaflow/tests/map/test_sync_mapper.py | 25 ++++++++- packages/pynumaflow/tests/map/utils.py | 8 +++ .../tests/mapstream/test_messages.py | 10 +++- .../tests/sourcetransform/test_messages.py | 11 +++- .../tests/sourcetransform/test_sync_server.py | 26 ++++++++- 17 files changed, 176 insertions(+), 14 deletions(-) diff --git a/packages/pynumaflow/pynumaflow/_constants.py b/packages/pynumaflow/pynumaflow/_constants.py index ed015671..63f09697 100644 --- a/packages/pynumaflow/pynumaflow/_constants.py +++ b/packages/pynumaflow/pynumaflow/_constants.py @@ -55,6 +55,7 @@ DELIMITER = ":" DROP = "U+005C__DROP__" NACK = "U+005C__NACK__" +FAIL = "U+005C__FAIL__" _PROCESS_COUNT = os.cpu_count() # Cap max value to 16 diff --git a/packages/pynumaflow/pynumaflow/batchmapper/__init__.py b/packages/pynumaflow/pynumaflow/batchmapper/__init__.py index 3374a0f1..ac17b6e7 100644 --- a/packages/pynumaflow/pynumaflow/batchmapper/__init__.py +++ b/packages/pynumaflow/pynumaflow/batchmapper/__init__.py @@ -1,4 +1,4 @@ -from pynumaflow._constants import DROP +from pynumaflow._constants import DROP, FAIL from pynumaflow.batchmapper._dtypes import ( Message, @@ -14,6 +14,7 @@ "Message", "Datum", "DROP", + "FAIL", "BatchMapAsyncServer", "BatchMapper", "BatchResponses", diff --git a/packages/pynumaflow/pynumaflow/batchmapper/_dtypes.py b/packages/pynumaflow/pynumaflow/batchmapper/_dtypes.py index 0cba0031..a22520a3 100644 --- a/packages/pynumaflow/pynumaflow/batchmapper/_dtypes.py +++ b/packages/pynumaflow/pynumaflow/batchmapper/_dtypes.py @@ -5,7 +5,7 @@ from typing import TypeAlias, TypeVar from collections.abc import AsyncIterable, Callable -from pynumaflow._constants import DROP, NACK +from pynumaflow._constants import DROP, FAIL, NACK from pynumaflow._nack import NackOptions from pynumaflow._validate import _validate_message_fields @@ -51,6 +51,10 @@ def to_nack(cls: type[M], opts: NackOptions | None = None) -> M: m._nack_options = opts return m + @classmethod + def to_fail(cls: type[M]) -> M: + return cls(b"", None, [FAIL]) + @property def nack_options(self) -> NackOptions | None: return self._nack_options diff --git a/packages/pynumaflow/pynumaflow/mapper/__init__.py b/packages/pynumaflow/pynumaflow/mapper/__init__.py index 477cd187..dd9c163a 100644 --- a/packages/pynumaflow/pynumaflow/mapper/__init__.py +++ b/packages/pynumaflow/pynumaflow/mapper/__init__.py @@ -5,12 +5,14 @@ from pynumaflow.mapper._dtypes import Message, Messages, Datum, DROP, Mapper from pynumaflow._metadata import UserMetadata, SystemMetadata from pynumaflow._nack import NackOptions +from pynumaflow._constants import FAIL __all__ = [ "Message", "Messages", "Datum", "DROP", + "FAIL", "Mapper", "MapServer", "MapAsyncServer", diff --git a/packages/pynumaflow/pynumaflow/mapper/_dtypes.py b/packages/pynumaflow/pynumaflow/mapper/_dtypes.py index 1f870e1b..20a7394f 100644 --- a/packages/pynumaflow/pynumaflow/mapper/_dtypes.py +++ b/packages/pynumaflow/pynumaflow/mapper/_dtypes.py @@ -6,7 +6,7 @@ from collections.abc import Callable from warnings import warn -from pynumaflow._constants import DROP, NACK +from pynumaflow._constants import DROP, FAIL, NACK from pynumaflow._nack import NackOptions from pynumaflow._metadata import UserMetadata, SystemMetadata from pynumaflow._validate import _validate_message_fields @@ -61,6 +61,10 @@ def to_nack(cls: type[M], opts: NackOptions | None = None) -> M: m._nack_options = opts return m + @classmethod + def to_fail(cls: type[M]) -> M: + return cls(b"", None, [FAIL]) + @property def nack_options(self) -> NackOptions | None: return self._nack_options diff --git a/packages/pynumaflow/pynumaflow/mapstreamer/__init__.py b/packages/pynumaflow/pynumaflow/mapstreamer/__init__.py index f993e607..06ad146f 100644 --- a/packages/pynumaflow/pynumaflow/mapstreamer/__init__.py +++ b/packages/pynumaflow/pynumaflow/mapstreamer/__init__.py @@ -1,4 +1,4 @@ -from pynumaflow._constants import DROP +from pynumaflow._constants import DROP, FAIL from pynumaflow.mapstreamer._dtypes import Message, Messages, Datum, MapStreamer from pynumaflow.mapstreamer.async_server import MapStreamAsyncServer @@ -9,6 +9,7 @@ "Messages", "Datum", "DROP", + "FAIL", "MapStreamAsyncServer", "MapStreamer", "NackOptions", diff --git a/packages/pynumaflow/pynumaflow/mapstreamer/_dtypes.py b/packages/pynumaflow/pynumaflow/mapstreamer/_dtypes.py index 01e7fac9..a70950e8 100644 --- a/packages/pynumaflow/pynumaflow/mapstreamer/_dtypes.py +++ b/packages/pynumaflow/pynumaflow/mapstreamer/_dtypes.py @@ -6,7 +6,7 @@ from collections.abc import AsyncIterable, Callable from warnings import warn -from pynumaflow._constants import DROP, NACK +from pynumaflow._constants import DROP, FAIL, NACK from pynumaflow._nack import NackOptions from pynumaflow._validate import _validate_message_fields @@ -51,6 +51,10 @@ def to_nack(cls: type[M], opts: NackOptions | None = None) -> M: m._nack_options = opts return m + @classmethod + def to_fail(cls: type[M]) -> M: + return cls(b"", None, [FAIL]) + @property def nack_options(self) -> NackOptions | None: return self._nack_options diff --git a/packages/pynumaflow/pynumaflow/sourcetransformer/__init__.py b/packages/pynumaflow/pynumaflow/sourcetransformer/__init__.py index d6dccf73..417be0ea 100644 --- a/packages/pynumaflow/pynumaflow/sourcetransformer/__init__.py +++ b/packages/pynumaflow/pynumaflow/sourcetransformer/__init__.py @@ -10,12 +10,14 @@ from pynumaflow.sourcetransformer.async_server import SourceTransformAsyncServer from pynumaflow._metadata import UserMetadata, SystemMetadata from pynumaflow._nack import NackOptions +from pynumaflow._constants import FAIL __all__ = [ "Message", "Messages", "Datum", "DROP", + "FAIL", "SourceTransformServer", "SourceTransformer", "SourceTransformMultiProcServer", diff --git a/packages/pynumaflow/pynumaflow/sourcetransformer/_dtypes.py b/packages/pynumaflow/pynumaflow/sourcetransformer/_dtypes.py index b9b2f07f..318fc92b 100644 --- a/packages/pynumaflow/pynumaflow/sourcetransformer/_dtypes.py +++ b/packages/pynumaflow/pynumaflow/sourcetransformer/_dtypes.py @@ -6,7 +6,7 @@ from collections.abc import Awaitable, Callable from warnings import warn -from pynumaflow._constants import DROP, NACK +from pynumaflow._constants import DROP, FAIL, NACK from pynumaflow._nack import NackOptions from pynumaflow._metadata import UserMetadata, SystemMetadata from pynumaflow._validate import _validate_message_fields @@ -70,6 +70,10 @@ def to_nack( m._nack_options = opts return m + @classmethod + def to_fail(cls: type[M], event_time: datetime) -> M: + return cls(b"", event_time, None, [FAIL]) + @property def nack_options(self) -> NackOptions | None: return self._nack_options diff --git a/packages/pynumaflow/tests/batchmap/test_messages.py b/packages/pynumaflow/tests/batchmap/test_messages.py index a0f81662..73d25aa1 100644 --- a/packages/pynumaflow/tests/batchmap/test_messages.py +++ b/packages/pynumaflow/tests/batchmap/test_messages.py @@ -1,7 +1,7 @@ import pytest from pynumaflow.batchmapper import Message, DROP, BatchResponse, BatchResponses, NackOptions -from pynumaflow._constants import NACK +from pynumaflow._constants import FAIL, NACK from tests.batchmap.test_datatypes import TEST_ID from tests.testing_utils import mock_message @@ -32,6 +32,14 @@ def test_message_default_nack_options(): assert msg.nack_options is None +def test_message_to_fail(): + msg = Message.to_fail() + assert type(msg) is Message + assert msg.keys == [] + assert msg.value == b"" + assert msg.tags == [FAIL] + + def test_batch_responses_init(): batch_responses = BatchResponses() batch_response1 = BatchResponse.from_id(TEST_ID) diff --git a/packages/pynumaflow/tests/map/test_async_mapper.py b/packages/pynumaflow/tests/map/test_async_mapper.py index 5c485ab9..a8483d11 100644 --- a/packages/pynumaflow/tests/map/test_async_mapper.py +++ b/packages/pynumaflow/tests/map/test_async_mapper.py @@ -16,8 +16,13 @@ from pynumaflow.proto.common import metadata_pb2 from pynumaflow.proto.mapper import map_pb2, map_pb2_grpc from tests.conftest import create_async_loop, start_async_server, teardown_async_server -from tests.map.utils import get_test_datums, async_nack_map_handler, NACK_TEST_OPTIONS -from pynumaflow._constants import NACK +from tests.map.utils import ( + get_test_datums, + async_nack_map_handler, + async_fail_map_handler, + NACK_TEST_OPTIONS, +) +from pynumaflow._constants import FAIL, NACK pytestmark = pytest.mark.integration @@ -28,6 +33,7 @@ SOCK_PATH = "unix:///tmp/async_map.sock" NACK_SOCK_PATH = "unix:///tmp/async_map_nack.sock" +FAIL_SOCK_PATH = "unix:///tmp/async_map_fail.sock" def request_generator(req): @@ -156,6 +162,51 @@ def test_map_nack(async_nack_map_server): assert result.nack_options.reason == NACK_TEST_OPTIONS.reason +async def _start_fail_server(udfs): + _server_options = [ + ("grpc.max_send_message_length", MAX_MESSAGE_SIZE), + ("grpc.max_receive_message_length", MAX_MESSAGE_SIZE), + ] + server = grpc.aio.server(options=_server_options) + map_pb2_grpc.add_MapServicer_to_server(udfs, server) + server.add_insecure_port(FAIL_SOCK_PATH) + logging.info("Starting fail server on %s", FAIL_SOCK_PATH) + await server.start() + return server, FAIL_SOCK_PATH + + +@pytest.fixture(scope="module") +def async_fail_map_server(): + """Module-scoped fixture: async gRPC map server whose handler fails every message.""" + loop = create_async_loop() + + server_obj = MapAsyncServer(mapper_instance=async_fail_map_handler) + udfs = server_obj.servicer + server = start_async_server(loop, _start_fail_server(udfs)) + + yield loop + + teardown_async_server(loop, server) + + +def test_map_fail(async_fail_map_server): + with grpc.insecure_channel(FAIL_SOCK_PATH) as channel: + stub = map_pb2_grpc.MapStub(channel) + request = get_test_datums() + generator_response = stub.MapFn(request_iterator=request_generator(request)) + + responses = list(generator_response) + + # 1 handshake + 3 data responses + assert len(responses) == 4 + assert responses[0].handshake.sot + + for resp in responses[1:]: + assert len(resp.results) == 1 + result = resp.results[0] + assert FAIL in result.tags + + def test_map(map_stub): request = get_test_datums() generator_response: Iterator[map_pb2.MapResponse] = map_stub.MapFn( diff --git a/packages/pynumaflow/tests/map/test_messages.py b/packages/pynumaflow/tests/map/test_messages.py index d430ca3d..61569b3c 100644 --- a/packages/pynumaflow/tests/map/test_messages.py +++ b/packages/pynumaflow/tests/map/test_messages.py @@ -1,7 +1,7 @@ import pytest from pynumaflow.mapper import Messages, Message, DROP, Mapper, Datum, NackOptions -from pynumaflow._constants import NACK +from pynumaflow._constants import FAIL, NACK from tests.testing_utils import mock_message @@ -52,6 +52,14 @@ def test_message_to_nack_default_opts(): assert msg.nack_options is None +def test_message_to_fail(): + msg = Message.to_fail() + assert type(msg) is Message + assert msg.keys == [] + assert msg.value == b"" + assert msg.tags == [FAIL] + + def test_message_default_nack_options(): msg = Message(mock_message()) assert msg.nack_options is None diff --git a/packages/pynumaflow/tests/map/test_sync_mapper.py b/packages/pynumaflow/tests/map/test_sync_mapper.py index 2d0d9057..c2960951 100644 --- a/packages/pynumaflow/tests/map/test_sync_mapper.py +++ b/packages/pynumaflow/tests/map/test_sync_mapper.py @@ -12,8 +12,9 @@ get_test_datums, nack_map_handler, NACK_TEST_OPTIONS, + fail_map_handler, ) -from pynumaflow._constants import NACK +from pynumaflow._constants import FAIL, NACK from tests.conftest import collect_responses, drain_responses, send_test_requests @@ -131,6 +132,28 @@ def test_map_nack(): assert code == StatusCode.OK +def test_map_fail(): + my_server = MapServer(mapper_instance=fail_map_handler) + services = {map_pb2.DESCRIPTOR.services_by_name["Map"]: my_server.servicer} + test_server = server_from_dictionary(services, strict_real_time()) + + test_datums = get_test_datums(handshake=True) + method = _invoke_map_fn(test_server) + send_test_requests(method, test_datums) + responses = collect_responses(method) + + metadata, code, details = method.termination() + # 1 handshake + 3 data responses + assert len(responses) == 4 + assert responses[0].handshake.sot + + for resp in responses[1:]: + assert len(resp.results) == 1 + result = resp.results[0] + assert FAIL in result.tags + assert code == StatusCode.OK + + def test_invalid_input(): with pytest.raises(TypeError): MapServer() diff --git a/packages/pynumaflow/tests/map/utils.py b/packages/pynumaflow/tests/map/utils.py index 38851185..9102b242 100644 --- a/packages/pynumaflow/tests/map/utils.py +++ b/packages/pynumaflow/tests/map/utils.py @@ -16,6 +16,14 @@ async def async_nack_map_handler(keys: list[str], datum: Datum) -> Messages: return Messages(Message.to_nack(NACK_TEST_OPTIONS)) +def fail_map_handler(keys: list[str], datum: Datum) -> Messages: + return Messages(Message.to_fail()) + + +async def async_fail_map_handler(keys: list[str], datum: Datum) -> Messages: + return Messages(Message.to_fail()) + + async def async_map_error_fn(keys: list[str], datum: Datum) -> Messages: raise ValueError("error invoking map") diff --git a/packages/pynumaflow/tests/mapstream/test_messages.py b/packages/pynumaflow/tests/mapstream/test_messages.py index 0ba8ed97..7b658299 100644 --- a/packages/pynumaflow/tests/mapstream/test_messages.py +++ b/packages/pynumaflow/tests/mapstream/test_messages.py @@ -1,7 +1,7 @@ import pytest from pynumaflow.mapstreamer import Messages, Message, DROP, NackOptions -from pynumaflow._constants import NACK +from pynumaflow._constants import FAIL, NACK from tests.testing_utils import mock_message @@ -26,6 +26,14 @@ def test_message_default_nack_options(): assert msg.nack_options is None +def test_message_to_fail(): + msg = Message.to_fail() + assert type(msg) is Message + assert msg.keys == [] + assert msg.value == b"" + assert msg.tags == [FAIL] + + def test_message_key(): mock_obj = {"Keys": ["test-key"], "Value": mock_message()} msg = Message(value=mock_obj["Value"], keys=mock_obj["Keys"]) diff --git a/packages/pynumaflow/tests/sourcetransform/test_messages.py b/packages/pynumaflow/tests/sourcetransform/test_messages.py index 061ae1a1..bd2cb88d 100644 --- a/packages/pynumaflow/tests/sourcetransform/test_messages.py +++ b/packages/pynumaflow/tests/sourcetransform/test_messages.py @@ -11,7 +11,7 @@ SystemMetadata, NackOptions, ) -from pynumaflow._constants import NACK +from pynumaflow._constants import FAIL, NACK from tests.testing_utils import mock_new_event_time @@ -83,6 +83,15 @@ def test_message_to_nack_default_opts(): assert msgt.nack_options is None +def test_message_to_fail(): + msgt = Message.to_fail(mock_event_time()) + assert isinstance(msgt, Message) + assert msgt.keys == [] + assert msgt.value == b"" + assert msgt.tags == [FAIL] + assert msgt.event_time == mock_event_time() + + def test_message_default_nack_options(): msgt = Message(mock_message_t(), mock_event_time()) assert msgt.nack_options is None diff --git a/packages/pynumaflow/tests/sourcetransform/test_sync_server.py b/packages/pynumaflow/tests/sourcetransform/test_sync_server.py index 85d74793..302614b9 100644 --- a/packages/pynumaflow/tests/sourcetransform/test_sync_server.py +++ b/packages/pynumaflow/tests/sourcetransform/test_sync_server.py @@ -13,7 +13,7 @@ Message, NackOptions, ) -from pynumaflow._constants import NACK +from pynumaflow._constants import FAIL, NACK from tests.sourcetransform.utils import transform_handler, err_transform_handler, get_test_datums from tests.conftest import collect_responses, drain_responses, send_test_requests from tests.testing_utils import mock_new_event_time @@ -155,6 +155,30 @@ def test_transform_nack(): assert code == StatusCode.OK +def fail_transform_handler(keys: list[str], datum: Datum) -> Messages: + return Messages(Message.to_fail(mock_new_event_time())) + + +def test_transform_fail(): + test_server = _make_transform_server(fail_transform_handler) + test_datums = get_test_datums() + method = _invoke_transform_fn(test_server) + + send_test_requests(method, test_datums) + responses = collect_responses(method) + + metadata, code, details = method.termination() + # 1 handshake + 3 data responses + assert len(responses) == 4 + assert responses[0].handshake.sot + + for resp in responses[1:]: + assert len(resp.results) == 1 + result = resp.results[0] + assert FAIL in result.tags + assert code == StatusCode.OK + + def test_invalid_input(): with pytest.raises(TypeError): SourceTransformServer()