From 4f41df0edded8e93b119c69557e7e30105948a66 Mon Sep 17 00:00:00 2001 From: Eric Hare Date: Mon, 24 Aug 2026 19:24:56 -0700 Subject: [PATCH 1/5] test: add reproducer for reranker truncate config --- .../NvidiaRerankingProviderTest.java | 19 ++++++++++++++++++- 1 file changed, 18 insertions(+), 1 deletion(-) diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java index 8abeec3a4d..10df0b6b48 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java @@ -28,7 +28,15 @@ public class NvidiaRerankingProviderTest { .RequestPropertiesImpl REQUEST_PROPERTIES = new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl - .RequestPropertiesImpl(3, 10, 100, 100, 0.5, 10); + .RequestPropertiesImpl( + 3, + 10, + 100, + 100, + 0.5, + 10, + RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties + .TruncateOption.END); private static final RerankingProvidersConfig.RerankingProviderConfig.ModelConfig MODEL_CONFIG = new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl( @@ -87,4 +95,13 @@ void testTenantIdIsExtractedFromCredentials() { .isNotNull() .isEqualTo(expectedTenantId); } + + @Test + void configuredTruncationIsSentInRequest() { + NvidiaRerankingProvider provider = new NvidiaRerankingProvider(MODEL_CONFIG); + + var request = provider.createRequest("test query", List.of("passage1", "passage2")); + + assertThat(request.truncate()).isEqualTo("END"); + } } From abd76bb357d974884e91e2c72e1f466979ab22a4 Mon Sep 17 00:00:00 2001 From: Eric Hare Date: Mon, 24 Aug 2026 19:29:17 -0700 Subject: [PATCH 2/5] fix: configure Nvidia reranker truncation --- .../RerankingProvidersConfig.java | 9 +++++++ .../RerankingProvidersConfigImpl.java | 23 ++++++++++++++-- .../operation/NvidiaRerankingProvider.java | 26 +++++++------------ .../resources/reranking-providers-config.yaml | 3 ++- .../test-reranking-providers-config.yaml | 3 ++- .../NvidiaRerankingProviderTest.java | 24 ++++++++++------- 6 files changed, 57 insertions(+), 31 deletions(-) diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProvidersConfig.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProvidersConfig.java index 4978debeff..c5a303b0cf 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProvidersConfig.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProvidersConfig.java @@ -73,6 +73,11 @@ interface ModelConfig { RequestProperties properties(); interface RequestProperties { + enum TruncateOption { + NONE, + END + } + /** * Specifies the maximum number of attempts before failing. Default is 3 (1 request + 2 * retries). @@ -120,6 +125,10 @@ interface RequestProperties { /** Maximum batch size supported by the provider. */ int maxBatchSize(); + + /** How the provider handles query and passage pairs that exceed the model token limit. */ + @WithDefault("NONE") + TruncateOption truncate(); } } } diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProvidersConfigImpl.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProvidersConfigImpl.java index cc5e0897f9..bd6a5197aa 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProvidersConfigImpl.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProvidersConfigImpl.java @@ -35,8 +35,27 @@ public record RequestPropertiesImpl( int readTimeoutMillis, int maxBackOffMillis, double jitter, - int maxBatchSize) - implements RerankingProviderConfig.ModelConfig.RequestProperties {} + int maxBatchSize, + TruncateOption truncate) + implements RerankingProviderConfig.ModelConfig.RequestProperties { + + public RequestPropertiesImpl( + int atMostRetries, + int initialBackOffMillis, + int readTimeoutMillis, + int maxBackOffMillis, + double jitter, + int maxBatchSize) { + this( + atMostRetries, + initialBackOffMillis, + readTimeoutMillis, + maxBackOffMillis, + jitter, + maxBatchSize, + TruncateOption.NONE); + } + } } } } diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProvider.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProvider.java index e2a9ff2979..0ea876725d 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProvider.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProvider.java @@ -60,17 +60,6 @@ public class NvidiaRerankingProvider extends RerankingProvider { private final NvidiaRerankingClient nvidiaClient; - /** - * Nvidia Reranking Service supports truncation or error behavior when the passage is too long. - * - *

The Data API uses {@code NONE} as the default, which means the reranking request will error - * out if there is a query and passage pair that exceeds the allowed token size of 8192. - * - *

See: - * https://docs.nvidia.com/nim/nemo-retriever/text-reranking/latest/using-reranking.html#token-limits-truncation - */ - private static final String TRUNCATE_PASSAGE = "NONE"; - public NvidiaRerankingProvider( RerankingProvidersConfig.RerankingProviderConfig.ModelConfig modelConfig) { super(ModelProvider.NVIDIA, modelConfig); @@ -97,12 +86,7 @@ public Uni rerank( } var accessToken = HttpConstants.BEARER_PREFIX_FOR_API_KEY + rerankingCredentials.apiKey(); - var nvidiaRequest = - new NvidiaRerankingRequest( - modelName(), - new NvidiaRerankingRequest.TextWrapper(query), - passages.stream().map(NvidiaRerankingRequest.TextWrapper::new).toList(), - TRUNCATE_PASSAGE); + var nvidiaRequest = createRequest(query, passages); final long callStartNano = System.nanoTime(); return retryHTTPCall( @@ -132,6 +116,14 @@ public Uni rerank( }); } + NvidiaRerankingRequest createRequest(String query, List passages) { + return new NvidiaRerankingRequest( + modelName(), + new NvidiaRerankingRequest.TextWrapper(query), + passages.stream().map(NvidiaRerankingRequest.TextWrapper::new).toList(), + modelConfig.properties().truncate().name()); + } + /** * REST client interface for the Nvidia Reranking Service. * diff --git a/src/main/resources/reranking-providers-config.yaml b/src/main/resources/reranking-providers-config.yaml index 67739fe0a2..a851a20eb1 100644 --- a/src/main/resources/reranking-providers-config.yaml +++ b/src/main/resources/reranking-providers-config.yaml @@ -14,4 +14,5 @@ stargate: is-default: true url: https://us-west-2.api-dev.ai.datastax.com/nvidia/v1/ranking properties: - max-batch-size: 10 \ No newline at end of file + max-batch-size: 10 + truncate: END diff --git a/src/main/resources/test-reranking-providers-config.yaml b/src/main/resources/test-reranking-providers-config.yaml index 7f1350470a..1ad6d9fa6e 100644 --- a/src/main/resources/test-reranking-providers-config.yaml +++ b/src/main/resources/test-reranking-providers-config.yaml @@ -15,6 +15,7 @@ stargate: url: https://us-west-2.api-dev.ai.datastax.com/nvidia/v1/ranking properties: max-batch-size: 10 + truncate: END - name: nvidia/a-random-deprecated-model api-model-support: status: DEPRECATED @@ -28,4 +29,4 @@ stargate: message: This model is at END_OF_LIFE status, it is not supported. url: https://us-west-2.api-dev.ai.datastax.com/nvidia/v1/ranking properties: - max-batch-size: 10 \ No newline at end of file + max-batch-size: 10 diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java index 10df0b6b48..11f18d5f9b 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java @@ -13,6 +13,7 @@ import io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig; import io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfigImpl; import io.stargate.sgv2.jsonapi.testresource.NoGlobalResourcesTestProfile; +import jakarta.inject.Inject; import java.util.List; import java.util.Optional; import org.junit.jupiter.api.Test; @@ -28,15 +29,7 @@ public class NvidiaRerankingProviderTest { .RequestPropertiesImpl REQUEST_PROPERTIES = new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl - .RequestPropertiesImpl( - 3, - 10, - 100, - 100, - 0.5, - 10, - RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties - .TruncateOption.END); + .RequestPropertiesImpl(3, 10, 100, 100, 0.5, 10); private static final RerankingProvidersConfig.RerankingProviderConfig.ModelConfig MODEL_CONFIG = new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl( @@ -47,6 +40,8 @@ public class NvidiaRerankingProviderTest { "https://us-west-2.api-dev.ai.datastax.com/nvidia/v1/ranking", REQUEST_PROPERTIES); + @Inject RerankingProvidersConfig rerankingProvidersConfig; + @Test void testEmptyApiKeyThrowsException() { NvidiaRerankingProvider provider = new NvidiaRerankingProvider(MODEL_CONFIG); @@ -98,10 +93,19 @@ void testTenantIdIsExtractedFromCredentials() { @Test void configuredTruncationIsSentInRequest() { - NvidiaRerankingProvider provider = new NvidiaRerankingProvider(MODEL_CONFIG); + var modelConfig = rerankingProvidersConfig.providers().get("nvidia").models().getFirst(); + NvidiaRerankingProvider provider = new NvidiaRerankingProvider(modelConfig); var request = provider.createRequest("test query", List.of("passage1", "passage2")); assertThat(request.truncate()).isEqualTo("END"); } + + @Test + void programmaticConfigDefaultsTruncationToNone() { + assertThat(REQUEST_PROPERTIES.truncate()) + .isEqualTo( + RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties + .TruncateOption.NONE); + } } From a5a72ce25860f02db0b3a5db4940bdb1c543c990 Mon Sep 17 00:00:00 2001 From: Eric Hare Date: Mon, 24 Aug 2026 19:41:58 -0700 Subject: [PATCH 3/5] test: verify Nvidia rerank truncate payload --- .../NvidiaRerankingProviderTest.java | 76 +++++++++++++++++++ 1 file changed, 76 insertions(+) diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java index 11f18d5f9b..6ba58a6b85 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java @@ -1,8 +1,16 @@ package io.stargate.sgv2.jsonapi.service.reranking.operation; +import static com.github.tomakehurst.wiremock.client.WireMock.aResponse; +import static com.github.tomakehurst.wiremock.client.WireMock.equalTo; +import static com.github.tomakehurst.wiremock.client.WireMock.matchingJsonPath; +import static com.github.tomakehurst.wiremock.client.WireMock.post; +import static com.github.tomakehurst.wiremock.client.WireMock.postRequestedFor; +import static com.github.tomakehurst.wiremock.client.WireMock.urlEqualTo; +import static com.github.tomakehurst.wiremock.client.WireMock.verify; import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatThrownBy; +import com.github.tomakehurst.wiremock.WireMockServer; import io.quarkus.test.junit.QuarkusTest; import io.quarkus.test.junit.TestProfile; import io.smallrye.mutiny.helpers.test.UniAssertSubscriber; @@ -14,8 +22,12 @@ import io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfigImpl; import io.stargate.sgv2.jsonapi.testresource.NoGlobalResourcesTestProfile; import jakarta.inject.Inject; +import jakarta.ws.rs.core.HttpHeaders; +import jakarta.ws.rs.core.MediaType; import java.util.List; import java.util.Optional; +import org.junit.jupiter.api.AfterAll; +import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Test; /** Tests for {@link NvidiaRerankingProvider} */ @@ -24,6 +36,9 @@ public class NvidiaRerankingProviderTest { private static final TestConstants testConstants = new TestConstants(); + private static final String NVIDIA_PATH = "/v1/ranking"; + private static final String NVIDIA_URL = "http://localhost:8080" + NVIDIA_PATH; + private static WireMockServer wireMockServer; private static final RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl .RequestPropertiesImpl @@ -40,8 +55,53 @@ public class NvidiaRerankingProviderTest { "https://us-west-2.api-dev.ai.datastax.com/nvidia/v1/ranking", REQUEST_PROPERTIES); + private static final RerankingCredentials RERANKING_CREDENTIALS = + new RerankingCredentials(testConstants.TENANT, "mocked reranking api key"); + + private static final RerankingProvidersConfig.RerankingProviderConfig.ModelConfig + END_MODEL_CONFIG = + new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl( + "nvidia/llama-3.2-nv-rerankqa-1b-v2", + new ApiModelSupport.ApiModelSupportImpl( + ApiModelSupport.SupportStatus.SUPPORTED, Optional.empty()), + false, + NVIDIA_URL, + new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl + .RequestPropertiesImpl( + 3, + 10, + 100, + 100, + 0.5, + 10, + RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties + .TruncateOption.END)); + @Inject RerankingProvidersConfig rerankingProvidersConfig; + @BeforeAll + static void startWireMock() { + wireMockServer = new WireMockServer(); + wireMockServer.start(); + wireMockServer.stubFor( + post(urlEqualTo(NVIDIA_PATH)) + .willReturn( + aResponse() + .withHeader(HttpHeaders.CONTENT_TYPE, MediaType.APPLICATION_JSON) + .withBody( + """ + { + "rankings": [{"index": 0, "logit": 0.75}], + "usage": {"prompt_tokens": 12, "total_tokens": 12} + } + """))); + } + + @AfterAll + static void stopWireMock() { + wireMockServer.stop(); + } + @Test void testEmptyApiKeyThrowsException() { NvidiaRerankingProvider provider = new NvidiaRerankingProvider(MODEL_CONFIG); @@ -108,4 +168,20 @@ void programmaticConfigDefaultsTruncationToNone() { RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties .TruncateOption.NONE); } + + @Test + void sendsConfiguredTruncation() { + NvidiaRerankingProvider provider = new NvidiaRerankingProvider(END_MODEL_CONFIG); + + provider + .rerank(1, "test query", List.of("test passage"), RERANKING_CREDENTIALS) + .subscribe() + .withSubscriber(UniAssertSubscriber.create()) + .awaitItem() + .getItem(); + + verify( + postRequestedFor(urlEqualTo(NVIDIA_PATH)) + .withRequestBody(matchingJsonPath("$.truncate", equalTo("END")))); + } } From 160a6e8ebe8f95eb3d24e0a2eb2549f48f8befa7 Mon Sep 17 00:00:00 2001 From: Eric Hare Date: Mon, 24 Aug 2026 19:44:52 -0700 Subject: [PATCH 4/5] refactor: make rerank truncation explicit --- .../RerankingProviderConfigProducer.java | 5 +++- .../RerankingProvidersConfigImpl.java | 20 +------------ .../operation/NvidiaRerankingProvider.java | 2 +- .../gateway/RerankingGatewayClientTest.java | 3 +- .../NvidiaRerankingProviderTest.java | 28 ++++++------------- .../operation/TestRerankingProvider.java | 6 ++-- .../FindAndRerankOperationBuilderTest.java | 3 +- 7 files changed, 22 insertions(+), 45 deletions(-) diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderConfigProducer.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderConfigProducer.java index eb447b3133..71d39e25d1 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderConfigProducer.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderConfigProducer.java @@ -1,5 +1,7 @@ package io.stargate.sgv2.jsonapi.service.reranking.configuration; +import static io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties.TruncateOption.NONE; + import io.quarkus.grpc.GrpcClient; import io.quarkus.runtime.Startup; import io.stargate.embedding.gateway.EmbeddingGateway; @@ -214,7 +216,8 @@ private RerankingProvidersConfig.RerankingProviderConfig createRerankingProvider model.getProperties().getReadTimeoutMillis(), model.getProperties().getMaxBackOffMillis(), model.getProperties().getJitter(), - model.getProperties().getMaxBatchSize()))) + model.getProperties().getMaxBatchSize(), + NONE))) .collect(Collectors.toList()); return new RerankingProvidersConfigImpl.RerankingProviderConfigImpl( diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProvidersConfigImpl.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProvidersConfigImpl.java index bd6a5197aa..96bfcc7d18 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProvidersConfigImpl.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProvidersConfigImpl.java @@ -37,25 +37,7 @@ public record RequestPropertiesImpl( double jitter, int maxBatchSize, TruncateOption truncate) - implements RerankingProviderConfig.ModelConfig.RequestProperties { - - public RequestPropertiesImpl( - int atMostRetries, - int initialBackOffMillis, - int readTimeoutMillis, - int maxBackOffMillis, - double jitter, - int maxBatchSize) { - this( - atMostRetries, - initialBackOffMillis, - readTimeoutMillis, - maxBackOffMillis, - jitter, - maxBatchSize, - TruncateOption.NONE); - } - } + implements RerankingProviderConfig.ModelConfig.RequestProperties {} } } } diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProvider.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProvider.java index 0ea876725d..4f630a5ae6 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProvider.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProvider.java @@ -116,7 +116,7 @@ public Uni rerank( }); } - NvidiaRerankingRequest createRequest(String query, List passages) { + private NvidiaRerankingRequest createRequest(String query, List passages) { return new NvidiaRerankingRequest( modelName(), new NvidiaRerankingRequest.TextWrapper(query), diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingGatewayClientTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingGatewayClientTest.java index 9d7d0da39f..8bb75183be 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingGatewayClientTest.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingGatewayClientTest.java @@ -1,5 +1,6 @@ package io.stargate.sgv2.jsonapi.service.reranking.gateway; +import static io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties.TruncateOption.NONE; import static org.assertj.core.api.Assertions.assertThat; import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.mock; @@ -46,7 +47,7 @@ public class RerankingGatewayClientTest { .RequestPropertiesImpl REQUEST_PROPERTIES = new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl - .RequestPropertiesImpl(3, 10, 100, 100, 0.5, 10); + .RequestPropertiesImpl(3, 10, 100, 100, 0.5, 10, NONE); private static final RerankingProvidersConfig.RerankingProviderConfig.ModelConfig MODEL_CONFIG = new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl( diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java index 6ba58a6b85..8d2eeda8b1 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java @@ -7,6 +7,8 @@ import static com.github.tomakehurst.wiremock.client.WireMock.postRequestedFor; import static com.github.tomakehurst.wiremock.client.WireMock.urlEqualTo; import static com.github.tomakehurst.wiremock.client.WireMock.verify; +import static io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties.TruncateOption.END; +import static io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties.TruncateOption.NONE; import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatThrownBy; @@ -44,7 +46,7 @@ public class NvidiaRerankingProviderTest { .RequestPropertiesImpl REQUEST_PROPERTIES = new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl - .RequestPropertiesImpl(3, 10, 100, 100, 0.5, 10); + .RequestPropertiesImpl(3, 10, 100, 100, 0.5, 10, NONE); private static final RerankingProvidersConfig.RerankingProviderConfig.ModelConfig MODEL_CONFIG = new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl( @@ -67,15 +69,7 @@ public class NvidiaRerankingProviderTest { false, NVIDIA_URL, new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl - .RequestPropertiesImpl( - 3, - 10, - 100, - 100, - 0.5, - 10, - RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties - .TruncateOption.END)); + .RequestPropertiesImpl(3, 10, 100, 100, 0.5, 10, END)); @Inject RerankingProvidersConfig rerankingProvidersConfig; @@ -152,21 +146,15 @@ void testTenantIdIsExtractedFromCredentials() { } @Test - void configuredTruncationIsSentInRequest() { + void mapsConfiguredTruncationFromYaml() { var modelConfig = rerankingProvidersConfig.providers().get("nvidia").models().getFirst(); - NvidiaRerankingProvider provider = new NvidiaRerankingProvider(modelConfig); - var request = provider.createRequest("test query", List.of("passage1", "passage2")); - - assertThat(request.truncate()).isEqualTo("END"); + assertThat(modelConfig.properties().truncate()).isEqualTo(END); } @Test - void programmaticConfigDefaultsTruncationToNone() { - assertThat(REQUEST_PROPERTIES.truncate()) - .isEqualTo( - RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties - .TruncateOption.NONE); + void programmaticConfigUsesExplicitNone() { + assertThat(REQUEST_PROPERTIES.truncate()).isEqualTo(NONE); } @Test diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/TestRerankingProvider.java b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/TestRerankingProvider.java index dc30b8ac70..dea314002c 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/TestRerankingProvider.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/TestRerankingProvider.java @@ -1,5 +1,7 @@ package io.stargate.sgv2.jsonapi.service.reranking.operation; +import static io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties.TruncateOption.NONE; + import io.smallrye.mutiny.Uni; import io.stargate.sgv2.jsonapi.TestConstants; import io.stargate.sgv2.jsonapi.api.request.RerankingCredentials; @@ -24,7 +26,7 @@ public class TestRerankingProvider extends RerankingProvider { .RequestPropertiesImpl REQUEST_PROPERTIES = new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl - .RequestPropertiesImpl(3, 10, 100, 100, 0.5, 10); + .RequestPropertiesImpl(3, 10, 100, 100, 0.5, 10, NONE); private static final RerankingProvidersConfig.RerankingProviderConfig.ModelConfig MODEL_CONFIG = new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl( @@ -49,7 +51,7 @@ protected TestRerankingProvider(int maxBatchSize) { false, "http://testing.com", new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl - .RequestPropertiesImpl(3, 100, 5000, 500, 0.5, maxBatchSize))); + .RequestPropertiesImpl(3, 100, 5000, 500, 0.5, maxBatchSize, NONE))); } @Override diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/resolver/FindAndRerankOperationBuilderTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/resolver/FindAndRerankOperationBuilderTest.java index 3ca81fc205..12802d4f85 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/resolver/FindAndRerankOperationBuilderTest.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/resolver/FindAndRerankOperationBuilderTest.java @@ -1,5 +1,6 @@ package io.stargate.sgv2.jsonapi.service.resolver; +import static io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties.TruncateOption.NONE; import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatCode; import static org.assertj.core.api.Assertions.assertThatThrownBy; @@ -50,7 +51,7 @@ class FindAndRerankOperationBuilderTest { .RequestPropertiesImpl REQUEST_PROPERTIES = new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl - .RequestPropertiesImpl(3, 10, 100, 100, 0.5, 10); + .RequestPropertiesImpl(3, 10, 100, 100, 0.5, 10, NONE); private static RerankingProvidersConfig.RerankingProviderConfig.ModelConfig modelConfig( String name, ApiModelSupport.SupportStatus status) { From b70bc49c0b11ac75fbefbcb741e9d7a68cab4a66 Mon Sep 17 00:00:00 2001 From: Eric Hare Date: Mon, 24 Aug 2026 19:54:24 -0700 Subject: [PATCH 5/5] refactor: simplify reranking truncate coverage --- .../operation/NvidiaRerankingProvider.java | 18 ++++----- .../NvidiaRerankingProviderTest.java | 39 ++++++------------- 2 files changed, 19 insertions(+), 38 deletions(-) diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProvider.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProvider.java index 4f630a5ae6..c1343d8160 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProvider.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProvider.java @@ -10,6 +10,7 @@ import io.stargate.sgv2.jsonapi.service.provider.ModelProvider; import io.stargate.sgv2.jsonapi.service.provider.ProviderBillingFilter; import io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig; +import io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties.TruncateOption; import jakarta.ws.rs.HeaderParam; import jakarta.ws.rs.POST; import jakarta.ws.rs.core.HttpHeaders; @@ -86,7 +87,12 @@ public Uni rerank( } var accessToken = HttpConstants.BEARER_PREFIX_FOR_API_KEY + rerankingCredentials.apiKey(); - var nvidiaRequest = createRequest(query, passages); + var nvidiaRequest = + new NvidiaRerankingRequest( + modelName(), + new NvidiaRerankingRequest.TextWrapper(query), + passages.stream().map(NvidiaRerankingRequest.TextWrapper::new).toList(), + modelConfig.properties().truncate()); final long callStartNano = System.nanoTime(); return retryHTTPCall( @@ -116,14 +122,6 @@ public Uni rerank( }); } - private NvidiaRerankingRequest createRequest(String query, List passages) { - return new NvidiaRerankingRequest( - modelName(), - new NvidiaRerankingRequest.TextWrapper(query), - passages.stream().map(NvidiaRerankingRequest.TextWrapper::new).toList(), - modelConfig.properties().truncate().name()); - } - /** * REST client interface for the Nvidia Reranking Service. * @@ -148,7 +146,7 @@ Uni rerank( *

.. */ public record NvidiaRerankingRequest( - String model, TextWrapper query, List passages, String truncate) { + String model, TextWrapper query, List passages, TruncateOption truncate) { /** * query and passage string needs to be wrapped in with text key for request to the Nvidia diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java index 8d2eeda8b1..5311b25f9a 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/NvidiaRerankingProviderTest.java @@ -6,8 +6,6 @@ import static com.github.tomakehurst.wiremock.client.WireMock.post; import static com.github.tomakehurst.wiremock.client.WireMock.postRequestedFor; import static com.github.tomakehurst.wiremock.client.WireMock.urlEqualTo; -import static com.github.tomakehurst.wiremock.client.WireMock.verify; -import static io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties.TruncateOption.END; import static io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig.RerankingProviderConfig.ModelConfig.RequestProperties.TruncateOption.NONE; import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatThrownBy; @@ -60,17 +58,6 @@ public class NvidiaRerankingProviderTest { private static final RerankingCredentials RERANKING_CREDENTIALS = new RerankingCredentials(testConstants.TENANT, "mocked reranking api key"); - private static final RerankingProvidersConfig.RerankingProviderConfig.ModelConfig - END_MODEL_CONFIG = - new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl( - "nvidia/llama-3.2-nv-rerankqa-1b-v2", - new ApiModelSupport.ApiModelSupportImpl( - ApiModelSupport.SupportStatus.SUPPORTED, Optional.empty()), - false, - NVIDIA_URL, - new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl - .RequestPropertiesImpl(3, 10, 100, 100, 0.5, 10, END)); - @Inject RerankingProvidersConfig rerankingProvidersConfig; @BeforeAll @@ -146,20 +133,16 @@ void testTenantIdIsExtractedFromCredentials() { } @Test - void mapsConfiguredTruncationFromYaml() { - var modelConfig = rerankingProvidersConfig.providers().get("nvidia").models().getFirst(); - - assertThat(modelConfig.properties().truncate()).isEqualTo(END); - } - - @Test - void programmaticConfigUsesExplicitNone() { - assertThat(REQUEST_PROPERTIES.truncate()).isEqualTo(NONE); - } - - @Test - void sendsConfiguredTruncation() { - NvidiaRerankingProvider provider = new NvidiaRerankingProvider(END_MODEL_CONFIG); + void sendsTruncationConfiguredInYaml() { + var configuredModel = rerankingProvidersConfig.providers().get("nvidia").models().getFirst(); + var localModel = + new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl( + configuredModel.name(), + configuredModel.apiModelSupport(), + configuredModel.isDefault(), + NVIDIA_URL, + configuredModel.properties()); + NvidiaRerankingProvider provider = new NvidiaRerankingProvider(localModel); provider .rerank(1, "test query", List.of("test passage"), RERANKING_CREDENTIALS) @@ -168,7 +151,7 @@ void sendsConfiguredTruncation() { .awaitItem() .getItem(); - verify( + wireMockServer.verify( postRequestedFor(urlEqualTo(NVIDIA_PATH)) .withRequestBody(matchingJsonPath("$.truncate", equalTo("END")))); }