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) {