Skip to content
Draft
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
Expand Up @@ -4,6 +4,7 @@
import io.smallrye.config.WithDefault;
import io.stargate.sgv2.jsonapi.service.provider.ApiModelSupport;
import io.stargate.sgv2.jsonapi.service.schema.collections.CollectionRerankDef;
import jakarta.validation.constraints.Positive;
import java.util.List;
import java.util.Map;
import java.util.Objects;
Expand Down Expand Up @@ -73,6 +74,8 @@ interface ModelConfig {
RequestProperties properties();

interface RequestProperties {
int DEFAULT_CONNECTION_POOL_SIZE = 50;

/**
* Specifies the maximum number of attempts before failing. Default is 3 (1 request + 2
* retries).
Expand Down Expand Up @@ -100,6 +103,11 @@ interface RequestProperties {
@WithDefault("5000")
int readTimeoutMillis();

/** Maximum number of HTTP/1.x connections in this model's shared REST client pool. */
@Positive
@WithDefault("50")
int connectionPoolSize();

/**
* The maximum delay between retries in milliseconds.
*
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -35,8 +35,26 @@ public record RequestPropertiesImpl(
int readTimeoutMillis,
int maxBackOffMillis,
double jitter,
int connectionPoolSize,
int maxBatchSize)
implements RerankingProviderConfig.ModelConfig.RequestProperties {}
implements RerankingProviderConfig.ModelConfig.RequestProperties {
public RequestPropertiesImpl(
int atMostRetries,
int initialBackOffMillis,
int readTimeoutMillis,
int maxBackOffMillis,
double jitter,
int maxBatchSize) {
this(
atMostRetries,
initialBackOffMillis,
readTimeoutMillis,
maxBackOffMillis,
jitter,
DEFAULT_CONNECTION_POOL_SIZE,
maxBatchSize);
}
}
}
}
}
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.vertx.core.http.HttpClientOptions;
import jakarta.ws.rs.HeaderParam;
import jakarta.ws.rs.POST;
import jakarta.ws.rs.core.HttpHeaders;
Expand All @@ -18,6 +19,7 @@
import java.net.URI;
import java.util.*;
import java.util.concurrent.TimeUnit;
import java.util.function.Consumer;
import org.eclipse.microprofile.rest.client.annotation.ClientHeaderParam;
import org.eclipse.microprofile.rest.client.annotation.RegisterProvider;
import org.eclipse.microprofile.rest.client.inject.RegisterRestClient;
Expand Down Expand Up @@ -56,7 +58,7 @@
* }
* }</pre>
*/
public class NvidiaRerankingProvider extends RerankingProvider {
public class NvidiaRerankingProvider extends RerankingProvider implements AutoCloseable {

private final NvidiaRerankingClient nvidiaClient;

Expand All @@ -73,13 +75,28 @@ public class NvidiaRerankingProvider extends RerankingProvider {

public NvidiaRerankingProvider(
RerankingProvidersConfig.RerankingProviderConfig.ModelConfig modelConfig) {
this(modelConfig, createClient(modelConfig));
}

NvidiaRerankingProvider(
RerankingProvidersConfig.RerankingProviderConfig.ModelConfig modelConfig,
NvidiaRerankingClient nvidiaClient) {
super(ModelProvider.NVIDIA, modelConfig);
this.nvidiaClient = nvidiaClient;
}

private static NvidiaRerankingClient createClient(
RerankingProvidersConfig.RerankingProviderConfig.ModelConfig modelConfig) {
return QuarkusRestClientBuilder.newBuilder()
.baseUri(URI.create(modelConfig.url()))
.readTimeout(modelConfig.properties().readTimeoutMillis(), TimeUnit.MILLISECONDS)
.httpClientOptionsCustomizer(clientOptionsCustomizer(modelConfig))
.build(NvidiaRerankingClient.class);
}

nvidiaClient =
QuarkusRestClientBuilder.newBuilder()
.baseUri(URI.create(modelConfig.url()))
.readTimeout(modelConfig.properties().readTimeoutMillis(), TimeUnit.MILLISECONDS)
.build(NvidiaRerankingClient.class);
static Consumer<HttpClientOptions> clientOptionsCustomizer(
RerankingProvidersConfig.RerankingProviderConfig.ModelConfig modelConfig) {
return options -> options.setMaxPoolSize(modelConfig.properties().connectionPoolSize());
}

@Override
Expand Down Expand Up @@ -132,6 +149,11 @@ public Uni<BatchedRerankingResponse> rerank(
});
}

@Override
public void close() throws Exception {
nvidiaClient.close();
}

/**
* REST client interface for the Nvidia Reranking Service.
*
Expand All @@ -140,7 +162,7 @@ public Uni<BatchedRerankingResponse> rerank(
@RegisterRestClient
@RegisterProvider(RerankingProviderContentTypeFilter.class)
@RegisterProvider(ProviderBillingFilter.class)
public interface NvidiaRerankingClient {
public interface NvidiaRerankingClient extends AutoCloseable {

@POST
@ClientHeaderParam(name = HttpHeaders.CONTENT_TYPE, value = MediaType.APPLICATION_JSON)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,9 +8,12 @@
import io.stargate.sgv2.jsonapi.service.provider.ModelProvider;
import io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig;
import io.stargate.sgv2.jsonapi.service.reranking.gateway.RerankingEGWClient;
import jakarta.annotation.PreDestroy;
import jakarta.enterprise.context.ApplicationScoped;
import jakarta.inject.Inject;
import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.ConcurrentMap;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;

Expand All @@ -33,6 +36,18 @@ RerankingProvider create(
private static final Map<ModelProvider, ProviderConstructor> RERANKING_PROVIDER_CTORS =
Map.ofEntries(Map.entry(ModelProvider.NVIDIA, NvidiaRerankingProvider::new));

private final Map<ModelProvider, ProviderConstructor> providerConstructors;
private final ConcurrentMap<ProviderKey, RerankingProvider> directProviders =
new ConcurrentHashMap<>();

public RerankingProviderFactory() {
this(RERANKING_PROVIDER_CTORS);
}

RerankingProviderFactory(Map<ModelProvider, ProviderConstructor> providerConstructors) {
this.providerConstructors = Map.copyOf(providerConstructors);
}

public RerankingProvider create(
Tenant tenant,
String authToken,
Expand All @@ -59,7 +74,7 @@ public RerankingProvider create(
return create(tenant, authToken, modelProvider, modelName, authentication, commandName);
}

private synchronized RerankingProvider create(
private RerankingProvider create(
Tenant tenant,
String authToken,
ModelProvider modelProvider,
Expand Down Expand Up @@ -105,16 +120,35 @@ private synchronized RerankingProvider create(
commandName);
}

RerankingProviderFactory.ProviderConstructor ctor = RERANKING_PROVIDER_CTORS.get(modelProvider);
RerankingProviderFactory.ProviderConstructor ctor = providerConstructors.get(modelProvider);
if (ctor == null) {
throw SchemaException.Code.RERANKING_SERVICE_TYPE_UNAVAILABLE.get(
Map.of(
"errorMessage", "unknown service provider '%s'".formatted(modelProvider.apiName())));
}
return ctor.create(modelConfig);
return directProviders.computeIfAbsent(
new ProviderKey(modelProvider, modelConfig.name()), ignored -> ctor.create(modelConfig));
}

@PreDestroy
void close() {
directProviders.values().stream()
.filter(AutoCloseable.class::isInstance)
.map(AutoCloseable.class::cast)
.forEach(
provider -> {
try {
provider.close();
} catch (Exception exception) {
LOGGER.warn("Failed to close a cached reranking provider", exception);
}
});
directProviders.clear();
}

public RerankingProvidersConfig getRerankingConfig() {
return rerankingConfig;
}

private record ProviderKey(ModelProvider modelProvider, String modelName) {}
}
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
connection-pool-size: 50
max-batch-size: 10
5 changes: 4 additions & 1 deletion src/main/resources/test-reranking-providers-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -14,18 +14,21 @@ stargate:
is-default: true
url: https://us-west-2.api-dev.ai.datastax.com/nvidia/v1/ranking
properties:
connection-pool-size: 50
max-batch-size: 10
- name: nvidia/a-random-deprecated-model
api-model-support:
status: DEPRECATED
message: This model has been deprecated, it will be removed in a future release. It is not supported for new Collections or Tables.
url: https://us-west-2.api-dev.ai.datastax.com/nvidia/v1/ranking
properties:
connection-pool-size: 50
max-batch-size: 10
- name: nvidia/a-random-EOL-model
api-model-support:
status: END_OF_LIFE
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
connection-pool-size: 50
max-batch-size: 10
Original file line number Diff line number Diff line change
Expand Up @@ -2,19 +2,32 @@

import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.assertThatThrownBy;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.eq;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;

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.api.request.tenant.Tenant;
import io.stargate.sgv2.jsonapi.config.DatabaseType;
import io.stargate.sgv2.jsonapi.config.constants.HttpConstants;
import io.stargate.sgv2.jsonapi.exception.SchemaException;
import io.stargate.sgv2.jsonapi.service.provider.ApiModelSupport;
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 io.vertx.core.http.HttpClientOptions;
import jakarta.inject.Inject;
import java.util.List;
import java.util.Optional;
import java.util.concurrent.CyclicBarrier;
import java.util.concurrent.Executors;
import org.junit.jupiter.api.Test;

/** Tests for {@link NvidiaRerankingProvider} */
Expand All @@ -28,7 +41,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, 37, 10);

private static final RerankingProvidersConfig.RerankingProviderConfig.ModelConfig MODEL_CONFIG =
new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl(
Expand All @@ -39,6 +52,79 @@ public class NvidiaRerankingProviderTest {
"https://us-west-2.api-dev.ai.datastax.com/nvidia/v1/ranking",
REQUEST_PROPERTIES);

@Inject RerankingProvidersConfig rerankingProvidersConfig;

@Test
void configuresConnectionPoolSizeFromModel() {
var options = new HttpClientOptions();

NvidiaRerankingProvider.clientOptionsCustomizer(MODEL_CONFIG).accept(options);

assertThat(options.getMaxPoolSize()).isEqualTo(37);
}

@Test
void loadsExplicitConnectionPoolSize() {
var configuredModel =
rerankingProvidersConfig.providers().values().stream()
.flatMap(provider -> provider.models().stream())
.filter(RerankingProvidersConfig.RerankingProviderConfig.ModelConfig::isDefault)
.findFirst()
.orElseThrow();

assertThat(configuredModel.properties().connectionPoolSize()).isEqualTo(50);
}

@Test
void sharedProviderUsesCredentialsFromConcurrentCalls() throws Exception {
var nvidiaClient = mock(NvidiaRerankingProvider.NvidiaRerankingClient.class);
when(nvidiaClient.rerank(any(), any(), any())).thenReturn(Uni.createFrom().nothing());
var provider = new NvidiaRerankingProvider(MODEL_CONFIG, nvidiaClient);
var firstTenant = Tenant.create(DatabaseType.ASTRA, "first-tenant");
var secondTenant = Tenant.create(DatabaseType.ASTRA, "second-tenant");
var barrier = new CyclicBarrier(2);

try (var executor = Executors.newFixedThreadPool(2)) {
var calls =
List.of(
executor.submit(
() -> {
barrier.await();
provider.rerank(
0,
"query",
List.of("first"),
new RerankingCredentials(firstTenant, "first-key"));
return null;
}),
executor.submit(
() -> {
barrier.await();
provider.rerank(
1,
"query",
List.of("second"),
new RerankingCredentials(secondTenant, "second-key"));
return null;
}));

for (var call : calls) {
call.get();
}
}

verify(nvidiaClient)
.rerank(
eq(HttpConstants.BEARER_PREFIX_FOR_API_KEY + "first-key"),
eq(firstTenant.toString()),
any(NvidiaRerankingProvider.NvidiaRerankingRequest.class));
verify(nvidiaClient)
.rerank(
eq(HttpConstants.BEARER_PREFIX_FOR_API_KEY + "second-key"),
eq(secondTenant.toString()),
any(NvidiaRerankingProvider.NvidiaRerankingRequest.class));
}

@Test
void testEmptyApiKeyThrowsException() {
NvidiaRerankingProvider provider = new NvidiaRerankingProvider(MODEL_CONFIG);
Expand Down
Loading