diff --git a/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/input/WorkerInputManager.java b/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/input/WorkerInputManager.java index 1a3cd2c86..5fc4971db 100644 --- a/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/input/WorkerInputManager.java +++ b/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/input/WorkerInputManager.java @@ -139,6 +139,7 @@ public void loadGraph() { "sending edges", e); }).join(); this.sendManager.finishSend(MessageType.EDGE); + this.sendManager.checkFatal(); this.sendManager.clearBuffer(); } diff --git a/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/master/MasterService.java b/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/master/MasterService.java index da01fa7b2..e1e6ff748 100644 --- a/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/master/MasterService.java +++ b/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/master/MasterService.java @@ -159,7 +159,7 @@ public synchronized void close() { LOG.error("Error occurred while closing master service", e); } - if (!failed && this.bsp4Master != null) { + if (this.inited && !failed && this.bsp4Master != null) { this.bsp4Master.waitWorkersCloseDone(); } diff --git a/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/MessageSendManager.java b/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/MessageSendManager.java index bda242ab0..8a53321d3 100644 --- a/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/MessageSendManager.java +++ b/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/MessageSendManager.java @@ -150,6 +150,7 @@ public void startSend(MessageType type) { .map(this.partitioner::workerId) .collect(Collectors.toSet()); this.sendControlMessageToWorkers(workerIds, MessageType.START); + this.sender.checkFatal(); LOG.info("Start sending message(type={})", type); } @@ -166,6 +167,7 @@ public void finishSend(MessageType type) { .map(this.partitioner::workerId) .collect(Collectors.toSet()); this.sendControlMessageToWorkers(workerIds, MessageType.FINISH); + this.sender.checkFatal(); LOG.info("Finish sending message(type={},count={},bytes={})", type, stat.messageCount(), stat.messageBytes()); } @@ -178,6 +180,10 @@ public void clearBuffer() { this.buffers.clear(); } + public void checkFatal() { + this.checkException(); + } + private void sortIfTargetBufferIsFull(WriteBuffers buffer, int partitionId, MessageType type) { @@ -286,6 +292,7 @@ private void sendControlMessageToWorkers(Set workerIds, } private void checkException() { + this.sender.checkFatal(); if (this.exception.get() != null) { throw new ComputerException("Failed to send message", this.exception.get()); diff --git a/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/MessageSender.java b/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/MessageSender.java index a700b22c9..893706900 100644 --- a/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/MessageSender.java +++ b/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/MessageSender.java @@ -45,4 +45,14 @@ CompletableFuture send(int workerId, MessageType type) * an exception is thrown processing message. */ void transportExceptionCaught(TransportException cause, ConnectionId connectionId); + + /** + * Check whether the sender has encountered a fatal error. Implementations + * that run background threads should propagate the first fatal error to + * callers so that the caller can fail fast instead of hanging on a future + * or barrier. + */ + default void checkFatal() { + // no-op by default + } } diff --git a/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/QueuedMessage.java b/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/QueuedMessage.java index 7401daea4..fdb1d2ddd 100644 --- a/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/QueuedMessage.java +++ b/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/QueuedMessage.java @@ -18,6 +18,7 @@ package org.apache.hugegraph.computer.core.sender; import java.nio.ByteBuffer; +import java.util.concurrent.CompletableFuture; import org.apache.hugegraph.computer.core.network.message.MessageType; @@ -26,11 +27,18 @@ public class QueuedMessage { private final int partitionId; private final MessageType type; private final ByteBuffer buffer; + private final CompletableFuture controlFuture; public QueuedMessage(int partitionId, MessageType type, ByteBuffer buffer) { + this(partitionId, type, buffer, null); + } + + QueuedMessage(int partitionId, MessageType type, ByteBuffer buffer, + CompletableFuture controlFuture) { this.partitionId = partitionId; this.type = type; this.buffer = buffer; + this.controlFuture = controlFuture; } public int partitionId() { @@ -44,4 +52,8 @@ public MessageType type() { public ByteBuffer buffer() { return this.buffer; } + + CompletableFuture controlFuture() { + return this.controlFuture; + } } diff --git a/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/QueuedMessageSender.java b/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/QueuedMessageSender.java index b2006b886..4e2410311 100644 --- a/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/QueuedMessageSender.java +++ b/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/sender/QueuedMessageSender.java @@ -44,6 +44,8 @@ public class QueuedMessageSender implements MessageSender { private final Thread sendExecutor; private final BarrierEvent anyQueueNotEmptyEvent; private final BarrierEvent anyClientNotBusyEvent; + private final AtomicReference fatalError; + private volatile boolean closed; public QueuedMessageSender(Config config) { int workerCount = config.get(ComputerOptions.JOB_WORKERS_COUNT); @@ -53,6 +55,7 @@ public QueuedMessageSender(Config config) { this.sendExecutor = new Thread(new Sender(), NAME); this.anyQueueNotEmptyEvent = new BarrierEvent(); this.anyClientNotBusyEvent = new BarrierEvent(); + this.fatalError = new AtomicReference<>(); } public void init() { @@ -63,15 +66,30 @@ public void init() { } public void close() { + this.closed = true; this.sendExecutor.interrupt(); try { this.sendExecutor.join(); } catch (InterruptedException e) { + Thread.currentThread().interrupt(); throw new ComputerException("Interrupted when waiting for " + "send-executor to stop", e); } } + @Override + public void checkFatal() { + Throwable error = this.fatalError.get(); + if (error != null) { + throw new ComputerException("Send-executor encountered fatal error", + error); + } + } + + private void recordFatal(Throwable error) { + this.fatalError.compareAndSet(null, error); + } + public void addWorkerClient(int workerId, TransportClient client) { MessageQueue queue = new MessageQueue( this.anyQueueNotEmptyEvent::signal); @@ -84,22 +102,37 @@ public void addWorkerClient(int workerId, TransportClient client) { @Override public CompletableFuture send(int workerId, MessageType type) throws InterruptedException { + this.checkFatal(); + E.checkArgument(type == MessageType.START || + type == MessageType.FINISH, + "The control message type must be START or FINISH, " + + "but got '%s'", type); WorkerChannel channel = this.channels[channelId(workerId)]; - CompletableFuture future = channel.newFuture(); - future.whenComplete((r, e) -> { - channel.resetFuture(future); - }); + CompletableFuture future = new CompletableFuture<>(); + if (!channel.setControlFuture(future)) { + return future; + } /* * Control message just need message type is enough, * partitionId = -1 and buffer = null represents a meaningless value */ - channel.queue.put(new QueuedMessage(-1, type, null)); + try { + channel.queue.put(new QueuedMessage(-1, type, null, future)); + } catch (InterruptedException e) { + channel.completeControlFuture(future, e); + throw e; + } return future; } @Override public void send(int workerId, QueuedMessage message) throws InterruptedException { + this.checkFatal(); + E.checkArgument(message.type() != null && + message.type().category() == MessageType.Category.DATA, + "The queued message type must be DATA, but got '%s'", + message.type()); WorkerChannel channel = this.channels[channelId(workerId)]; channel.queue.put(message); } @@ -108,7 +141,7 @@ public void send(int workerId, QueuedMessage message) public void transportExceptionCaught(TransportException cause, ConnectionId connectionId) { for (WorkerChannel channel : this.channels) { if (channel.client.connectionId().equals(connectionId)) { - channel.futureRef.get().completeExceptionally(cause); + channel.transportExceptionCaught(cause); } } } @@ -127,57 +160,75 @@ private class Sender implements Runnable { public void run() { LOG.info("The send-executor is running"); Thread thread = Thread.currentThread(); - while (!thread.isInterrupted()) { - try { - int emptyQueueCount = 0; - int busyClientCount = 0; - for (WorkerChannel channel : channels) { - QueuedMessage message = channel.queue.peek(); - if (message == null) { - ++emptyQueueCount; - continue; + try { + while (!thread.isInterrupted()) { + try { + int emptyQueueCount = 0; + int busyClientCount = 0; + for (WorkerChannel channel : channels) { + QueuedMessage message = channel.queue.peek(); + if (message == null) { + ++emptyQueueCount; + continue; + } + try { + if (channel.doSend(message)) { + // Only consume the message after it is sent + channel.queue.take(); + } else { + ++busyClientCount; + } + } catch (TransportException | RuntimeException e) { + channel.failDataSend(e); + // Discard the failed data message to keep sending + channel.queue.take(); + LOG.warn("Failed to send {} message to {}, " + + "discard it", message.type(), channel, e); + } } - if (channel.doSend(message)) { - // Only consume the message after it is sent - channel.queue.take(); - } else { - ++busyClientCount; + int channelCount = channels.length; + /* + * If all queues are empty, let send thread wait + * until any queue is available + */ + if (emptyQueueCount >= channelCount) { + LOG.debug("The send executor was blocked " + + "to wait any queue not empty"); + QueuedMessageSender.this.waitAnyQueueNotEmpty(); + } + /* + * If all clients are busy, let send thread wait + * until any client is available + */ + if (busyClientCount >= channelCount) { + LOG.debug("The send executor was blocked " + + "to wait any client not busy"); + QueuedMessageSender.this.waitAnyClientNotBusy(); + } + } catch (InterruptedException e) { + // Reset interrupted flag + thread.interrupt(); + if (QueuedMessageSender.this.closed) { + // Normal shutdown path + return; + } + // Any client is active means that sending task in running + if (QueuedMessageSender.this.activeClientCount() > 0) { + throw new ComputerException( + "Interrupted when waiting for message " + + "queue not empty"); } } - int channelCount = channels.length; - /* - * If all queues are empty, let send thread wait - * until any queue is available - */ - if (emptyQueueCount >= channelCount) { - LOG.debug("The send executor was blocked " + - "to wait any queue not empty"); - QueuedMessageSender.this.waitAnyQueueNotEmpty(); - } - /* - * If all clients are busy, let send thread wait - * until any client is available - */ - if (busyClientCount >= channelCount) { - LOG.debug("The send executor was blocked " + - "to wait any client not busy"); - QueuedMessageSender.this.waitAnyClientNotBusy(); - } - } catch (InterruptedException e) { - // Reset interrupted flag - thread.interrupt(); - // Any client is active means that sending task in running - if (QueuedMessageSender.this.activeClientCount() > 0) { - throw new ComputerException( - "Interrupted when waiting for message " + - "queue not empty"); - } - } catch (TransportException e) { - // TODO: should handle this in main workflow thread - throw new ComputerException("Failed to send message", e); } + } catch (Throwable t) { + if (!QueuedMessageSender.this.closed) { + QueuedMessageSender.this.recordFatal(t); + LOG.error("The send-executor terminated unexpectedly", t); + } + return; + } finally { + LOG.info("The send-executor is terminated"); } - LOG.info("The send-executor is terminated"); } } @@ -198,8 +249,14 @@ private void waitAnyClientNotBusy() { } catch (InterruptedException e) { // Reset interrupted flag Thread.currentThread().interrupt(); - throw new ComputerException("Interrupted when waiting any client " + - "not busy"); + if (this.closed) { + // Normal shutdown, do not treat as error + return; + } + ComputerException error = new ComputerException( + "Interrupted when waiting any client not busy"); + this.recordFatal(error); + throw error; } finally { this.anyClientNotBusyEvent.reset(); } @@ -227,75 +284,139 @@ private static class WorkerChannel { private final MessageQueue queue; // Each target worker has a TransportClient private final TransportClient client; - private final AtomicReference> futureRef; + private final AtomicReference> controlFutureRef; + private final AtomicReference dataFailureRef; public WorkerChannel(int workerId, MessageQueue queue, TransportClient client) { this.workerId = workerId; this.queue = queue; this.client = client; - this.futureRef = new AtomicReference<>(); - } - - public CompletableFuture newFuture() { - CompletableFuture future = new CompletableFuture<>(); - if (!this.futureRef.compareAndSet(null, future)) { - throw new ComputerException("The origin future must be null"); - } - return future; - } - - public void resetFuture(CompletableFuture future) { - if (!this.futureRef.compareAndSet(future, null)) { - throw new ComputerException("Failed to reset futureRef, " + - "expect future object is %s, " + - "but some thread modified it", - future); - } + this.controlFutureRef = new AtomicReference<>(); + this.dataFailureRef = new AtomicReference<>(); } public boolean doSend(QueuedMessage message) throws TransportException, InterruptedException { switch (message.type()) { case START: - this.sendStartMessage(); + this.sendStartMessage(this.controlFuture(message)); return true; case FINISH: - this.sendFinishMessage(); + this.sendFinishMessage(this.controlFuture(message)); return true; default: return this.sendDataMessage(message); } } - public void sendStartMessage() throws TransportException { - this.client.startSessionAsync().whenComplete((r, e) -> { - CompletableFuture future = this.futureRef.get(); - assert future != null; - - if (e != null) { - LOG.info("Failed to start session connected to {}", this); - future.completeExceptionally(e); - } else { - LOG.info("Start session connected to {}", this); - future.complete(null); + private CompletableFuture controlFuture(QueuedMessage message) { + CompletableFuture future = message.controlFuture(); + E.checkState(future != null, + "The control future can't be null for message '%s'", + message.type()); + return future; + } + + public void sendStartMessage(CompletableFuture future) { + if (!this.controlFutureInFlight(future)) { + return; + } + try { + this.client.startSessionAsync().whenComplete((r, e) -> { + if (e != null) { + LOG.info("Failed to start session connected to {}", this); + } else { + LOG.info("Start session connected to {}", this); + } + this.completeControlFuture(future, e); + }); + } catch (TransportException e) { + this.completeControlFuture(future, e); + } catch (RuntimeException e) { + this.completeControlFuture(future, e); + } + } + + public void sendFinishMessage(CompletableFuture future) { + if (!this.controlFutureInFlight(future)) { + return; + } + try { + this.client.finishSessionAsync().whenComplete((r, e) -> { + if (e != null) { + LOG.info("Failed to finish session connected to {}", this); + } else { + LOG.info("Finish session connected to {}", this); + } + this.completeControlFuture(future, e); + }); + } catch (TransportException e) { + this.completeControlFuture(future, e); + } catch (RuntimeException e) { + this.completeControlFuture(future, e); + } + } + + public void transportExceptionCaught(TransportException cause) { + CompletableFuture future = this.controlFutureRef.get(); + if (future == null) { + this.failDataSend(cause); + } else { + this.completeControlFuture(future, cause); + } + } + + public void failControlFuture(Throwable cause) { + CompletableFuture future = this.controlFutureRef.getAndSet(null); + if (future != null) { + future.completeExceptionally(cause); + } + } + + public void failDataSend(Throwable cause) { + this.dataFailureRef.compareAndSet(null, cause); + this.failControlFuture(this.dataFailureRef.get()); + } + + private boolean setControlFuture(CompletableFuture future) { + Throwable failure = this.dataFailureRef.get(); + if (failure != null) { + future.completeExceptionally(failure); + return false; + } + if (this.controlFutureRef.compareAndSet(null, future)) { + failure = this.dataFailureRef.get(); + if (failure == null) { + return true; } - }); + this.completeControlFuture(future, failure); + return false; + } + ComputerException e = new ComputerException( + "The origin future must be null"); + future.completeExceptionally(e); + return false; + } + + private boolean controlFutureInFlight(CompletableFuture future) { + return future != null && this.controlFutureRef.get() == future; } - public void sendFinishMessage() throws TransportException { - this.client.finishSessionAsync().whenComplete((r, e) -> { - CompletableFuture future = this.futureRef.get(); - assert future != null; - - if (e != null) { - LOG.info("Failed to finish session connected to {}", this); - future.completeExceptionally(e); - } else { - LOG.info("Finish session connected to {}", this); - future.complete(null); + private void completeControlFuture(CompletableFuture future, + Throwable cause) { + E.checkState(future != null, "The control future can't be null"); + if (!this.controlFutureRef.compareAndSet(future, null)) { + if (cause != null) { + this.failDataSend(cause); } - }); + return; + } + if (cause == null) { + future.complete(null); + } else { + future.completeExceptionally(cause); + } } public boolean sendDataMessage(QueuedMessage message) diff --git a/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/worker/WorkerService.java b/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/worker/WorkerService.java index 8d776b327..5a85f82c2 100644 --- a/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/worker/WorkerService.java +++ b/computer/computer-core/src/main/java/org/apache/hugegraph/computer/core/worker/WorkerService.java @@ -63,6 +63,7 @@ public class WorkerService implements Closeable { private volatile boolean inited; private volatile boolean closed; + private volatile boolean registered; private final ComputerContext context; private final Map workers; @@ -83,9 +84,15 @@ public WorkerService() { this.workers = new HashMap<>(); this.inited = false; this.closed = false; + this.registered = false; this.shutdownHook = new ShutdownHook(); } + WorkerService(Bsp4Worker bsp4Worker) { + this(); + this.bsp4Worker = bsp4Worker; + } + /** * Init worker service, create the managers used by worker service. */ @@ -101,7 +108,9 @@ public synchronized void init(Config config) { this.workerInfo = new ContainerInfo(); LOG.info("{} Start to initialize worker", this); - this.bsp4Worker = new Bsp4Worker(this.config, this.workerInfo); + if (this.bsp4Worker == null) { + this.bsp4Worker = new Bsp4Worker(this.config, this.workerInfo); + } /* * Keep the waitMasterInitDone() called before initManagers(), * in order to ensure master init() before worker managers init() @@ -113,6 +122,7 @@ public synchronized void init(Config config) { LOG.info("{} register WorkerService", this); this.bsp4Worker.workerInitDone(); + this.registered = true; this.connectToWorkers(); this.computeManager = new ComputeManager(this.workerInfo.id(), this.context, @@ -176,7 +186,6 @@ public synchronized void close() { this.computeManager.close(); } else { LOG.warn("The computeManager is null"); - return; } } catch (Exception e) { LOG.error("Error when closing ComputeManager", e); @@ -194,8 +203,12 @@ public synchronized void close() { } try { - this.bsp4Worker.workerCloseDone(); - this.bsp4Worker.close(); + if (this.bsp4Worker != null) { + if (this.registered) { + this.bsp4Worker.workerCloseDone(); + } + this.bsp4Worker.close(); + } } catch (Exception e) { LOG.error("Error while closing bsp4Worker", e); } @@ -304,7 +317,7 @@ public String toString() { return String.format("[worker %s]", id); } - private InetSocketAddress initManagers(ContainerInfo masterInfo) { + InetSocketAddress initManagers(ContainerInfo masterInfo) { // Create managers WorkerRpcManager rpcManager = new WorkerRpcManager(); this.managers.add(rpcManager); @@ -375,6 +388,10 @@ private SuperstepStat inputstep() { WorkerInputManager manager = this.managers.get(WorkerInputManager.NAME); manager.loadGraph(); + // Fail fast if the sender thread died, before signaling workerInputDone + MessageSendManager sendManager = this.managers.get(MessageSendManager.NAME); + sendManager.checkFatal(); + this.bsp4Worker.workerInputDone(); this.bsp4Worker.waitMasterInputDone(); diff --git a/computer/computer-test/src/main/java/org/apache/hugegraph/computer/core/network/netty/NettyTransportClientTest.java b/computer/computer-test/src/main/java/org/apache/hugegraph/computer/core/network/netty/NettyTransportClientTest.java index fac4046b6..6f254bdad 100644 --- a/computer/computer-test/src/main/java/org/apache/hugegraph/computer/core/network/netty/NettyTransportClientTest.java +++ b/computer/computer-test/src/main/java/org/apache/hugegraph/computer/core/network/netty/NettyTransportClientTest.java @@ -127,7 +127,6 @@ public void testSend() throws IOException { @Test public void testDataUniformity() throws IOException { - NettyTransportClient client = (NettyTransportClient) this.oneClient(); byte[] sourceBytes1 = StringEncodeUtil.encode("test data message"); byte[] sourceBytes2 = StringEncodeUtil.encode("test data edge"); byte[] sourceBytes3 = StringEncodeUtil.encode("test data vertex"); @@ -165,6 +164,7 @@ public void testDataUniformity() throws IOException { return null; }).when(serverHandler).handle(Mockito.any(), Mockito.eq(1), Mockito.any()); + NettyTransportClient client = (NettyTransportClient) this.oneClient(); client.startSession(); client.send(MessageType.MSG, 1, ByteBuffer.wrap(sourceBytes1)); client.send(MessageType.EDGE, 1, ByteBuffer.wrap(sourceBytes2)); @@ -274,12 +274,12 @@ public void testFlowControl() throws IOException { @Test public void testHandlerException() throws IOException { - NettyTransportClient client = (NettyTransportClient) this.oneClient(); - client.startSession(); - Mockito.doThrow(new RuntimeException("test exception")).when(serverHandler) .handle(Mockito.any(), Mockito.anyInt(), Mockito.any()); + NettyTransportClient client = (NettyTransportClient) this.oneClient(); + client.startSession(); + ByteBuffer buffer = ByteBuffer.wrap(StringEncodeUtil.encode("test data")); boolean send = client.send(MessageType.MSG, 1, buffer); Assert.assertTrue(send); diff --git a/computer/computer-test/src/main/java/org/apache/hugegraph/computer/core/sender/QueuedMessageSenderTest.java b/computer/computer-test/src/main/java/org/apache/hugegraph/computer/core/sender/QueuedMessageSenderTest.java index 07fd15929..328080712 100644 --- a/computer/computer-test/src/main/java/org/apache/hugegraph/computer/core/sender/QueuedMessageSenderTest.java +++ b/computer/computer-test/src/main/java/org/apache/hugegraph/computer/core/sender/QueuedMessageSenderTest.java @@ -17,8 +17,21 @@ package org.apache.hugegraph.computer.core.sender; +import java.net.InetSocketAddress; +import java.nio.ByteBuffer; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; +import java.util.concurrent.atomic.AtomicReference; + +import org.apache.hugegraph.computer.core.common.exception.TransportException; import org.apache.hugegraph.computer.core.config.ComputerOptions; import org.apache.hugegraph.computer.core.config.Config; +import org.apache.hugegraph.computer.core.network.ConnectionId; +import org.apache.hugegraph.computer.core.network.TransportClient; +import org.apache.hugegraph.computer.core.network.message.MessageType; import org.apache.hugegraph.computer.core.worker.MockComputation2; import org.apache.hugegraph.computer.suite.unit.UnitTestBase; import org.apache.hugegraph.testutil.Assert; @@ -47,21 +60,557 @@ public void setup() { ); } + private QueuedMessageSender newSender(TransportClient first, TransportClient second) { + QueuedMessageSender sender = new QueuedMessageSender(this.config); + sender.addWorkerClient(1, first); + sender.addWorkerClient(2, second); + sender.init(); + return sender; + } + @Test public void testInitAndClose() { + QueuedMessageSender sender = this.newSender(new MockTransportClient(), + new MockTransportClient()); + Thread sendExecutor = Whitebox.getInternalState(sender, + "sendExecutor"); + try { + Assert.assertTrue(ImmutableSet.of(Thread.State.NEW, + Thread.State.RUNNABLE, + Thread.State.WAITING) + .contains(sendExecutor.getState())); + } finally { + sender.close(); + } + Assert.assertEquals(Thread.State.TERMINATED, sendExecutor.getState()); + } + + @Test + public void testRejectsMessageTypeFromWrongOverload() { QueuedMessageSender sender = new QueuedMessageSender(this.config); - sender.addWorkerClient(1, new MockTransportClient()); - sender.addWorkerClient(2, new MockTransportClient()); - sender.init(); + sender.addWorkerClient(1, new ControlFutureClient()); + + Assert.assertThrows(IllegalArgumentException.class, () -> { + sender.send(1, MessageType.MSG); + }); + Assert.assertThrows(IllegalArgumentException.class, () -> { + sender.send(1, MessageType.PING); + }); + Assert.assertThrows(IllegalArgumentException.class, () -> { + sender.send(1, new QueuedMessage(-1, MessageType.START, null)); + }); + Assert.assertThrows(IllegalArgumentException.class, () -> { + sender.send(1, new QueuedMessage(-1, MessageType.FINISH, null)); + }); + } + + @Test + public void testControlBeforeCompletionFinishes() throws Exception { + ControlFutureClient client = new ControlFutureClient(); + QueuedMessageSender sender = this.newSender(client, new MockTransportClient()); + + CountDownLatch completionStarted = new CountDownLatch(1); + CountDownLatch allowCompletion = new CountDownLatch(1); + Thread completionThread = null; + try { + CompletableFuture startFuture = sender.send(1, MessageType.START); + Assert.assertTrue(await(client.startCalled)); + startFuture.whenComplete((r, e) -> { + completionStarted.countDown(); + try { + allowCompletion.await(); + } catch (InterruptedException exception) { + Thread.currentThread().interrupt(); + throw new AssertionError(exception); + } + }); + + completionThread = new Thread( + () -> client.startFuture.complete(null)); + completionThread.start(); + Assert.assertTrue(completionStarted.await(1, TimeUnit.SECONDS)); + + CompletableFuture finishFuture = sender.send(1, MessageType.FINISH); + Assert.assertTrue(await(client.finishCalled)); + allowCompletion.countDown(); + completionThread.join(TimeUnit.SECONDS.toMillis(1)); + Assert.assertFalse(completionThread.isAlive()); + client.finishFuture.complete(null); + finishFuture.get(1, TimeUnit.SECONDS); + } finally { + allowCompletion.countDown(); + if (completionThread != null) { + completionThread.join(TimeUnit.SECONDS.toMillis(1)); + } + sender.close(); + } + } + + @Test + public void testTransportExceptionControlFuture() throws Exception { + ControlFutureClient client = new ControlFutureClient(); + QueuedMessageSender sender = this.newSender(client, new MockTransportClient()); + + try { + CompletableFuture startFuture = sender.send(1, MessageType.START); + Assert.assertTrue(await(client.startCalled)); + + TransportException cause = new TransportException("connection failed"); + sender.transportExceptionCaught(cause, client.connectionId()); + assertFutureFailedWith(startFuture, cause); + + CompletableFuture finishFuture = sender.send(1, MessageType.FINISH); + Assert.assertTrue(await(client.finishCalled)); + client.startFuture.complete(null); + assertFutureFailedWith(startFuture, cause); + Assert.assertFalse(finishFuture.isDone()); + + client.finishFuture.complete(null); + finishFuture.get(1, TimeUnit.SECONDS); + } finally { + sender.close(); + } + } + + @Test + public void testExceptionalCompletionCasLossFailsNextControl() + throws Exception { + ControlFutureClient client = new ControlFutureClient(); + QueuedMessageSender sender = this.newSender( + client, new MockTransportClient()); + CountDownLatch failureObserved = new CountDownLatch(1); + CountDownLatch resumeFailure = new CountDownLatch(1); + Thread failureThread = null; + + try { + CompletableFuture startFuture = sender.send( + 1, MessageType.START); + Assert.assertTrue(await(client.startCalled)); + + Object[] channels = Whitebox.getInternalState(sender, "channels"); + Object channel = channels[0]; + AtomicReference> controlFutureRef = + Whitebox.getInternalState(channel, "controlFutureRef"); + AtomicReference> observedFuture = + new AtomicReference<>(); + TransportException cause = + new TransportException("connection failed"); + failureThread = new Thread(() -> { + observedFuture.set(controlFutureRef.get()); + failureObserved.countDown(); + try { + resumeFailure.await(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + throw new AssertionError(e); + } + Whitebox.invoke(channel.getClass(), new Class[] { + CompletableFuture.class, + Throwable.class}, + "completeControlFuture", channel, + observedFuture.get(), cause); + }); + failureThread.start(); + Assert.assertTrue(await(failureObserved)); + + client.startFuture.complete(null); + startFuture.get(1, TimeUnit.SECONDS); + CompletableFuture finishFuture = sender.send( + 1, MessageType.FINISH); + Assert.assertTrue(await(client.finishCalled)); + + resumeFailure.countDown(); + assertFutureFailedWith(finishFuture, cause); + client.finishFuture.complete(null); + } finally { + resumeFailure.countDown(); + if (failureThread != null) { + failureThread.join(TimeUnit.SECONDS.toMillis(1L)); + } + sender.close(); + } + } + + @Test + public void testTransportExceptionDispatch() throws Exception { + ControlFutureClient client = new ControlFutureClient(); + QueuedMessageSender sender = this.newSender(client, new MockTransportClient()); + + client.blockDataSend = true; + try { + sender.send(1, new QueuedMessage(0, MessageType.MSG, ByteBuffer.allocate(1))); + Assert.assertTrue(await(client.dataSendCalled)); + + CompletableFuture startFuture = sender.send(1, MessageType.START); + TransportException cause = new TransportException("connection failed before start"); + sender.transportExceptionCaught(cause, client.connectionId()); + assertFutureFailedWith(startFuture, cause); + + client.allowDataSend.countDown(); + Assert.assertFalse(await(client.startCalled)); + } finally { + client.allowDataSend.countDown(); + sender.close(); + } + } + + @Test + public void testExecutorAlive() throws Exception { + ControlFutureClient client = new ControlFutureClient(); + QueuedMessageSender sender = this.newSender(client, new MockTransportClient()); + + try { + RuntimeException startCause = new IllegalArgumentException("start session failed"); + client.startFailure = startCause; + CompletableFuture startFuture = sender.send(1, MessageType.START); + assertFutureFailedWith(startFuture, startCause); + + RuntimeException finishCause = new IllegalArgumentException("finish session failed"); + client.finishFailure = finishCause; + CompletableFuture finishFuture = sender.send(1, MessageType.FINISH); + assertFutureFailedWith(finishFuture, finishCause); + + Thread sendExecutor = Whitebox.getInternalState(sender, "sendExecutor"); + sendExecutor.join(TimeUnit.SECONDS.toMillis(1)); + Assert.assertTrue(sendExecutor.isAlive()); + + client.startFailure = null; + CompletableFuture nextStartFuture = sender.send(1, MessageType.START); + Assert.assertTrue(await(client.startCalled)); + client.startFuture.complete(null); + nextStartFuture.get(1, TimeUnit.SECONDS); + } finally { + sender.close(); + } + } + + @Test + public void testAsyncControlFutureFailures() throws Exception { + ControlFutureClient client = new ControlFutureClient(); + QueuedMessageSender sender = this.newSender(client, + new MockTransportClient()); + + try { + CompletableFuture startFuture = sender.send( + 1, MessageType.START); + Assert.assertTrue(await(client.startCalled)); + TransportException startCause = + new TransportException("async start failed"); + client.startFuture.completeExceptionally(startCause); + assertFutureFailedWith(startFuture, startCause); + + CompletableFuture finishFuture = sender.send( + 1, MessageType.FINISH); + Assert.assertTrue(await(client.finishCalled)); + TransportException finishCause = + new TransportException("async finish failed"); + client.finishFuture.completeExceptionally(finishCause); + assertFutureFailedWith(finishFuture, finishCause); + } finally { + sender.close(); + } + } + + @Test + public void testOtherClients() throws Exception { + ControlFutureClient failedClient = new ControlFutureClient(1); + ControlFutureClient activeClient = new ControlFutureClient(2); + QueuedMessageSender sender = this.newSender(failedClient, activeClient); + + try { + Assert.assertFalse(failedClient.connectionId() + .equals(activeClient.connectionId())); + TransportException startCause = + new TransportException("start session failed"); + failedClient.startFailure = startCause; + CompletableFuture failedStart = sender.send( + 1, MessageType.START); + assertFutureFailedWith(failedStart, startCause); + + CompletableFuture activeStart = sender.send(2, MessageType.START); + Assert.assertTrue(await(activeClient.startCalled)); + activeClient.startFuture.complete(null); + activeStart.get(1, TimeUnit.SECONDS); + + failedClient.startFailure = null; + CompletableFuture callbackStart = sender.send( + 1, MessageType.START); + Assert.assertTrue(await(failedClient.startCalled)); + TransportException callbackCause = + new TransportException("connection failed"); + sender.transportExceptionCaught(callbackCause, + failedClient.connectionId()); + assertFutureFailedWith(callbackStart, callbackCause); + + CompletableFuture activeStartAfterCallback = sender.send( + 2, MessageType.START); + activeStartAfterCallback.get(1, TimeUnit.SECONDS); + + TransportException finishCause = + new TransportException("finish session failed"); + failedClient.finishFailure = finishCause; + CompletableFuture failedFinish = sender.send( + 1, MessageType.FINISH); + assertFutureFailedWith(failedFinish, finishCause); + + CompletableFuture activeFinish = sender.send(2, MessageType.FINISH); + Assert.assertTrue(await(activeClient.finishCalled)); + activeClient.finishFuture.complete(null); + activeFinish.get(1, TimeUnit.SECONDS); + + Thread sendExecutor = Whitebox.getInternalState(sender, "sendExecutor"); + Assert.assertTrue(sendExecutor.isAlive()); + } finally { + sender.close(); + } + } + + @Test + public void testQueuedFinish() throws Exception { + this.assertSynchronousDataFailureCompletesQueuedFinish( + new TransportException("data send failed")); + } + + @Test + public void testCompletesQueuedFinish() throws Exception { + this.assertSynchronousDataFailureCompletesQueuedFinish( + new IllegalStateException("data send failed")); + } + + @Test + public void testFinishFailsFinish() throws Exception { + this.assertSynchronousDataFailureBeforeFinishFailsFinish( + new TransportException("data send failed before finish")); + } + + @Test + public void testDataRuntimeFinish() throws Exception { + this.assertSynchronousDataFailureBeforeFinishFailsFinish( + new IllegalStateException("data send failed before finish")); + } + + @Test + public void testConflictKeepsSendExecutorAlive() throws Exception { + ControlFutureClient client = new ControlFutureClient(); + QueuedMessageSender sender = this.newSender(client, new MockTransportClient()); + + try { + CompletableFuture startFuture = sender.send(1, MessageType.START); + Assert.assertTrue(await(client.startCalled)); + + CompletableFuture conflictingFinishFuture = sender.send(1, MessageType.FINISH); + assertFutureFailedWithMessage(conflictingFinishFuture, "The origin future must be null"); + + Thread sendExecutor = Whitebox.getInternalState(sender, "sendExecutor"); + sendExecutor.join(TimeUnit.SECONDS.toMillis(1)); + Assert.assertTrue(sendExecutor.isAlive()); + + client.startFuture.complete(null); + startFuture.get(1, TimeUnit.SECONDS); + + CompletableFuture finishFuture = sender.send(1, MessageType.FINISH); + Assert.assertTrue(await(client.finishCalled)); + client.finishFuture.complete(null); + finishFuture.get(1, TimeUnit.SECONDS); + } finally { + sender.close(); + } + } + + private static void assertFutureFailedWith(CompletableFuture future, Throwable cause) + throws InterruptedException, TimeoutException { + try { + future.get(1, TimeUnit.SECONDS); + Assert.fail("Expected control future to fail"); + } catch (ExecutionException exception) { + Assert.assertSame(cause, exception.getCause()); + } + } + + private static boolean await(CountDownLatch latch) throws InterruptedException { + return latch.await(1, TimeUnit.SECONDS); + } + + private void assertSynchronousDataFailureCompletesQueuedFinish( + Throwable cause) throws Exception { + ControlFutureClient failedClient = new ControlFutureClient(); + ControlFutureClient activeClient = new ControlFutureClient(2); + QueuedMessageSender sender = this.newSender(failedClient, activeClient); + + failedClient.blockDataSend = true; + try { + sender.send(1, new QueuedMessage(0, MessageType.MSG, + ByteBuffer.allocate(1))); + Assert.assertTrue(await(failedClient.dataSendCalled)); + + CompletableFuture finishFuture = sender.send(1, MessageType.FINISH); + failedClient.dataFailure = cause; + failedClient.allowDataSend.countDown(); + assertFutureFailedWith(finishFuture, cause); + waitForQueueEmpty(sender, 1); + Assert.assertEquals(1L, failedClient.finishCalled.getCount()); + + CompletableFuture activeStart = sender.send(2, MessageType.START); + Assert.assertTrue(await(activeClient.startCalled)); + activeClient.startFuture.complete(null); + activeStart.get(1, TimeUnit.SECONDS); + } finally { + failedClient.allowDataSend.countDown(); + sender.close(); + } + } + + @Test + public void testTransportExceptionDuringDataSendFailsLaterFinish() + throws Exception { + ControlFutureClient client = new ControlFutureClient(); + QueuedMessageSender sender = this.newSender( + client, new MockTransportClient()); + + client.blockDataSend = true; + try { + sender.send(1, new QueuedMessage(0, MessageType.MSG, + ByteBuffer.allocate(1))); + Assert.assertTrue(await(client.dataSendCalled)); + + TransportException cause = + new TransportException("connection failed during data"); + sender.transportExceptionCaught(cause, client.connectionId()); + client.allowDataSend.countDown(); + waitForQueueEmpty(sender, 1); + + CompletableFuture finishFuture = sender.send( + 1, MessageType.FINISH); + assertFutureFailedWith(finishFuture, cause); + Assert.assertFalse(await(client.finishCalled)); + } finally { + client.allowDataSend.countDown(); + sender.close(); + } + } + + private void assertSynchronousDataFailureBeforeFinishFailsFinish( + Throwable cause) throws Exception { + ControlFutureClient failedClient = new ControlFutureClient(); + QueuedMessageSender sender = this.newSender( + failedClient, new MockTransportClient()); + + failedClient.dataFailure = cause; + try { + sender.send(1, new QueuedMessage(0, MessageType.MSG, + ByteBuffer.allocate(1))); + waitForQueueEmpty(sender, 1); + + CompletableFuture finishFuture = sender.send(1, MessageType.FINISH); + assertFutureFailedWith(finishFuture, cause); + Assert.assertFalse(await(failedClient.finishCalled)); + } finally { + sender.close(); + } + } + + private static void waitForQueueEmpty(QueuedMessageSender sender, + int workerId) + throws InterruptedException { + Object[] channels = Whitebox.getInternalState(sender, "channels"); + MessageQueue queue = Whitebox.getInternalState(channels[workerId - 1], + "queue"); + long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(1L); + while (queue.peek() != null && System.nanoTime() < deadline) { + Thread.sleep(10L); + } + Assert.assertTrue("Timed out to wait for sender queue to be empty", + queue.peek() == null); + } + + private static void assertFutureFailedWithMessage(CompletableFuture future, + String message) + throws InterruptedException, TimeoutException { + try { + future.get(1, TimeUnit.SECONDS); + Assert.fail("Expected control future to fail"); + } catch (ExecutionException exception) { + Assert.assertContains(message, exception.getCause().getMessage()); + } + } + + private static class ControlFutureClient extends MockTransportClient { + + private final ConnectionId connectionId; + private final CountDownLatch startCalled = new CountDownLatch(1); + private final CountDownLatch finishCalled = new CountDownLatch(1); + private final CountDownLatch dataSendCalled = new CountDownLatch(1); + private final CountDownLatch allowDataSend = new CountDownLatch(1); + private final CompletableFuture startFuture = new CompletableFuture<>(); + private final CompletableFuture finishFuture = new CompletableFuture<>(); + private Throwable startFailure; + private Throwable finishFailure; + private Throwable dataFailure; + private boolean blockDataSend; + + private ControlFutureClient() { + this(1); + } + + private ControlFutureClient(int clientIndex) { + this.connectionId = new ConnectionId( + new InetSocketAddress("localhost", 8080), + clientIndex); + } + + @Override + public ConnectionId connectionId() { + return this.connectionId; + } + + @Override + public CompletableFuture startSessionAsync() throws TransportException { + throwFailure(this.startFailure); + this.startCalled.countDown(); + return this.startFuture; + } + + @Override + public CompletableFuture finishSessionAsync() throws TransportException { + throwFailure(this.finishFailure); + this.finishCalled.countDown(); + return this.finishFuture; + } + + @Override + public boolean send(MessageType messageType, int partition, + ByteBuffer buffer) throws TransportException { + if (this.blockDataSend) { + this.dataSendCalled.countDown(); + try { + this.allowDataSend.await(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + throw new TransportException("Interrupted data send", e); + } + } + throwFailure(this.dataFailure); + return true; + } + + private static void throwFailure(Throwable failure) throws TransportException { + if (failure instanceof TransportException) { + throw (TransportException) failure; + } + if (failure instanceof RuntimeException) { + throw (RuntimeException) failure; + } + } + + @Override + public boolean sessionActive() { + return false; + } - Thread sendExecutor = Whitebox.getInternalState(sender, "sendExecutor"); - Assert.assertTrue(ImmutableSet.of(Thread.State.NEW, - Thread.State.RUNNABLE, - Thread.State.WAITING) - .contains(sendExecutor.getState())); + @Override + public InetSocketAddress remoteAddress() { + return new InetSocketAddress("127.0.0.1", 8080); + } - sender.close(); - Assert.assertTrue(ImmutableSet.of(Thread.State.TERMINATED) - .contains(sendExecutor.getState())); } } diff --git a/computer/computer-test/src/main/java/org/apache/hugegraph/computer/core/worker/WorkerServiceTest.java b/computer/computer-test/src/main/java/org/apache/hugegraph/computer/core/worker/WorkerServiceTest.java index 3e7d1c0ba..005fcfaeb 100644 --- a/computer/computer-test/src/main/java/org/apache/hugegraph/computer/core/worker/WorkerServiceTest.java +++ b/computer/computer-test/src/main/java/org/apache/hugegraph/computer/core/worker/WorkerServiceTest.java @@ -17,11 +17,14 @@ package org.apache.hugegraph.computer.core.worker; +import java.net.InetSocketAddress; import java.util.Arrays; import java.util.concurrent.CountDownLatch; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; +import org.apache.hugegraph.computer.core.bsp.Bsp4Worker; +import org.apache.hugegraph.computer.core.common.ContainerInfo; import org.apache.hugegraph.computer.core.common.exception.ComputerException; import org.apache.hugegraph.computer.core.config.ComputerOptions; import org.apache.hugegraph.computer.core.config.Config; @@ -33,6 +36,8 @@ import org.apache.hugegraph.testutil.Assert; import org.apache.hugegraph.util.Log; import org.junit.Test; +import org.mockito.InOrder; +import org.mockito.Mockito; import org.slf4j.Logger; public class WorkerServiceTest extends UnitTestBase { @@ -232,6 +237,49 @@ public void testFailToConnectEtcd() { } } + @Test + public void testInitFailsAfterRegistration() { + Bsp4Worker bsp4Worker = Mockito.mock(Bsp4Worker.class); + ContainerInfo masterInfo = new ContainerInfo(ContainerInfo.MASTER_ID, + "localhost", 8099); + Mockito.when(bsp4Worker.waitMasterInitDone()).thenReturn(masterInfo); + Mockito.when(bsp4Worker.waitMasterAllInitDone()) + .thenThrow(new ComputerException( + "Mocked failure to connect to workers after " + + "registration")); + + Config config = UnitTestBase.updateWithRequiredOptions( + ComputerOptions.JOB_ID, "local_005", + ComputerOptions.JOB_WORKERS_COUNT, "1", + ComputerOptions.WORKER_COMPUTATION_CLASS, + MockComputation.class.getName(), + ComputerOptions.ALGORITHM_RESULT_CLASS, + DoubleValue.class.getName(), + ComputerOptions.ALGORITHM_MESSAGE_CLASS, + DoubleValue.class.getName() + ); + + try (WorkerService service = new WorkerService(bsp4Worker) { + @Override + InetSocketAddress initManagers(ContainerInfo masterInfo) { + return new InetSocketAddress(masterInfo.hostname(), + masterInfo.rpcPort()); + } + }) { + Assert.assertThrows(ComputerException.class, () -> { + service.init(config); + }); + } + + InOrder inOrder = Mockito.inOrder(bsp4Worker); + inOrder.verify(bsp4Worker).waitMasterInitDone(); + inOrder.verify(bsp4Worker).workerInitDone(); + inOrder.verify(bsp4Worker).waitMasterAllInitDone(); + inOrder.verify(bsp4Worker).workerCloseDone(); + inOrder.verify(bsp4Worker).close(); + Mockito.verifyNoMoreInteractions(bsp4Worker); + } + @Test public void testDataTransportManagerFail() { /* diff --git a/computer/computer-test/src/main/java/org/apache/hugegraph/computer/suite/integrate/SenderIntegrateTest.java b/computer/computer-test/src/main/java/org/apache/hugegraph/computer/suite/integrate/SenderIntegrateTest.java index 6ae0008be..a81301873 100644 --- a/computer/computer-test/src/main/java/org/apache/hugegraph/computer/suite/integrate/SenderIntegrateTest.java +++ b/computer/computer-test/src/main/java/org/apache/hugegraph/computer/suite/integrate/SenderIntegrateTest.java @@ -17,6 +17,7 @@ package org.apache.hugegraph.computer.suite.integrate; +import java.io.Closeable; import java.io.IOException; import java.util.ArrayList; import java.util.HashMap; @@ -24,8 +25,11 @@ import java.util.Map; import java.util.Set; import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CompletionException; import java.util.concurrent.Future; +import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Consumer; import java.util.function.Function; import org.apache.hugegraph.computer.algorithm.centrality.pagerank.PageRankParams; @@ -43,6 +47,7 @@ import org.apache.hugegraph.computer.core.util.ComputerContextUtil; import org.apache.hugegraph.computer.core.worker.WorkerService; import org.apache.hugegraph.config.RpcOptions; +import org.apache.hugegraph.testutil.Assert; import org.apache.hugegraph.testutil.Whitebox; import org.apache.hugegraph.util.Log; import org.junit.AfterClass; @@ -57,6 +62,9 @@ public class SenderIntegrateTest { public static final Logger LOG = Log.logger(SenderIntegrateTest.class); private static final Class COMPUTATION = MockComputation.class; + private static final long BSP_TIMEOUT_MS = TimeUnit.SECONDS.toMillis(90L); + private static final long SERVICE_WAIT_TIMEOUT = + BSP_TIMEOUT_MS + TimeUnit.SECONDS.toMillis(30L); @BeforeClass public static void init() { @@ -83,17 +91,20 @@ public void testOneWorker() { .withMaxSuperStep(3) .withComputationClass(COMPUTATION) .withWorkerCount(1) - .withBufferThreshold(50) - .withBufferCapacity(60) + // 4KB keeps the sort+send path exercised without + // flooding the ACK-throttled transport (50B stalled) + .withBufferThreshold(4096) + .withBufferCapacity(8192) .withRpcServerHost("127.0.0.1") .withRpcServerPort(8611) - .withRpcServerPort(0) - .build(); - try (MasterService service = initMaster(args)) { - masterServiceRef.set(service); + .withRpcServerPort(0) + .withTestBspTimeouts() + .build(); + try (MasterService service = initMaster( + args, masterServiceRef::set)) { service.execute(); masterFuture.complete(null); - } catch (Exception e) { + } catch (Throwable e) { LOG.error("Failed to execute master service", e); masterFuture.completeExceptionally(e); } @@ -111,12 +122,13 @@ public void testOneWorker() { .withMaxSuperStep(3) .withComputationClass(COMPUTATION) .withWorkerCount(1) - .withBufferThreshold(50) - .withBufferCapacity(60) - .withTransoprtServerPort(0) - .build(); - try (WorkerService service = initWorker(args)) { - workerServiceRef.set(service); + .withBufferThreshold(4096) + .withBufferCapacity(8192) + .withTransportServerPort(0) + .withTestBspTimeouts() + .build(); + try (WorkerService service = initWorker( + args, workerServiceRef::set)) { service.execute(); workerFuture.complete(null); } catch (Throwable e) { @@ -128,11 +140,15 @@ public void testOneWorker() { masterThread.start(); workerThread.start(); + Throwable failure = null; try { - CompletableFuture.allOf(workerFuture, masterFuture).join(); + awaitServices(workerFuture, masterFuture); + } catch (RuntimeException | Error e) { + failure = e; + throw e; } finally { - workerServiceRef.get().close(); - masterServiceRef.get().close(); + closeServices(failure, workerServiceRef.get(), + masterServiceRef.get()); } } @@ -155,11 +171,12 @@ public void testMultiWorkers() throws IOException { .withWorkerCount(workerCount) .withPartitionCount(partitionCount) .withRpcServerHost("127.0.0.1") - .withRpcServerPort(0) - .build(); + .withRpcServerPort(0) + .withTestBspTimeouts() + .build(); try { - MasterService service = initMaster(args); - masterServiceRef.set(service); + MasterService service = initMaster( + args, masterServiceRef::set); service.execute(); masterFuture.complete(null); } catch (Throwable e) { @@ -186,12 +203,13 @@ public void testMultiWorkers() throws IOException { .withComputationClass(COMPUTATION) .withWorkerCount(workerCount) .withPartitionCount(partitionCount) - .withTransoprtServerPort(0) - .withDataDirs(dir) - .build(); + .withTransportServerPort(0) + .withDataDirs(dir) + .withTestBspTimeouts() + .build(); try { - WorkerService service = initWorker(args); - workerServices.add(service); + WorkerService service = initWorker( + args, workerServices::add); service.execute(); workerFuture.complete(null); } catch (Throwable e) { @@ -210,13 +228,16 @@ public void testMultiWorkers() throws IOException { List> futures = new ArrayList<>(workers.values()); futures.add(masterFuture); + Throwable failure = null; try { - CompletableFuture.allOf(futures.toArray(new CompletableFuture[0])).join(); + awaitServices(futures.toArray(new CompletableFuture[0])); + } catch (RuntimeException | Error e) { + failure = e; + throw e; } finally { - for (WorkerService workerService : workerServices) { - workerService.close(); - } - masterServiceRef.get().close(); + List services = new ArrayList<>(workerServices); + services.add(masterServiceRef.get()); + closeServices(failure, services.toArray(new Closeable[0])); } } @@ -238,10 +259,11 @@ public void testOneWorkerWithBusyClient() { .withWriteBufferHighMark(10) .withWriteBufferLowMark(5) .withRpcServerHost("127.0.0.1") - .withRpcServerPort(0) - .build(); - try (MasterService service = initMaster(args)) { - masterServiceRef.set(service); + .withRpcServerPort(0) + .withTestBspTimeouts() + .build(); + try (MasterService service = initMaster( + args, masterServiceRef::set)) { service.execute(); masterFuture.complete(null); } catch (Throwable e) { @@ -252,7 +274,7 @@ public void testOneWorkerWithBusyClient() { masterThread.setDaemon(true); CompletableFuture workerFuture = new CompletableFuture<>(); - int transoprtServerPort = 8998; + int transportServerPort = 8998; Thread workerThread = new Thread(() -> { String[] args = OptionsBuilder.newInstance() .withJobId("local_002") @@ -265,12 +287,14 @@ public void testOneWorkerWithBusyClient() { .withWorkerCount(1) .withWriteBufferHighMark(20) .withWriteBufferLowMark(10) - .withTransoprtServerPort(transoprtServerPort) - .build(); - try (WorkerService service = initWorker(args)) { - workerServiceRef.set(service); + .withTransportServerPort( + transportServerPort) + .withTestBspTimeouts() + .build(); + try (WorkerService service = initWorker( + args, workerServiceRef::set)) { // Let send rate slowly - this.slowSendFunc(service, transoprtServerPort); + this.slowSendFunc(service, transportServerPort); service.execute(); workerFuture.complete(null); } catch (Throwable e) { @@ -282,15 +306,42 @@ public void testOneWorkerWithBusyClient() { masterThread.start(); workerThread.start(); + Throwable failure = null; try { - CompletableFuture.allOf(workerFuture, masterFuture).join(); + awaitServices(workerFuture, masterFuture); + } catch (RuntimeException | Error e) { + failure = e; + throw e; } finally { - workerServiceRef.get().close(); - masterServiceRef.get().close(); + closeServices(failure, workerServiceRef.get(), + masterServiceRef.get()); + } + } + + @Test + public void testFailedServiceReleasesPeerWaitFast() { + CompletableFuture failed = new CompletableFuture<>(); + CompletableFuture unfinished = new CompletableFuture<>(); + IllegalStateException cause = + new IllegalStateException("service failed"); + failed.completeExceptionally(cause); + + long start = System.currentTimeMillis(); + try { + awaitServices(TimeUnit.SECONDS.toMillis(1L), failed, unfinished); + Assert.fail("Expected CompletionException when a service fails"); + } catch (CompletionException e) { + long elapsed = System.currentTimeMillis() - start; + Assert.assertTrue( + "Wait should fail fast on first failure, but waited " + + elapsed + " ms", + elapsed < TimeUnit.SECONDS.toMillis(5L)); + Assert.assertSame(cause, e.getCause()); } } - private void slowSendFunc(WorkerService service, int port) throws TransportException { + private void slowSendFunc(WorkerService service, int port) + throws TransportException { Managers managers = Whitebox.getInternalState(service, "managers"); DataClientManager clientManager = managers.get( DataClientManager.NAME); @@ -306,7 +357,7 @@ private void slowSendFunc(WorkerService service, int port) throws TransportExcep "sendFunction"); Function> sendFunc = message -> { try { - Thread.sleep(100); + Thread.sleep(20); } catch (InterruptedException e) { e.printStackTrace(); } @@ -315,22 +366,77 @@ private void slowSendFunc(WorkerService service, int port) throws TransportExcep Whitebox.setInternalState(clientSession, "sendFunction", sendFunc); } - private MasterService initMaster(String[] args) { + private MasterService initMaster(String[] args, + Consumer register) { Config config = ComputerContextUtil.initContext( ComputerContextUtil.convertToMap(args)); MasterService service = new MasterService(); + register.accept(service); service.init(config); return service; } - private WorkerService initWorker(String[] args) { + private WorkerService initWorker(String[] args, + Consumer register) { Config config = ComputerContextUtil.initContext( ComputerContextUtil.convertToMap(args)); WorkerService service = new WorkerService(); + register.accept(service); service.init(config); return service; } + private static void closeServices(Throwable primaryFailure, + Closeable... services) { + Throwable cleanupFailure = null; + for (Closeable service : services) { + if (service == null) { + continue; + } + try { + service.close(); + } catch (Throwable failure) { + if (primaryFailure != null) { + primaryFailure.addSuppressed(failure); + } else if (cleanupFailure == null) { + cleanupFailure = failure; + } else { + cleanupFailure.addSuppressed(failure); + } + } + } + if (cleanupFailure instanceof RuntimeException) { + throw (RuntimeException) cleanupFailure; + } + if (cleanupFailure instanceof Error) { + throw (Error) cleanupFailure; + } + if (cleanupFailure != null) { + throw new RuntimeException("Failed to close services", + cleanupFailure); + } + } + + private static void awaitServices(CompletableFuture... futures) { + awaitServices(SERVICE_WAIT_TIMEOUT, futures); + } + + private static void awaitServices(long timeout, + CompletableFuture... futures) { + CompletableFuture allDone = CompletableFuture.allOf(futures); + CompletableFuture firstFailure = new CompletableFuture<>(); + for (CompletableFuture future : futures) { + future.whenComplete((result, error) -> { + if (error != null) { + firstFailure.completeExceptionally(error); + } + }); + } + CompletableFuture.anyOf(firstFailure, allDone) + .orTimeout(timeout, TimeUnit.MILLISECONDS) + .join(); + } + private static class OptionsBuilder { private final List options; @@ -415,7 +521,7 @@ public OptionsBuilder withBufferCapacity(int sizeInByte) { return this; } - public OptionsBuilder withTransoprtServerPort(int dataPort) { + public OptionsBuilder withTransportServerPort(int dataPort) { this.options.add(ComputerOptions.TRANSPORT_SERVER_PORT.name()); this.options.add(String.valueOf(dataPort)); return this; @@ -452,5 +558,14 @@ public OptionsBuilder withDataDirs(String dataDirs) { this.options.add(String.valueOf(dataDirs)); return this; } + + public OptionsBuilder withTestBspTimeouts() { + this.options.add(ComputerOptions.BSP_WAIT_WORKERS_TIMEOUT.name()); + this.options.add(String.valueOf(BSP_TIMEOUT_MS)); + this.options.add(ComputerOptions.BSP_WAIT_MASTER_TIMEOUT.name()); + this.options.add(String.valueOf(BSP_TIMEOUT_MS)); + return this; + } + } }