diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java index 9a9f22ffff..ca8ab183bc 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java @@ -79,16 +79,40 @@ public ModelUsage createEmptyModelUsage(RerankingCredentials rerankingCredential public Uni rerank( String query, List passages, RerankingCredentials rerankingCredentials) { - // TODO: what to do if passages is empty? List> passageBatches = createPassageBatches(passages); - List> batchRerankings = new ArrayList<>(); - - for (int batchId = 0; batchId < passageBatches.size(); batchId++) { - batchRerankings.add( - rerank(batchId, query, passageBatches.get(batchId), rerankingCredentials)); + if (passageBatches.isEmpty()) { + return Uni.createFrom().item(aggregateRanks(List.of())); } - return Uni.join().all(batchRerankings).andFailFast().map(this::aggregateRanks); + // Complete the first batch before starting the remaining batches. This primes upstream + // authorization caches without sacrificing concurrency for the rest of the request. + return rerank(0, query, passageBatches.getFirst(), rerankingCredentials) + .onItem() + .transformToUni( + firstBatchResponse -> { + if (passageBatches.size() == 1) { + return Uni.createFrom().item(List.of(firstBatchResponse)); + } + + List> remainingBatchRerankings = new ArrayList<>(); + for (int batchId = 1; batchId < passageBatches.size(); batchId++) { + remainingBatchRerankings.add( + rerank(batchId, query, passageBatches.get(batchId), rerankingCredentials)); + } + + return Uni.join() + .all(remainingBatchRerankings) + .andFailFast() + .map( + remainingBatchResponses -> { + List batchResponses = + new ArrayList<>(passageBatches.size()); + batchResponses.add(firstBatchResponse); + batchResponses.addAll(remainingBatchResponses); + return batchResponses; + }); + }) + .map(this::aggregateRanks); } /** diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderTest.java index 3c9add3bbf..b6335e0a0b 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderTest.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderTest.java @@ -4,11 +4,16 @@ import io.quarkus.test.junit.QuarkusTest; import io.quarkus.test.junit.TestProfile; +import io.smallrye.mutiny.Uni; import io.smallrye.mutiny.helpers.test.UniAssertSubscriber; import io.stargate.sgv2.jsonapi.TestConstants; import io.stargate.sgv2.jsonapi.api.request.RerankingCredentials; import io.stargate.sgv2.jsonapi.testresource.NoGlobalResourcesTestProfile; import java.util.List; +import java.util.Map; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.CopyOnWriteArrayList; import java.util.stream.IntStream; import org.junit.jupiter.api.Test; @@ -99,4 +104,141 @@ void microBatchingTest() { IntStream.range(0, 15) .forEach(i -> assertThat(finalResult.ranks().get(i).index()).isEqualTo(i)); } + + @Test + void primesFirstBatchBeforeStartingRemainingBatches() { + ControlledRerankingProvider provider = new ControlledRerankingProvider(); + + UniAssertSubscriber subscriber = + provider + .rerank("query", List.of("first", "second", "third"), RERANK_CREDENTIALS) + .subscribe() + .withSubscriber(UniAssertSubscriber.create()); + + assertThat(provider.invokedBatchIds()).containsExactly(0); + assertThat(provider.subscribedBatchIds()).containsExactly(0); + + provider.completeBatch(0); + + assertThat(provider.invokedBatchIds()).containsExactly(0, 1, 2); + assertThat(provider.subscribedBatchIds()).containsExactly(0, 1, 2); + assertThat(subscriber.getItem()).isNull(); + + // Completing one tail batch must not complete the result: tail batches stay concurrent. + provider.completeBatch(2); + assertThat(subscriber.getItem()).isNull(); + provider.completeBatch(1); + + RerankingProvider.RerankingResponse result = subscriber.awaitItem().getItem(); + assertThat(result.ranks()).extracting(RerankingProvider.Rank::index).containsExactly(0, 1, 2); + assertThat(result.modelUsage().batchCount()).isEqualTo(3); + } + + @Test + void doesNotStartRemainingBatchesWhenFirstBatchFails() { + ControlledRerankingProvider provider = new ControlledRerankingProvider(); + RuntimeException expected = new RuntimeException("first batch failed"); + + UniAssertSubscriber subscriber = + provider + .rerank("query", List.of("first", "second", "third"), RERANK_CREDENTIALS) + .subscribe() + .withSubscriber(UniAssertSubscriber.create()); + + provider.failBatch(0, expected); + + assertThat(subscriber.awaitFailure().getFailure()).isSameAs(expected); + assertThat(provider.invokedBatchIds()).containsExactly(0); + assertThat(provider.subscribedBatchIds()).containsExactly(0); + } + + @Test + void handlesSingleBatchWithoutStartingTailBatches() { + ControlledRerankingProvider provider = new ControlledRerankingProvider(); + + UniAssertSubscriber subscriber = + provider + .rerank("query", List.of("only"), RERANK_CREDENTIALS) + .subscribe() + .withSubscriber(UniAssertSubscriber.create()); + + assertThat(provider.invokedBatchIds()).containsExactly(0); + assertThat(provider.subscribedBatchIds()).containsExactly(0); + + provider.completeBatch(0); + + RerankingProvider.RerankingResponse result = subscriber.awaitItem().getItem(); + assertThat(result.ranks()).extracting(RerankingProvider.Rank::index).containsExactly(0); + assertThat(result.modelUsage().batchCount()).isEqualTo(1); + } + + @Test + void preservesEmptyPassagesResult() { + ControlledRerankingProvider provider = new ControlledRerankingProvider(); + + RerankingProvider.RerankingResponse result = + provider + .rerank("query", List.of(), RERANK_CREDENTIALS) + .subscribe() + .withSubscriber(UniAssertSubscriber.create()) + .awaitItem() + .getItem(); + + assertThat(provider.invokedBatchIds()).isEmpty(); + assertThat(provider.subscribedBatchIds()).isEmpty(); + assertThat(result.ranks()).isEmpty(); + assertThat(result.modelUsage()).isNull(); + } + + private static final class ControlledRerankingProvider extends TestRerankingProvider { + + private final Map> batchResults = + new ConcurrentHashMap<>(); + private final List invokedBatchIds = new CopyOnWriteArrayList<>(); + private final List subscribedBatchIds = new CopyOnWriteArrayList<>(); + + private ControlledRerankingProvider() { + super(1); + } + + @Override + public Uni rerank( + int batchId, + String query, + List passages, + RerankingCredentials rerankingCredentials) { + invokedBatchIds.add(batchId); + CompletableFuture batchResult = new CompletableFuture<>(); + batchResults.put(batchId, batchResult); + + return Uni.createFrom() + .deferred( + () -> { + subscribedBatchIds.add(batchId); + return Uni.createFrom().completionStage(batchResult); + }); + } + + private void completeBatch(int batchId) { + batchResults + .get(batchId) + .complete( + new BatchedRerankingResponse( + batchId, + List.of(new Rank(0, batchId)), + createEmptyModelUsage(RERANK_CREDENTIALS))); + } + + private void failBatch(int batchId, Throwable failure) { + batchResults.get(batchId).completeExceptionally(failure); + } + + private List invokedBatchIds() { + return List.copyOf(invokedBatchIds); + } + + private List subscribedBatchIds() { + return List.copyOf(subscribedBatchIds); + } + } }