diff --git a/runners/google-cloud-dataflow-java/src/main/java/org/apache/beam/runners/dataflow/options/DataflowStreamingPipelineOptions.java b/runners/google-cloud-dataflow-java/src/main/java/org/apache/beam/runners/dataflow/options/DataflowStreamingPipelineOptions.java index 4c1a82418848..ffb2e27e55b2 100644 --- a/runners/google-cloud-dataflow-java/src/main/java/org/apache/beam/runners/dataflow/options/DataflowStreamingPipelineOptions.java +++ b/runners/google-cloud-dataflow-java/src/main/java/org/apache/beam/runners/dataflow/options/DataflowStreamingPipelineOptions.java @@ -310,6 +310,9 @@ public Integer create(PipelineOptions options) { class EnableWindmillServiceDirectPathFactory implements DefaultValueFactory { @Override public Boolean create(PipelineOptions options) { + if (ExperimentalOptions.hasExperiment(options, "disable_windmill_service_direct_path")) { + return false; + } return ExperimentalOptions.hasExperiment(options, "enable_windmill_service_direct_path"); } } diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorker.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorker.java index 2a4b111af225..b0d6cb7b13d3 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorker.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorker.java @@ -33,6 +33,7 @@ import java.util.concurrent.ScheduledExecutorService; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicReference; import java.util.function.Consumer; import java.util.function.Function; import java.util.function.Supplier; @@ -65,6 +66,7 @@ import org.apache.beam.runners.dataflow.worker.util.MemoryMonitor; import org.apache.beam.runners.dataflow.worker.windmill.ApplianceWindmillClient; import org.apache.beam.runners.dataflow.worker.windmill.Windmill; +import org.apache.beam.runners.dataflow.worker.windmill.Windmill.ConnectivityType; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.JobHeader; import org.apache.beam.runners.dataflow.worker.windmill.WindmillServerStub; import org.apache.beam.runners.dataflow.worker.windmill.appliance.JniWindmillApplianceServer; @@ -170,11 +172,12 @@ public final class StreamingDataflowWorker { "windmill_bounded_queue_executor_use_fair_monitor"; private final WindmillStateCache stateCache; - private final StreamingWorkerStatusPages statusPages; + private AtomicReference statusPages = new AtomicReference<>(); private final ComputationConfig.Fetcher configFetcher; private final ComputationStateCache computationStateCache; private final BoundedQueueExecutor workUnitExecutor; - private final StreamingWorkerHarness streamingWorkerHarness; + private final AtomicReference streamingWorkerHarness = + new AtomicReference<>(); private final AtomicBoolean running = new AtomicBoolean(); private final DataflowWorkerHarnessOptions options; private final BackgroundMemoryMonitor memoryMonitor; @@ -183,6 +186,14 @@ public final class StreamingDataflowWorker { private final ActiveWorkRefresher activeWorkRefresher; private final StreamingWorkerStatusReporter workerStatusReporter; private final int numCommitThreads; + private final Supplier clock; + private final GrpcDispatcherClient dispatcherClient; + private final ExecutorService harnessSwitchExecutor; + private final long clientId; + private final WindmillServerStub windmillServer; + private final GrpcWindmillStreamFactory windmillStreamFactory; + private final StreamingWorkScheduler streamingWorkScheduler; + private final ThrottlingGetDataMetricTracker getDataMetricTracker; private StreamingDataflowWorker( WindmillServerStub windmillServer, @@ -215,150 +226,71 @@ private StreamingDataflowWorker( Executors.newCachedThreadPool()); this.options = options; this.workUnitExecutor = workUnitExecutor; + this.harnessSwitchExecutor = + Executors.newSingleThreadExecutor( + new ThreadFactoryBuilder().setNameFormat("HarnessSwitchExecutor").build()); + this.clock = clock; this.memoryMonitor = BackgroundMemoryMonitor.create(memoryMonitor); this.numCommitThreads = options.isEnableStreamingEngine() ? Math.max(options.getWindmillServiceCommitThreads(), 1) : 1; - - StreamingWorkScheduler streamingWorkScheduler = + this.dispatcherClient = dispatcherClient; + this.clientId = clientId; + this.windmillServer = windmillServer; + this.windmillStreamFactory = windmillStreamFactory; + this.streamingWorkScheduler = StreamingWorkScheduler.create( options, clock, readerCache, mapTaskExecutorFactory, workUnitExecutor, - stateCache::forComputation, + this.stateCache::forComputation, failureTracker, workFailureProcessor, streamingCounters, hotKeyLogger, sampler, ID_GENERATOR, - configFetcher.getGlobalConfigHandle(), + this.configFetcher.getGlobalConfigHandle(), stageInfoMap); - ThrottlingGetDataMetricTracker getDataMetricTracker = - new ThrottlingGetDataMetricTracker(memoryMonitor); - // Status page members. Different implementations on whether the harness is streaming engine + this.getDataMetricTracker = new ThrottlingGetDataMetricTracker(memoryMonitor); + StreamingWorkerHarnessFactoryOutput harnessFactoryOutput; + // Different implementations on whether the harness is streaming engine // direct path, streaming engine cloud path, or streaming appliance. - @Nullable ChannelzServlet channelzServlet = null; - Consumer getDataStatusProvider; - Supplier currentActiveCommitBytesProvider; - ChannelCache channelCache = null; - if (options.isEnableStreamingEngine() && options.getIsWindmillServiceDirectPathEnabled()) { - // Direct path pipelines. - WeightedSemaphore maxCommitByteSemaphore = Commits.maxCommitByteSemaphore(); - channelCache = createChannelCache(options, configFetcher); - FanOutStreamingEngineWorkerHarness fanOutStreamingEngineWorkerHarness = - FanOutStreamingEngineWorkerHarness.create( - createJobHeader(options, clientId), - GetWorkBudget.builder() - .setItems(chooseMaxBundlesOutstanding(options)) - .setBytes(MAX_GET_WORK_FETCH_BYTES) - .build(), - windmillStreamFactory, - (workItem, - serializedWorkItemSize, - watermarks, - processingContext, - getWorkStreamLatencies) -> - computationStateCache - .get(processingContext.computationId()) - .ifPresent( - computationState -> { - memoryMonitor.waitForResources("GetWork"); - streamingWorkScheduler.scheduleWork( - computationState, - workItem, - serializedWorkItemSize, - watermarks, - processingContext, - getWorkStreamLatencies); - }), - ChannelCachingRemoteStubFactory.create(options.getGcpCredential(), channelCache), - GetWorkBudgetDistributors.distributeEvenly(), - Preconditions.checkNotNull(dispatcherClient), - commitWorkStream -> - StreamingEngineWorkCommitter.builder() - // Share the commitByteSemaphore across all created workCommitters. - .setCommitByteSemaphore(maxCommitByteSemaphore) - .setBackendWorkerToken(commitWorkStream.backendWorkerToken()) - .setOnCommitComplete(this::onCompleteCommit) - .setNumCommitSenders(Math.max(options.getWindmillServiceCommitThreads(), 1)) - .setCommitWorkStreamFactory( - () -> CloseableStream.create(commitWorkStream, () -> {})) - .build(), - getDataMetricTracker); - getDataStatusProvider = getDataMetricTracker::printHtml; - currentActiveCommitBytesProvider = - fanOutStreamingEngineWorkerHarness::currentActiveCommitBytes; - channelzServlet = - createChannelzServlet( - options, fanOutStreamingEngineWorkerHarness::currentWindmillEndpoints); - this.streamingWorkerHarness = fanOutStreamingEngineWorkerHarness; - } else { - // Non-direct path pipelines. - Windmill.GetWorkRequest request = - Windmill.GetWorkRequest.newBuilder() - .setClientId(clientId) - .setMaxItems(chooseMaxBundlesOutstanding(options)) - .setMaxBytes(MAX_GET_WORK_FETCH_BYTES) - .build(); - GetDataClient getDataClient; - HeartbeatSender heartbeatSender; - WorkCommitter workCommitter; - GetWorkSender getWorkSender; - if (options.isEnableStreamingEngine()) { - WindmillStreamPool getDataStreamPool = - WindmillStreamPool.create( - Math.max(1, options.getWindmillGetDataStreamCount()), - GET_DATA_STREAM_TIMEOUT, - windmillServer::getDataStream); - getDataClient = new StreamPoolGetDataClient(getDataMetricTracker, getDataStreamPool); - heartbeatSender = - createStreamingEngineHeartbeatSender( - options, windmillServer, getDataStreamPool, configFetcher.getGlobalConfigHandle()); - channelzServlet = - createChannelzServlet(options, windmillServer::getWindmillServiceEndpoints); - workCommitter = - StreamingEngineWorkCommitter.builder() - .setCommitWorkStreamFactory( - WindmillStreamPool.create( - numCommitThreads, - COMMIT_STREAM_TIMEOUT, - windmillServer::commitWorkStream) - ::getCloseableStream) - .setCommitByteSemaphore(Commits.maxCommitByteSemaphore()) - .setNumCommitSenders(numCommitThreads) - .setOnCommitComplete(this::onCompleteCommit) - .build(); - getWorkSender = - GetWorkSender.forStreamingEngine( - receiver -> windmillServer.getWorkStream(request, receiver)); + if (options.isEnableStreamingEngine()) { + if (options.getIsWindmillServiceDirectPathEnabled()) { + harnessFactoryOutput = + createFanOutStreamingEngineWorkerHarness( + clientId, + options, + windmillStreamFactory, + streamingWorkScheduler, + getDataMetricTracker, + memoryMonitor, + this.dispatcherClient); } else { - getDataClient = new ApplianceGetDataClient(windmillServer, getDataMetricTracker); - heartbeatSender = new ApplianceHeartbeatSender(windmillServer::getData); - workCommitter = - StreamingApplianceWorkCommitter.create( - windmillServer::commitWork, this::onCompleteCommit); - getWorkSender = GetWorkSender.forAppliance(() -> windmillServer.getWork(request)); + harnessFactoryOutput = + createSingleSourceWorkerHarness( + clientId, + options, + windmillServer, + streamingWorkScheduler, + getDataMetricTracker, + memoryMonitor); } - - getDataStatusProvider = getDataClient::printHtml; - currentActiveCommitBytesProvider = workCommitter::currentActiveCommitBytes; - - this.streamingWorkerHarness = - SingleSourceWorkerHarness.builder() - .setStreamingWorkScheduler(streamingWorkScheduler) - .setWorkCommitter(workCommitter) - .setGetDataClient(getDataClient) - .setComputationStateFetcher(this.computationStateCache::get) - .setWaitForResources(() -> memoryMonitor.waitForResources("GetWork")) - .setHeartbeatSender(heartbeatSender) - .setGetWorkSender(getWorkSender) - .build(); + } else { // Appliance + harnessFactoryOutput = + createApplianceWorkerHarness( + clientId, + options, + windmillServer, + streamingWorkScheduler, + getDataMetricTracker, + memoryMonitor); } - + this.streamingWorkerHarness.set(harnessFactoryOutput.streamingWorkerHarness()); this.workerStatusReporter = streamingWorkerStatusReporter; this.activeWorkRefresher = new ActiveWorkRefresher( @@ -372,20 +304,21 @@ private StreamingDataflowWorker( activeWorkRefreshExecutorFn, getDataMetricTracker::trackHeartbeats); - this.statusPages = - createStatusPageBuilder(options, windmillStreamFactory, memoryMonitor) - .setClock(clock) - .setClientId(clientId) - .setIsRunning(running) - .setStateCache(stateCache) + this.statusPages.set( + createStatusPageBuilder( + this.options, this.windmillStreamFactory, this.memoryMonitor.memoryMonitor()) + .setClock(this.clock) + .setClientId(this.clientId) + .setIsRunning(this.running) + .setStateCache(this.stateCache) .setComputationStateCache(this.computationStateCache) - .setWorkUnitExecutor(workUnitExecutor) - .setGlobalConfigHandle(configFetcher.getGlobalConfigHandle()) - .setChannelzServlet(channelzServlet) - .setGetDataStatusProvider(getDataStatusProvider) - .setCurrentActiveCommitBytes(currentActiveCommitBytesProvider) - .setChannelCache(channelCache) - .build(); + .setWorkUnitExecutor(this.workUnitExecutor) + .setGlobalConfigHandle(this.configFetcher.getGlobalConfigHandle()) + .setChannelzServlet(harnessFactoryOutput.channelzServlet()) + .setGetDataStatusProvider(harnessFactoryOutput.getDataStatusProvider()) + .setCurrentActiveCommitBytes(harnessFactoryOutput.currentActiveCommitBytesProvider()) + .setChannelCache(harnessFactoryOutput.channelCache()) + .build()); LOG.debug("isDirectPathEnabled: {}", options.getIsWindmillServiceDirectPathEnabled()); LOG.debug("windmillServiceEnabled: {}", options.isEnableStreamingEngine()); @@ -394,6 +327,238 @@ private StreamingDataflowWorker( LOG.debug("LocalWindmillHostport: {}", options.getLocalWindmillHostport()); } + private StreamingWorkerHarnessFactoryOutput createApplianceWorkerHarness( + long clientId, + DataflowWorkerHarnessOptions options, + WindmillServerStub windmillServer, + StreamingWorkScheduler streamingWorkScheduler, + ThrottlingGetDataMetricTracker getDataMetricTracker, + MemoryMonitor memoryMonitor) { + Windmill.GetWorkRequest request = + Windmill.GetWorkRequest.newBuilder() + .setClientId(clientId) + .setMaxItems(chooseMaxBundlesOutstanding(options)) + .setMaxBytes(MAX_GET_WORK_FETCH_BYTES) + .build(); + + GetDataClient getDataClient = new ApplianceGetDataClient(windmillServer, getDataMetricTracker); + HeartbeatSender heartbeatSender = new ApplianceHeartbeatSender(windmillServer::getData); + WorkCommitter workCommitter = + StreamingApplianceWorkCommitter.create(windmillServer::commitWork, this::onCompleteCommit); + GetWorkSender getWorkSender = GetWorkSender.forAppliance(() -> windmillServer.getWork(request)); + + return StreamingWorkerHarnessFactoryOutput.builder() + .setStreamingWorkerHarness( + SingleSourceWorkerHarness.builder() + .setStreamingWorkScheduler(streamingWorkScheduler) + .setWorkCommitter(workCommitter) + .setGetDataClient(getDataClient) + .setComputationStateFetcher(this.computationStateCache::get) + .setWaitForResources(() -> memoryMonitor.waitForResources("GetWork")) + .setHeartbeatSender(heartbeatSender) + .setGetWorkSender(getWorkSender) + .build()) + .setGetDataStatusProvider(getDataClient::printHtml) + .setCurrentActiveCommitBytesProvider(workCommitter::currentActiveCommitBytes) + .setChannelzServlet(null) // Appliance doesn't use ChannelzServlet + .setChannelCache(null) // Appliance doesn't use ChannelCache + .build(); + } + + private StreamingWorkerHarnessFactoryOutput createFanOutStreamingEngineWorkerHarness( + long clientId, + DataflowWorkerHarnessOptions options, + GrpcWindmillStreamFactory windmillStreamFactory, + StreamingWorkScheduler streamingWorkScheduler, + ThrottlingGetDataMetricTracker getDataMetricTracker, + MemoryMonitor memoryMonitor, + GrpcDispatcherClient dispatcherClient) { + WeightedSemaphore maxCommitByteSemaphore = Commits.maxCommitByteSemaphore(); + ChannelCache channelCache = createChannelCache(options, configFetcher); + FanOutStreamingEngineWorkerHarness fanOutStreamingEngineWorkerHarness = + FanOutStreamingEngineWorkerHarness.create( + createJobHeader(options, clientId), + GetWorkBudget.builder() + .setItems(chooseMaxBundlesOutstanding(options)) + .setBytes(MAX_GET_WORK_FETCH_BYTES) + .build(), + windmillStreamFactory, + (workItem, + serializedWorkItemSize, + watermarks, + processingContext, + getWorkStreamLatencies) -> + computationStateCache + .get(processingContext.computationId()) + .ifPresent( + computationState -> { + memoryMonitor.waitForResources("GetWork"); + streamingWorkScheduler.scheduleWork( + computationState, + workItem, + serializedWorkItemSize, + watermarks, + processingContext, + getWorkStreamLatencies); + }), + ChannelCachingRemoteStubFactory.create(options.getGcpCredential(), channelCache), + GetWorkBudgetDistributors.distributeEvenly(), + Preconditions.checkNotNull(dispatcherClient), + commitWorkStream -> + StreamingEngineWorkCommitter.builder() + // Share the commitByteSemaphore across all created workCommitters. + .setCommitByteSemaphore(maxCommitByteSemaphore) + .setBackendWorkerToken(commitWorkStream.backendWorkerToken()) + .setOnCommitComplete(this::onCompleteCommit) + .setNumCommitSenders(Math.max(options.getWindmillServiceCommitThreads(), 1)) + .setCommitWorkStreamFactory( + () -> CloseableStream.create(commitWorkStream, () -> {})) + .build(), + getDataMetricTracker); + ChannelzServlet channelzServlet = + createChannelzServlet( + options, fanOutStreamingEngineWorkerHarness::currentWindmillEndpoints); + return StreamingWorkerHarnessFactoryOutput.builder() + .setStreamingWorkerHarness(fanOutStreamingEngineWorkerHarness) + .setGetDataStatusProvider(getDataMetricTracker::printHtml) + .setCurrentActiveCommitBytesProvider( + fanOutStreamingEngineWorkerHarness::currentActiveCommitBytes) + .setChannelzServlet(channelzServlet) + .setChannelCache(channelCache) + .build(); + } + + private StreamingWorkerHarnessFactoryOutput createSingleSourceWorkerHarness( + long clientId, + DataflowWorkerHarnessOptions options, + WindmillServerStub windmillServer, + StreamingWorkScheduler streamingWorkScheduler, + ThrottlingGetDataMetricTracker getDataMetricTracker, + MemoryMonitor memoryMonitor) { + Windmill.GetWorkRequest request = + Windmill.GetWorkRequest.newBuilder() + .setClientId(clientId) + .setMaxItems(chooseMaxBundlesOutstanding(options)) + .setMaxBytes(MAX_GET_WORK_FETCH_BYTES) + .build(); + WindmillStreamPool getDataStreamPool = + WindmillStreamPool.create( + Math.max(1, options.getWindmillGetDataStreamCount()), + GET_DATA_STREAM_TIMEOUT, + windmillServer::getDataStream); + GetDataClient getDataClient = + new StreamPoolGetDataClient(getDataMetricTracker, getDataStreamPool); + HeartbeatSender heartbeatSender = + createStreamingEngineHeartbeatSender( + options, windmillServer, getDataStreamPool, configFetcher.getGlobalConfigHandle()); + WorkCommitter workCommitter = + StreamingEngineWorkCommitter.builder() + .setCommitWorkStreamFactory( + WindmillStreamPool.create( + numCommitThreads, COMMIT_STREAM_TIMEOUT, windmillServer::commitWorkStream) + ::getCloseableStream) + .setCommitByteSemaphore(Commits.maxCommitByteSemaphore()) + .setNumCommitSenders(numCommitThreads) + .setOnCommitComplete(this::onCompleteCommit) + .build(); + GetWorkSender getWorkSender = + GetWorkSender.forStreamingEngine( + receiver -> windmillServer.getWorkStream(request, receiver)); + ChannelzServlet channelzServlet = + createChannelzServlet(options, windmillServer::getWindmillServiceEndpoints); + return StreamingWorkerHarnessFactoryOutput.builder() + .setStreamingWorkerHarness( + SingleSourceWorkerHarness.builder() + .setStreamingWorkScheduler(streamingWorkScheduler) + .setWorkCommitter(workCommitter) + .setGetDataClient(getDataClient) + .setComputationStateFetcher(this.computationStateCache::get) + .setWaitForResources(() -> memoryMonitor.waitForResources("GetWork")) + .setHeartbeatSender(heartbeatSender) + .setGetWorkSender(getWorkSender) + .build()) + .setGetDataStatusProvider(getDataClient::printHtml) + .setCurrentActiveCommitBytesProvider(workCommitter::currentActiveCommitBytes) + .setChannelzServlet(channelzServlet) + .setChannelCache(null) // SingleSourceWorkerHarness doesn't use ChannelCache + .build(); + } + + private void switchStreamingWorkerHarness(ConnectivityType connectivityType) { + if ((connectivityType == ConnectivityType.CONNECTIVITY_TYPE_DIRECTPATH + && this.streamingWorkerHarness.get() instanceof FanOutStreamingEngineWorkerHarness) + || (connectivityType == ConnectivityType.CONNECTIVITY_TYPE_CLOUDPATH + && streamingWorkerHarness.get() instanceof SingleSourceWorkerHarness)) { + return; + } + // Stop the current status pages before switching the harness. + this.statusPages.get().stop(); + LOG.debug("Stopped StreamingWorkerStatusPages before switching connectivity type."); + StreamingWorkerHarnessFactoryOutput newHarnessFactoryOutput = null; + if (connectivityType == ConnectivityType.CONNECTIVITY_TYPE_DIRECTPATH) { + // If dataflow experiment `enable_windmill_service_direct_path` is not set for + // the job, do not switch to FanOutStreamingEngineWorkerHarness. This is because + // `enable_windmill_service_direct_path` is tied to SDK version and is only + // enabled for job running with SDK above the cut off version, + // and we do not want jobs below the cutoff to switch to + // FanOutStreamingEngineWorkerHarness + if (!options.getIsWindmillServiceDirectPathEnabled()) { + LOG.info( + "Dataflow experiment `enable_windmill_service_direct_path` is not set for the job. Job" + + " cannot switch to connectivity type DIRECTPATH. Job will continue running on" + + " CLOUDPATH"); + return; + } + LOG.info("Switching connectivity type from CLOUDPATH to DIRECTPATH"); + LOG.debug("Shutting down to SingleSourceWorkerHarness"); + this.streamingWorkerHarness.get().shutdown(); + newHarnessFactoryOutput = + createFanOutStreamingEngineWorkerHarness( + this.clientId, + this.options, + this.windmillStreamFactory, + this.streamingWorkScheduler, + this.getDataMetricTracker, + this.memoryMonitor.memoryMonitor(), + this.dispatcherClient); + this.streamingWorkerHarness.set(newHarnessFactoryOutput.streamingWorkerHarness()); + streamingWorkerHarness.get().start(); + LOG.debug("Started FanOutStreamingEngineWorkerHarness"); + } else if (connectivityType == ConnectivityType.CONNECTIVITY_TYPE_CLOUDPATH) { + LOG.info("Switching connectivity type from DIRECTPATH to CLOUDPATH"); + LOG.debug("Shutting down FanOutStreamingEngineWorkerHarness"); + streamingWorkerHarness.get().shutdown(); + newHarnessFactoryOutput = + createSingleSourceWorkerHarness( + this.clientId, + this.options, + this.windmillServer, + this.streamingWorkScheduler, + this.getDataMetricTracker, + this.memoryMonitor.memoryMonitor()); + this.streamingWorkerHarness.set(newHarnessFactoryOutput.streamingWorkerHarness()); + streamingWorkerHarness.get().start(); + LOG.debug("Started SingleSourceWorkerHarness"); + } + this.statusPages.set( + createStatusPageBuilder( + this.options, this.windmillStreamFactory, this.memoryMonitor.memoryMonitor()) + .setClock(this.clock) + .setClientId(this.clientId) + .setIsRunning(this.running) + .setStateCache(this.stateCache) + .setComputationStateCache(this.computationStateCache) + .setWorkUnitExecutor(this.workUnitExecutor) + .setGlobalConfigHandle(this.configFetcher.getGlobalConfigHandle()) + .setChannelzServlet(newHarnessFactoryOutput.channelzServlet()) + .setGetDataStatusProvider(newHarnessFactoryOutput.getDataStatusProvider()) + .setCurrentActiveCommitBytes(newHarnessFactoryOutput.currentActiveCommitBytesProvider()) + .setChannelCache(newHarnessFactoryOutput.channelCache()) + .build()); + this.statusPages.get().start(this.options); + LOG.info("Started new StreamingWorkerStatusPages instance."); + } + private static StreamingWorkerStatusPages.Builder createStatusPageBuilder( DataflowWorkerHarnessOptions options, GrpcWindmillStreamFactory windmillStreamFactory, @@ -736,6 +901,11 @@ static StreamingDataflowWorker forTesting( createGrpcwindmillStreamFactoryBuilder(options, 1) .setProcessHeartbeatResponses( new WorkHeartbeatResponseProcessor(computationStateCache::get)); + GrpcDispatcherClient grpcDispatcherClient = GrpcDispatcherClient.create(options, stubFactory); + grpcDispatcherClient.consumeWindmillDispatcherEndpoints( + ImmutableSet.builder() + .add(HostAndPort.fromHost("StreamingDataflowWorkerTest")) + .build()); return new StreamingDataflowWorker( windmillServer, @@ -761,7 +931,7 @@ static StreamingDataflowWorker forTesting( : windmillStreamFactory.build(), executorSupplier.apply("RefreshWork"), stageInfo, - GrpcDispatcherClient.create(options, stubFactory)); + grpcDispatcherClient); } private static GrpcWindmillStreamFactory.Builder createGrpcwindmillStreamFactoryBuilder( @@ -889,15 +1059,36 @@ public void start() { running.set(true); configFetcher.start(); memoryMonitor.start(); - streamingWorkerHarness.start(); + streamingWorkerHarness.get().start(); sampler.start(); workerStatusReporter.start(); activeWorkRefresher.start(); + configFetcher + .getGlobalConfigHandle() + .registerConfigObserver( + streamingGlobalConfig -> { + ConnectivityType connectivityType = + streamingGlobalConfig.userWorkerJobSettings().getConnectivityType(); + if (connectivityType != ConnectivityType.CONNECTIVITY_TYPE_DEFAULT) { + LOG.debug("Switching to connectivityType: {}.", connectivityType); + harnessSwitchExecutor.execute(() -> switchStreamingWorkerHarness(connectivityType)); + } + }); } /** Starts the status page server for debugging. May be omitted for lighter weight testing. */ private void startStatusPages() { - statusPages.start(options); + statusPages.get().start(options); + } + + @VisibleForTesting + StreamingWorkerHarness getStreamingWorkerHarness() { + return streamingWorkerHarness.get(); + } + + @VisibleForTesting + ExecutorService getHarnessSwitchExecutor() { + return harnessSwitchExecutor; } @VisibleForTesting @@ -905,9 +1096,10 @@ void stop() { try { configFetcher.stop(); activeWorkRefresher.stop(); - statusPages.stop(); + statusPages.get().stop(); running.set(false); - streamingWorkerHarness.shutdown(); + harnessSwitchExecutor.shutdown(); + streamingWorkerHarness.get().shutdown(); memoryMonitor.shutdown(); workUnitExecutor.shutdown(); computationStateCache.closeAndInvalidateAll(); @@ -1000,4 +1192,40 @@ private void shutdown() { executor().shutdown(); } } + + /** + * Holds the {@link StreamingWorkerHarness} and its associated dependencies that are created + * together. + */ + @AutoValue + abstract static class StreamingWorkerHarnessFactoryOutput { + static Builder builder() { + return new AutoValue_StreamingDataflowWorker_StreamingWorkerHarnessFactoryOutput.Builder(); + } + + abstract StreamingWorkerHarness streamingWorkerHarness(); + + abstract Consumer getDataStatusProvider(); + + abstract Supplier currentActiveCommitBytesProvider(); + + abstract @Nullable ChannelzServlet channelzServlet(); + + abstract @Nullable ChannelCache channelCache(); + + @AutoValue.Builder + abstract static class Builder { + abstract Builder setStreamingWorkerHarness(StreamingWorkerHarness value); + + abstract Builder setGetDataStatusProvider(Consumer value); + + abstract Builder setCurrentActiveCommitBytesProvider(Supplier value); + + abstract Builder setChannelzServlet(@Nullable ChannelzServlet value); + + abstract Builder setChannelCache(@Nullable ChannelCache value); + + abstract StreamingWorkerHarnessFactoryOutput build(); + } + } } diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/SingleSourceWorkerHarness.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/SingleSourceWorkerHarness.java index 95023d117299..0de9d130b650 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/SingleSourceWorkerHarness.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/SingleSourceWorkerHarness.java @@ -27,6 +27,7 @@ import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; import java.util.function.Function; +import javax.annotation.Nullable; import org.apache.beam.runners.dataflow.worker.WindmillTimeUtils; import org.apache.beam.runners.dataflow.worker.streaming.ComputationState; import org.apache.beam.runners.dataflow.worker.streaming.Watermarks; @@ -66,6 +67,7 @@ public final class SingleSourceWorkerHarness implements StreamingWorkerHarness { private final Function> computationStateFetcher; private final ExecutorService workProviderExecutor; private final GetWorkSender getWorkSender; + @Nullable private WindmillStream.GetWorkStream getWorkStream; SingleSourceWorkerHarness( WorkCommitter workCommitter, @@ -140,12 +142,15 @@ public void shutdown() { LOG.warn("Unable to shutdown {}", getClass()); } workCommitter.stop(); + if (getWorkStream != null) { + getWorkStream.shutdown(); + } } private void streamingEngineDispatchLoop( Function getWorkStreamFactory) { while (isRunning.get()) { - WindmillStream.GetWorkStream stream = + getWorkStream = getWorkStreamFactory.apply( (computationId, inputDataWatermark, @@ -179,8 +184,10 @@ private void streamingEngineDispatchLoop( // Reconnect every now and again to enable better load balancing. // If at any point the server closes the stream, we will reconnect immediately; otherwise // we half-close the stream after some time and create a new one. - if (!stream.awaitTermination(GET_WORK_STREAM_TIMEOUT_MINUTES, TimeUnit.MINUTES)) { - stream.halfClose(); + if (getWorkStream != null) { + if (!getWorkStream.awaitTermination(GET_WORK_STREAM_TIMEOUT_MINUTES, TimeUnit.MINUTES)) { + Preconditions.checkNotNull(getWorkStream).halfClose(); + } } } catch (InterruptedException e) { // Continue processing until !running.get() diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/FakeWindmillServer.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/FakeWindmillServer.java index dd13d5b55930..a5c8909b8d07 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/FakeWindmillServer.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/FakeWindmillServer.java @@ -47,6 +47,7 @@ import org.apache.beam.runners.dataflow.worker.streaming.ComputationState; import org.apache.beam.runners.dataflow.worker.streaming.WorkHeartbeatResponseProcessor; import org.apache.beam.runners.dataflow.worker.streaming.WorkId; +import org.apache.beam.runners.dataflow.worker.windmill.CloudWindmillMetadataServiceV1Alpha1Grpc; import org.apache.beam.runners.dataflow.worker.windmill.Windmill; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.CommitWorkResponse; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.ComputationCommitWorkRequest; @@ -60,12 +61,15 @@ import org.apache.beam.runners.dataflow.worker.windmill.Windmill.LatencyAttribution; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.LatencyAttribution.State; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.WorkItemCommitRequest; +import org.apache.beam.runners.dataflow.worker.windmill.Windmill.WorkerMetadataRequest; +import org.apache.beam.runners.dataflow.worker.windmill.Windmill.WorkerMetadataResponse; import org.apache.beam.runners.dataflow.worker.windmill.WindmillServerStub; import org.apache.beam.runners.dataflow.worker.windmill.client.WindmillStream.CommitWorkStream; import org.apache.beam.runners.dataflow.worker.windmill.client.WindmillStream.GetDataStream; import org.apache.beam.runners.dataflow.worker.windmill.client.WindmillStream.GetWorkStream; import org.apache.beam.runners.dataflow.worker.windmill.work.WorkItemReceiver; import org.apache.beam.runners.dataflow.worker.windmill.work.budget.GetWorkBudget; +import org.apache.beam.vendor.grpc.v1p69p0.io.grpc.stub.StreamObserver; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableList; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableSet; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.net.HostAndPort; @@ -92,6 +96,8 @@ public final class FakeWindmillServer extends WindmillServerStub { private final ConcurrentHashMap> droppedStreamingCommits; private final List getDataRequests = new ArrayList<>(); private final Consumer> processHeartbeatResponses; + private StreamObserver workerMetadataObserver = null; + private int commitsRequested = 0; private boolean dropStreamingCommits = false; @@ -553,6 +559,47 @@ public synchronized void setWindmillServiceEndpoints(Set endpoints) this.dispatcherEndpoints = ImmutableSet.copyOf(endpoints); } + public void injectWorkerMetadata(WorkerMetadataResponse response) { + if (workerMetadataObserver != null) { + workerMetadataObserver.onNext(response); + } + } + + private void setWorkerMetadataObserver( + StreamObserver workerMetadataObserver) { + this.workerMetadataObserver = workerMetadataObserver; + } + + public static class FakeWindmillMetadataService + extends CloudWindmillMetadataServiceV1Alpha1Grpc + .CloudWindmillMetadataServiceV1Alpha1ImplBase { + private final FakeWindmillServer server; + + public FakeWindmillMetadataService(FakeWindmillServer server) { + this.server = server; + } + + @Override + public StreamObserver getWorkerMetadata( + StreamObserver responseObserver) { + server.setWorkerMetadataObserver(responseObserver); + return new StreamObserver() { + @Override + public void onNext(WorkerMetadataRequest value) {} + + @Override + public void onError(Throwable t) { + responseObserver.onError(t); + } + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + } + public static class ResponseQueue { private final Queue> responses = new ConcurrentLinkedQueue<>(); Duration sleep = Duration.ZERO; diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java index a60535dfbd69..b21b8e830ae8 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java @@ -60,6 +60,7 @@ import com.google.auto.value.AutoValue; import java.io.IOException; import java.io.InputStream; +import java.net.ServerSocket; import java.util.ArrayList; import java.util.Arrays; import java.util.Collection; @@ -75,6 +76,7 @@ import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.CountDownLatch; import java.util.concurrent.ExecutionException; +import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; import java.util.concurrent.Future; import java.util.concurrent.ScheduledExecutorService; @@ -106,17 +108,21 @@ import org.apache.beam.runners.dataflow.worker.streaming.Work; import org.apache.beam.runners.dataflow.worker.streaming.config.StreamingGlobalConfig; import org.apache.beam.runners.dataflow.worker.streaming.config.StreamingGlobalConfigHandleImpl; +import org.apache.beam.runners.dataflow.worker.streaming.harness.FanOutStreamingEngineWorkerHarness; +import org.apache.beam.runners.dataflow.worker.streaming.harness.SingleSourceWorkerHarness; import org.apache.beam.runners.dataflow.worker.streaming.harness.StreamingCounters; import org.apache.beam.runners.dataflow.worker.testing.RestoreDataflowLoggingMDC; import org.apache.beam.runners.dataflow.worker.testing.TestCountingSource; import org.apache.beam.runners.dataflow.worker.util.BoundedQueueExecutor; import org.apache.beam.runners.dataflow.worker.util.WorkerPropertyNames; +import org.apache.beam.runners.dataflow.worker.windmill.CloudWindmillServiceV1Alpha1Grpc; import org.apache.beam.runners.dataflow.worker.windmill.Windmill; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.CommitStatus; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.ComputationGetDataRequest; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.ComputationGetDataResponse; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.ComputationHeartbeatRequest; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.ComputationHeartbeatResponse; +import org.apache.beam.runners.dataflow.worker.windmill.Windmill.ConnectivityType; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.GetDataRequest; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.GetDataResponse; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.GetWorkResponse; @@ -131,6 +137,7 @@ import org.apache.beam.runners.dataflow.worker.windmill.Windmill.Timer.Type; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.WatermarkHold; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.WorkItemCommitRequest; +import org.apache.beam.runners.dataflow.worker.windmill.Windmill.WorkerMetadataResponse; import org.apache.beam.runners.dataflow.worker.windmill.client.getdata.FakeGetDataClient; import org.apache.beam.runners.dataflow.worker.windmill.client.grpc.stubs.WindmillChannels; import org.apache.beam.runners.dataflow.worker.windmill.testing.FakeWindmillStubFactory; @@ -182,6 +189,9 @@ import org.apache.beam.sdk.values.WindowingStrategy.AccumulationMode; import org.apache.beam.vendor.grpc.v1p69p0.com.google.protobuf.ByteString; import org.apache.beam.vendor.grpc.v1p69p0.com.google.protobuf.TextFormat; +import org.apache.beam.vendor.grpc.v1p69p0.io.grpc.Server; +import org.apache.beam.vendor.grpc.v1p69p0.io.grpc.ServerBuilder; +import org.apache.beam.vendor.grpc.v1p69p0.io.grpc.testing.GrpcCleanupRule; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.cache.CacheStats; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableList; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableMap; @@ -284,6 +294,7 @@ public Long get() { @Rule public transient Timeout globalTimeout = Timeout.seconds(600); @Rule public BlockingFn blockingFn = new BlockingFn(); @Rule public TestRule restoreMDC = new RestoreDataflowLoggingMDC(); + @Rule public final GrpcCleanupRule grpcCleanup = new GrpcCleanupRule(); @Rule public ErrorCollector errorCollector = new ErrorCollector(); WorkUnitClient mockWorkUnitClient = mock(WorkUnitClient.class); StreamingGlobalConfigHandleImpl mockGlobalConfigHandle = @@ -4058,6 +4069,143 @@ public void testStuckCommit() throws Exception { removeDynamicFields(result.get(1L))); } + @Test + public void testSwitchStreamingWorkerHarness() throws Exception { + if (!streamingEngine) { + return; + } + + int port = -1; + try (ServerSocket socket = new ServerSocket(0)) { + port = socket.getLocalPort(); + } + String serverEndpoint = "localhost:" + port; + Server fakeServer = + grpcCleanup + .register( + ServerBuilder.forPort(port) + .directExecutor() + .addService(new FakeWindmillServer.FakeWindmillMetadataService(server)) + .addService( + new CloudWindmillServiceV1Alpha1Grpc + .CloudWindmillServiceV1Alpha1ImplBase() {}) + .build()) + .start(); + List instructions = + Arrays.asList( + makeSourceInstruction(StringUtf8Coder.of()), + makeSinkInstruction(StringUtf8Coder.of(), 0)); + + // Start with Directpath. + DataflowWorkerHarnessOptions options = + createTestingPipelineOptions("--isWindmillServiceDirectPathEnabled=true"); + options.setWindmillServiceEndpoint(serverEndpoint); + + StreamingDataflowWorker worker = + makeWorker( + defaultWorkerParams() + .setOptions(options) + .setInstructions(instructions) + .publishCounters() + .build()); + + ArgumentCaptor> observerCaptor = + ArgumentCaptor.forClass(Consumer.class); + + worker.start(); + + verify(mockGlobalConfigHandle, atLeastOnce()).registerConfigObserver(observerCaptor.capture()); + + List> observers = observerCaptor.getAllValues(); + + assertTrue( + "Worker should start with FanOutStreamingEngineWorkerHarness", + worker.getStreamingWorkerHarness() instanceof FanOutStreamingEngineWorkerHarness); + + // Prepare WorkerMetadataResponse + server.injectWorkerMetadata( + WorkerMetadataResponse.newBuilder() + .setMetadataVersion(1) + .addWorkEndpoints( + WorkerMetadataResponse.Endpoint.newBuilder() + .setBackendWorkerToken("workerToken1") + .setDirectEndpoint(serverEndpoint) + .build()) + .build()); + + // Switch to Cloudpath. + StreamingGlobalConfig cloudPathConfig = + StreamingGlobalConfig.builder() + .setUserWorkerJobSettings( + Windmill.UserWorkerRunnerV1Settings.newBuilder() + .setConnectivityType(ConnectivityType.CONNECTIVITY_TYPE_CLOUDPATH) + .build()) + .build(); + for (Consumer observer : observers) { + observer.accept(cloudPathConfig); + } + + ExecutorService harnessSwitchExecutor = worker.getHarnessSwitchExecutor(); + Future cloudPathSwitchFuture = harnessSwitchExecutor.submit(() -> {}); + cloudPathSwitchFuture.get(30, TimeUnit.SECONDS); + assertTrue( + "Worker should switch to SingleSourceWorkerHarness", + worker.getStreamingWorkerHarness() instanceof SingleSourceWorkerHarness); + + // Process some work with CloudPath. + server.whenGetWorkCalled().thenReturn(makeInput(1, 1000)); + Map result = server.waitForAndGetCommits(1); + assertEquals(1, result.size()); + assertTrue(result.containsKey(1L)); + + // Switch to Directpath. + StreamingGlobalConfig directPathConfig = + StreamingGlobalConfig.builder() + .setUserWorkerJobSettings( + Windmill.UserWorkerRunnerV1Settings.newBuilder() + .setConnectivityType(ConnectivityType.CONNECTIVITY_TYPE_DIRECTPATH) + .build()) + .build(); + + for (Consumer observer : observers) { + observer.accept(directPathConfig); + } + + // Wait for the harnessSwitchExecutor to complete the switch. + Future directPathSwitchFuture = harnessSwitchExecutor.submit(() -> {}); + // Wait for the dummy task to complete. The dummy task will be executed after + // switchStreamingWorkerHarness has completed. + directPathSwitchFuture.get(30, TimeUnit.SECONDS); + assertTrue( + "Worker should switch to FanOutStreamingEngineWorkerHarness", + worker.getStreamingWorkerHarness() instanceof FanOutStreamingEngineWorkerHarness); + + // Switch to Cloudpath again. + cloudPathConfig = + StreamingGlobalConfig.builder() + .setUserWorkerJobSettings( + Windmill.UserWorkerRunnerV1Settings.newBuilder() + .setConnectivityType(ConnectivityType.CONNECTIVITY_TYPE_CLOUDPATH) + .build()) + .build(); + for (Consumer observer : observers) { + observer.accept(cloudPathConfig); + } + + cloudPathSwitchFuture = harnessSwitchExecutor.submit(() -> {}); + cloudPathSwitchFuture.get(30, TimeUnit.SECONDS); + assertTrue( + "Worker should switch back to SingleSourceWorkerHarness", + worker.getStreamingWorkerHarness() instanceof SingleSourceWorkerHarness); + // Process some work with CloudPath again. + server.whenGetWorkCalled().thenReturn(makeInput(2, 2000)); + result = server.waitForAndGetCommits(1); + assertEquals(2, result.size()); + assertTrue(result.containsKey(2L)); + + worker.stop(); + } + private void runNumCommitThreadsTest(int configNumCommitThreads, int expectedNumCommitThreads) { List instructions = Arrays.asList( diff --git a/runners/google-cloud-dataflow-java/worker/windmill/src/main/proto/windmill.proto b/runners/google-cloud-dataflow-java/worker/windmill/src/main/proto/windmill.proto index 77401be4ac77..a4b3df906dd9 100644 --- a/runners/google-cloud-dataflow-java/worker/windmill/src/main/proto/windmill.proto +++ b/runners/google-cloud-dataflow-java/worker/windmill/src/main/proto/windmill.proto @@ -958,6 +958,12 @@ message UserWorkerGrpcFlowControlSettings { optional int32 on_ready_threshold_bytes = 3; } +enum ConnectivityType { + CONNECTIVITY_TYPE_DEFAULT = 0; + CONNECTIVITY_TYPE_CLOUDPATH = 1; + CONNECTIVITY_TYPE_DIRECTPATH = 2; +} + // Settings to control runtime behavior of the java runner v1 user worker. message UserWorkerRunnerV1Settings { // If true, use separate channels for each windmill RPC. @@ -967,6 +973,9 @@ message UserWorkerRunnerV1Settings { optional bool use_separate_windmill_heartbeat_streams = 2 [default = true]; optional UserWorkerGrpcFlowControlSettings flow_control_settings = 3; + + optional ConnectivityType connectivity_type = 4 + [default = CONNECTIVITY_TYPE_DEFAULT]; } service WindmillAppliance {