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/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..42532f559 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 @@ -84,22 +84,35 @@ public void addWorkerClient(int workerId, TransportClient client) { @Override public CompletableFuture send(int workerId, MessageType type) throws InterruptedException { + 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 { + 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 +121,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); } } } @@ -137,11 +150,19 @@ public void run() { ++emptyQueueCount; continue; } - if (channel.doSend(message)) { - // Only consume the message after it is sent + 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(); - } else { - ++busyClientCount; + LOG.warn("Failed to send {} message to {}, " + + "discard it", message.type(), channel, e); } } int channelCount = channels.length; @@ -172,9 +193,6 @@ public void run() { "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); } } LOG.info("The send-executor is terminated"); @@ -227,75 +245,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..183ef5fa9 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 @@ -176,7 +176,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 +193,12 @@ public synchronized void close() { } try { - this.bsp4Worker.workerCloseDone(); - this.bsp4Worker.close(); + if (this.bsp4Worker != null) { + if (this.inited) { + this.bsp4Worker.workerCloseDone(); + } + this.bsp4Worker.close(); + } } catch (Exception e) { LOG.error("Error while closing bsp4Worker", e); } 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..5fdcb01f9 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()); + + try { + Thread sendExecutor = Whitebox.getInternalState(sender, + "sendExecutor"); + Assert.assertTrue(ImmutableSet.of(Thread.State.NEW, + Thread.State.RUNNABLE, + Thread.State.WAITING) + .contains(sendExecutor.getState())); + } finally { + sender.close(); + } + } + + @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/suite/integrate/SenderIntegrateTest.java b/computer/computer-test/src/main/java/org/apache/hugegraph/computer/suite/integrate/SenderIntegrateTest.java index 6ae0008be..e7a13339d 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; @@ -26,6 +27,7 @@ import java.util.concurrent.CompletableFuture; import java.util.concurrent.Future; 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; @@ -57,7 +59,6 @@ public class SenderIntegrateTest { public static final Logger LOG = Log.logger(SenderIntegrateTest.class); private static final Class COMPUTATION = MockComputation.class; - @BeforeClass public static void init() { // pass @@ -87,13 +88,13 @@ public void testOneWorker() { .withBufferCapacity(60) .withRpcServerHost("127.0.0.1") .withRpcServerPort(8611) - .withRpcServerPort(0) - .build(); - try (MasterService service = initMaster(args)) { - masterServiceRef.set(service); + .withRpcServerPort(0) + .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); } @@ -113,10 +114,10 @@ public void testOneWorker() { .withWorkerCount(1) .withBufferThreshold(50) .withBufferCapacity(60) - .withTransoprtServerPort(0) - .build(); - try (WorkerService service = initWorker(args)) { - workerServiceRef.set(service); + .withTransportServerPort(0) + .build(); + try (WorkerService service = initWorker( + args, workerServiceRef::set)) { service.execute(); workerFuture.complete(null); } catch (Throwable e) { @@ -128,11 +129,15 @@ public void testOneWorker() { masterThread.start(); workerThread.start(); + Throwable failure = null; try { CompletableFuture.allOf(workerFuture, masterFuture).join(); + } catch (RuntimeException | Error e) { + failure = e; + throw e; } finally { - workerServiceRef.get().close(); - masterServiceRef.get().close(); + closeServices(failure, workerServiceRef.get(), + masterServiceRef.get()); } } @@ -155,11 +160,11 @@ public void testMultiWorkers() throws IOException { .withWorkerCount(workerCount) .withPartitionCount(partitionCount) .withRpcServerHost("127.0.0.1") - .withRpcServerPort(0) - .build(); + .withRpcServerPort(0) + .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 +191,12 @@ public void testMultiWorkers() throws IOException { .withComputationClass(COMPUTATION) .withWorkerCount(workerCount) .withPartitionCount(partitionCount) - .withTransoprtServerPort(0) - .withDataDirs(dir) - .build(); + .withTransportServerPort(0) + .withDataDirs(dir) + .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 +215,17 @@ 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(); + CompletableFuture.allOf(futures.toArray(new CompletableFuture[0])) + .join(); + } 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 +247,10 @@ 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) + .build(); + try (MasterService service = initMaster( + args, masterServiceRef::set)) { service.execute(); masterFuture.complete(null); } catch (Throwable e) { @@ -252,7 +261,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 +274,13 @@ public void testOneWorkerWithBusyClient() { .withWorkerCount(1) .withWriteBufferHighMark(20) .withWriteBufferLowMark(10) - .withTransoprtServerPort(transoprtServerPort) - .build(); - try (WorkerService service = initWorker(args)) { - workerServiceRef.set(service); + .withTransportServerPort( + transportServerPort) + .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 +292,20 @@ public void testOneWorkerWithBusyClient() { masterThread.start(); workerThread.start(); + Throwable failure = null; try { CompletableFuture.allOf(workerFuture, masterFuture).join(); + } catch (RuntimeException | Error e) { + failure = e; + throw e; } finally { - workerServiceRef.get().close(); - masterServiceRef.get().close(); + closeServices(failure, workerServiceRef.get(), + masterServiceRef.get()); } } - 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); @@ -315,22 +330,57 @@ 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 class OptionsBuilder { private final List options; @@ -415,7 +465,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 +502,6 @@ public OptionsBuilder withDataDirs(String dataDirs) { this.options.add(String.valueOf(dataDirs)); return this; } + } }