diff --git a/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/CancellationSharer.java b/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/CancellationSharer.java new file mode 100644 index 000000000000..59ecf13f5891 --- /dev/null +++ b/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/CancellationSharer.java @@ -0,0 +1,140 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.cloud.pubsub.v1; + +import com.google.api.core.AbstractApiFuture; +import com.google.api.core.ApiFuture; +import com.google.api.core.ApiFutureCallback; +import com.google.api.core.ApiFutures; +import com.google.common.util.concurrent.MoreExecutors; +import com.google.pubsub.v1.PublishResponse; +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicReference; + +/** + * Coordinates multiple publish attempts for a single batch of messages. + * + *

Implements {@link ApiFuture} to act as the single future returned to the publisher's client. + * It manages the lifecycle of the original attempt and any subsequent hedged attempts. + */ +class CancellationSharer extends AbstractApiFuture { + private final Publisher.OutstandingBatch batch; + private final Publisher publisher; + private final Map> runningAttempts = + new ConcurrentHashMap<>(); + private final AtomicBoolean done = new AtomicBoolean(false); + final AtomicBoolean isInQueue = new AtomicBoolean(false); + private final AtomicReference lastError = new AtomicReference<>(); + + CancellationSharer(Publisher.OutstandingBatch batch, Publisher publisher) { + this.batch = batch; + this.publisher = publisher; + } + + /** + * Adds an attempt to be tracked by this coordinator. + * + * @param attemptNumber the 1-based index of the attempt (1 is original, 2+ are hedged) + * @param future the future representing the gRPC call for this attempt + */ + void addAttempt(final int attemptNumber, ApiFuture future) { + runningAttempts.put(attemptNumber, future); + + if (done.get()) { + future.cancel(true); + runningAttempts.remove(attemptNumber); + return; + } + + ApiFutures.addCallback( + future, + new ApiFutureCallback() { + @Override + public void onSuccess(PublishResponse result) { + handleAttemptSuccess(attemptNumber, result); + } + + @Override + public void onFailure(Throwable t) { + handleAttemptFailure(attemptNumber, t); + } + }, + MoreExecutors.directExecutor()); + } + + private void handleAttemptSuccess(int attemptNumber, PublishResponse response) { + if (done.compareAndSet(false, true)) { + set(response); // Resolve parent future + cancelAllExcept(attemptNumber); + publisher.refillTokenBucket(); + } + } + + private void handleAttemptFailure(int attemptNumber, Throwable t) { + runningAttempts.remove(attemptNumber); + + if (done.get()) { + return; + } + lastError.set(t); + if (runningAttempts.isEmpty() && done.compareAndSet(false, true)) { + setException(lastError.get()); + if (isInQueue.get()) { + publisher.removeFromHedgingQueue(this); + } + } + } + + void checkCompletionOnQueueExit() { + if (!done.get() && runningAttempts.isEmpty() && !isInQueue.get()) { + if (done.compareAndSet(false, true)) { + Throwable error = lastError.get(); + setException( + error != null ? error : new RuntimeException("Hedging failed with no active attempts")); + } + } + } + + private void cancelAllExcept(int successfulAttempt) { + for (Map.Entry> entry : runningAttempts.entrySet()) { + if (entry.getKey() != successfulAttempt) { + entry.getValue().cancel(true); + } + } + } + + @Override + public boolean cancel(boolean mayInterruptIfRunning) { + if (super.cancel(mayInterruptIfRunning)) { + done.set(true); + if (isInQueue.get()) { + publisher.removeFromHedgingQueue(this); + } + for (ApiFuture future : runningAttempts.values()) { + future.cancel(mayInterruptIfRunning); + } + return true; + } + return false; + } + + Publisher.OutstandingBatch getBatch() { + return batch; + } +} diff --git a/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/HedgeSettings.java b/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/HedgeSettings.java new file mode 100644 index 000000000000..47efec2e42ae --- /dev/null +++ b/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/HedgeSettings.java @@ -0,0 +1,143 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.cloud.pubsub.v1; + +import com.google.common.base.Preconditions; +import java.time.Duration; + +/** Settings for configuring publish hedging. */ +public final class HedgeSettings { + /** Default hedging delay. */ + private static final Duration DEFAULT_DELAY = Duration.ofMillis(1000); + + /** Default maximum number of tokens in the bucket. */ + private static final int DEFAULT_MAX_TOKENS = 50; + + /** Default refill rate (tokens per successful request). */ + private static final float DEFAULT_REFILL_RATIO = 0.1f; + + /** Hedging delay. */ + private final Duration hedgeDelay; + + /** Maximum tokens. */ + private final int maxTokens; + + /** Refill rate. */ + private final float refillRatio; + + private HedgeSettings(final Builder builder) { + this.hedgeDelay = builder.hedgeDelay; + this.maxTokens = builder.maxTokens; + this.refillRatio = builder.refillRatio; + } + + /** + * Returns the configured hedging delay. + * + * @return the hedging delay. + */ + Duration getHedgeDelay() { + return hedgeDelay; + } + + int getMaxTokens() { + return maxTokens; + } + + float getRefillRatio() { + return refillRatio; + } + + /** + * Returns a new builder for {@code HedgeSettings}. + * + * @return a new builder. + */ + public static Builder newBuilder() { + return new Builder(); + } + + /** Builder for {@code HedgeSettings}. */ + public static final class Builder { + /** Hedging delay. */ + private Duration hedgeDelay = DEFAULT_DELAY; + + /** Maximum tokens. */ + private int maxTokens = DEFAULT_MAX_TOKENS; + + /** Refill rate. */ + private float refillRatio = DEFAULT_REFILL_RATIO; + + private Builder() {} + + /** + * Allows hedging delay to be configurable. + * + * @param delay the hedging delay, must be 0.1s <= HedgeDelay <= 10s. + * @return this builder. + */ + public Builder setHedgeDelay(final Duration delay) { + Preconditions.checkNotNull(delay); + if (delay.toMillis() < 100 || delay.toMillis() > 10000) { + throw new IllegalArgumentException( + "hedgeDelay must be greater than or equal to 100ms and less than or equal to 10s"); + } + this.hedgeDelay = delay; + return this; + } + + /** + * Allows the maximum number of tokens in the bucket to be configurable. + * + * @param maxTokens the maximum number of tokens, must be 0 < MaxTokens <= 250. + * @return this builder. + */ + public Builder setMaxTokens(final int maxTokens) { + if (maxTokens <= 0 || maxTokens > 250) { + throw new IllegalArgumentException( + "maxTokens must be greater than 0 and less than or equal to 250"); + } + this.maxTokens = maxTokens; + return this; + } + + /** + * Allows the token bucket refill rate to be configurable. + * + * @param refill the refill rate (tokens per successful request), must be 0 < RefillRatio <= + * 0.2. + * @return this builder. + */ + public Builder setRefillRatio(final float refillRatio) { + if (refillRatio <= 0.0f || refillRatio > 0.2f) { + throw new IllegalArgumentException( + "refillRatio must be greater than 0.0 and less than or equal to 0.2"); + } + this.refillRatio = refillRatio; + return this; + } + + /** + * Builds an instance of {@code HedgeSettings}. + * + * @return the built {@code HedgeSettings} instance. + */ + public HedgeSettings build() { + return new HedgeSettings(this); + } + } +} diff --git a/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/HedgedRequest.java b/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/HedgedRequest.java new file mode 100644 index 000000000000..ad07de9a8a90 --- /dev/null +++ b/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/HedgedRequest.java @@ -0,0 +1,42 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.cloud.pubsub.v1; + +/** Represents a pending hedging check in the publisher's queue. */ +class HedgedRequest { + private final CancellationSharer coordinator; + private final int attemptNumber; + private final long sendAfterMs; + + HedgedRequest(CancellationSharer coordinator, int attemptNumber, long sendAfterMs) { + this.coordinator = coordinator; + this.attemptNumber = attemptNumber; + this.sendAfterMs = sendAfterMs; + } + + CancellationSharer getCoordinator() { + return coordinator; + } + + int getAttemptNumber() { + return attemptNumber; + } + + long getSendAfterMs() { + return sendAfterMs; + } +} diff --git a/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/OpenTelemetryPubsubTracer.java b/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/OpenTelemetryPubsubTracer.java index 3de4484586d0..a811e30e1452 100644 --- a/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/OpenTelemetryPubsubTracer.java +++ b/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/OpenTelemetryPubsubTracer.java @@ -179,6 +179,11 @@ void endPublishBatchingSpan(PubsubMessageWrapper message) { * links with the publisher parent span are created for sampled messages in the batch. */ Span startPublishRpcSpan(TopicName topicName, List messages) { + return startPublishRpcSpan(topicName, messages, 1); + } + + Span startPublishRpcSpan( + TopicName topicName, List messages, int attemptNumber) { if (!enabled) { return null; } @@ -203,7 +208,7 @@ Span startPublishRpcSpan(TopicName topicName, List message for (PubsubMessageWrapper message : messages) { if (publishRpcSpan.getSpanContext().isSampled()) { message.getPublisherSpan().addLink(publishRpcSpan.getSpanContext(), linkAttributes); - message.addPublishStartEvent(); + message.addPublishStartEvent(attemptNumber); } } return publishRpcSpan; diff --git a/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/Publisher.java b/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/Publisher.java index 56c920bcfdc1..85ff8c95c78e 100644 --- a/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/Publisher.java +++ b/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/Publisher.java @@ -18,11 +18,13 @@ import static com.google.common.util.concurrent.MoreExecutors.directExecutor; +import com.google.api.core.ApiClock; import com.google.api.core.ApiFunction; import com.google.api.core.ApiFuture; import com.google.api.core.ApiFutureCallback; import com.google.api.core.ApiFutures; import com.google.api.core.BetaApi; +import com.google.api.core.CurrentMillisClock; import com.google.api.core.SettableApiFuture; import com.google.api.gax.batching.BatchingSettings; import com.google.api.gax.batching.FlowControlSettings; @@ -46,6 +48,8 @@ import com.google.cloud.pubsub.v1.stub.PublisherStub; import com.google.cloud.pubsub.v1.stub.PublisherStubSettings; import com.google.common.base.Preconditions; +import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableMap; import com.google.protobuf.CodedOutputStream; import com.google.pubsub.v1.PublishRequest; import com.google.pubsub.v1.PublishResponse; @@ -59,18 +63,21 @@ import java.io.IOException; import java.time.Duration; import java.util.ArrayList; +import java.util.Collections; import java.util.HashMap; import java.util.Iterator; import java.util.LinkedList; import java.util.List; import java.util.Map; import java.util.concurrent.Callable; +import java.util.concurrent.ConcurrentLinkedQueue; import java.util.concurrent.CountDownLatch; import java.util.concurrent.Executor; import java.util.concurrent.ScheduledExecutorService; import java.util.concurrent.ScheduledFuture; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; import java.util.concurrent.locks.Lock; import java.util.concurrent.locks.ReentrantLock; import java.util.logging.Level; @@ -138,6 +145,26 @@ public class Publisher implements PublisherInterface { private final OpenTelemetry openTelemetry; private OpenTelemetryPubsubTracer tracer = new OpenTelemetryPubsubTracer(null, false); + private final HedgeSettings hedgeSettings; + private final Map> hedgingMetadata; + + /** + * Scale factor to represent decimal token values (e.g. 0.1 refill ratio) as integers inside the + * AtomicInteger token bucket. A scale of 100 allows representing decimal ratios down to 0.01 + * (1%). For example, 1.0 logical token is represented as 100. + */ + private static final int HEDGE_TOKEN_SCALE = 100; + + private final AtomicInteger hedgeTokenBucket = new AtomicInteger(); + private int scaledMaxHedgeTokens; + private int scaledHedgeRefillAmount; + private final ApiClock clock; + + private final ConcurrentLinkedQueue hedgingQueue; + private final AtomicBoolean isQueueProcessingScheduled; + private ScheduledFuture queueProcessingFuture; + private final Lock queueLock; + /** The maximum number of messages in one request. Defined by the API. */ public static long getApiMaxRequestElementCount() { return 1000L; @@ -230,10 +257,34 @@ private Publisher(Builder builder) throws IOException { backgroundResources = new BackgroundResourceAggregation(backgroundResourceList); shutdown = new AtomicBoolean(false); messagesWaiter = new Waiter(); + this.hedgeSettings = builder.hedgeSettings; + if (this.hedgeSettings != null) { + // Verify that the hedge delay is strictly less than the initial RPC timeout. + Duration hedgeDelay = this.hedgeSettings.getHedgeDelay(); + Duration initialRpcTimeout = builder.retrySettings.getInitialRpcTimeoutDuration(); + if (hedgeDelay.compareTo(initialRpcTimeout) >= 0) { + throw new IllegalArgumentException( + "hedgeDelay (" + + hedgeDelay.toMillis() + + "ms) must be strictly less than the initial RPC timeout duration (" + + initialRpcTimeout.toMillis() + + "ms)"); + } + this.scaledMaxHedgeTokens = this.hedgeSettings.getMaxTokens() * HEDGE_TOKEN_SCALE; + this.scaledHedgeRefillAmount = + (int) (this.hedgeSettings.getRefillRatio() * HEDGE_TOKEN_SCALE); + this.hedgeTokenBucket.set(0); + } + this.clock = builder.clock != null ? builder.clock : CurrentMillisClock.getDefaultClock(); this.publishContext = GrpcCallContext.createDefault(); + this.hedgingMetadata = ImmutableMap.of("x-goog-pubsub-hedged", ImmutableList.of("true")); this.publishContextWithCompression = GrpcCallContext.createDefault() .withCallOptions(CallOptions.DEFAULT.withCompression(GZIP_COMPRESSION)); + this.hedgingQueue = new ConcurrentLinkedQueue<>(); + this.isQueueProcessingScheduled = new AtomicBoolean(false); + this.queueLock = new ReentrantLock(); + this.queueProcessingFuture = null; } /** Topic which the publisher publishes to. */ @@ -246,6 +297,18 @@ public String getTopicNameString() { return topicName; } + /** Returns the configured hedging settings, or null if hedging is disabled. */ + public HedgeSettings getHedgeSettings() { + return hedgeSettings; + } + + Float getHedgeTokenBalance() { + if (hedgeSettings == null) { + return null; + } + return (float) hedgeTokenBucket.get() / HEDGE_TOKEN_SCALE; + } + /** * Schedules the publishing of a message. The publishing of the message may occur immediately or * be delayed based on the publisher batching options. @@ -481,10 +544,22 @@ private void publishAllWithoutInflightForKey(final String orderingKey) { } private ApiFuture publishCall(OutstandingBatch outstandingBatch) { + return publishCall(outstandingBatch, 1); + } + + private ApiFuture publishCall( + OutstandingBatch outstandingBatch, int attemptNumber) { GrpcCallContext context = publishContext; if (enableCompression && outstandingBatch.batchSizeBytes >= compressionBytesThreshold) { context = publishContextWithCompression; } + if (attemptNumber > 1) { + logger.log(Level.FINER, "Publishing hedged attempt {0}", attemptNumber); + context = + context + .withExtraHeaders(hedgingMetadata) + .withRetryableCodes(Collections.emptySet()); + } int numMessagesInBatch = outstandingBatch.size(); List pubsubMessagesList = new ArrayList(numMessagesInBatch); @@ -494,7 +569,8 @@ private ApiFuture publishCall(OutstandingBatch outstandingBatch pubsubMessagesList.add(messageWrapper.getPubsubMessage()); } - outstandingBatch.publishRpcSpan = tracer.startPublishRpcSpan(topicNameObject, messageWrappers); + outstandingBatch.publishRpcSpan = + tracer.startPublishRpcSpan(topicNameObject, messageWrappers, attemptNumber); return publisherStub .publishCallable() @@ -572,7 +648,11 @@ public void onFailure(Throwable t) { ApiFuture future; Executor callbackExecutor = directExecutor(); if (outstandingBatch.orderingKey == null || outstandingBatch.orderingKey.isEmpty()) { - future = publishCall(outstandingBatch); + if (hedgeSettings != null) { + future = startHedgedCall(outstandingBatch); + } else { + future = publishCall(outstandingBatch); + } } else { // If ordering key is specified, publish the batch using the sequential executor. future = @@ -588,7 +668,148 @@ public ApiFuture call() { ApiFutures.addCallback(future, futureCallback, callbackExecutor); } - private final class OutstandingBatch { + void refillTokenBucket() { + if (hedgeSettings != null) { + while (true) { + int current = hedgeTokenBucket.get(); + if (current >= scaledMaxHedgeTokens) { + return; + } + int next = Math.min(scaledMaxHedgeTokens, current + scaledHedgeRefillAmount); + if (hedgeTokenBucket.compareAndSet(current, next)) { + return; + } + } + } + } + + boolean tryAcquireHedgeToken() { + if (hedgeSettings == null) { + return false; + } + while (true) { + int current = hedgeTokenBucket.get(); + if (current < HEDGE_TOKEN_SCALE) { + return false; + } + int next = current - HEDGE_TOKEN_SCALE; + if (hedgeTokenBucket.compareAndSet(current, next)) { + return true; + } + } + } + + private ApiFuture startHedgedCall(final OutstandingBatch outstandingBatch) { + final CancellationSharer coordinator = new CancellationSharer(outstandingBatch, this); + + // Register cancellation listeners on client futures to propagate cancel to coordinator + final AtomicInteger cancelledCount = new AtomicInteger(0); + final int batchSize = outstandingBatch.outstandingPublishes.size(); + for (final OutstandingPublish outstanding : outstandingBatch.outstandingPublishes) { + outstanding.publishResult.addListener( + new Runnable() { + @Override + public void run() { + if (outstanding.publishResult.isCancelled()) { + if (cancelledCount.incrementAndGet() == batchSize) { + coordinator.cancel(true); + } + } + } + }, + directExecutor()); + } + + ApiFuture firstAttemptFuture = publishCall(outstandingBatch); + coordinator.addAttempt(1, firstAttemptFuture); + long delayMs = hedgeSettings.getHedgeDelay().toMillis(); + HedgedRequest item = new HedgedRequest(coordinator, 2, clock.millisTime() + delayMs); + hedgingQueue.add(item); + coordinator.isInQueue.set(true); + scheduleQueueProcessing(); + + return coordinator; + } + + private void scheduleQueueProcessing() { + if (isQueueProcessingScheduled.compareAndSet(false, true)) { + HedgedRequest nextItem = hedgingQueue.peek(); + if (nextItem == null) { + isQueueProcessingScheduled.set(false); + return; + } + + long delay = Math.max(0, nextItem.getSendAfterMs() - clock.millisTime()); + + queueProcessingFuture = + executor.schedule( + new Runnable() { + @Override + public void run() { + processQueue(); + } + }, + delay, + TimeUnit.MILLISECONDS); + } + } + + void removeFromHedgingQueue(CancellationSharer coordinator) { + queueLock.lock(); + try { + Iterator iterator = hedgingQueue.iterator(); + while (iterator.hasNext()) { + if (iterator.next().getCoordinator() == coordinator) { + iterator.remove(); + coordinator.isInQueue.set(false); + } + } + } finally { + queueLock.unlock(); + } + } + + private void processQueue() { + queueLock.lock(); + try { + isQueueProcessingScheduled.set(false); + long now = clock.millisTime(); + + HedgedRequest item; + while ((item = hedgingQueue.peek()) != null && item.getSendAfterMs() <= now) { + hedgingQueue.poll(); + + CancellationSharer coordinator = item.getCoordinator(); + if (coordinator.isDone()) { + coordinator.isInQueue.set(false); + continue; + } + + if (tryAcquireHedgeToken()) { + // Clone and schedule next attempt check (Attempt + 1) + long delayMs = hedgeSettings.getHedgeDelay().toMillis(); + HedgedRequest nextItem = + new HedgedRequest(coordinator, item.getAttemptNumber() + 1, now + delayMs); + hedgingQueue.add(nextItem); + + // Start Hedged Attempt + ApiFuture hedgedFuture = + publishCall(coordinator.getBatch(), item.getAttemptNumber()); + coordinator.addAttempt(item.getAttemptNumber(), hedgedFuture); + } else { + coordinator.isInQueue.set(false); + coordinator.checkCompletionOnQueueExit(); + } + } + + // Reschedule for next items + scheduleQueueProcessing(); + } finally { + queueLock.unlock(); + } + } + + final class OutstandingBatch { final List outstandingPublishes; final long creationTime; int attempt; @@ -600,7 +821,7 @@ private final class OutstandingBatch { List outstandingPublishes, int batchSizeBytes, String orderingKey) { this.outstandingPublishes = outstandingPublishes; attempt = 1; - creationTime = System.currentTimeMillis(); + creationTime = clock.millisTime(); this.batchSizeBytes = batchSizeBytes; this.orderingKey = orderingKey; } @@ -677,6 +898,9 @@ public void shutdown() { if (currentAlarmFuture != null && activeAlarm.getAndSet(false)) { currentAlarmFuture.cancel(false); } + if (queueProcessingFuture != null) { + queueProcessingFuture.cancel(false); + } publishAllOutstanding(); messagesWaiter.waitComplete(); backgroundResources.shutdown(); @@ -814,6 +1038,8 @@ public PubsubMessage apply(PubsubMessage input) { private boolean enableOpenTelemetryTracing = false; private OpenTelemetry openTelemetry = null; + private HedgeSettings hedgeSettings = null; + ApiClock clock = null; private Builder(String topic) { this.topicName = Preconditions.checkNotNull(topic); @@ -966,12 +1192,26 @@ public Builder setOpenTelemetry(OpenTelemetry openTelemetry) { return this; } + /** Configures the Publisher's hedging parameters. */ + public Builder setHedgeSettings(HedgeSettings hedgeSettings) { + this.hedgeSettings = hedgeSettings; + return this; + } + + Builder setClock(ApiClock clock) { + this.clock = clock; + return this; + } + /** Returns the default BatchingSettings used by the client if settings are not provided. */ public static BatchingSettings getDefaultBatchingSettings() { return DEFAULT_BATCHING_SETTINGS; } public Publisher build() throws IOException { + Preconditions.checkState( + !(enableMessageOrdering && hedgeSettings != null), + "Publish hedging and message ordering cannot be enabled at the same time."); return new Publisher(this); } } diff --git a/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/PubsubMessageWrapper.java b/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/PubsubMessageWrapper.java index 19864a26f5a1..91009aa135d5 100644 --- a/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/PubsubMessageWrapper.java +++ b/java-pubsub/google-cloud-pubsub/src/main/java/com/google/cloud/pubsub/v1/PubsubMessageWrapper.java @@ -42,6 +42,7 @@ public class PubsubMessageWrapper { private final int deliveryAttempt; private static final String PUBLISH_START_EVENT = "publish start"; + private static final String HEDGED_PUBLISH_START_EVENT = "publish start (hedged)"; private static final String PUBLISH_END_EVENT = "publish end"; private static final String MODACK_START_EVENT = "modack start"; @@ -183,8 +184,20 @@ void setSubscribeProcessSpan(Span span) { /** Creates a publish start event that is tied to the publish RPC span time. */ void addPublishStartEvent() { + addPublishStartEvent(1); + } + + /** + * Creates a publish start event that is tied to the publish RPC span time, marking hedged + * attempts explicitly. + */ + void addPublishStartEvent(int attemptNumber) { if (publisherSpan != null) { - publisherSpan.addEvent(PUBLISH_START_EVENT); + if (attemptNumber > 1) { + publisherSpan.addEvent(HEDGED_PUBLISH_START_EVENT); + } else { + publisherSpan.addEvent(PUBLISH_START_EVENT); + } } } diff --git a/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/FakePublisherServiceImpl.java b/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/FakePublisherServiceImpl.java index 9ab1dec73471..315c01a0ff55 100644 --- a/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/FakePublisherServiceImpl.java +++ b/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/FakePublisherServiceImpl.java @@ -20,6 +20,7 @@ import com.google.pubsub.v1.PublishRequest; import com.google.pubsub.v1.PublishResponse; import com.google.pubsub.v1.PublisherGrpc.PublisherImplBase; +import io.grpc.Metadata; import io.grpc.stub.StreamObserver; import java.time.Duration; import java.util.ArrayList; @@ -36,6 +37,7 @@ class FakePublisherServiceImpl extends PublisherImplBase { private final LinkedBlockingQueue requests = new LinkedBlockingQueue<>(); + private final LinkedBlockingQueue capturedHeaders = new LinkedBlockingQueue<>(); private final LinkedBlockingQueue publishResponses = new LinkedBlockingQueue<>(); private final AtomicInteger nextMessageId = new AtomicInteger(1); private boolean autoPublishResponse; @@ -81,7 +83,6 @@ public String toString() { @Override public void publish( PublishRequest request, final StreamObserver responseObserver) { - requests.add(request); Response response; try { if (autoPublishResponse) { @@ -98,6 +99,7 @@ public void publish( } if (responseDelay == Duration.ZERO) { sendResponse(response, responseObserver); + requests.add(request); } else { final Response responseToSend = response; executor.schedule( @@ -109,6 +111,7 @@ public void run() { }, responseDelay.toMillis(), TimeUnit.MILLISECONDS); + requests.add(request); } } @@ -160,4 +163,17 @@ public FakePublisherServiceImpl addPublishError(Throwable error) { public List getCapturedRequests() { return new ArrayList(requests); } + + public void recordHeaders(Metadata headers) { + capturedHeaders.add(headers); + } + + public List getCapturedHeaders() { + return new ArrayList<>(capturedHeaders); + } + + public void clearRequests() { + requests.clear(); + capturedHeaders.clear(); + } } diff --git a/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/HedgeSettingsTest.java b/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/HedgeSettingsTest.java new file mode 100644 index 000000000000..8ca202ed30ad --- /dev/null +++ b/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/HedgeSettingsTest.java @@ -0,0 +1,112 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.cloud.pubsub.v1; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertNotNull; +import static org.junit.Assert.assertThrows; + +import java.time.Duration; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.junit.runners.JUnit4; + +@RunWith(JUnit4.class) +public class HedgeSettingsTest { + + @Test + public void testDefaultSettings() { + HedgeSettings settings = HedgeSettings.newBuilder().build(); + assertNotNull(settings); + assertEquals(Duration.ofMillis(1000), settings.getHedgeDelay()); + assertEquals(50, settings.getMaxTokens()); + assertEquals(0.1f, settings.getRefillRatio(), 0.0001f); + } + + @Test + public void testCustomDelay() { + Duration customDelay = Duration.ofMillis(200); + HedgeSettings settings = HedgeSettings.newBuilder().setHedgeDelay(customDelay).build(); + assertNotNull(settings); + assertEquals(customDelay, settings.getHedgeDelay()); + } + + @Test + public void testDelayTooSmallThrows() { + assertThrows( + IllegalArgumentException.class, + () -> HedgeSettings.newBuilder().setHedgeDelay(Duration.ofMillis(99))); + } + + @Test + public void testDelayTooLargeThrows() { + assertThrows( + IllegalArgumentException.class, + () -> HedgeSettings.newBuilder().setHedgeDelay(Duration.ofMillis(10001))); + } + + @Test + public void testNullDelayThrows() { + assertThrows(NullPointerException.class, () -> HedgeSettings.newBuilder().setHedgeDelay(null)); + } + + @Test + public void testCustomMaxTokens() { + HedgeSettings settings = HedgeSettings.newBuilder().setMaxTokens(10).build(); + assertEquals(10, settings.getMaxTokens()); + } + + @Test + public void testNegativeMaxTokensThrows() { + assertThrows(IllegalArgumentException.class, () -> HedgeSettings.newBuilder().setMaxTokens(-5)); + } + + @Test + public void testZeroMaxTokensThrows() { + assertThrows(IllegalArgumentException.class, () -> HedgeSettings.newBuilder().setMaxTokens(0)); + } + + @Test + public void testMaxTokensTooLargeThrows() { + assertThrows( + IllegalArgumentException.class, () -> HedgeSettings.newBuilder().setMaxTokens(251)); + } + + @Test + public void testCustomRefill() { + HedgeSettings settings = HedgeSettings.newBuilder().setRefillRatio(0.15f).build(); + assertEquals(0.15f, settings.getRefillRatio(), 0.0001f); + } + + @Test + public void testNegativeRefillThrows() { + assertThrows( + IllegalArgumentException.class, () -> HedgeSettings.newBuilder().setRefillRatio(-0.1f)); + } + + @Test + public void testZeroRefillThrows() { + assertThrows( + IllegalArgumentException.class, () -> HedgeSettings.newBuilder().setRefillRatio(0.0f)); + } + + @Test + public void testRefillTooLargeThrows() { + assertThrows( + IllegalArgumentException.class, () -> HedgeSettings.newBuilder().setRefillRatio(0.21f)); + } +} diff --git a/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/OpenTelemetryTest.java b/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/OpenTelemetryTest.java index 52351ddef466..af3d70d79990 100644 --- a/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/OpenTelemetryTest.java +++ b/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/OpenTelemetryTest.java @@ -57,6 +57,7 @@ public class OpenTelemetryTest { private static final String PUBLISH_BATCHING_SPAN_NAME = "publisher batching"; private static final String PUBLISH_RPC_SPAN_NAME = FULL_TOPIC_NAME.getTopic() + " publish"; private static final String PUBLISH_START_EVENT = "publish start"; + private static final String HEDGED_PUBLISH_START_EVENT = "publish start (hedged)"; private static final String PUBLISH_END_EVENT = "publish end"; private static final String SUBSCRIBER_SPAN_NAME = @@ -656,6 +657,55 @@ public void testSubscribeRpcSpanFailures() { .hasEnded(); } + @Test + public void testHedgedPublishSpanEvents() { + PubsubMessage message = getPubsubMessage(); + PubsubMessageWrapper messageWrapper = + PubsubMessageWrapper.newBuilder(message, FULL_TOPIC_NAME).build(); + List messageWrappers = + java.util.Collections.singletonList(messageWrapper); + + Tracer openTelemetryTracer = openTelemetryTesting.getOpenTelemetry().getTracer("test"); + OpenTelemetryPubsubTracer tracer = new OpenTelemetryPubsubTracer(openTelemetryTracer, true); + + // Start Publisher span + tracer.startPublisherSpan(messageWrapper); + + // Original Attempt 1 + Span publishRpcSpan1 = tracer.startPublishRpcSpan(FULL_TOPIC_NAME, messageWrappers, 1); + tracer.endPublishRpcSpan(publishRpcSpan1); + + // Hedged Attempt 2 + Span publishRpcSpan2 = tracer.startPublishRpcSpan(FULL_TOPIC_NAME, messageWrappers, 2); + tracer.endPublishRpcSpan(publishRpcSpan2); + + // End Publisher span + tracer.endPublisherSpan(messageWrapper); + + List allSpans = openTelemetryTesting.getSpans(); + // 3 Spans: publishRpcSpan1, publishRpcSpan2, publisherSpan + assertEquals(3, allSpans.size()); + SpanData publisherSpanData = allSpans.get(2); + + // The publisher parent span should have 3 events: + // 1. "publish start" (from attempt 1) + // 2. "publish start (hedged)" (from attempt 2) + // 3. "publish end" (when publisher span ends) + assertEquals(3, publisherSpanData.getEvents().size()); + + EventDataAssert startEvent1Assert = + OpenTelemetryAssertions.assertThat(publisherSpanData.getEvents().get(0)); + startEvent1Assert.hasName(PUBLISH_START_EVENT); + + EventDataAssert startEvent2Assert = + OpenTelemetryAssertions.assertThat(publisherSpanData.getEvents().get(1)); + startEvent2Assert.hasName(HEDGED_PUBLISH_START_EVENT); + + EventDataAssert endEventAssert = + OpenTelemetryAssertions.assertThat(publisherSpanData.getEvents().get(2)); + endEventAssert.hasName(PUBLISH_END_EVENT); + } + private PubsubMessage getPubsubMessage() { return PubsubMessage.newBuilder() .setData(ByteString.copyFromUtf8("test-data")) diff --git a/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/PublisherImplTest.java b/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/PublisherImplTest.java index 8e6efaf372c9..aa2f2e660a09 100644 --- a/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/PublisherImplTest.java +++ b/java-pubsub/google-cloud-pubsub/src/test/java/com/google/cloud/pubsub/v1/PublisherImplTest.java @@ -35,6 +35,7 @@ import com.google.api.gax.grpc.testing.LocalChannelProvider; import com.google.api.gax.rpc.DataLossException; import com.google.api.gax.rpc.FixedTransportChannelProvider; +import com.google.api.gax.rpc.InvalidArgumentException; import com.google.api.gax.rpc.TransportChannelProvider; import com.google.cloud.pubsub.v1.Publisher.Builder; import com.google.protobuf.ByteString; @@ -43,7 +44,12 @@ import com.google.pubsub.v1.PublishResponse; import com.google.pubsub.v1.PubsubMessage; import io.grpc.ManagedChannel; +import io.grpc.Metadata; import io.grpc.Server; +import io.grpc.ServerCall; +import io.grpc.ServerCallHandler; +import io.grpc.ServerInterceptor; +import io.grpc.ServerInterceptors; import io.grpc.Status; import io.grpc.StatusException; import io.grpc.inprocess.InProcessChannelBuilder; @@ -98,13 +104,25 @@ public class PublisherImplTest { public void setUp() throws Exception { testPublisherServiceImpl = new FakePublisherServiceImpl(); + ServerInterceptor headerInterceptor = + new ServerInterceptor() { + @Override + public ServerCall.Listener interceptCall( + ServerCall call, Metadata headers, ServerCallHandler next) { + testPublisherServiceImpl.recordHeaders(headers); + return next.startCall(call, headers); + } + }; + InProcessServerBuilder serverBuilder = InProcessServerBuilder.forName("test-server"); - serverBuilder.addService(testPublisherServiceImpl); + serverBuilder.addService( + ServerInterceptors.intercept(testPublisherServiceImpl, headerInterceptor)); testServer = serverBuilder.build(); testChannel = InProcessChannelBuilder.forName("test-server").build(); testServer.start(); fakeExecutor = new FakeScheduledExecutorService(); + testPublisherServiceImpl.setExecutor(fakeExecutor); } @After @@ -1339,6 +1357,309 @@ public void testPublishOpenTelemetryTracing() throws Exception { .hasEnded(); } + @Test + public void testPublisherWithHedgeSettings() throws Exception { + HedgeSettings hedgeSettings = + HedgeSettings.newBuilder().setHedgeDelay(Duration.ofMillis(100)).build(); + Publisher publisher = getTestPublisherBuilder().setHedgeSettings(hedgeSettings).build(); + + assertThat(publisher.getHedgeSettings()).isEqualTo(hedgeSettings); + assertThat(publisher.getHedgeTokenBalance()).isNotNull(); + assertThat(publisher.getHedgeTokenBalance()).isWithin(0.0001f).of(0.0f); + + shutdownTestPublisher(publisher); + } + + @Test + public void testPublisherThrowsIfHedgeDelayGtRpcTimeout() throws Exception { + HedgeSettings hedgeSettings = + HedgeSettings.newBuilder().setHedgeDelay(Duration.ofMillis(500)).build(); + com.google.api.gax.retrying.RetrySettings retrySettings = + com.google.api.gax.retrying.RetrySettings.newBuilder() + .setInitialRpcTimeoutDuration(Duration.ofMillis(400)) + .setMaxRpcTimeoutDuration(Duration.ofMillis(400)) + .setTotalTimeoutDuration(Duration.ofSeconds(10)) + .build(); + + try { + getTestPublisherBuilder() + .setHedgeSettings(hedgeSettings) + .setRetrySettings(retrySettings) + .build(); + fail( + "Should have thrown IllegalArgumentException because hedgeDelay (500ms) > RPC timeout (400ms)"); + } catch (IllegalArgumentException e) { + assertThat(e.getMessage()) + .contains("must be strictly less than the initial RPC timeout duration"); + } + } + + @Test + public void testPublisherThrowsIfHedgeDelayEqRpcTimeout() throws Exception { + HedgeSettings hedgeSettings = + HedgeSettings.newBuilder().setHedgeDelay(Duration.ofMillis(500)).build(); + com.google.api.gax.retrying.RetrySettings retrySettings = + com.google.api.gax.retrying.RetrySettings.newBuilder() + .setInitialRpcTimeoutDuration(Duration.ofMillis(500)) + .setMaxRpcTimeoutDuration(Duration.ofMillis(500)) + .setTotalTimeoutDuration(Duration.ofSeconds(10)) + .build(); + + try { + getTestPublisherBuilder() + .setHedgeSettings(hedgeSettings) + .setRetrySettings(retrySettings) + .build(); + fail( + "Should have thrown IllegalArgumentException because hedgeDelay (500ms) == RPC timeout (500ms)"); + } catch (IllegalArgumentException e) { + assertThat(e.getMessage()) + .contains("must be strictly less than the initial RPC timeout duration"); + } + } + + @Test + public void testPublisherWithoutHedgeSettings() throws Exception { + Publisher publisher = getTestPublisherBuilder().build(); + + assertThat(publisher.getHedgeSettings()).isNull(); + assertThat(publisher.getHedgeTokenBalance()).isNull(); + + shutdownTestPublisher(publisher); + } + + private Publisher getPublisherWithHedge(Duration delay) throws Exception { + return getPublisherWithHedge(delay, 0.1f, 20); + } + + private Publisher getPublisherWithHedge(Duration delay, float refillRatio, int maxTokens) + throws Exception { + HedgeSettings hedgeSettings = + HedgeSettings.newBuilder() + .setHedgeDelay(delay) + .setRefillRatio(refillRatio) + .setMaxTokens(maxTokens) + .build(); + return getTestPublisherBuilder() + .setHedgeSettings(hedgeSettings) + .setClock(fakeExecutor.getClock()) + .setBatchingSettings( + Publisher.Builder.DEFAULT_BATCHING_SETTINGS.toBuilder() + .setElementCountThreshold(1L) + .build()) + .build(); + } + + private void fillTokenBucket(Publisher publisher, int tokensToFill) throws Exception { + testPublisherServiceImpl.setAutoPublishResponse(true); + for (int i = 0; i < tokensToFill; i++) { + ApiFuture future = sendTestMessage(publisher, "warmup-msg-" + i); + future.get(); + } + testPublisherServiceImpl.clearRequests(); + } + + private void waitForRequests(FakePublisherServiceImpl service, int expectedCount) + throws InterruptedException { + long timeout = System.currentTimeMillis() + 5000; + while (service.getCapturedRequests().size() < expectedCount + && System.currentTimeMillis() < timeout) { + Thread.sleep(5); + } + if (service.getCapturedRequests().size() < expectedCount) { + throw new AssertionError( + String.format( + "Timed out waiting for requests. Expected: %d, Got: %d", + expectedCount, service.getCapturedRequests().size())); + } + } + + @Test + public void testHedgingNotTriggeredIfFast() throws Exception { + Publisher publisher = getPublisherWithHedge(Duration.ofMillis(100)); + + // Prepare fast response (10ms delay) + testPublisherServiceImpl.setAutoPublishResponse(false); + testPublisherServiceImpl.setPublishResponseDelay(Duration.ofMillis(10)); + testPublisherServiceImpl.addPublishResponse(PublishResponse.newBuilder().addMessageIds("1")); + + ApiFuture future = sendTestMessage(publisher, "msg-fast"); + waitForRequests(testPublisherServiceImpl, 1); + + // Advance time past response but before hedge delay (e.g. 50ms) + fakeExecutor.advanceTime(Duration.ofMillis(50)); + + // Future should be completed + assertEquals("1", future.get()); + + // Only 1 request should be received by server + assertThat(testPublisherServiceImpl.getCapturedRequests()).hasSize(1); + + shutdownTestPublisher(publisher); + } + + @Test + public void testHedgingTriggeredIfSlow() throws Exception { + Publisher publisher = getPublisherWithHedge(Duration.ofMillis(100), 0.2f, 20); + fillTokenBucket(publisher, 5); + + // Set response delay to 200ms (greater than 100ms hedge delay) + testPublisherServiceImpl.setAutoPublishResponse(false); + testPublisherServiceImpl.setPublishResponseDelay(Duration.ofMillis(200)); + // Add two responses (one for main, one for hedge) + testPublisherServiceImpl.addPublishResponse(PublishResponse.newBuilder().addMessageIds("1")); + testPublisherServiceImpl.addPublishResponse(PublishResponse.newBuilder().addMessageIds("2")); + + ApiFuture future = sendTestMessage(publisher, "msg-slow"); + waitForRequests(testPublisherServiceImpl, 1); + + // Advance time to 80ms (before hedge delay) + fakeExecutor.advanceTime(Duration.ofMillis(80)); + assertThat(testPublisherServiceImpl.getCapturedRequests()).hasSize(1); // Only original sent + + // Advance time to 120ms (past 100ms hedge delay) + fakeExecutor.advanceTime(Duration.ofMillis(40)); + waitForRequests(testPublisherServiceImpl, 2); + + // Now attempt 2 should have been triggered + assertThat(testPublisherServiceImpl.getCapturedRequests()).hasSize(2); + + // Advance to 220ms to let responses complete + fakeExecutor.advanceTime(Duration.ofMillis(100)); + fakeExecutor.advanceTime(Duration.ZERO); // Drain pending tasks + assertEquals("1", future.get(5, TimeUnit.SECONDS)); + + List capturedHeaders = testPublisherServiceImpl.getCapturedHeaders(); + assertThat(capturedHeaders).hasSize(2); + Metadata.Key hedgedHeaderKey = + Metadata.Key.of("x-goog-pubsub-hedged", Metadata.ASCII_STRING_MARSHALLER); + // Original request should NOT have the header + assertThat(capturedHeaders.get(0).get(hedgedHeaderKey)).isNull(); + // Hedged request SHOULD have the header set to "true" + assertThat(capturedHeaders.get(1).get(hedgedHeaderKey)).isEqualTo("true"); + + shutdownTestPublisher(publisher); + } + + @Test + public void testMultipleHedging() throws Exception { + Publisher publisher = getPublisherWithHedge(Duration.ofMillis(100), 0.2f, 20); + fillTokenBucket(publisher, 10); + + // Set delay to 400ms + testPublisherServiceImpl.setAutoPublishResponse(false); + testPublisherServiceImpl.setPublishResponseDelay(Duration.ofMillis(400)); + // Add responses for 3 attempts + testPublisherServiceImpl.addPublishResponse(PublishResponse.newBuilder().addMessageIds("1")); + testPublisherServiceImpl.addPublishResponse(PublishResponse.newBuilder().addMessageIds("2")); + testPublisherServiceImpl.addPublishResponse(PublishResponse.newBuilder().addMessageIds("3")); + + ApiFuture future = sendTestMessage(publisher, "msg-very-slow"); + waitForRequests(testPublisherServiceImpl, 1); + + // T=0: Attempt 1 sent. + // T=120 (Hedge 1): Attempt 2 sent. + fakeExecutor.advanceTime(Duration.ofMillis(120)); + waitForRequests(testPublisherServiceImpl, 2); + assertThat(testPublisherServiceImpl.getCapturedRequests()).hasSize(2); + + // T=240 (Hedge 2): Attempt 3 sent. + fakeExecutor.advanceTime(Duration.ofMillis(120)); + waitForRequests(testPublisherServiceImpl, 3); + assertThat(testPublisherServiceImpl.getCapturedRequests()).hasSize(3); + + // Advance to complete + fakeExecutor.advanceTime(Duration.ofMillis(200)); + assertEquals("1", future.get(5, TimeUnit.SECONDS)); + + shutdownTestPublisher(publisher); + } + + @Test + public void testHedgingBypassedIfNoTokens() throws Exception { + Publisher publisher = getPublisherWithHedge(Duration.ofMillis(100)); + + // Drain the token bucket completely (since it starts full) + while (publisher.tryAcquireHedgeToken()) {} + assertThat(publisher.getHedgeTokenBalance()).isEqualTo(0.0f); + + testPublisherServiceImpl.setPublishResponseDelay(Duration.ofMillis(200)); + testPublisherServiceImpl.addPublishResponse(PublishResponse.newBuilder().addMessageIds("1")); + + ApiFuture future = sendTestMessage(publisher, "msg-slow-no-tokens"); + waitForRequests(testPublisherServiceImpl, 1); + + // Advance past hedge delay + fakeExecutor.advanceTime(Duration.ofMillis(120)); + + // Should NOT trigger hedge because token bucket is empty + assertThat(testPublisherServiceImpl.getCapturedRequests()).hasSize(1); + + fakeExecutor.advanceTime(Duration.ofMillis(100)); + assertEquals("1", future.get(5, TimeUnit.SECONDS)); + + shutdownTestPublisher(publisher); + } + + @Test + public void testHedgingCancellationPropagates() throws Exception { + Publisher publisher = getPublisherWithHedge(Duration.ofMillis(100), 0.2f, 20); + fillTokenBucket(publisher, 5); + + testPublisherServiceImpl.setAutoPublishResponse(false); + testPublisherServiceImpl.setPublishResponseDelay(Duration.ofMillis(200)); + testPublisherServiceImpl.addPublishResponse(PublishResponse.newBuilder().addMessageIds("1")); + testPublisherServiceImpl.addPublishResponse(PublishResponse.newBuilder().addMessageIds("2")); + + ApiFuture future = sendTestMessage(publisher, "msg-cancel"); + waitForRequests(testPublisherServiceImpl, 1); + + // Trigger hedge + fakeExecutor.advanceTime(Duration.ofMillis(120)); + waitForRequests(testPublisherServiceImpl, 2); + assertThat(testPublisherServiceImpl.getCapturedRequests()).hasSize(2); + + // Cancel the future + future.cancel(true); + + // Verify cancellation propagates to overall future + assertTrue(future.isCancelled()); + + shutdownTestPublisher(publisher); + } + + @Test + public void testNoHedgingIfOriginalFailsImmediately() throws Exception { + Publisher publisher = getPublisherWithHedge(Duration.ofMillis(100), 0.2f, 20); + fillTokenBucket(publisher, 5); + + // Configure the fake to immediately return an INVALID_ARGUMENT error + testPublisherServiceImpl.setAutoPublishResponse(false); + testPublisherServiceImpl.addPublishError(new StatusException(Status.INVALID_ARGUMENT)); + + ApiFuture future = sendTestMessage(publisher, "msg-fail-fast"); + + // The request should fail immediately without waiting or advancing time + try { + future.get(1, TimeUnit.SECONDS); + fail("Should have failed with ExecutionException"); + } catch (ExecutionException e) { + // expected + assertThat(e.getCause()).isInstanceOf(InvalidArgumentException.class); + } + + // Server should receive exactly 1 request (the original attempt) + assertThat(testPublisherServiceImpl.getCapturedRequests()).hasSize(1); + + // Advance time past the 100ms hedge delay and check that no hedge was sent + fakeExecutor.advanceTime(Duration.ofMillis(200)); + + // Captured requests should still be 1 (no hedge triggered) + assertThat(testPublisherServiceImpl.getCapturedRequests()).hasSize(1); + + shutdownTestPublisher(publisher); + } + private Builder getTestPublisherBuilder() { return Publisher.newBuilder(TEST_TOPIC) .setExecutorProvider(FixedExecutorProvider.create(fakeExecutor))