Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,11 @@ boolean commitWorkItem(
Windmill.WorkItemCommitRequest request,
Consumer<Windmill.CommitStatus> onDone);

boolean commitMultiKeyWorkItem(
String computation,
Windmill.MultiKeyWorkItemCommitRequest request,
Consumer<Windmill.CommitStatus> onDone);

/** Flushes any pending work items to the wire. */
void flush();

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,35 +17,89 @@
*/
package org.apache.beam.runners.dataflow.worker.windmill.client.commits;

import com.google.auto.value.AutoValue;
import static org.apache.beam.sdk.util.Preconditions.checkStateNotNull;

import org.apache.beam.runners.dataflow.worker.streaming.ComputationState;
import org.apache.beam.runners.dataflow.worker.streaming.Work;
import org.apache.beam.runners.dataflow.worker.windmill.Windmill.MultiKeyWorkItemCommitRequest;
import org.apache.beam.runners.dataflow.worker.windmill.Windmill.WorkItemCommitRequest;
import org.apache.beam.sdk.annotations.Internal;
import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Preconditions;
import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableList;
import org.checkerframework.checker.nullness.qual.Nullable;

/** Value class for a queued commit. */
@Internal
@AutoValue
public abstract class Commit {
public class Commit {

private final ComputationState computationState;
private final ImmutableList<Work> workBatch;
private final @Nullable WorkItemCommitRequest singleKeyRequest;
private final @Nullable MultiKeyWorkItemCommitRequest multiKeyRequest;

public static Commit create(
WorkItemCommitRequest request, ComputationState computationState, Work work) {
Preconditions.checkArgument(request.getSerializedSize() > 0);
return new AutoValue_Commit(request, computationState, work);
return new Commit(computationState, ImmutableList.of(work), request, null);
}

public static Commit createMultiKey(
MultiKeyWorkItemCommitRequest multiKeyRequest,
ComputationState computationState,
ImmutableList<Work> workBatch) {
Preconditions.checkArgument(!workBatch.isEmpty());
return new Commit(computationState, workBatch, null, multiKeyRequest);
}

private Commit(
ComputationState computationState,
ImmutableList<Work> workBatch,
@Nullable WorkItemCommitRequest singleKeyRequest,
@Nullable MultiKeyWorkItemCommitRequest multiKeyRequest) {
this.computationState = computationState;
this.workBatch = workBatch;
this.singleKeyRequest = singleKeyRequest;
this.multiKeyRequest = multiKeyRequest;
}

public final String computationId() {
return computationState().getComputationId();
}

public abstract WorkItemCommitRequest request();
public @Nullable WorkItemCommitRequest singleKeyRequest() {
return singleKeyRequest;
};

public abstract ComputationState computationState();
public ComputationState computationState() {
return computationState;
}

public @Nullable MultiKeyWorkItemCommitRequest multiKeyRequest() {
return multiKeyRequest;
}

public abstract Work work();
public ImmutableList<Work> workBatch() {
return workBatch;
}

public final int getSerializedByteSize() {
if (multiKeyRequest() != null) {
return checkStateNotNull(multiKeyRequest()).getSerializedSize();
}
return checkStateNotNull(singleKeyRequest()).getSerializedSize();
}

public final int getSize() {
return request().getSerializedSize();
@Override
public String toString() {
Work work = workBatch.get(0);
return "[computationId="
+ computationId()
+ ", shardingKey="
+ work.getShardedKey()
+ ", workId="
+ work.id()
+ ", workBatchSize="
+ workBatch.size()
+ "]";
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,6 @@ public final class Commits {
private Commits() {}

public static WeightedSemaphore<Commit> maxCommitByteSemaphore() {
return WeightedSemaphore.create(MAX_QUEUED_COMMITS_BYTES, Commit::getSize);
return WeightedSemaphore.create(MAX_QUEUED_COMMITS_BYTES, Commit::getSerializedByteSize);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -37,26 +37,11 @@
@AutoValue
public abstract class CompleteCommit {

public static CompleteCommit create(Commit commit, CommitStatus commitStatus) {
return new AutoValue_CompleteCommit(
commit.computationId(),
ShardedKey.create(commit.request().getKey(), commit.request().getShardingKey()),
WorkId.builder()
.setWorkToken(commit.request().getWorkToken())
.setCacheToken(commit.request().getCacheToken())
.build(),
commitStatus);
}

public static CompleteCommit create(
String computationId, ShardedKey shardedKey, WorkId workId, CommitStatus status) {
return new AutoValue_CompleteCommit(computationId, shardedKey, workId, status);
}

public static CompleteCommit forFailedWork(Commit commit) {
return create(commit, CommitStatus.ABORTED);
}

public abstract String computationId();

public abstract ShardedKey shardedKey();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,9 @@
*/
package org.apache.beam.runners.dataflow.worker.windmill.client.commits;

import static org.apache.beam.sdk.util.Preconditions.checkStateNotNull;
import static org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Preconditions.checkState;

import java.util.HashMap;
import java.util.Map;
import java.util.concurrent.ExecutorService;
Expand Down Expand Up @@ -112,18 +115,19 @@ private void commitLoop() {
}
while (commit != null) {
ComputationState computationState = commit.computationState();
commit.work().setState(Work.State.COMMITTING);
checkState(commit.workBatch().size() == 1);
commit.workBatch().get(0).setState(Work.State.COMMITTING);
Windmill.ComputationCommitWorkRequest.Builder computationRequestBuilder =
computationRequestMap.get(computationState);
if (computationRequestBuilder == null) {
computationRequestBuilder = commitRequestBuilder.addRequestsBuilder();
computationRequestBuilder.setComputationId(computationState.getComputationId());
computationRequestMap.put(computationState, computationRequestBuilder);
}
computationRequestBuilder.addRequests(commit.request());
computationRequestBuilder.addRequests(checkStateNotNull(commit.singleKeyRequest()));
// Send the request if we've exceeded the bytes or there is no more
// pending work. commitBytes is a long, so this cannot overflow.
commitBytes += commit.getSize();
commitBytes += commit.getSerializedByteSize();
if (commitBytes >= TARGET_COMMIT_BUNDLE_BYTES) {
break;
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,8 @@
*/
package org.apache.beam.runners.dataflow.worker.windmill.client.commits;

import static org.apache.beam.sdk.util.Preconditions.checkStateNotNull;

import com.google.auto.value.AutoBuilder;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
Expand All @@ -30,6 +32,7 @@
import org.apache.beam.runners.dataflow.worker.streaming.WeightedBoundedQueue;
import org.apache.beam.runners.dataflow.worker.streaming.WeightedSemaphore;
import org.apache.beam.runners.dataflow.worker.streaming.Work;
import org.apache.beam.runners.dataflow.worker.windmill.Windmill.CommitStatus;
import org.apache.beam.runners.dataflow.worker.windmill.client.CloseableStream;
import org.apache.beam.runners.dataflow.worker.windmill.client.WindmillStream.CommitWorkStream;
import org.apache.beam.sdk.annotations.Internal;
Expand Down Expand Up @@ -100,21 +103,16 @@ public void start() {

@Override
public void commit(Commit commit) {
if (commit.work().isFailed()) {
failCommit(commit);
if (shouldFailCommit(commit)) {
failQueuedCommit(commit);
} else {
commitQueue.put(commit);
}

// Do this check after adding to commitQueue, else commitQueue.put() can race with
// drainCommitQueue() in stop() and leave commits orphaned in the queue.
if (!this.isRunning.get()) {
LOG.debug(
"Trying to queue commit on shutdown, failing commit=[computationId={}, shardingKey={},"
+ " workId={} ].",
commit.computationId(),
commit.work().getShardedKey(),
commit.work().id());
LOG.debug("Trying to queue commit on shutdown, failing commit={}", commit);
drainCommitQueue();
}
}
Expand All @@ -141,14 +139,18 @@ public void stop() {
private void drainCommitQueue() {
Commit queuedCommit = commitQueue.poll();
while (queuedCommit != null) {
failCommit(queuedCommit);
failQueuedCommit(queuedCommit);
queuedCommit = commitQueue.poll();
}
}

private void failCommit(Commit commit) {
commit.work().setFailed();
onCommitComplete.accept(CompleteCommit.forFailedWork(commit));
private void failQueuedCommit(Commit commit) {
for (Work w : commit.workBatch()) {
w.setFailed();
onCommitComplete.accept(
CompleteCommit.create(
commit.computationId(), w.getShardedKey(), w.id(), CommitStatus.ABORTED));
}
}

@Override
Expand All @@ -173,8 +175,8 @@ private void streamingCommitLoop() {
// take() blocks until a value is available in the commitQueue.
Preconditions.checkNotNull(initialCommit);

if (initialCommit.work().isFailed()) {
onCommitComplete.accept(CompleteCommit.forFailedWork(initialCommit));
if (shouldFailCommit(initialCommit)) {
failQueuedCommit(initialCommit);
initialCommit = null;
continue;
}
Expand All @@ -194,29 +196,61 @@ private void streamingCommitLoop() {
}
} finally {
if (initialCommit != null) {
failCommit(initialCommit);
failQueuedCommit(initialCommit);
}
}
}

boolean shouldFailCommit(Commit commit) {
for (Work w : commit.workBatch()) {
if (w.isFailed()) {
return true;
}
}
return false;
}

/** Adds the commit to the batch if it fits, returning true if it is consumed. */
private boolean tryAddToCommitBatch(Commit commit, CommitWorkStream.RequestBatcher batcher) {
Preconditions.checkNotNull(commit);
commit.work().setState(Work.State.COMMITTING);
activeCommitBytes.addAndGet(commit.getSize());
boolean isCommitAccepted =
batcher.commitWorkItem(
commit.computationId(),
commit.request(),
commitStatus -> {
onCommitComplete.accept(CompleteCommit.create(commit, commitStatus));
activeCommitBytes.addAndGet(-commit.getSize());
});
for (Work w : commit.workBatch()) {
w.setState(Work.State.COMMITTING);
}
activeCommitBytes.addAndGet(commit.getSerializedByteSize());
boolean isCommitAccepted;
if (commit.multiKeyRequest() != null) {
isCommitAccepted =
batcher.commitMultiKeyWorkItem(
commit.computationId(),
checkStateNotNull(commit.multiKeyRequest()),
commitStatus -> {
for (Work w : commit.workBatch()) {
onCommitComplete.accept(
CompleteCommit.create(
commit.computationId(), w.getShardedKey(), w.id(), commitStatus));
}
activeCommitBytes.addAndGet(-commit.getSerializedByteSize());
});
} else {
isCommitAccepted =
batcher.commitWorkItem(
commit.computationId(),
checkStateNotNull(commit.singleKeyRequest()),
commitStatus -> {
Work w = commit.workBatch().get(0);
onCommitComplete.accept(
CompleteCommit.create(
commit.computationId(), w.getShardedKey(), w.id(), commitStatus));
activeCommitBytes.addAndGet(-commit.getSerializedByteSize());
});
}

// Since the commit was not accepted, revert the changes made above.
if (!isCommitAccepted) {
commit.work().setState(Work.State.COMMIT_QUEUED);
activeCommitBytes.addAndGet(-commit.getSize());
for (Work w : commit.workBatch()) {
w.setState(Work.State.COMMIT_QUEUED);
}
activeCommitBytes.addAndGet(-commit.getSerializedByteSize());
}

return isCommitAccepted;
Expand Down Expand Up @@ -246,8 +280,8 @@ private boolean tryAddToCommitBatch(Commit commit, CommitWorkStream.RequestBatch
}

// Drop commits for failed work. Such commits will be dropped by Windmill anyway.
if (commit.work().isFailed()) {
onCommitComplete.accept(CompleteCommit.forFailedWork(commit));
if (shouldFailCommit(commit)) {
failQueuedCommit(commit);
continue;
}

Expand Down
Loading
Loading