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/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..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
@@ -35,7 +35,8 @@ public record RequestPropertiesImpl(
int readTimeoutMillis,
int maxBackOffMillis,
double jitter,
- int maxBatchSize)
+ int maxBatchSize,
+ TruncateOption truncate)
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 e2a9ff2979..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;
@@ -60,17 +61,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);
@@ -102,7 +92,7 @@ public Uni rerank(
modelName(),
new NvidiaRerankingRequest.TextWrapper(query),
passages.stream().map(NvidiaRerankingRequest.TextWrapper::new).toList(),
- TRUNCATE_PASSAGE);
+ modelConfig.properties().truncate());
final long callStartNano = System.nanoTime();
return retryHTTPCall(
@@ -156,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/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/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 8abeec3a4d..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
@@ -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 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;
+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;
@@ -13,8 +21,13 @@
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 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} */
@@ -23,12 +36,15 @@
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
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(
@@ -39,6 +55,34 @@ 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");
+
+ @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);
@@ -87,4 +131,28 @@ void testTenantIdIsExtractedFromCredentials() {
.isNotNull()
.isEqualTo(expectedTenantId);
}
+
+ @Test
+ 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)
+ .subscribe()
+ .withSubscriber(UniAssertSubscriber.create())
+ .awaitItem()
+ .getItem();
+
+ wireMockServer.verify(
+ postRequestedFor(urlEqualTo(NVIDIA_PATH))
+ .withRequestBody(matchingJsonPath("$.truncate", equalTo("END"))));
+ }
}
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) {