From 937747f5174b07a265801241536d4622455b0ea1 Mon Sep 17 00:00:00 2001 From: Vaibhav Tiwari Date: Fri, 14 Aug 2026 14:26:10 -0400 Subject: [PATCH] feat: introduce API to fail a message in udf/transformer Signed-off-by: Vaibhav Tiwari --- .../numaflow/batchmapper/Message.java | 10 ++ .../io/numaproj/numaflow/mapper/Message.java | 10 ++ .../numaflow/mapstreamer/Message.java | 10 ++ .../numaflow/sourcetransformer/Message.java | 13 +++ .../numaflow/batchmapper/ServerFailTest.java | 97 +++++++++++++++++++ .../numaflow/mapper/ServerFailTest.java | 79 +++++++++++++++ .../numaflow/mapstreamer/ServiceFailTest.java | 77 +++++++++++++++ .../sourcetransformer/ServerFailTest.java | 83 ++++++++++++++++ 8 files changed, 379 insertions(+) create mode 100644 src/test/java/io/numaproj/numaflow/batchmapper/ServerFailTest.java create mode 100644 src/test/java/io/numaproj/numaflow/mapper/ServerFailTest.java create mode 100644 src/test/java/io/numaproj/numaflow/mapstreamer/ServiceFailTest.java create mode 100644 src/test/java/io/numaproj/numaflow/sourcetransformer/ServerFailTest.java diff --git a/src/main/java/io/numaproj/numaflow/batchmapper/Message.java b/src/main/java/io/numaproj/numaflow/batchmapper/Message.java index 7f54a159..a837babd 100644 --- a/src/main/java/io/numaproj/numaflow/batchmapper/Message.java +++ b/src/main/java/io/numaproj/numaflow/batchmapper/Message.java @@ -8,6 +8,7 @@ public class Message { private static final String[] DROP_TAGS = {"U+005C__DROP__"}; private static final String[] NACK_TAGS = {"U+005C__NACK__"}; + private static final String[] FAIL_TAGS = {"U+005C__FAIL__"}; private final String[] keys; private final byte[] value; private final String[] tags; @@ -69,4 +70,13 @@ public static Message toDrop() { public static Message toNack(NackOptions nackOptions) { return new Message(new byte[0], null, NACK_TAGS, nackOptions); } + + /** + * creates a Message that marks the input message as failed, triggering a retry by numaflow-core. + * + * @return the Message which will be failed + */ + public static Message toFail() { + return new Message(new byte[0], null, FAIL_TAGS); + } } diff --git a/src/main/java/io/numaproj/numaflow/mapper/Message.java b/src/main/java/io/numaproj/numaflow/mapper/Message.java index 83c7cfef..772e12de 100644 --- a/src/main/java/io/numaproj/numaflow/mapper/Message.java +++ b/src/main/java/io/numaproj/numaflow/mapper/Message.java @@ -13,6 +13,7 @@ public class Message { private static final String[] DROP_TAGS = {"U+005C__DROP__"}; private static final String[] NACK_TAGS = {"U+005C__NACK__"}; + private static final String[] FAIL_TAGS = {"U+005C__FAIL__"}; private final String[] keys; private final byte[] value; private final String[] tags; @@ -89,4 +90,13 @@ public static Message toDrop() { public static Message toNack(NackOptions nackOptions) { return new Message(new byte[0], null, NACK_TAGS, null, nackOptions); } + + /** + * creates a Message that marks the input message as failed, triggering a retry by numaflow-core. + * + * @return the Message which will be failed + */ + public static Message toFail() { + return new Message(new byte[0], null, FAIL_TAGS, null); + } } diff --git a/src/main/java/io/numaproj/numaflow/mapstreamer/Message.java b/src/main/java/io/numaproj/numaflow/mapstreamer/Message.java index 4632d9e5..98bf8e32 100644 --- a/src/main/java/io/numaproj/numaflow/mapstreamer/Message.java +++ b/src/main/java/io/numaproj/numaflow/mapstreamer/Message.java @@ -8,6 +8,7 @@ public class Message { private static final String[] DROP_TAGS = {"U+005C__DROP__"}; private static final String[] NACK_TAGS = {"U+005C__NACK__"}; + private static final String[] FAIL_TAGS = {"U+005C__FAIL__"}; private final String[] keys; private final byte[] value; private final String[] tags; @@ -69,4 +70,13 @@ public static Message toDrop() { public static Message toNack(NackOptions nackOptions) { return new Message(new byte[0], null, NACK_TAGS, nackOptions); } + + /** + * creates a Message that marks the input message as failed, triggering a retry by numaflow-core. + * + * @return the Message which will be failed + */ + public static Message toFail() { + return new Message(new byte[0], null, FAIL_TAGS); + } } diff --git a/src/main/java/io/numaproj/numaflow/sourcetransformer/Message.java b/src/main/java/io/numaproj/numaflow/sourcetransformer/Message.java index 183572dc..34608881 100644 --- a/src/main/java/io/numaproj/numaflow/sourcetransformer/Message.java +++ b/src/main/java/io/numaproj/numaflow/sourcetransformer/Message.java @@ -11,6 +11,7 @@ public class Message { private static final String[] DROP_TAGS = {"U+005C__DROP__"}; private static final String[] NACK_TAGS = {"U+005C__NACK__"}; + private static final String[] FAIL_TAGS = {"U+005C__FAIL__"}; private final String[] keys; private final byte[] value; private final Instant eventTime; @@ -99,4 +100,16 @@ public static Message toDrop(Instant eventTime) { public static Message toNack(Instant eventTime, NackOptions nackOptions) { return new Message(new byte[0], eventTime, null, NACK_TAGS, null, nackOptions); } + + /** + * creates a Message that marks the input message as failed, triggering a retry by numaflow-core. + * eventTime is required: even though the message is failed, it is considered processed for + * watermark purposes. + * + * @param eventTime message eventTime + * @return the Message which will be failed + */ + public static Message toFail(Instant eventTime) { + return new Message(new byte[0], eventTime, null, FAIL_TAGS, null); + } } diff --git a/src/test/java/io/numaproj/numaflow/batchmapper/ServerFailTest.java b/src/test/java/io/numaproj/numaflow/batchmapper/ServerFailTest.java new file mode 100644 index 00000000..75af50c9 --- /dev/null +++ b/src/test/java/io/numaproj/numaflow/batchmapper/ServerFailTest.java @@ -0,0 +1,97 @@ +package io.numaproj.numaflow.batchmapper; + +import com.google.protobuf.ByteString; +import io.grpc.ManagedChannel; +import io.grpc.inprocess.InProcessChannelBuilder; +import io.grpc.inprocess.InProcessServerBuilder; +import io.grpc.stub.StreamObserver; +import io.grpc.testing.GrpcCleanupRule; +import io.numaproj.numaflow.map.v1.MapGrpc; +import io.numaproj.numaflow.map.v1.MapOuterClass; +import org.junit.After; +import org.junit.Before; +import org.junit.Rule; +import org.junit.Test; + +import java.util.Arrays; +import java.util.List; +import java.util.concurrent.ExecutionException; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.fail; + +public class ServerFailTest { + @Rule + public final GrpcCleanupRule grpcCleanup = new GrpcCleanupRule(); + private Server server; + private ManagedChannel inProcessChannel; + + @Before + public void setUp() throws Exception { + String serverName = InProcessServerBuilder.generateName(); + GRPCConfig grpcServerConfig = GRPCConfig.newBuilder() + .maxMessageSize(Constants.DEFAULT_MESSAGE_SIZE) + .socketPath(Constants.DEFAULT_SOCKET_PATH) + .infoFilePath("/tmp/numaflow-test-server-info)") + .build(); + server = new Server(grpcServerConfig, new FailBatchMapFn(), null, serverName); + server.start(); + inProcessChannel = grpcCleanup.register( + InProcessChannelBuilder.forName(serverName).directExecutor().build()); + } + + @After + public void tearDown() throws Exception { + server.stop(); + } + + @Test + public void batchMapFail() { + // expect: handshake resp + 1 per-id response + 1 EOT = 3 + BatchMapOutputStreamObserver outputStreamObserver = new BatchMapOutputStreamObserver(3); + StreamObserver in = MapGrpc.newStub(inProcessChannel) + .mapFn(outputStreamObserver); + in.onNext(MapOuterClass.MapRequest.newBuilder() + .setHandshake(MapOuterClass.Handshake.newBuilder().setSot(true)).build()); + in.onNext(MapOuterClass.MapRequest.newBuilder() + .setRequest(MapOuterClass.MapRequest.Request.newBuilder() + .setValue(ByteString.copyFromUtf8("x")).addKeys("k").build()) + .setId("id-1").build()); + in.onNext(MapOuterClass.MapRequest.newBuilder() + .setStatus(MapOuterClass.TransmissionStatus.newBuilder().setEot(true)).build()); + in.onCompleted(); + try { + outputStreamObserver.done.get(); + } catch (InterruptedException | ExecutionException e) { + fail("Error in getting done signal " + e.getMessage()); + } + List result = outputStreamObserver.getMapResponses(); + MapOuterClass.MapResponse.Result r = result.stream() + .filter(resp -> resp.getResultsCount() > 0) + .findFirst().orElseThrow(() -> new AssertionError("no result")).getResults(0); + assertEquals(Arrays.asList("U+005C__FAIL__"), r.getTagsList()); + } + + private static class FailBatchMapFn extends BatchMapper { + @Override + public BatchResponses processMessage(DatumIterator datumStream) { + BatchResponses batchResponses = new BatchResponses(); + while (true) { + Datum datum; + try { + datum = datumStream.next(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + continue; + } + if (datum == null) { + break; + } + BatchResponse batchResponse = new BatchResponse(datum.getId()); + batchResponse.append(Message.toFail()); + batchResponses.append(batchResponse); + } + return batchResponses; + } + } +} diff --git a/src/test/java/io/numaproj/numaflow/mapper/ServerFailTest.java b/src/test/java/io/numaproj/numaflow/mapper/ServerFailTest.java new file mode 100644 index 00000000..ee0479df --- /dev/null +++ b/src/test/java/io/numaproj/numaflow/mapper/ServerFailTest.java @@ -0,0 +1,79 @@ +package io.numaproj.numaflow.mapper; + +import com.google.protobuf.ByteString; +import io.grpc.ManagedChannel; +import io.grpc.inprocess.InProcessChannelBuilder; +import io.grpc.inprocess.InProcessServerBuilder; +import io.grpc.testing.GrpcCleanupRule; +import io.numaproj.numaflow.map.v1.MapGrpc; +import io.numaproj.numaflow.map.v1.MapOuterClass; +import org.junit.After; +import org.junit.Before; +import org.junit.Rule; +import org.junit.Test; + +import java.util.Arrays; +import java.util.List; +import java.util.concurrent.ExecutionException; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.fail; + +public class ServerFailTest { + @Rule + public final GrpcCleanupRule grpcCleanup = new GrpcCleanupRule(); + private Server server; + private ManagedChannel inProcessChannel; + + @Before + public void setUp() throws Exception { + String serverName = InProcessServerBuilder.generateName(); + GRPCConfig grpcServerConfig = GRPCConfig.newBuilder() + .maxMessageSize(Constants.DEFAULT_MESSAGE_SIZE) + .socketPath(Constants.DEFAULT_SOCKET_PATH) + .infoFilePath("/tmp/numaflow-test-server-info)") + .build(); + server = new Server(grpcServerConfig, new FailMapFn(), null, serverName); + server.start(); + inProcessChannel = grpcCleanup.register( + InProcessChannelBuilder.forName(serverName).directExecutor().build()); + } + + @After + public void tearDown() throws Exception { + server.stop(); + } + + @Test + public void mapperFail() { + MapOuterClass.MapRequest handshake = MapOuterClass.MapRequest.newBuilder() + .setHandshake(MapOuterClass.Handshake.newBuilder().setSot(true)).build(); + MapOuterClass.MapRequest inDatum = MapOuterClass.MapRequest.newBuilder() + .setRequest(MapOuterClass.MapRequest.Request.newBuilder() + .setValue(ByteString.copyFromUtf8("x")).addKeys("k").build()).build(); + + MapOutputStreamObserver responseObserver = new MapOutputStreamObserver(2); + var stub = MapGrpc.newStub(inProcessChannel); + var requestStreamObserver = stub.mapFn(responseObserver); + requestStreamObserver.onNext(handshake); + requestStreamObserver.onNext(inDatum); + try { + responseObserver.done.get(); + } catch (InterruptedException | ExecutionException e) { + fail("Error while waiting for response" + e.getMessage()); + } + List responses = responseObserver.getMapResponses().subList(1, 2); + MapOuterClass.MapResponse.Result r = responses.get(0).getResults(0); + assertEquals(Arrays.asList("U+005C__FAIL__"), r.getTagsList()); + requestStreamObserver.onCompleted(); + } + + private static class FailMapFn extends Mapper { + @Override + public MessageList processMessage(String[] keys, Datum datum) { + return MessageList.newBuilder() + .addMessage(Message.toFail()) + .build(); + } + } +} diff --git a/src/test/java/io/numaproj/numaflow/mapstreamer/ServiceFailTest.java b/src/test/java/io/numaproj/numaflow/mapstreamer/ServiceFailTest.java new file mode 100644 index 00000000..78cb62b8 --- /dev/null +++ b/src/test/java/io/numaproj/numaflow/mapstreamer/ServiceFailTest.java @@ -0,0 +1,77 @@ +package io.numaproj.numaflow.mapstreamer; + +import com.google.protobuf.ByteString; +import io.grpc.ManagedChannel; +import io.grpc.inprocess.InProcessChannelBuilder; +import io.grpc.inprocess.InProcessServerBuilder; +import io.grpc.testing.GrpcCleanupRule; +import io.numaproj.numaflow.map.v1.MapGrpc; +import io.numaproj.numaflow.map.v1.MapOuterClass; +import org.junit.After; +import org.junit.Before; +import org.junit.Rule; +import org.junit.Test; + +import java.util.Arrays; +import java.util.List; +import java.util.concurrent.CompletableFuture; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.fail; + +public class ServiceFailTest { + @Rule + public final GrpcCleanupRule grpcCleanup = new GrpcCleanupRule(); + private Service service; + private ManagedChannel inProcessChannel; + + @Before + public void setUp() throws Exception { + String serverName = InProcessServerBuilder.generateName(); + CompletableFuture shutdownSignal = new CompletableFuture<>(); + service = new Service(new FailMapStreamer(), shutdownSignal); + grpcCleanup.register(InProcessServerBuilder.forName(serverName).directExecutor() + .addService(service).build().start()); + inProcessChannel = grpcCleanup.register( + InProcessChannelBuilder.forName(serverName).directExecutor().build()); + } + + @After + public void tearDown() { + inProcessChannel.shutdownNow(); + } + + @Test + public void mapStreamerFail() { + MapOuterClass.MapRequest handshake = MapOuterClass.MapRequest.newBuilder() + .setHandshake(MapOuterClass.Handshake.newBuilder().setSot(true)).build(); + MapOuterClass.MapRequest inDatum = MapOuterClass.MapRequest.newBuilder() + .setRequest(MapOuterClass.MapRequest.Request.newBuilder() + .setValue(ByteString.copyFromUtf8("x")).addKeys("k").build()).build(); + + // expect: handshake resp + 1 result + 1 EOT = 3 responses + MapStreamOutputStreamObserver responseObserver = new MapStreamOutputStreamObserver(3); + var stub = MapGrpc.newStub(inProcessChannel); + var requestStreamObserver = stub.mapFn(responseObserver); + requestStreamObserver.onNext(handshake); + requestStreamObserver.onNext(inDatum); + try { + responseObserver.done.get(); + } catch (Exception e) { + fail("Error while waiting for response" + e.getMessage()); + } + List responses = responseObserver.getMapResponses(); + MapOuterClass.MapResponse.Result r = responses.stream() + .filter(resp -> resp.getResultsCount() > 0) + .findFirst().orElseThrow(() -> new AssertionError("no result")).getResults(0); + assertEquals(Arrays.asList("U+005C__FAIL__"), r.getTagsList()); + requestStreamObserver.onCompleted(); + } + + private static class FailMapStreamer extends MapStreamer { + @Override + public void processMessage(String[] keys, Datum datum, OutputObserver outputObserver) { + outputObserver.send(Message.toFail()); + } + } +} diff --git a/src/test/java/io/numaproj/numaflow/sourcetransformer/ServerFailTest.java b/src/test/java/io/numaproj/numaflow/sourcetransformer/ServerFailTest.java new file mode 100644 index 00000000..c0a57fa7 --- /dev/null +++ b/src/test/java/io/numaproj/numaflow/sourcetransformer/ServerFailTest.java @@ -0,0 +1,83 @@ +package io.numaproj.numaflow.sourcetransformer; + +import com.google.protobuf.ByteString; +import io.grpc.ManagedChannel; +import io.grpc.inprocess.InProcessChannelBuilder; +import io.grpc.inprocess.InProcessServerBuilder; +import io.grpc.testing.GrpcCleanupRule; +import io.numaproj.numaflow.sourcetransformer.v1.SourceTransformGrpc; +import io.numaproj.numaflow.sourcetransformer.v1.Sourcetransformer; +import org.junit.After; +import org.junit.Before; +import org.junit.Rule; +import org.junit.Test; + +import java.time.Instant; +import java.util.Arrays; +import java.util.List; +import java.util.concurrent.ExecutionException; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.fail; + +public class ServerFailTest { + private static final Instant TEST_EVENT_TIME = Instant.ofEpochMilli(1000L); + + @Rule + public final GrpcCleanupRule grpcCleanup = new GrpcCleanupRule(); + private Server server; + private ManagedChannel inProcessChannel; + + @Before + public void setUp() throws Exception { + String serverName = InProcessServerBuilder.generateName(); + GRPCConfig grpcServerConfig = GRPCConfig.newBuilder() + .maxMessageSize(Constants.DEFAULT_MESSAGE_SIZE) + .socketPath(Constants.DEFAULT_SOCKET_PATH) + .infoFilePath("/tmp/numaflow-test-server-info)") + .build(); + server = new Server(grpcServerConfig, new FailTransformer(), null, serverName); + server.start(); + inProcessChannel = grpcCleanup.register( + InProcessChannelBuilder.forName(serverName).directExecutor().build()); + } + + @After + public void tearDown() throws Exception { + server.stop(); + } + + @Test + public void transformerFail() { + Sourcetransformer.SourceTransformRequest handshake = Sourcetransformer.SourceTransformRequest.newBuilder() + .setHandshake(Sourcetransformer.Handshake.newBuilder().setSot(true).build()).build(); + Sourcetransformer.SourceTransformRequest req = Sourcetransformer.SourceTransformRequest.newBuilder() + .setRequest(Sourcetransformer.SourceTransformRequest.Request.newBuilder() + .setValue(ByteString.copyFromUtf8("x")).addKeys("k").build()).build(); + + TransformerOutputStreamObserver responseObserver = new TransformerOutputStreamObserver(2); + var stub = SourceTransformGrpc.newStub(inProcessChannel); + var requestStreamObserver = stub.sourceTransformFn(responseObserver); + requestStreamObserver.onNext(handshake); + requestStreamObserver.onNext(req); + try { + responseObserver.done.get(); + } catch (InterruptedException | ExecutionException e) { + fail("Error while waiting for response" + e.getMessage()); + } + List responses = responseObserver.getResponses().subList(1, 2); + Sourcetransformer.SourceTransformResponse.Result r = responses.get(0).getResults(0); + assertEquals(Arrays.asList("U+005C__FAIL__"), r.getTagsList()); + assertEquals(TEST_EVENT_TIME.getEpochSecond(), r.getEventTime().getSeconds()); + requestStreamObserver.onCompleted(); + } + + private static class FailTransformer extends SourceTransformer { + @Override + public MessageList processMessage(String[] keys, Datum datum) { + return MessageList.newBuilder() + .addMessage(Message.toFail(TEST_EVENT_TIME)) + .build(); + } + } +}