ziting-openai commented on code in PR #5513:
URL: https://github.com/apache/datafusion-comet/pull/5513#discussion_r3877184418
##########
spark/src/main/java/org/apache/comet/shuffle/CelebornShufflePartitionPusher.java:
##########
@@ -207,30 +479,503 @@ public void pushPartitionData(int partitionId, byte[]
data, int length) throws I
numPartitions,
true,
true);
- } catch (IllegalAccessException e) {
- throw new IOException("Cannot invoke the public Celeborn raw-push API",
e);
- } catch (InvocationTargetException e) {
- Throwable cause = e.getCause();
- if (cause instanceof IOException) {
- throw (IOException) cause;
+ submitted = accepted > 0;
+ if (submitted && observed != null) {
+ if (clientPushStates != null) {
+ Object current = ((Map<?, ?>)
clientPushStates.get(shuffleClient)).get(mapKey());
+ if (current != null && current != pushState) {
+ observed = observePushState(current);
+ }
+ }
+ markSubmitted(reservation, observed, accepted);
}
- if (cause instanceof RuntimeException) {
- throw (RuntimeException) cause;
+
+ int minimumAccepted = length + CELEBORN_BATCH_HEADER_BYTES;
+ if (accepted < minimumAccepted) {
+ throw new IOException(
+ "Celeborn raw shuffle push accepted "
+ + accepted
+ + " bytes; expected at least "
+ + minimumAccepted
+ + " including its transport header");
+ }
+ if (accepted > reservation.bytes) {
+ throw new IOException(
+ "Celeborn encrypted shuffle request exceeds its reserved in-flight
byte limit");
}
- if (cause instanceof Error) {
- throw (Error) cause;
+ throwIfAsyncFailure();
+ if (isAborted()) {
+ throw new IOException("Celeborn shuffle map attempt was aborted during
its push");
+ }
+ partitionLengths.addAndGet(partitionId, length);
+ } catch (IllegalAccessException cause) {
+ failure = new IOException("Cannot invoke the public Celeborn raw-push
API", cause);
+ abortAndSuppress(failure);
+ throw (IOException) failure;
+ } catch (InvocationTargetException cause) {
+ failure = unwrapFailure("Celeborn raw shuffle push failed", cause);
+ abortAndSuppress(failure);
+ throwFailure(failure);
+ throw new AssertionError("unreachable");
+ } catch (IOException | RuntimeException | Error cause) {
+ failure = cause;
+ abortAndSuppress(cause);
+ throw cause;
+ } finally {
+ if (insideClient) {
+ endClientPush();
+ }
+ if (reservation != null) {
+ if (registered && !submitted) {
+ releaseUnsubmittedPush(reservation);
+ } else if (!registered) {
+ admission.release(reservation.bytes);
+ }
+ }
+ try {
+ endPush();
+ } catch (IOException cleanupFailure) {
+ if (failure == null) {
+ throw cleanupFailure;
+ }
+ if (cleanupFailure != failure) {
+ failure.addSuppressed(cleanupFailure);
+ }
}
- throw new IOException("Celeborn raw shuffle push failed", cause);
}
+ }
- int minimumAccepted = length + CELEBORN_BATCH_HEADER_BYTES;
- if (accepted < minimumAccepted) {
+ private void validateFrame(int partitionId, byte[] data, int length) throws
IOException {
+ if (partitionId < 0 || partitionId >= numPartitions) {
+ throw new IOException("Celeborn output partition is outside this task's
partition count");
+ }
+ if (data == null) {
+ throw new IOException("Celeborn shuffle frame must not be null");
+ }
+ if (length > Integer.MAX_VALUE - CELEBORN_BATCH_HEADER_BYTES) {
+ throw new IOException("Celeborn shuffle frame and transport header
exceed the byte limit");
+ }
+ if (length < MINIMUM_COMET_FRAME_BYTES || length > data.length) {
+ throw new IOException("Celeborn shuffle frame length must describe one
complete frame");
+ }
+ if (length > maxFrameBytes) {
+ throw new IOException("Celeborn shuffle frame exceeds its configured
maximum frame size");
+ }
+
+ long declaredBodyLength =
ByteBuffer.wrap(data).order(ByteOrder.LITTLE_ENDIAN).getLong();
+ if (declaredBodyLength != (long) length - Long.BYTES) {
throw new IOException(
- "Celeborn raw shuffle push accepted "
- + accepted
- + " bytes; expected at least "
- + minimumAccepted
- + " including its transport header");
+ "Celeborn shuffle frame declares "
+ + declaredBodyLength
+ + " body bytes, but contains "
+ + (length - Long.BYTES));
+ }
+ }
+
+ private PushReservation claimEncodingReservation(int frameBytes) throws
IOException {
+ int required = Math.addExact(Math.multiplyExact(frameBytes, 3),
CELEBORN_BATCH_HEADER_BYTES);
+ PushReservation reservation = encodingReservation.get();
+ if (reservation != null) {
+ if (required > reservation.bytes) {
+ throw new IOException("Celeborn shuffle push exceeds its native
encoding reservation");
+ }
+ encodingReservation.remove();
+ synchronized (lifecycleLock) {
+ activeEncoders--;
+ lifecycleLock.notifyAll();
+ }
+ if (required < reservation.bytes) {
+ admission.release(reservation.bytes - required);
+ reservation.bytes = required;
+ }
+ return reservation;
+ }
+ admission.acquire(required, this::isAborted);
+ return new PushReservation(required);
+ }
+
+ private ObservedPushState observePushState(Object pushState) throws
IllegalAccessException {
+ synchronized (lifecycleLock) {
+ ObservedPushState observed = observedPushStates.get(pushState);
+ if (observed != null) {
+ return observed;
+ }
+ Object tracker = inFlightRequestTracker.get(pushState);
+ ObservedPushState created =
+ new ObservedPushState(
+ (LongAdder) totalInFlightRequests.get(tracker),
+ (AtomicReference<?>) pushStateException.get(pushState));
+ observedPushStates.put(pushState, created);
+ return created;
+ }
+ }
+
+ private void beginPush() throws IOException {
+ synchronized (lifecycleLock) {
+ if (asynchronousFailure != null) {
+ throw asynchronousFailure;
+ }
+ if (state != State.OPEN) {
+ throw new IOException("Celeborn shuffle map attempt no longer accepts
partition data");
+ }
+ activePushes++;
+ }
+ }
+
+ private void beginClientPush() throws IOException {
+ synchronized (lifecycleLock) {
+ if (state == State.ABORTED) {
+ throw new IOException("Celeborn shuffle map attempt was aborted before
its client push");
+ }
+ activeClientPushes++;
+ }
+ }
+
+ private void registerPendingPush(PushReservation reservation,
ObservedPushState pushState)
+ throws IOException {
+ synchronized (lifecycleLock) {
+ if (state == State.ABORTED) {
+ throw new IOException("Celeborn shuffle map attempt was aborted before
submission");
+ }
+ reservation.pushState = pushState;
+ pendingPushes.addLast(reservation);
+ if (completionReconciliation == null ||
completionReconciliation.isDone()) {
+ completionReconciliation =
+ COMPLETION_RECONCILER.scheduleWithFixedDelay(
+ this::safelyReconcileAcceptedPushes,
+ RECONCILIATION_INTERVAL_MILLIS,
+ RECONCILIATION_INTERVAL_MILLIS,
+ TimeUnit.MILLISECONDS);
+ }
+ }
+ }
+
+ private void markSubmitted(
+ PushReservation reservation, ObservedPushState pushState, int
acceptedBytes) {
+ int released;
+ synchronized (lifecycleLock) {
+ reservation.pushState = pushState;
+ reservation.submitted = true;
+ pushState.submittedPushes++;
+ int retained = Math.min(reservation.bytes, acceptedBytes);
+ released = reservation.bytes - retained;
+ reservation.bytes = retained;
+ }
+ if (released > 0) {
+ admission.release(released);
+ }
+ reconcileAcceptedPushes();
+ }
+
+ private void safelyReconcileAcceptedPushes() {
+ try {
+ reconcileAcceptedPushes();
+ } catch (RuntimeException | Error cause) {
+ IOException failure =
+ new IOException("Celeborn push completion reconciliation failed",
cause);
+ synchronized (lifecycleLock) {
+ if (asynchronousFailure == null) {
+ asynchronousFailure = failure;
+ } else {
+ failure = asynchronousFailure;
+ }
+ }
+ abortAndSuppress(failure);
+ }
+ }
+
+ private void reconcileAcceptedPushes() {
+ ArrayList<PushReservation> completed = new ArrayList<>();
+ IOException detectedFailure = null;
+ synchronized (lifecycleLock) {
+ for (ObservedPushState observed : observedPushStates.values()) {
+ Object failure = observed.exception.get();
+ if (failure instanceof IOException
+ && !"Cleaned Up".equals(((IOException) failure).getMessage())) {
+ IOException observedFailure = (IOException) failure;
+ // Only the raw transport callback's pinned message proves request
termination.
+ // PushState can also publish interruption/backpressure exceptions
while an older
+ // request is still live; those failures must never free that
request's admission.
+ if (!observed.terminalFailureCredited &&
isTerminalRawPushFailure(observedFailure)) {
+ observed.terminalFailureCredited = true;
+ }
+ if (asynchronousFailure == null) {
+ asynchronousFailure = observedFailure;
+ detectedFailure = asynchronousFailure;
+ }
+ }
+ long retainedRequests =
+ Math.max(
+ 0L, observed.inFlightRequests.sum() -
(observed.terminalFailureCredited ? 1L : 0L));
+ long completions = Math.max(0L, observed.submittedPushes -
retainedRequests);
Review Comment:
Rechecked at be5f4af7d7876e862af9be172b95e2e704acf049 with three independent
reviews. The callback tracker now retains ownership before request publication
and releases it only after the original callback returns, including retry
handoffs. This addresses the cleanup/failure leak and the
removal-before-callback interleaving described above; map membership is no
longer the completion signal. I traced the implementation and inspected the
added regressions; I did not rerun the Celeborn integration tests.
_Posted by Codex on behalf of ziting-openai._
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]