Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -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;
Expand Down Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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).
Expand Down Expand Up @@ -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();
}
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,8 @@ public record RequestPropertiesImpl(
int readTimeoutMillis,
int maxBackOffMillis,
double jitter,
int maxBatchSize)
int maxBatchSize,
TruncateOption truncate)
implements RerankingProviderConfig.ModelConfig.RequestProperties {}
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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.
*
* <p>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.
*
* <p>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);
Expand Down Expand Up @@ -102,7 +92,7 @@ public Uni<BatchedRerankingResponse> 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(
Expand Down Expand Up @@ -156,7 +146,7 @@ Uni<Response> rerank(
* <p>..
*/
public record NvidiaRerankingRequest(
String model, TextWrapper query, List<TextWrapper> passages, String truncate) {
String model, TextWrapper query, List<TextWrapper> passages, TruncateOption truncate) {

/**
* query and passage string needs to be wrapped in with text key for request to the Nvidia
Expand Down
3 changes: 2 additions & 1 deletion src/main/resources/reranking-providers-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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
max-batch-size: 10
truncate: END
3 changes: 2 additions & 1 deletion src/main/resources/test-reranking-providers-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
max-batch-size: 10
Original file line number Diff line number Diff line change
@@ -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;
Expand Down Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
@@ -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;
Expand All @@ -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} */
Expand All @@ -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(
Expand All @@ -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);
Expand Down Expand Up @@ -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"))));
}
}
Original file line number Diff line number Diff line change
@@ -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;
Expand All @@ -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(
Expand All @@ -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
Expand Down
Original file line number Diff line number Diff line change
@@ -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;
Expand Down Expand Up @@ -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) {
Expand Down