From 6c4898ff158042e3f2de908e15fafd8368dc637b Mon Sep 17 00:00:00 2001 From: Eric Hare Date: Wed, 5 Aug 2026 10:45:20 -0700 Subject: [PATCH 1/4] fix: handle transient reranking provider failures --- .../exception/RerankingProviderException.java | 18 ++ .../RerankingProviderResponseValidation.java | 6 + .../RerankingProvidersConfig.java | 6 +- .../reranking/gateway/RerankingEGWClient.java | 74 +++++--- .../operation/RerankingProvider.java | 6 + .../resources/reranking-providers-config.yaml | 3 +- ...rankingProviderResponseValidationTest.java | 68 +++++++ .../gateway/RerankingGatewayClientTest.java | 172 +++++++++++++++++- .../operation/RerankingProviderRetryTest.java | 151 +++++++++++++++ 9 files changed, 476 insertions(+), 28 deletions(-) create mode 100644 src/test/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidationTest.java create mode 100644 src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderRetryTest.java diff --git a/src/main/java/io/stargate/sgv2/jsonapi/exception/RerankingProviderException.java b/src/main/java/io/stargate/sgv2/jsonapi/exception/RerankingProviderException.java index 7c9ee3da7a..2495bcfe1e 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/exception/RerankingProviderException.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/exception/RerankingProviderException.java @@ -1,5 +1,9 @@ package io.stargate.sgv2.jsonapi.exception; +import java.util.EnumSet; +import java.util.Optional; +import java.util.UUID; + public class RerankingProviderException extends ServerException { public static final Scope SCOPE = Scope.RERANKING_PROVIDER; @@ -8,6 +12,20 @@ public RerankingProviderException(ErrorInstance errorInstance) { super(errorInstance); } + /** Constructs a reranking provider exception from an unrecognized EGW error response. */ + public RerankingProviderException(String code, String title, String body) { + this( + new ErrorInstance( + UUID.randomUUID(), + FAMILY, + SCOPE, + code, + title, + body, + Optional.empty(), + EnumSet.noneOf(ExceptionFlags.class))); + } + public enum Code implements ErrorCode { RERANKING_PROVIDER_TIMEOUT; diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidation.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidation.java index aefcd1618e..68896190c1 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidation.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidation.java @@ -42,6 +42,12 @@ public void filter(ClientRequestContext requestContext, ClientResponseContext re if (responseContext.getStatus() == 0) { return; } + + // HTTP errors may be empty or non-JSON; ProviderBase maps statuses 400 and above. + if (responseContext.getStatus() >= 400) { + return; + } + // Throw error if there is no response body if (!responseContext.hasEntity()) { throw SchemaException.Code.RERANKING_PROVIDER_UNEXPECTED_RESPONSE.get( 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..349a6c4102 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 @@ -74,10 +74,10 @@ interface ModelConfig { interface RequestProperties { /** - * Specifies the maximum number of attempts before failing. Default is 3 (1 request + 2 - * retries). + * Specifies the maximum number of retries after the initial request. Default is 3 (up to 4 + * total attempts). * - * @return The maximum number of attempts before failing. + * @return The maximum number of retries after the initial request. */ @WithDefault("3") int atMostRetries(); diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java index f3af9fca2f..ba395a032d 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java @@ -9,6 +9,7 @@ import io.stargate.sgv2.jsonapi.api.request.tenant.Tenant; import io.stargate.sgv2.jsonapi.exception.RerankingProviderException; import io.stargate.sgv2.jsonapi.exception.SchemaException; +import io.stargate.sgv2.jsonapi.exception.ServerException; import io.stargate.sgv2.jsonapi.service.provider.ModelProvider; import io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig; import io.stargate.sgv2.jsonapi.service.reranking.operation.RerankingProvider; @@ -79,33 +80,15 @@ public Uni rerank( .setProviderContext(contextBuilder.build()) .build(); - // TODO: Why is this error handling here not part of the uni pipeline? - Uni gatewayRerankingUni; - try { - gatewayRerankingUni = grpcGatewayService.rerank(gatewayRequest); - } catch (StatusRuntimeException e) { - if (e.getStatus().getCode().equals(Status.Code.DEADLINE_EXCEEDED)) { - throw RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.get( - Map.of( - "modelProvider", - modelProvider().apiName(), - "httpStatus", - String.valueOf(e.getStatus().getCode()), - "errorMessage", - e.getMessage())); - } - throw e; - } - - return gatewayRerankingUni + return Uni.createFrom() + .deferred(() -> grpcGatewayService.rerank(gatewayRequest)) + .onFailure(StatusRuntimeException.class) + .transform(this::mapStatusFailure) .onItem() .transform( gatewayResponse -> { if (gatewayResponse.hasError()) { - // 22-Jan-2026, tatu: This is ugly. But has to be done to work around fragility - // of exception mapping - throw SchemaException.Code.valueOf(gatewayResponse.getError().getErrorCode()) - .withPreformattedMessage(gatewayResponse.getError().getErrorBody()); + throw mapGatewayError(gatewayResponse.getError()); } return new BatchedRerankingResponse( @@ -116,4 +99,49 @@ public Uni rerank( createModelUsage(gatewayResponse.getModelUsage())); }); } + + private Throwable mapStatusFailure(Throwable failure) { + var statusException = (StatusRuntimeException) failure; + if (statusException.getStatus().getCode().equals(Status.Code.DEADLINE_EXCEEDED)) { + return RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.get( + Map.of( + "modelProvider", + modelProvider().apiName(), + "httpStatus", + String.valueOf(statusException.getStatus().getCode()), + "errorMessage", + statusException.getMessage())); + } + return failure; + } + + private RuntimeException mapGatewayError(EmbeddingGateway.RerankingResponse.ErrorResponse error) { + String errorCode = error.getErrorCode(); + + var schemaCode = + Arrays.stream(SchemaException.Code.values()) + .filter(code -> code.name().equals(errorCode)) + .findFirst(); + if (schemaCode.isPresent()) { + return schemaCode.get().withPreformattedMessage(error.getErrorBody()); + } + + var serverCode = + Arrays.stream(ServerException.Code.values()) + .filter(code -> code.name().equals(errorCode)) + .findFirst(); + if (serverCode.isPresent()) { + return serverCode.get().withPreformattedMessage(error.getErrorBody()); + } + + var rerankingProviderCode = + Arrays.stream(RerankingProviderException.Code.values()) + .filter(code -> code.name().equals(errorCode)) + .findFirst(); + if (rerankingProviderCode.isPresent()) { + return rerankingProviderCode.get().withPreformattedMessage(error.getErrorBody()); + } + + return new RerankingProviderException(errorCode, error.getErrorTitle(), error.getErrorBody()); + } } diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java index 9a9f22ffff..a4450d3d11 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java @@ -125,6 +125,12 @@ protected boolean decideRetry(Throwable throwable) { boolean retry = throwable instanceof RerankingProviderException rpe && RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.name().equals(rpe.code); + retry = + retry + || throwable instanceof SchemaException schemaException + && SchemaException.Code.RERANKING_PROVIDER_SERVER_ERROR + .name() + .equals(schemaException.code); return retry || super.decideRetry(throwable); } diff --git a/src/main/resources/reranking-providers-config.yaml b/src/main/resources/reranking-providers-config.yaml index 67739fe0a2..6a17984287 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 + at-most-retries: 1 + max-batch-size: 10 diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidationTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidationTest.java new file mode 100644 index 0000000000..96e280d23e --- /dev/null +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidationTest.java @@ -0,0 +1,68 @@ +package io.stargate.sgv2.jsonapi.service.reranking.configuration; + +import static org.assertj.core.api.Assertions.assertThatCode; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import io.stargate.sgv2.jsonapi.exception.SchemaException; +import jakarta.ws.rs.client.ClientResponseContext; +import jakarta.ws.rs.core.MediaType; +import jakarta.ws.rs.core.Response; +import org.junit.jupiter.api.Test; + +class RerankingProviderResponseValidationTest { + + private final RerankingProviderResponseValidation validation = + new RerankingProviderResponseValidation(); + + @Test + void ignoresEmptyServerErrorResponse() { + ClientResponseContext responseContext = responseContext(Response.Status.INTERNAL_SERVER_ERROR); + when(responseContext.hasEntity()).thenReturn(false); + + assertThatCode(() -> validation.filter(null, responseContext)).doesNotThrowAnyException(); + + verify(responseContext, never()).hasEntity(); + } + + @Test + void ignoresNonJsonServerErrorResponse() { + ClientResponseContext responseContext = responseContext(Response.Status.INTERNAL_SERVER_ERROR); + when(responseContext.hasEntity()).thenReturn(true); + when(responseContext.getMediaType()).thenReturn(MediaType.TEXT_PLAIN_TYPE); + + assertThatCode(() -> validation.filter(null, responseContext)).doesNotThrowAnyException(); + + verify(responseContext, never()).getMediaType(); + } + + @Test + void rejectsEmptySuccessfulResponse() { + ClientResponseContext responseContext = responseContext(Response.Status.OK); + when(responseContext.hasEntity()).thenReturn(false); + + assertThatThrownBy(() -> validation.filter(null, responseContext)) + .isInstanceOf(SchemaException.class) + .hasMessageContaining("No response body from the reranking provider"); + } + + @Test + void rejectsEmptyRedirectResponse() { + ClientResponseContext responseContext = responseContext(Response.Status.TEMPORARY_REDIRECT); + when(responseContext.hasEntity()).thenReturn(false); + + assertThatThrownBy(() -> validation.filter(null, responseContext)) + .isInstanceOf(SchemaException.class) + .hasMessageContaining("No response body from the reranking provider"); + } + + private static ClientResponseContext responseContext(Response.Status status) { + ClientResponseContext responseContext = mock(ClientResponseContext.class); + when(responseContext.getStatus()).thenReturn(status.getStatusCode()); + when(responseContext.getStatusInfo()).thenReturn(status); + return responseContext; + } +} 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 7f88c80ef6..67dc4c8a0a 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 @@ -5,6 +5,8 @@ import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; +import io.grpc.Status; +import io.grpc.StatusRuntimeException; import io.quarkus.test.junit.QuarkusTest; import io.quarkus.test.junit.TestProfile; import io.smallrye.mutiny.Uni; @@ -13,7 +15,9 @@ import io.stargate.embedding.gateway.RerankingService; import io.stargate.sgv2.jsonapi.TestConstants; import io.stargate.sgv2.jsonapi.api.request.RerankingCredentials; +import io.stargate.sgv2.jsonapi.exception.RerankingProviderException; import io.stargate.sgv2.jsonapi.exception.SchemaException; +import io.stargate.sgv2.jsonapi.exception.ServerException; import io.stargate.sgv2.jsonapi.service.provider.ApiModelSupport; import io.stargate.sgv2.jsonapi.service.provider.ModelProvider; import io.stargate.sgv2.jsonapi.service.provider.ModelType; @@ -21,6 +25,7 @@ import io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfigImpl; import io.stargate.sgv2.jsonapi.service.reranking.operation.RerankingProvider; import io.stargate.sgv2.jsonapi.testresource.NoGlobalResourcesTestProfile; +import jakarta.inject.Inject; import java.util.List; import java.util.Map; import java.util.Optional; @@ -61,6 +66,16 @@ public class RerankingGatewayClientTest { new RerankingProvidersConfigImpl.RerankingProviderConfigImpl( false, "test", true, Map.of(), List.of()); + @Inject RerankingProvidersConfig rerankingProvidersConfig; + + @Test + void productionNvidiaConfigUsesOneRetry() { + assertThat(rerankingProvidersConfig.providers()).containsKey("nvidia"); + assertThat(rerankingProvidersConfig.providers().get("nvidia").models()) + .singleElement() + .satisfies(model -> assertThat(model.properties().atMostRetries()).isEqualTo(1)); + } + @Test void handleValidResponse() { RerankingService rerankService = mock(RerankingService.class); @@ -133,7 +148,7 @@ void handleValidResponse() { } @Test - void handleError() { + void mapsSchemaErrorFromGateway() { RerankingService rerankService = mock(RerankingService.class); final EmbeddingGateway.RerankingResponse.Builder builder = @@ -175,4 +190,159 @@ void handleError() { assertThat(exception.code).isEqualTo(apiException.code); }); } + + @Test + void mapsServerErrorFromGateway() { + RerankingService rerankService = mock(RerankingService.class); + when(rerankService.rerank(any())) + .thenReturn( + Uni.createFrom() + .item( + gatewayErrorResponse( + ServerException.Code.UNEXPECTED_SERVER_ERROR.name(), + "Gateway server error", + "Gateway server error body"))); + + Throwable result = rerankAndAwaitFailure(rerankService); + + assertThat(result) + .isInstanceOf(ServerException.class) + .satisfies( + failure -> { + ServerException exception = (ServerException) failure; + assertThat(exception.code) + .isEqualTo(ServerException.Code.UNEXPECTED_SERVER_ERROR.name()); + assertThat(exception.body).isEqualTo("Gateway server error body"); + }); + } + + @Test + void mapsRerankingProviderErrorFromGateway() { + RerankingService rerankService = mock(RerankingService.class); + when(rerankService.rerank(any())) + .thenReturn( + Uni.createFrom() + .item( + gatewayErrorResponse( + RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.name(), + "Gateway timeout", + "Gateway timeout body"))); + + Throwable result = rerankAndAwaitFailure(rerankService); + + assertThat(result) + .isInstanceOf(RerankingProviderException.class) + .satisfies( + failure -> { + RerankingProviderException exception = (RerankingProviderException) failure; + assertThat(exception.code) + .isEqualTo(RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.name()); + assertThat(exception.body).isEqualTo("Gateway timeout body"); + }); + } + + @Test + void preservesUnknownGatewayError() { + RerankingService rerankService = mock(RerankingService.class); + when(rerankService.rerank(any())) + .thenReturn( + Uni.createFrom() + .item( + gatewayErrorResponse( + "FUTURE_GATEWAY_ERROR", + "Future gateway error", + "Future gateway error body"))); + + Throwable result = rerankAndAwaitFailure(rerankService); + + assertThat(result) + .isInstanceOf(RerankingProviderException.class) + .satisfies( + failure -> { + RerankingProviderException exception = (RerankingProviderException) failure; + assertThat(exception.code).isEqualTo("FUTURE_GATEWAY_ERROR"); + assertThat(exception.title).isEqualTo("Future gateway error"); + assertThat(exception.body).isEqualTo("Future gateway error body"); + }); + } + + @Test + void mapsAsyncDeadlineExceeded() { + RerankingService rerankService = mock(RerankingService.class); + when(rerankService.rerank(any())) + .thenReturn(Uni.createFrom().failure(Status.DEADLINE_EXCEEDED.asRuntimeException())); + + Throwable result = rerankAndAwaitFailure(rerankService); + + assertThat(result) + .isInstanceOf(RerankingProviderException.class) + .satisfies( + failure -> { + RerankingProviderException exception = (RerankingProviderException) failure; + assertThat(exception.code) + .isEqualTo(RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.name()); + assertThat(exception.body).contains(ModelProvider.NVIDIA.apiName()); + assertThat(exception.body).contains(Status.Code.DEADLINE_EXCEEDED.name()); + }); + } + + @Test + void mapsSynchronousDeadlineExceeded() { + RerankingService rerankService = mock(RerankingService.class); + when(rerankService.rerank(any())).thenThrow(Status.DEADLINE_EXCEEDED.asRuntimeException()); + + Throwable result = rerankAndAwaitFailure(rerankService); + + assertThat(result) + .isInstanceOf(RerankingProviderException.class) + .satisfies( + failure -> { + RerankingProviderException exception = (RerankingProviderException) failure; + assertThat(exception.code) + .isEqualTo(RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.name()); + assertThat(exception.body).contains(Status.Code.DEADLINE_EXCEEDED.name()); + }); + } + + @Test + void preservesAsyncNonDeadlineStatusFailure() { + RerankingService rerankService = mock(RerankingService.class); + StatusRuntimeException unavailable = Status.UNAVAILABLE.asRuntimeException(); + when(rerankService.rerank(any())).thenReturn(Uni.createFrom().failure(unavailable)); + + Throwable result = rerankAndAwaitFailure(rerankService); + + assertThat(result).isSameAs(unavailable); + } + + private static Throwable rerankAndAwaitFailure(RerankingService rerankService) { + return createClient(rerankService) + .rerank(1, "apple", List.of("orange", "apple"), RERANK_CREDENTIALS) + .subscribe() + .withSubscriber(UniAssertSubscriber.create()) + .awaitFailure() + .getFailure(); + } + + private static RerankingEGWClient createClient(RerankingService rerankService) { + return new RerankingEGWClient( + ModelProvider.NVIDIA, + MODEL_CONFIG, + testConstants.TENANT, + "default", + rerankService, + Map.of(), + TESTING_COMMAND_NAME); + } + + private static EmbeddingGateway.RerankingResponse gatewayErrorResponse( + String code, String title, String body) { + return EmbeddingGateway.RerankingResponse.newBuilder() + .setError( + EmbeddingGateway.RerankingResponse.ErrorResponse.newBuilder() + .setErrorCode(code) + .setErrorTitle(title) + .setErrorBody(body)) + .build(); + } } diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderRetryTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderRetryTest.java new file mode 100644 index 0000000000..8a039d0fa2 --- /dev/null +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderRetryTest.java @@ -0,0 +1,151 @@ +package io.stargate.sgv2.jsonapi.service.reranking.operation; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +import io.smallrye.mutiny.Uni; +import io.smallrye.mutiny.helpers.test.UniAssertSubscriber; +import io.stargate.sgv2.jsonapi.api.request.RerankingCredentials; +import io.stargate.sgv2.jsonapi.exception.SchemaException; +import io.stargate.sgv2.jsonapi.service.provider.ApiModelSupport; +import io.stargate.sgv2.jsonapi.service.provider.ModelProvider; +import io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig; +import io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfigImpl; +import jakarta.ws.rs.core.MediaType; +import jakarta.ws.rs.core.Response; +import java.util.List; +import java.util.Optional; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.api.Test; + +class RerankingProviderRetryTest { + + @Test + void retriesServerErrorOnceThenReturnsSuccess() { + RetryingTestProvider provider = new RetryingTestProvider(); + AtomicInteger calls = new AtomicInteger(); + Response successfulResponse = response(Response.Status.OK); + + Response result = + provider + .execute( + Uni.createFrom() + .deferred( + () -> + Uni.createFrom() + .item( + calls.incrementAndGet() == 1 + ? response(Response.Status.INTERNAL_SERVER_ERROR) + : successfulResponse))) + .subscribe() + .withSubscriber(UniAssertSubscriber.create()) + .awaitItem() + .getItem(); + + assertThat(result).isSameAs(successfulResponse); + assertThat(calls).hasValue(2); + } + + @Test + void stopsAfterOneServerErrorRetry() { + RetryingTestProvider provider = new RetryingTestProvider(); + AtomicInteger calls = new AtomicInteger(); + + Throwable failure = + provider + .execute( + Uni.createFrom() + .deferred( + () -> { + calls.incrementAndGet(); + return Uni.createFrom() + .item(response(Response.Status.INTERNAL_SERVER_ERROR)); + })) + .subscribe() + .withSubscriber(UniAssertSubscriber.create()) + .awaitFailure() + .getFailure(); + + assertThat(calls).hasValue(2); + assertThat(failure) + .isInstanceOf(SchemaException.class) + .satisfies( + throwable -> + assertThat(((SchemaException) throwable).code) + .isEqualTo(SchemaException.Code.RERANKING_PROVIDER_SERVER_ERROR.name())); + } + + @Test + void doesNotRetryClientError() { + RetryingTestProvider provider = new RetryingTestProvider(); + AtomicInteger calls = new AtomicInteger(); + + Throwable failure = + provider + .execute( + Uni.createFrom() + .deferred( + () -> { + calls.incrementAndGet(); + return Uni.createFrom().item(response(Response.Status.BAD_REQUEST)); + })) + .subscribe() + .withSubscriber(UniAssertSubscriber.create()) + .awaitFailure() + .getFailure(); + + assertThat(calls).hasValue(1); + assertThat(failure) + .isInstanceOf(SchemaException.class) + .satisfies( + throwable -> + assertThat(((SchemaException) throwable).code) + .isEqualTo(SchemaException.Code.RERANKING_PROVIDER_CLIENT_ERROR.name())); + } + + private static Response response(Response.Status status) { + Response response = mock(Response.class); + when(response.getStatus()).thenReturn(status.getStatusCode()); + when(response.getStatusInfo()).thenReturn(status); + when(response.getMediaType()).thenReturn(MediaType.TEXT_PLAIN_TYPE); + when(response.readEntity(String.class)).thenReturn("provider response"); + return response; + } + + private static final class RetryingTestProvider extends RerankingProvider { + + private RetryingTestProvider() { + super(ModelProvider.NVIDIA, modelConfig()); + } + + private Uni execute(Uni request) { + return retryHTTPCall(request); + } + + @Override + protected String errorMessageJsonPtr() { + return "/message"; + } + + @Override + public Uni rerank( + int batchId, + String query, + List passages, + RerankingCredentials rerankingCredentials) { + throw new UnsupportedOperationException("Not used by retry tests"); + } + + private static RerankingProvidersConfig.RerankingProviderConfig.ModelConfig modelConfig() { + return new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl( + "test-model", + new ApiModelSupport.ApiModelSupportImpl( + ApiModelSupport.SupportStatus.SUPPORTED, Optional.empty()), + false, + "http://testing.com", + new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl + .RequestPropertiesImpl(1, 1, 100, 1, 0.0, 10)); + } + } +} From 4cc105c8ccc247b0d8b06129ae1bc58db8840169 Mon Sep 17 00:00:00 2001 From: Eric Hare Date: Wed, 5 Aug 2026 11:30:38 -0700 Subject: [PATCH 2/4] fix: refine provider failure recovery --- .../gateway/EmbeddingGatewayClient.java | 43 ++++---- .../AwsBedrockEmbeddingProvider.java | 2 +- .../operation/EmbeddingProvider.java | 2 +- .../EmbeddingProviderExceptionHandler.java | 2 +- .../service/provider/ProviderBase.java | 50 ++++++---- .../RerankingProviderExceptionHandler.java | 2 +- .../reranking/gateway/RerankingEGWClient.java | 30 +++--- .../operation/RerankingProvider.java | 35 ++++--- src/main/resources/errors.yaml | 4 +- .../resources/reranking-providers-config.yaml | 3 +- .../operation/EmbeddingGatewayClientTest.java | 79 +++++++++++++++ .../EmbeddingProviderErrorMessageTest.java | 2 +- ...rankingProviderResponseValidationTest.java | 15 +-- .../gateway/RerankingGatewayClientTest.java | 36 +++---- .../operation/RerankingProviderRetryTest.java | 98 ++++++++++++++++++- 15 files changed, 295 insertions(+), 108 deletions(-) diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/gateway/EmbeddingGatewayClient.java b/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/gateway/EmbeddingGatewayClient.java index 275c945723..5ab7a11759 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/gateway/EmbeddingGatewayClient.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/gateway/EmbeddingGatewayClient.java @@ -164,26 +164,12 @@ public Uni vectorize( .setProviderContext(contextBuilder.build()) .build(); - // aaron 17 June 2025 - unsure why this error handled was not in the uni pipeline below - // kept it as is when refactoring - Uni embeddingResponse; - try { - embeddingResponse = grpcGatewayClient.embed(gatewayRequest); - } catch (StatusRuntimeException e) { - if (e.getStatus().getCode().equals(Status.Code.DEADLINE_EXCEEDED)) { - throw EmbeddingProviderException.Code.EMBEDDING_PROVIDER_TIMEOUT.get( - Map.of( - "modelProvider", - modelProvider().apiName(), - "httpStatus", - String.valueOf(e.getStatus().getCode()), - "errorMessage", - e.getMessage())); - } - throw e; - } - - return embeddingResponse + // Defer the gRPC invocation so synchronous stub failures and asynchronous Uni failures use the + // same status-mapping path. + return Uni.createFrom() + .deferred(() -> grpcGatewayClient.embed(gatewayRequest)) + .onFailure(StatusRuntimeException.class) + .transform(this::mapStatusFailure) .onItem() .transform( gatewayResponse -> { @@ -212,6 +198,23 @@ public Uni vectorize( }); } + private Throwable mapStatusFailure(Throwable failure) { + var statusException = (StatusRuntimeException) failure; + if (statusException.getStatus().getCode().equals(Status.Code.DEADLINE_EXCEEDED)) { + return EmbeddingProviderException.Code.EMBEDDING_PROVIDER_TIMEOUT.get( + Map.of( + "modelProvider", + modelProvider().apiName(), + "providerStatus", + String.valueOf(statusException.getStatus().getCode()), + "errorMessage", + statusException.getMessage())); + } + + // Preserve non-deadline gRPC failures so downstream handling retains their status and cause. + return failure; + } + /** Return MAX_VALUE because the batching is done inside EGW */ @Override public int maxBatchSize() { diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/AwsBedrockEmbeddingProvider.java b/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/AwsBedrockEmbeddingProvider.java index 28c1b6cc54..84965d3c37 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/AwsBedrockEmbeddingProvider.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/AwsBedrockEmbeddingProvider.java @@ -197,7 +197,7 @@ private Throwable mapBedrockException(BedrockRuntimeException bedrockException) Map.of( "modelProvider", modelProvider().apiName(), - "httpStatus", + "providerStatus", String.valueOf(bedrockException.statusCode()), "errorMessage", bedrockException.getMessage())); diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProvider.java b/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProvider.java index 0f05ce34f1..9308650a0e 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProvider.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProvider.java @@ -223,7 +223,7 @@ protected RuntimeException mapHTTPError(Response jakartaResponse, String errorMe Map.of( "modelProvider", modelProvider().apiName(), - "httpStatus", + "providerStatus", String.valueOf(jakartaResponse.getStatus()), "errorMessage", errorMessage)); diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/provider/EmbeddingProviderExceptionHandler.java b/src/main/java/io/stargate/sgv2/jsonapi/service/provider/EmbeddingProviderExceptionHandler.java index 03efc7f439..fbf8c60164 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/provider/EmbeddingProviderExceptionHandler.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/provider/EmbeddingProviderExceptionHandler.java @@ -17,7 +17,7 @@ public RuntimeException handle(TimeoutException exception) { return EmbeddingProviderException.Code.EMBEDDING_PROVIDER_TIMEOUT.get( Map.of( "modelProvider", modelProvider.apiName(), - "httpStatus", "", + "providerStatus", "", "errorMessage", "")); } diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/provider/ProviderBase.java b/src/main/java/io/stargate/sgv2/jsonapi/service/provider/ProviderBase.java index 662df4d74d..44b2792c4c 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/provider/ProviderBase.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/provider/ProviderBase.java @@ -99,24 +99,38 @@ public ModelProvider modelProvider() { */ protected Uni retryHTTPCall(Uni uni) { - return uni - // Catch *any* web exception from jakarta rest client - .onFailure(WebApplicationException.class) - // and recover with the jakarta response, so we can translate to API exception - .recoverWithItem(ex -> ex.getResponse()) - .onItem() - // handle the response, throws if there is an error - .transform(this::handleHTTPResponse) - // decide if we want to retry - .onFailure(this::decideRetry) - .retry() - .withBackOff(initialBackOffDuration(), maxBackOffDuration()) - .withJitter(jitter()) - .atMost(atMostRetries()) - // after all retry logic, if we have an error we will need to handle it into an API - // exception - .onFailure() - .transform(exceptionHandler::maybeHandle); + return Uni.createFrom() + .deferred( + () -> + uni + // Catch *any* web exception from jakarta rest client + .onFailure(WebApplicationException.class) + // and recover with the jakarta response, so we can translate to API exception + .recoverWithItem(ex -> ex.getResponse()) + .onItem() + // handle the response, throws if there is an error + .transform(this::handleHTTPResponse) + // decide if we want to retry + .onFailure(newRetryPredicate()) + .retry() + .withBackOff(initialBackOffDuration(), maxBackOffDuration()) + .withJitter(jitter()) + .atMost(atMostRetries()) + // after all retry logic, if we have an error we will need to handle it into an + // API exception + .onFailure() + .transform(exceptionHandler::maybeHandle)); + } + + /** + * Creates the retry predicate for one subscription. + * + *

Subclasses may override this method when retry decisions need per-request state. The + * predicate must not retain state on the provider instance because providers are shared across + * requests. + */ + protected Predicate newRetryPredicate() { + return this::decideRetry; } /** diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/provider/RerankingProviderExceptionHandler.java b/src/main/java/io/stargate/sgv2/jsonapi/service/provider/RerankingProviderExceptionHandler.java index 2e4e5dc105..50d2dd4a85 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/provider/RerankingProviderExceptionHandler.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/provider/RerankingProviderExceptionHandler.java @@ -16,7 +16,7 @@ public RuntimeException handle(TimeoutException exception) { return RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.get( Map.of( "modelProvider", modelProvider.apiName(), - "httpStatus", "", + "providerStatus", "", "errorMessage", "")); } } diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java index ba395a032d..1e5d182f95 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java @@ -7,6 +7,8 @@ import io.stargate.embedding.gateway.RerankingService; import io.stargate.sgv2.jsonapi.api.request.RerankingCredentials; import io.stargate.sgv2.jsonapi.api.request.tenant.Tenant; +import io.stargate.sgv2.jsonapi.exception.APIException; +import io.stargate.sgv2.jsonapi.exception.ErrorCode; import io.stargate.sgv2.jsonapi.exception.RerankingProviderException; import io.stargate.sgv2.jsonapi.exception.SchemaException; import io.stargate.sgv2.jsonapi.exception.ServerException; @@ -107,41 +109,43 @@ private Throwable mapStatusFailure(Throwable failure) { Map.of( "modelProvider", modelProvider().apiName(), - "httpStatus", + "providerStatus", String.valueOf(statusException.getStatus().getCode()), "errorMessage", statusException.getMessage())); } + + // Only DEADLINE_EXCEEDED has a defined Data API mapping. Preserve other gRPC statuses so + // upstream handlers retain the original status instead of misclassifying it as a timeout. return failure; } private RuntimeException mapGatewayError(EmbeddingGateway.RerankingResponse.ErrorResponse error) { String errorCode = error.getErrorCode(); - var schemaCode = - Arrays.stream(SchemaException.Code.values()) - .filter(code -> code.name().equals(errorCode)) - .findFirst(); + // Preserve API compatibility for known gateway codes. This precedence is intentional: + // Schema (REQUEST/SCHEMA), then unscoped Server, then RerankingProvider. Unknown gateway + // codes pass through with their gateway-supplied title and body rather than failing lookup. + var schemaCode = findErrorCode(errorCode, SchemaException.Code.values()); if (schemaCode.isPresent()) { return schemaCode.get().withPreformattedMessage(error.getErrorBody()); } - var serverCode = - Arrays.stream(ServerException.Code.values()) - .filter(code -> code.name().equals(errorCode)) - .findFirst(); + var serverCode = findErrorCode(errorCode, ServerException.Code.values()); if (serverCode.isPresent()) { return serverCode.get().withPreformattedMessage(error.getErrorBody()); } - var rerankingProviderCode = - Arrays.stream(RerankingProviderException.Code.values()) - .filter(code -> code.name().equals(errorCode)) - .findFirst(); + var rerankingProviderCode = findErrorCode(errorCode, RerankingProviderException.Code.values()); if (rerankingProviderCode.isPresent()) { return rerankingProviderCode.get().withPreformattedMessage(error.getErrorBody()); } return new RerankingProviderException(errorCode, error.getErrorTitle(), error.getErrorBody()); } + + private static & ErrorCode> + Optional findErrorCode(String errorCode, C[] codes) { + return Arrays.stream(codes).filter(code -> code.name().equals(errorCode)).findFirst(); + } } diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java index a4450d3d11..9b35f19782 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java @@ -14,6 +14,8 @@ import java.util.Comparator; import java.util.List; import java.util.Map; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.function.Predicate; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @@ -122,16 +124,27 @@ protected int atMostRetries() { @Override protected boolean decideRetry(Throwable throwable) { - boolean retry = - throwable instanceof RerankingProviderException rpe - && RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.name().equals(rpe.code); - retry = - retry - || throwable instanceof SchemaException schemaException - && SchemaException.Code.RERANKING_PROVIDER_SERVER_ERROR - .name() - .equals(schemaException.code); - return retry || super.decideRetry(throwable); + if (throwable instanceof RerankingProviderException rpe + && RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.name().equals(rpe.code)) { + return true; + } + return super.decideRetry(throwable); + } + + @Override + protected Predicate newRetryPredicate() { + AtomicBoolean serverErrorRetryAvailable = new AtomicBoolean(true); + return throwable -> { + // Treat every mapped HTTP 5xx as transient once. The configured retry count remains the + // global cap across server errors and timeouts, so mixed failures cannot amplify attempts. + if (throwable instanceof SchemaException schemaException + && SchemaException.Code.RERANKING_PROVIDER_SERVER_ERROR + .name() + .equals(schemaException.code)) { + return serverErrorRetryAvailable.getAndSet(false); + } + return decideRetry(throwable); + }; } @Override @@ -142,7 +155,7 @@ protected RuntimeException mapHTTPError(Response jakartaResponse, String errorMe return RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.get( Map.of( "modelProvider", modelProvider().apiName(), - "httpStatus", String.valueOf(jakartaResponse.getStatus()), + "providerStatus", String.valueOf(jakartaResponse.getStatus()), "errorMessage", errorMessage)); } diff --git a/src/main/resources/errors.yaml b/src/main/resources/errors.yaml index 6d267f0457..b83d4f924a 100644 --- a/src/main/resources/errors.yaml +++ b/src/main/resources/errors.yaml @@ -2770,7 +2770,7 @@ server-errors: The command called an embedding provider to vectorize text, but the provider did not respond within the allowed time. The embedding provider was: ${modelProvider}. - The HTTP status code was: ${httpStatus}. + The provider status was: ${providerStatus}. The error message was: ${errorMessage}. ${SNIPPET.RETRY} @@ -2801,7 +2801,7 @@ server-errors: The command called an reranking provider to rerank results, but the provider did not respond within the allowed time. The reranking provider was: ${modelProvider}. - The HTTP status code was: ${httpStatus}. + The provider status was: ${providerStatus}. The error message was: ${errorMessage}. ${SNIPPET.RETRY} diff --git a/src/main/resources/reranking-providers-config.yaml b/src/main/resources/reranking-providers-config.yaml index 6a17984287..67739fe0a2 100644 --- a/src/main/resources/reranking-providers-config.yaml +++ b/src/main/resources/reranking-providers-config.yaml @@ -14,5 +14,4 @@ stargate: is-default: true url: https://us-west-2.api-dev.ai.datastax.com/nvidia/v1/ranking properties: - at-most-retries: 1 - max-batch-size: 10 + max-batch-size: 10 \ No newline at end of file diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingGatewayClientTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingGatewayClientTest.java index fc42e8ae8c..982c885ab7 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingGatewayClientTest.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingGatewayClientTest.java @@ -5,6 +5,8 @@ import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; +import io.grpc.Status; +import io.grpc.StatusRuntimeException; import io.quarkus.test.junit.QuarkusTest; import io.quarkus.test.junit.TestProfile; import io.smallrye.mutiny.Uni; @@ -239,4 +241,81 @@ void handleError() { assertThat(exception.code).isEqualTo(apiException.code); }); } + + @Test + void mapsAsyncDeadlineExceeded() { + EmbeddingService embeddingService = mock(EmbeddingService.class); + when(embeddingService.embed(any())) + .thenReturn(Uni.createFrom().failure(Status.DEADLINE_EXCEEDED.asRuntimeException())); + + Throwable result = vectorizeAndAwaitFailure(embeddingService); + + assertThat(result) + .isInstanceOf(EmbeddingProviderException.class) + .satisfies( + failure -> { + EmbeddingProviderException exception = (EmbeddingProviderException) failure; + assertThat(exception.code) + .isEqualTo(EmbeddingProviderException.Code.EMBEDDING_PROVIDER_TIMEOUT.name()); + assertThat(exception.body).contains(ModelProvider.OPENAI.apiName()); + assertThat(exception.body).contains(Status.Code.DEADLINE_EXCEEDED.name()); + }); + } + + @Test + void mapsSynchronousDeadlineExceeded() { + EmbeddingService embeddingService = mock(EmbeddingService.class); + when(embeddingService.embed(any())).thenThrow(Status.DEADLINE_EXCEEDED.asRuntimeException()); + + Throwable result = vectorizeAndAwaitFailure(embeddingService); + + assertThat(result) + .isInstanceOf(EmbeddingProviderException.class) + .satisfies( + failure -> { + EmbeddingProviderException exception = (EmbeddingProviderException) failure; + assertThat(exception.code) + .isEqualTo(EmbeddingProviderException.Code.EMBEDDING_PROVIDER_TIMEOUT.name()); + assertThat(exception.body).contains(Status.Code.DEADLINE_EXCEEDED.name()); + }); + } + + @Test + void preservesAsyncUnavailableStatusFailure() { + EmbeddingService embeddingService = mock(EmbeddingService.class); + StatusRuntimeException unavailable = Status.UNAVAILABLE.asRuntimeException(); + when(embeddingService.embed(any())).thenReturn(Uni.createFrom().failure(unavailable)); + + Throwable result = vectorizeAndAwaitFailure(embeddingService); + + assertThat(result).isSameAs(unavailable); + } + + private Throwable vectorizeAndAwaitFailure(EmbeddingService embeddingService) { + return createClient(embeddingService) + .vectorize( + 1, + List.of("data 1", "data 2"), + testConstants.EMBEDDING_CREDENTIALS, + EmbeddingGatewayClient.EmbeddingRequestType.INDEX) + .subscribe() + .withSubscriber(UniAssertSubscriber.create()) + .awaitFailure() + .getFailure(); + } + + private EmbeddingGatewayClient createClient(EmbeddingService embeddingService) { + return new EmbeddingGatewayClient( + ModelProvider.OPENAI, + PROVIDER_CONFIG, + MODEL_CONFIG, + SERVICE_CONFIG, + 1536, + Map.of(), + testConstants.TENANT, + "default", + embeddingService, + Map.of(), + TESTING_COMMAND_NAME); + } } diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProviderErrorMessageTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProviderErrorMessageTest.java index 93f67d0974..c5cfcf085d 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProviderErrorMessageTest.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProviderErrorMessageTest.java @@ -151,7 +151,7 @@ public void testRetryError() throws Exception { assertApiException( exception, EmbeddingProviderException.Code.EMBEDDING_PROVIDER_TIMEOUT, - "The HTTP status code was: 408."); + "The provider status was: 408."); } @Test diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidationTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidationTest.java index 96e280d23e..2bac3a561c 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidationTest.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidationTest.java @@ -9,7 +9,6 @@ import io.stargate.sgv2.jsonapi.exception.SchemaException; import jakarta.ws.rs.client.ClientResponseContext; -import jakarta.ws.rs.core.MediaType; import jakarta.ws.rs.core.Response; import org.junit.jupiter.api.Test; @@ -19,23 +18,12 @@ class RerankingProviderResponseValidationTest { new RerankingProviderResponseValidation(); @Test - void ignoresEmptyServerErrorResponse() { + void skipsBodyValidationForServerErrors() { ClientResponseContext responseContext = responseContext(Response.Status.INTERNAL_SERVER_ERROR); - when(responseContext.hasEntity()).thenReturn(false); assertThatCode(() -> validation.filter(null, responseContext)).doesNotThrowAnyException(); verify(responseContext, never()).hasEntity(); - } - - @Test - void ignoresNonJsonServerErrorResponse() { - ClientResponseContext responseContext = responseContext(Response.Status.INTERNAL_SERVER_ERROR); - when(responseContext.hasEntity()).thenReturn(true); - when(responseContext.getMediaType()).thenReturn(MediaType.TEXT_PLAIN_TYPE); - - assertThatCode(() -> validation.filter(null, responseContext)).doesNotThrowAnyException(); - verify(responseContext, never()).getMediaType(); } @@ -62,7 +50,6 @@ void rejectsEmptyRedirectResponse() { private static ClientResponseContext responseContext(Response.Status status) { ClientResponseContext responseContext = mock(ClientResponseContext.class); when(responseContext.getStatus()).thenReturn(status.getStatusCode()); - when(responseContext.getStatusInfo()).thenReturn(status); return responseContext; } } 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 67dc4c8a0a..2e93ed8133 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 @@ -15,6 +15,7 @@ import io.stargate.embedding.gateway.RerankingService; import io.stargate.sgv2.jsonapi.TestConstants; import io.stargate.sgv2.jsonapi.api.request.RerankingCredentials; +import io.stargate.sgv2.jsonapi.exception.ErrorFamily; import io.stargate.sgv2.jsonapi.exception.RerankingProviderException; import io.stargate.sgv2.jsonapi.exception.SchemaException; import io.stargate.sgv2.jsonapi.exception.ServerException; @@ -25,7 +26,6 @@ import io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfigImpl; import io.stargate.sgv2.jsonapi.service.reranking.operation.RerankingProvider; import io.stargate.sgv2.jsonapi.testresource.NoGlobalResourcesTestProfile; -import jakarta.inject.Inject; import java.util.List; import java.util.Map; import java.util.Optional; @@ -62,20 +62,6 @@ public class RerankingGatewayClientTest { "http://testing.com", REQUEST_PROPERTIES); - private static final RerankingProvidersConfigImpl.RerankingProviderConfigImpl PROVIDER_CONFIG = - new RerankingProvidersConfigImpl.RerankingProviderConfigImpl( - false, "test", true, Map.of(), List.of()); - - @Inject RerankingProvidersConfig rerankingProvidersConfig; - - @Test - void productionNvidiaConfigUsesOneRetry() { - assertThat(rerankingProvidersConfig.providers()).containsKey("nvidia"); - assertThat(rerankingProvidersConfig.providers().get("nvidia").models()) - .singleElement() - .satisfies(model -> assertThat(model.properties().atMostRetries()).isEqualTo(1)); - } - @Test void handleValidResponse() { RerankingService rerankService = mock(RerankingService.class); @@ -158,7 +144,10 @@ void mapsSchemaErrorFromGateway() { final SchemaException apiException = SchemaException.Code.RERANKING_PROVIDER_SERVER_ERROR.get( Map.of("errorMessage", "Test fail")); - errorResponseBuilder.setErrorCode(apiException.code).setErrorBody(apiException.getMessage()); + errorResponseBuilder + .setErrorCode(apiException.code) + .setErrorTitle("Gateway schema title") + .setErrorBody(apiException.getMessage()); builder.setError(errorResponseBuilder.build()); when(rerankService.rerank(any())).thenReturn(Uni.createFrom().item(builder.build())); @@ -186,8 +175,11 @@ void mapsSchemaErrorFromGateway() { .satisfies( e -> { SchemaException exception = (SchemaException) e; - assertThat(exception.getMessage()).isEqualTo(apiException.getMessage()); + assertThat(exception.family).isEqualTo(ErrorFamily.REQUEST); + assertThat(exception.scope).isEqualTo(SchemaException.SCOPE.scope()); assertThat(exception.code).isEqualTo(apiException.code); + assertThat(exception.title).isEqualTo("Reranking provider server error"); + assertThat(exception.body).isEqualTo(apiException.body); }); } @@ -210,8 +202,11 @@ void mapsServerErrorFromGateway() { .satisfies( failure -> { ServerException exception = (ServerException) failure; + assertThat(exception.family).isEqualTo(ErrorFamily.SERVER); + assertThat(exception.scope).isEmpty(); assertThat(exception.code) .isEqualTo(ServerException.Code.UNEXPECTED_SERVER_ERROR.name()); + assertThat(exception.title).isEqualTo("Unexpected server error"); assertThat(exception.body).isEqualTo("Gateway server error body"); }); } @@ -235,8 +230,11 @@ void mapsRerankingProviderErrorFromGateway() { .satisfies( failure -> { RerankingProviderException exception = (RerankingProviderException) failure; + assertThat(exception.family).isEqualTo(ErrorFamily.SERVER); + assertThat(exception.scope).isEqualTo(RerankingProviderException.SCOPE.scope()); assertThat(exception.code) .isEqualTo(RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.name()); + assertThat(exception.title).isEqualTo("Reranking Provider timed out"); assertThat(exception.body).isEqualTo("Gateway timeout body"); }); } @@ -260,6 +258,8 @@ void preservesUnknownGatewayError() { .satisfies( failure -> { RerankingProviderException exception = (RerankingProviderException) failure; + assertThat(exception.family).isEqualTo(ErrorFamily.SERVER); + assertThat(exception.scope).isEqualTo(RerankingProviderException.SCOPE.scope()); assertThat(exception.code).isEqualTo("FUTURE_GATEWAY_ERROR"); assertThat(exception.title).isEqualTo("Future gateway error"); assertThat(exception.body).isEqualTo("Future gateway error body"); @@ -305,7 +305,7 @@ void mapsSynchronousDeadlineExceeded() { } @Test - void preservesAsyncNonDeadlineStatusFailure() { + void preservesAsyncUnavailableStatusFailure() { RerankingService rerankService = mock(RerankingService.class); StatusRuntimeException unavailable = Status.UNAVAILABLE.asRuntimeException(); when(rerankService.rerank(any())).thenReturn(Uni.createFrom().failure(unavailable)); diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderRetryTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderRetryTest.java index 8a039d0fa2..b18244609f 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderRetryTest.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderRetryTest.java @@ -16,6 +16,7 @@ import jakarta.ws.rs.core.Response; import java.util.List; import java.util.Optional; +import java.util.concurrent.TimeoutException; import java.util.concurrent.atomic.AtomicInteger; import org.junit.jupiter.api.Test; @@ -49,7 +50,7 @@ void retriesServerErrorOnceThenReturnsSuccess() { @Test void stopsAfterOneServerErrorRetry() { - RetryingTestProvider provider = new RetryingTestProvider(); + RetryingTestProvider provider = new RetryingTestProvider(3); AtomicInteger calls = new AtomicInteger(); Throwable failure = @@ -76,9 +77,85 @@ void stopsAfterOneServerErrorRetry() { .isEqualTo(SchemaException.Code.RERANKING_PROVIDER_SERVER_ERROR.name())); } + @Test + void usesConfiguredRetryBudgetForTimeouts() { + RetryingTestProvider provider = new RetryingTestProvider(3); + AtomicInteger calls = new AtomicInteger(); + + provider + .execute( + Uni.createFrom() + .deferred( + () -> { + calls.incrementAndGet(); + return Uni.createFrom() + .failure(new TimeoutException("provider timed out")); + })) + .subscribe() + .withSubscriber(UniAssertSubscriber.create()) + .awaitFailure(); + + assertThat(calls).hasValue(4); + } + + @Test + void mixedRetryableFailuresShareConfiguredRetryBudget() { + RetryingTestProvider provider = new RetryingTestProvider(3); + AtomicInteger calls = new AtomicInteger(); + + provider + .execute( + Uni.createFrom() + .deferred( + () -> { + int attempt = calls.incrementAndGet(); + if (attempt == 2) { + return Uni.createFrom() + .item(response(Response.Status.INTERNAL_SERVER_ERROR)); + } + return Uni.createFrom() + .failure(new TimeoutException("provider timed out")); + })) + .subscribe() + .withSubscriber(UniAssertSubscriber.create()) + .awaitFailure(); + + assertThat(calls).hasValue(4); + } + + @Test + void serverErrorRetryBudgetIsFreshForEachSubscription() { + RetryingTestProvider provider = new RetryingTestProvider(3); + AtomicInteger calls = new AtomicInteger(); + Response successfulResponse = response(Response.Status.OK); + Uni operation = + provider.execute( + Uni.createFrom() + .deferred( + () -> + Uni.createFrom() + .item( + calls.incrementAndGet() % 2 == 1 + ? response(Response.Status.INTERNAL_SERVER_ERROR) + : successfulResponse))); + + operation + .subscribe() + .withSubscriber(UniAssertSubscriber.create()) + .awaitItem() + .assertItem(successfulResponse); + operation + .subscribe() + .withSubscriber(UniAssertSubscriber.create()) + .awaitItem() + .assertItem(successfulResponse); + + assertThat(calls).hasValue(4); + } + @Test void doesNotRetryClientError() { - RetryingTestProvider provider = new RetryingTestProvider(); + RetryingTestProvider provider = new RetryingTestProvider(3); AtomicInteger calls = new AtomicInteger(); Throwable failure = @@ -116,7 +193,11 @@ private static Response response(Response.Status status) { private static final class RetryingTestProvider extends RerankingProvider { private RetryingTestProvider() { - super(ModelProvider.NVIDIA, modelConfig()); + this(3); + } + + private RetryingTestProvider(int atMostRetries) { + super(ModelProvider.NVIDIA, modelConfig(atMostRetries)); } private Uni execute(Uni request) { @@ -137,7 +218,8 @@ public Uni rerank( throw new UnsupportedOperationException("Not used by retry tests"); } - private static RerankingProvidersConfig.RerankingProviderConfig.ModelConfig modelConfig() { + private static RerankingProvidersConfig.RerankingProviderConfig.ModelConfig modelConfig( + int atMostRetries) { return new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl( "test-model", new ApiModelSupport.ApiModelSupportImpl( @@ -145,7 +227,13 @@ private static RerankingProvidersConfig.RerankingProviderConfig.ModelConfig mode false, "http://testing.com", new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl - .RequestPropertiesImpl(1, 1, 100, 1, 0.0, 10)); + .RequestPropertiesImpl( + /* atMostRetries= */ atMostRetries, + /* initialBackOffMillis= */ 1, + /* readTimeoutMillis= */ 100, + /* maxBackOffMillis= */ 1, + /* jitter= */ 0.0, + /* maxBatchSize= */ 10)); } } } From 4e051d5a16755e9e82378faf917baf7ce7c4428d Mon Sep 17 00:00:00 2001 From: Eric Hare Date: Wed, 5 Aug 2026 12:06:44 -0700 Subject: [PATCH 3/4] fix: narrow reranking gateway error handling --- .../exception/RerankingProviderException.java | 18 -- .../gateway/EmbeddingGatewayClient.java | 43 ++-- .../AwsBedrockEmbeddingProvider.java | 2 +- .../operation/EmbeddingProvider.java | 2 +- .../EmbeddingProviderExceptionHandler.java | 2 +- .../service/provider/ProviderBase.java | 50 ++-- .../RerankingProviderExceptionHandler.java | 2 +- .../RerankingProviderResponseValidation.java | 6 - .../RerankingProvidersConfig.java | 6 +- .../reranking/gateway/RerankingEGWClient.java | 83 +++--- .../operation/RerankingProvider.java | 29 +-- src/main/resources/errors.yaml | 4 +- .../operation/EmbeddingGatewayClientTest.java | 79 ------ .../EmbeddingProviderErrorMessageTest.java | 2 +- ...rankingProviderResponseValidationTest.java | 55 ---- .../gateway/RerankingGatewayClientTest.java | 203 +++------------ .../operation/RerankingProviderRetryTest.java | 239 ------------------ 17 files changed, 121 insertions(+), 704 deletions(-) delete mode 100644 src/test/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidationTest.java delete mode 100644 src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderRetryTest.java diff --git a/src/main/java/io/stargate/sgv2/jsonapi/exception/RerankingProviderException.java b/src/main/java/io/stargate/sgv2/jsonapi/exception/RerankingProviderException.java index 2495bcfe1e..7c9ee3da7a 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/exception/RerankingProviderException.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/exception/RerankingProviderException.java @@ -1,9 +1,5 @@ package io.stargate.sgv2.jsonapi.exception; -import java.util.EnumSet; -import java.util.Optional; -import java.util.UUID; - public class RerankingProviderException extends ServerException { public static final Scope SCOPE = Scope.RERANKING_PROVIDER; @@ -12,20 +8,6 @@ public RerankingProviderException(ErrorInstance errorInstance) { super(errorInstance); } - /** Constructs a reranking provider exception from an unrecognized EGW error response. */ - public RerankingProviderException(String code, String title, String body) { - this( - new ErrorInstance( - UUID.randomUUID(), - FAMILY, - SCOPE, - code, - title, - body, - Optional.empty(), - EnumSet.noneOf(ExceptionFlags.class))); - } - public enum Code implements ErrorCode { RERANKING_PROVIDER_TIMEOUT; diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/gateway/EmbeddingGatewayClient.java b/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/gateway/EmbeddingGatewayClient.java index 5ab7a11759..275c945723 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/gateway/EmbeddingGatewayClient.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/gateway/EmbeddingGatewayClient.java @@ -164,12 +164,26 @@ public Uni vectorize( .setProviderContext(contextBuilder.build()) .build(); - // Defer the gRPC invocation so synchronous stub failures and asynchronous Uni failures use the - // same status-mapping path. - return Uni.createFrom() - .deferred(() -> grpcGatewayClient.embed(gatewayRequest)) - .onFailure(StatusRuntimeException.class) - .transform(this::mapStatusFailure) + // aaron 17 June 2025 - unsure why this error handled was not in the uni pipeline below + // kept it as is when refactoring + Uni embeddingResponse; + try { + embeddingResponse = grpcGatewayClient.embed(gatewayRequest); + } catch (StatusRuntimeException e) { + if (e.getStatus().getCode().equals(Status.Code.DEADLINE_EXCEEDED)) { + throw EmbeddingProviderException.Code.EMBEDDING_PROVIDER_TIMEOUT.get( + Map.of( + "modelProvider", + modelProvider().apiName(), + "httpStatus", + String.valueOf(e.getStatus().getCode()), + "errorMessage", + e.getMessage())); + } + throw e; + } + + return embeddingResponse .onItem() .transform( gatewayResponse -> { @@ -198,23 +212,6 @@ public Uni vectorize( }); } - private Throwable mapStatusFailure(Throwable failure) { - var statusException = (StatusRuntimeException) failure; - if (statusException.getStatus().getCode().equals(Status.Code.DEADLINE_EXCEEDED)) { - return EmbeddingProviderException.Code.EMBEDDING_PROVIDER_TIMEOUT.get( - Map.of( - "modelProvider", - modelProvider().apiName(), - "providerStatus", - String.valueOf(statusException.getStatus().getCode()), - "errorMessage", - statusException.getMessage())); - } - - // Preserve non-deadline gRPC failures so downstream handling retains their status and cause. - return failure; - } - /** Return MAX_VALUE because the batching is done inside EGW */ @Override public int maxBatchSize() { diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/AwsBedrockEmbeddingProvider.java b/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/AwsBedrockEmbeddingProvider.java index 84965d3c37..28c1b6cc54 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/AwsBedrockEmbeddingProvider.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/AwsBedrockEmbeddingProvider.java @@ -197,7 +197,7 @@ private Throwable mapBedrockException(BedrockRuntimeException bedrockException) Map.of( "modelProvider", modelProvider().apiName(), - "providerStatus", + "httpStatus", String.valueOf(bedrockException.statusCode()), "errorMessage", bedrockException.getMessage())); diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProvider.java b/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProvider.java index 9308650a0e..0f05ce34f1 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProvider.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProvider.java @@ -223,7 +223,7 @@ protected RuntimeException mapHTTPError(Response jakartaResponse, String errorMe Map.of( "modelProvider", modelProvider().apiName(), - "providerStatus", + "httpStatus", String.valueOf(jakartaResponse.getStatus()), "errorMessage", errorMessage)); diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/provider/EmbeddingProviderExceptionHandler.java b/src/main/java/io/stargate/sgv2/jsonapi/service/provider/EmbeddingProviderExceptionHandler.java index fbf8c60164..03efc7f439 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/provider/EmbeddingProviderExceptionHandler.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/provider/EmbeddingProviderExceptionHandler.java @@ -17,7 +17,7 @@ public RuntimeException handle(TimeoutException exception) { return EmbeddingProviderException.Code.EMBEDDING_PROVIDER_TIMEOUT.get( Map.of( "modelProvider", modelProvider.apiName(), - "providerStatus", "", + "httpStatus", "", "errorMessage", "")); } diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/provider/ProviderBase.java b/src/main/java/io/stargate/sgv2/jsonapi/service/provider/ProviderBase.java index 44b2792c4c..662df4d74d 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/provider/ProviderBase.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/provider/ProviderBase.java @@ -99,38 +99,24 @@ public ModelProvider modelProvider() { */ protected Uni retryHTTPCall(Uni uni) { - return Uni.createFrom() - .deferred( - () -> - uni - // Catch *any* web exception from jakarta rest client - .onFailure(WebApplicationException.class) - // and recover with the jakarta response, so we can translate to API exception - .recoverWithItem(ex -> ex.getResponse()) - .onItem() - // handle the response, throws if there is an error - .transform(this::handleHTTPResponse) - // decide if we want to retry - .onFailure(newRetryPredicate()) - .retry() - .withBackOff(initialBackOffDuration(), maxBackOffDuration()) - .withJitter(jitter()) - .atMost(atMostRetries()) - // after all retry logic, if we have an error we will need to handle it into an - // API exception - .onFailure() - .transform(exceptionHandler::maybeHandle)); - } - - /** - * Creates the retry predicate for one subscription. - * - *

Subclasses may override this method when retry decisions need per-request state. The - * predicate must not retain state on the provider instance because providers are shared across - * requests. - */ - protected Predicate newRetryPredicate() { - return this::decideRetry; + return uni + // Catch *any* web exception from jakarta rest client + .onFailure(WebApplicationException.class) + // and recover with the jakarta response, so we can translate to API exception + .recoverWithItem(ex -> ex.getResponse()) + .onItem() + // handle the response, throws if there is an error + .transform(this::handleHTTPResponse) + // decide if we want to retry + .onFailure(this::decideRetry) + .retry() + .withBackOff(initialBackOffDuration(), maxBackOffDuration()) + .withJitter(jitter()) + .atMost(atMostRetries()) + // after all retry logic, if we have an error we will need to handle it into an API + // exception + .onFailure() + .transform(exceptionHandler::maybeHandle); } /** diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/provider/RerankingProviderExceptionHandler.java b/src/main/java/io/stargate/sgv2/jsonapi/service/provider/RerankingProviderExceptionHandler.java index 50d2dd4a85..2e4e5dc105 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/provider/RerankingProviderExceptionHandler.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/provider/RerankingProviderExceptionHandler.java @@ -16,7 +16,7 @@ public RuntimeException handle(TimeoutException exception) { return RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.get( Map.of( "modelProvider", modelProvider.apiName(), - "providerStatus", "", + "httpStatus", "", "errorMessage", "")); } } diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidation.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidation.java index 68896190c1..aefcd1618e 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidation.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidation.java @@ -42,12 +42,6 @@ public void filter(ClientRequestContext requestContext, ClientResponseContext re if (responseContext.getStatus() == 0) { return; } - - // HTTP errors may be empty or non-JSON; ProviderBase maps statuses 400 and above. - if (responseContext.getStatus() >= 400) { - return; - } - // Throw error if there is no response body if (!responseContext.hasEntity()) { throw SchemaException.Code.RERANKING_PROVIDER_UNEXPECTED_RESPONSE.get( 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 349a6c4102..4978debeff 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 @@ -74,10 +74,10 @@ interface ModelConfig { interface RequestProperties { /** - * Specifies the maximum number of retries after the initial request. Default is 3 (up to 4 - * total attempts). + * Specifies the maximum number of attempts before failing. Default is 3 (1 request + 2 + * retries). * - * @return The maximum number of retries after the initial request. + * @return The maximum number of attempts before failing. */ @WithDefault("3") int atMostRetries(); diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java index 1e5d182f95..301726c35c 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java @@ -7,8 +7,6 @@ import io.stargate.embedding.gateway.RerankingService; import io.stargate.sgv2.jsonapi.api.request.RerankingCredentials; import io.stargate.sgv2.jsonapi.api.request.tenant.Tenant; -import io.stargate.sgv2.jsonapi.exception.APIException; -import io.stargate.sgv2.jsonapi.exception.ErrorCode; import io.stargate.sgv2.jsonapi.exception.RerankingProviderException; import io.stargate.sgv2.jsonapi.exception.SchemaException; import io.stargate.sgv2.jsonapi.exception.ServerException; @@ -82,15 +80,39 @@ public Uni rerank( .setProviderContext(contextBuilder.build()) .build(); - return Uni.createFrom() - .deferred(() -> grpcGatewayService.rerank(gatewayRequest)) - .onFailure(StatusRuntimeException.class) - .transform(this::mapStatusFailure) + // TODO: Why is this error handling here not part of the uni pipeline? + Uni gatewayRerankingUni; + try { + gatewayRerankingUni = grpcGatewayService.rerank(gatewayRequest); + } catch (StatusRuntimeException e) { + if (e.getStatus().getCode().equals(Status.Code.DEADLINE_EXCEEDED)) { + throw RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.get( + Map.of( + "modelProvider", + modelProvider().apiName(), + "httpStatus", + String.valueOf(e.getStatus().getCode()), + "errorMessage", + e.getMessage())); + } + throw e; + } + + return gatewayRerankingUni .onItem() .transform( gatewayResponse -> { if (gatewayResponse.hasError()) { - throw mapGatewayError(gatewayResponse.getError()); + var error = gatewayResponse.getError(); + var unexpectedServerError = ServerException.Code.UNEXPECTED_SERVER_ERROR; + if (unexpectedServerError.name().equals(error.getErrorCode())) { + throw unexpectedServerError.withPreformattedMessage(error.getErrorBody()); + } + + // 22-Jan-2026, tatu: This is ugly. But has to be done to work around fragility + // of exception mapping + throw SchemaException.Code.valueOf(error.getErrorCode()) + .withPreformattedMessage(error.getErrorBody()); } return new BatchedRerankingResponse( @@ -101,51 +123,4 @@ public Uni rerank( createModelUsage(gatewayResponse.getModelUsage())); }); } - - private Throwable mapStatusFailure(Throwable failure) { - var statusException = (StatusRuntimeException) failure; - if (statusException.getStatus().getCode().equals(Status.Code.DEADLINE_EXCEEDED)) { - return RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.get( - Map.of( - "modelProvider", - modelProvider().apiName(), - "providerStatus", - String.valueOf(statusException.getStatus().getCode()), - "errorMessage", - statusException.getMessage())); - } - - // Only DEADLINE_EXCEEDED has a defined Data API mapping. Preserve other gRPC statuses so - // upstream handlers retain the original status instead of misclassifying it as a timeout. - return failure; - } - - private RuntimeException mapGatewayError(EmbeddingGateway.RerankingResponse.ErrorResponse error) { - String errorCode = error.getErrorCode(); - - // Preserve API compatibility for known gateway codes. This precedence is intentional: - // Schema (REQUEST/SCHEMA), then unscoped Server, then RerankingProvider. Unknown gateway - // codes pass through with their gateway-supplied title and body rather than failing lookup. - var schemaCode = findErrorCode(errorCode, SchemaException.Code.values()); - if (schemaCode.isPresent()) { - return schemaCode.get().withPreformattedMessage(error.getErrorBody()); - } - - var serverCode = findErrorCode(errorCode, ServerException.Code.values()); - if (serverCode.isPresent()) { - return serverCode.get().withPreformattedMessage(error.getErrorBody()); - } - - var rerankingProviderCode = findErrorCode(errorCode, RerankingProviderException.Code.values()); - if (rerankingProviderCode.isPresent()) { - return rerankingProviderCode.get().withPreformattedMessage(error.getErrorBody()); - } - - return new RerankingProviderException(errorCode, error.getErrorTitle(), error.getErrorBody()); - } - - private static & ErrorCode> - Optional findErrorCode(String errorCode, C[] codes) { - return Arrays.stream(codes).filter(code -> code.name().equals(errorCode)).findFirst(); - } } diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java index 9b35f19782..9a9f22ffff 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProvider.java @@ -14,8 +14,6 @@ import java.util.Comparator; import java.util.List; import java.util.Map; -import java.util.concurrent.atomic.AtomicBoolean; -import java.util.function.Predicate; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @@ -124,27 +122,10 @@ protected int atMostRetries() { @Override protected boolean decideRetry(Throwable throwable) { - if (throwable instanceof RerankingProviderException rpe - && RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.name().equals(rpe.code)) { - return true; - } - return super.decideRetry(throwable); - } - - @Override - protected Predicate newRetryPredicate() { - AtomicBoolean serverErrorRetryAvailable = new AtomicBoolean(true); - return throwable -> { - // Treat every mapped HTTP 5xx as transient once. The configured retry count remains the - // global cap across server errors and timeouts, so mixed failures cannot amplify attempts. - if (throwable instanceof SchemaException schemaException - && SchemaException.Code.RERANKING_PROVIDER_SERVER_ERROR - .name() - .equals(schemaException.code)) { - return serverErrorRetryAvailable.getAndSet(false); - } - return decideRetry(throwable); - }; + boolean retry = + throwable instanceof RerankingProviderException rpe + && RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.name().equals(rpe.code); + return retry || super.decideRetry(throwable); } @Override @@ -155,7 +136,7 @@ protected RuntimeException mapHTTPError(Response jakartaResponse, String errorMe return RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.get( Map.of( "modelProvider", modelProvider().apiName(), - "providerStatus", String.valueOf(jakartaResponse.getStatus()), + "httpStatus", String.valueOf(jakartaResponse.getStatus()), "errorMessage", errorMessage)); } diff --git a/src/main/resources/errors.yaml b/src/main/resources/errors.yaml index b83d4f924a..6d267f0457 100644 --- a/src/main/resources/errors.yaml +++ b/src/main/resources/errors.yaml @@ -2770,7 +2770,7 @@ server-errors: The command called an embedding provider to vectorize text, but the provider did not respond within the allowed time. The embedding provider was: ${modelProvider}. - The provider status was: ${providerStatus}. + The HTTP status code was: ${httpStatus}. The error message was: ${errorMessage}. ${SNIPPET.RETRY} @@ -2801,7 +2801,7 @@ server-errors: The command called an reranking provider to rerank results, but the provider did not respond within the allowed time. The reranking provider was: ${modelProvider}. - The provider status was: ${providerStatus}. + The HTTP status code was: ${httpStatus}. The error message was: ${errorMessage}. ${SNIPPET.RETRY} diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingGatewayClientTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingGatewayClientTest.java index 982c885ab7..fc42e8ae8c 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingGatewayClientTest.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingGatewayClientTest.java @@ -5,8 +5,6 @@ import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; -import io.grpc.Status; -import io.grpc.StatusRuntimeException; import io.quarkus.test.junit.QuarkusTest; import io.quarkus.test.junit.TestProfile; import io.smallrye.mutiny.Uni; @@ -241,81 +239,4 @@ void handleError() { assertThat(exception.code).isEqualTo(apiException.code); }); } - - @Test - void mapsAsyncDeadlineExceeded() { - EmbeddingService embeddingService = mock(EmbeddingService.class); - when(embeddingService.embed(any())) - .thenReturn(Uni.createFrom().failure(Status.DEADLINE_EXCEEDED.asRuntimeException())); - - Throwable result = vectorizeAndAwaitFailure(embeddingService); - - assertThat(result) - .isInstanceOf(EmbeddingProviderException.class) - .satisfies( - failure -> { - EmbeddingProviderException exception = (EmbeddingProviderException) failure; - assertThat(exception.code) - .isEqualTo(EmbeddingProviderException.Code.EMBEDDING_PROVIDER_TIMEOUT.name()); - assertThat(exception.body).contains(ModelProvider.OPENAI.apiName()); - assertThat(exception.body).contains(Status.Code.DEADLINE_EXCEEDED.name()); - }); - } - - @Test - void mapsSynchronousDeadlineExceeded() { - EmbeddingService embeddingService = mock(EmbeddingService.class); - when(embeddingService.embed(any())).thenThrow(Status.DEADLINE_EXCEEDED.asRuntimeException()); - - Throwable result = vectorizeAndAwaitFailure(embeddingService); - - assertThat(result) - .isInstanceOf(EmbeddingProviderException.class) - .satisfies( - failure -> { - EmbeddingProviderException exception = (EmbeddingProviderException) failure; - assertThat(exception.code) - .isEqualTo(EmbeddingProviderException.Code.EMBEDDING_PROVIDER_TIMEOUT.name()); - assertThat(exception.body).contains(Status.Code.DEADLINE_EXCEEDED.name()); - }); - } - - @Test - void preservesAsyncUnavailableStatusFailure() { - EmbeddingService embeddingService = mock(EmbeddingService.class); - StatusRuntimeException unavailable = Status.UNAVAILABLE.asRuntimeException(); - when(embeddingService.embed(any())).thenReturn(Uni.createFrom().failure(unavailable)); - - Throwable result = vectorizeAndAwaitFailure(embeddingService); - - assertThat(result).isSameAs(unavailable); - } - - private Throwable vectorizeAndAwaitFailure(EmbeddingService embeddingService) { - return createClient(embeddingService) - .vectorize( - 1, - List.of("data 1", "data 2"), - testConstants.EMBEDDING_CREDENTIALS, - EmbeddingGatewayClient.EmbeddingRequestType.INDEX) - .subscribe() - .withSubscriber(UniAssertSubscriber.create()) - .awaitFailure() - .getFailure(); - } - - private EmbeddingGatewayClient createClient(EmbeddingService embeddingService) { - return new EmbeddingGatewayClient( - ModelProvider.OPENAI, - PROVIDER_CONFIG, - MODEL_CONFIG, - SERVICE_CONFIG, - 1536, - Map.of(), - testConstants.TENANT, - "default", - embeddingService, - Map.of(), - TESTING_COMMAND_NAME); - } } diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProviderErrorMessageTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProviderErrorMessageTest.java index c5cfcf085d..93f67d0974 100644 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProviderErrorMessageTest.java +++ b/src/test/java/io/stargate/sgv2/jsonapi/service/embedding/operation/EmbeddingProviderErrorMessageTest.java @@ -151,7 +151,7 @@ public void testRetryError() throws Exception { assertApiException( exception, EmbeddingProviderException.Code.EMBEDDING_PROVIDER_TIMEOUT, - "The provider status was: 408."); + "The HTTP status code was: 408."); } @Test diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidationTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidationTest.java deleted file mode 100644 index 2bac3a561c..0000000000 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/configuration/RerankingProviderResponseValidationTest.java +++ /dev/null @@ -1,55 +0,0 @@ -package io.stargate.sgv2.jsonapi.service.reranking.configuration; - -import static org.assertj.core.api.Assertions.assertThatCode; -import static org.assertj.core.api.Assertions.assertThatThrownBy; -import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.never; -import static org.mockito.Mockito.verify; -import static org.mockito.Mockito.when; - -import io.stargate.sgv2.jsonapi.exception.SchemaException; -import jakarta.ws.rs.client.ClientResponseContext; -import jakarta.ws.rs.core.Response; -import org.junit.jupiter.api.Test; - -class RerankingProviderResponseValidationTest { - - private final RerankingProviderResponseValidation validation = - new RerankingProviderResponseValidation(); - - @Test - void skipsBodyValidationForServerErrors() { - ClientResponseContext responseContext = responseContext(Response.Status.INTERNAL_SERVER_ERROR); - - assertThatCode(() -> validation.filter(null, responseContext)).doesNotThrowAnyException(); - - verify(responseContext, never()).hasEntity(); - verify(responseContext, never()).getMediaType(); - } - - @Test - void rejectsEmptySuccessfulResponse() { - ClientResponseContext responseContext = responseContext(Response.Status.OK); - when(responseContext.hasEntity()).thenReturn(false); - - assertThatThrownBy(() -> validation.filter(null, responseContext)) - .isInstanceOf(SchemaException.class) - .hasMessageContaining("No response body from the reranking provider"); - } - - @Test - void rejectsEmptyRedirectResponse() { - ClientResponseContext responseContext = responseContext(Response.Status.TEMPORARY_REDIRECT); - when(responseContext.hasEntity()).thenReturn(false); - - assertThatThrownBy(() -> validation.filter(null, responseContext)) - .isInstanceOf(SchemaException.class) - .hasMessageContaining("No response body from the reranking provider"); - } - - private static ClientResponseContext responseContext(Response.Status status) { - ClientResponseContext responseContext = mock(ClientResponseContext.class); - when(responseContext.getStatus()).thenReturn(status.getStatusCode()); - return responseContext; - } -} 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 2e93ed8133..46ae175196 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 @@ -5,8 +5,6 @@ import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; -import io.grpc.Status; -import io.grpc.StatusRuntimeException; import io.quarkus.test.junit.QuarkusTest; import io.quarkus.test.junit.TestProfile; import io.smallrye.mutiny.Uni; @@ -15,8 +13,6 @@ import io.stargate.embedding.gateway.RerankingService; import io.stargate.sgv2.jsonapi.TestConstants; import io.stargate.sgv2.jsonapi.api.request.RerankingCredentials; -import io.stargate.sgv2.jsonapi.exception.ErrorFamily; -import io.stargate.sgv2.jsonapi.exception.RerankingProviderException; import io.stargate.sgv2.jsonapi.exception.SchemaException; import io.stargate.sgv2.jsonapi.exception.ServerException; import io.stargate.sgv2.jsonapi.service.provider.ApiModelSupport; @@ -62,6 +58,10 @@ public class RerankingGatewayClientTest { "http://testing.com", REQUEST_PROPERTIES); + private static final RerankingProvidersConfigImpl.RerankingProviderConfigImpl PROVIDER_CONFIG = + new RerankingProvidersConfigImpl.RerankingProviderConfigImpl( + false, "test", true, Map.of(), List.of()); + @Test void handleValidResponse() { RerankingService rerankService = mock(RerankingService.class); @@ -134,7 +134,7 @@ void handleValidResponse() { } @Test - void mapsSchemaErrorFromGateway() { + void handleError() { RerankingService rerankService = mock(RerankingService.class); final EmbeddingGateway.RerankingResponse.Builder builder = @@ -144,10 +144,7 @@ void mapsSchemaErrorFromGateway() { final SchemaException apiException = SchemaException.Code.RERANKING_PROVIDER_SERVER_ERROR.get( Map.of("errorMessage", "Test fail")); - errorResponseBuilder - .setErrorCode(apiException.code) - .setErrorTitle("Gateway schema title") - .setErrorBody(apiException.getMessage()); + errorResponseBuilder.setErrorCode(apiException.code).setErrorBody(apiException.getMessage()); builder.setError(errorResponseBuilder.build()); when(rerankService.rerank(any())).thenReturn(Uni.createFrom().item(builder.build())); @@ -175,174 +172,52 @@ void mapsSchemaErrorFromGateway() { .satisfies( e -> { SchemaException exception = (SchemaException) e; - assertThat(exception.family).isEqualTo(ErrorFamily.REQUEST); - assertThat(exception.scope).isEqualTo(SchemaException.SCOPE.scope()); + assertThat(exception.getMessage()).isEqualTo(apiException.getMessage()); assertThat(exception.code).isEqualTo(apiException.code); - assertThat(exception.title).isEqualTo("Reranking provider server error"); - assertThat(exception.body).isEqualTo(apiException.body); }); } @Test - void mapsServerErrorFromGateway() { + void handleUnexpectedServerError() { RerankingService rerankService = mock(RerankingService.class); - when(rerankService.rerank(any())) - .thenReturn( - Uni.createFrom() - .item( - gatewayErrorResponse( - ServerException.Code.UNEXPECTED_SERVER_ERROR.name(), - "Gateway server error", - "Gateway server error body"))); + final ServerException apiException = + ServerException.Code.UNEXPECTED_SERVER_ERROR.withPreformattedMessage("Test fail"); + final EmbeddingGateway.RerankingResponse gatewayResponse = + EmbeddingGateway.RerankingResponse.newBuilder() + .setError( + EmbeddingGateway.RerankingResponse.ErrorResponse.newBuilder() + .setErrorCode(apiException.code) + .setErrorBody(apiException.getMessage())) + .build(); + when(rerankService.rerank(any())).thenReturn(Uni.createFrom().item(gatewayResponse)); + + RerankingEGWClient rerankEGWClient = + new RerankingEGWClient( + ModelProvider.NVIDIA, + MODEL_CONFIG, + testConstants.TENANT, + "default", + rerankService, + Map.of(), + TESTING_COMMAND_NAME); - Throwable result = rerankAndAwaitFailure(rerankService); + Throwable result = + rerankEGWClient + .rerank(1, "apple", List.of("orange", "apple"), RERANK_CREDENTIALS) + .subscribe() + .withSubscriber(UniAssertSubscriber.create()) + .awaitFailure() + .getFailure(); assertThat(result) .isInstanceOf(ServerException.class) .satisfies( failure -> { ServerException exception = (ServerException) failure; - assertThat(exception.family).isEqualTo(ErrorFamily.SERVER); - assertThat(exception.scope).isEmpty(); - assertThat(exception.code) - .isEqualTo(ServerException.Code.UNEXPECTED_SERVER_ERROR.name()); - assertThat(exception.title).isEqualTo("Unexpected server error"); - assertThat(exception.body).isEqualTo("Gateway server error body"); - }); - } - - @Test - void mapsRerankingProviderErrorFromGateway() { - RerankingService rerankService = mock(RerankingService.class); - when(rerankService.rerank(any())) - .thenReturn( - Uni.createFrom() - .item( - gatewayErrorResponse( - RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.name(), - "Gateway timeout", - "Gateway timeout body"))); - - Throwable result = rerankAndAwaitFailure(rerankService); - - assertThat(result) - .isInstanceOf(RerankingProviderException.class) - .satisfies( - failure -> { - RerankingProviderException exception = (RerankingProviderException) failure; - assertThat(exception.family).isEqualTo(ErrorFamily.SERVER); - assertThat(exception.scope).isEqualTo(RerankingProviderException.SCOPE.scope()); - assertThat(exception.code) - .isEqualTo(RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.name()); - assertThat(exception.title).isEqualTo("Reranking Provider timed out"); - assertThat(exception.body).isEqualTo("Gateway timeout body"); - }); - } - - @Test - void preservesUnknownGatewayError() { - RerankingService rerankService = mock(RerankingService.class); - when(rerankService.rerank(any())) - .thenReturn( - Uni.createFrom() - .item( - gatewayErrorResponse( - "FUTURE_GATEWAY_ERROR", - "Future gateway error", - "Future gateway error body"))); - - Throwable result = rerankAndAwaitFailure(rerankService); - - assertThat(result) - .isInstanceOf(RerankingProviderException.class) - .satisfies( - failure -> { - RerankingProviderException exception = (RerankingProviderException) failure; - assertThat(exception.family).isEqualTo(ErrorFamily.SERVER); - assertThat(exception.scope).isEqualTo(RerankingProviderException.SCOPE.scope()); - assertThat(exception.code).isEqualTo("FUTURE_GATEWAY_ERROR"); - assertThat(exception.title).isEqualTo("Future gateway error"); - assertThat(exception.body).isEqualTo("Future gateway error body"); - }); - } - - @Test - void mapsAsyncDeadlineExceeded() { - RerankingService rerankService = mock(RerankingService.class); - when(rerankService.rerank(any())) - .thenReturn(Uni.createFrom().failure(Status.DEADLINE_EXCEEDED.asRuntimeException())); - - Throwable result = rerankAndAwaitFailure(rerankService); - - assertThat(result) - .isInstanceOf(RerankingProviderException.class) - .satisfies( - failure -> { - RerankingProviderException exception = (RerankingProviderException) failure; - assertThat(exception.code) - .isEqualTo(RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.name()); - assertThat(exception.body).contains(ModelProvider.NVIDIA.apiName()); - assertThat(exception.body).contains(Status.Code.DEADLINE_EXCEEDED.name()); - }); - } - - @Test - void mapsSynchronousDeadlineExceeded() { - RerankingService rerankService = mock(RerankingService.class); - when(rerankService.rerank(any())).thenThrow(Status.DEADLINE_EXCEEDED.asRuntimeException()); - - Throwable result = rerankAndAwaitFailure(rerankService); - - assertThat(result) - .isInstanceOf(RerankingProviderException.class) - .satisfies( - failure -> { - RerankingProviderException exception = (RerankingProviderException) failure; - assertThat(exception.code) - .isEqualTo(RerankingProviderException.Code.RERANKING_PROVIDER_TIMEOUT.name()); - assertThat(exception.body).contains(Status.Code.DEADLINE_EXCEEDED.name()); + assertThat(exception.code).isEqualTo(apiException.code); + assertThat(exception.getMessage()).isEqualTo(apiException.getMessage()); + assertThat(exception.fullyQualifiedCode()) + .isEqualTo("SERVER_UNEXPECTED_SERVER_ERROR"); }); } - - @Test - void preservesAsyncUnavailableStatusFailure() { - RerankingService rerankService = mock(RerankingService.class); - StatusRuntimeException unavailable = Status.UNAVAILABLE.asRuntimeException(); - when(rerankService.rerank(any())).thenReturn(Uni.createFrom().failure(unavailable)); - - Throwable result = rerankAndAwaitFailure(rerankService); - - assertThat(result).isSameAs(unavailable); - } - - private static Throwable rerankAndAwaitFailure(RerankingService rerankService) { - return createClient(rerankService) - .rerank(1, "apple", List.of("orange", "apple"), RERANK_CREDENTIALS) - .subscribe() - .withSubscriber(UniAssertSubscriber.create()) - .awaitFailure() - .getFailure(); - } - - private static RerankingEGWClient createClient(RerankingService rerankService) { - return new RerankingEGWClient( - ModelProvider.NVIDIA, - MODEL_CONFIG, - testConstants.TENANT, - "default", - rerankService, - Map.of(), - TESTING_COMMAND_NAME); - } - - private static EmbeddingGateway.RerankingResponse gatewayErrorResponse( - String code, String title, String body) { - return EmbeddingGateway.RerankingResponse.newBuilder() - .setError( - EmbeddingGateway.RerankingResponse.ErrorResponse.newBuilder() - .setErrorCode(code) - .setErrorTitle(title) - .setErrorBody(body)) - .build(); - } } diff --git a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderRetryTest.java b/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderRetryTest.java deleted file mode 100644 index b18244609f..0000000000 --- a/src/test/java/io/stargate/sgv2/jsonapi/service/reranking/operation/RerankingProviderRetryTest.java +++ /dev/null @@ -1,239 +0,0 @@ -package io.stargate.sgv2.jsonapi.service.reranking.operation; - -import static org.assertj.core.api.Assertions.assertThat; -import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.when; - -import io.smallrye.mutiny.Uni; -import io.smallrye.mutiny.helpers.test.UniAssertSubscriber; -import io.stargate.sgv2.jsonapi.api.request.RerankingCredentials; -import io.stargate.sgv2.jsonapi.exception.SchemaException; -import io.stargate.sgv2.jsonapi.service.provider.ApiModelSupport; -import io.stargate.sgv2.jsonapi.service.provider.ModelProvider; -import io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfig; -import io.stargate.sgv2.jsonapi.service.reranking.configuration.RerankingProvidersConfigImpl; -import jakarta.ws.rs.core.MediaType; -import jakarta.ws.rs.core.Response; -import java.util.List; -import java.util.Optional; -import java.util.concurrent.TimeoutException; -import java.util.concurrent.atomic.AtomicInteger; -import org.junit.jupiter.api.Test; - -class RerankingProviderRetryTest { - - @Test - void retriesServerErrorOnceThenReturnsSuccess() { - RetryingTestProvider provider = new RetryingTestProvider(); - AtomicInteger calls = new AtomicInteger(); - Response successfulResponse = response(Response.Status.OK); - - Response result = - provider - .execute( - Uni.createFrom() - .deferred( - () -> - Uni.createFrom() - .item( - calls.incrementAndGet() == 1 - ? response(Response.Status.INTERNAL_SERVER_ERROR) - : successfulResponse))) - .subscribe() - .withSubscriber(UniAssertSubscriber.create()) - .awaitItem() - .getItem(); - - assertThat(result).isSameAs(successfulResponse); - assertThat(calls).hasValue(2); - } - - @Test - void stopsAfterOneServerErrorRetry() { - RetryingTestProvider provider = new RetryingTestProvider(3); - AtomicInteger calls = new AtomicInteger(); - - Throwable failure = - provider - .execute( - Uni.createFrom() - .deferred( - () -> { - calls.incrementAndGet(); - return Uni.createFrom() - .item(response(Response.Status.INTERNAL_SERVER_ERROR)); - })) - .subscribe() - .withSubscriber(UniAssertSubscriber.create()) - .awaitFailure() - .getFailure(); - - assertThat(calls).hasValue(2); - assertThat(failure) - .isInstanceOf(SchemaException.class) - .satisfies( - throwable -> - assertThat(((SchemaException) throwable).code) - .isEqualTo(SchemaException.Code.RERANKING_PROVIDER_SERVER_ERROR.name())); - } - - @Test - void usesConfiguredRetryBudgetForTimeouts() { - RetryingTestProvider provider = new RetryingTestProvider(3); - AtomicInteger calls = new AtomicInteger(); - - provider - .execute( - Uni.createFrom() - .deferred( - () -> { - calls.incrementAndGet(); - return Uni.createFrom() - .failure(new TimeoutException("provider timed out")); - })) - .subscribe() - .withSubscriber(UniAssertSubscriber.create()) - .awaitFailure(); - - assertThat(calls).hasValue(4); - } - - @Test - void mixedRetryableFailuresShareConfiguredRetryBudget() { - RetryingTestProvider provider = new RetryingTestProvider(3); - AtomicInteger calls = new AtomicInteger(); - - provider - .execute( - Uni.createFrom() - .deferred( - () -> { - int attempt = calls.incrementAndGet(); - if (attempt == 2) { - return Uni.createFrom() - .item(response(Response.Status.INTERNAL_SERVER_ERROR)); - } - return Uni.createFrom() - .failure(new TimeoutException("provider timed out")); - })) - .subscribe() - .withSubscriber(UniAssertSubscriber.create()) - .awaitFailure(); - - assertThat(calls).hasValue(4); - } - - @Test - void serverErrorRetryBudgetIsFreshForEachSubscription() { - RetryingTestProvider provider = new RetryingTestProvider(3); - AtomicInteger calls = new AtomicInteger(); - Response successfulResponse = response(Response.Status.OK); - Uni operation = - provider.execute( - Uni.createFrom() - .deferred( - () -> - Uni.createFrom() - .item( - calls.incrementAndGet() % 2 == 1 - ? response(Response.Status.INTERNAL_SERVER_ERROR) - : successfulResponse))); - - operation - .subscribe() - .withSubscriber(UniAssertSubscriber.create()) - .awaitItem() - .assertItem(successfulResponse); - operation - .subscribe() - .withSubscriber(UniAssertSubscriber.create()) - .awaitItem() - .assertItem(successfulResponse); - - assertThat(calls).hasValue(4); - } - - @Test - void doesNotRetryClientError() { - RetryingTestProvider provider = new RetryingTestProvider(3); - AtomicInteger calls = new AtomicInteger(); - - Throwable failure = - provider - .execute( - Uni.createFrom() - .deferred( - () -> { - calls.incrementAndGet(); - return Uni.createFrom().item(response(Response.Status.BAD_REQUEST)); - })) - .subscribe() - .withSubscriber(UniAssertSubscriber.create()) - .awaitFailure() - .getFailure(); - - assertThat(calls).hasValue(1); - assertThat(failure) - .isInstanceOf(SchemaException.class) - .satisfies( - throwable -> - assertThat(((SchemaException) throwable).code) - .isEqualTo(SchemaException.Code.RERANKING_PROVIDER_CLIENT_ERROR.name())); - } - - private static Response response(Response.Status status) { - Response response = mock(Response.class); - when(response.getStatus()).thenReturn(status.getStatusCode()); - when(response.getStatusInfo()).thenReturn(status); - when(response.getMediaType()).thenReturn(MediaType.TEXT_PLAIN_TYPE); - when(response.readEntity(String.class)).thenReturn("provider response"); - return response; - } - - private static final class RetryingTestProvider extends RerankingProvider { - - private RetryingTestProvider() { - this(3); - } - - private RetryingTestProvider(int atMostRetries) { - super(ModelProvider.NVIDIA, modelConfig(atMostRetries)); - } - - private Uni execute(Uni request) { - return retryHTTPCall(request); - } - - @Override - protected String errorMessageJsonPtr() { - return "/message"; - } - - @Override - public Uni rerank( - int batchId, - String query, - List passages, - RerankingCredentials rerankingCredentials) { - throw new UnsupportedOperationException("Not used by retry tests"); - } - - private static RerankingProvidersConfig.RerankingProviderConfig.ModelConfig modelConfig( - int atMostRetries) { - return new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl( - "test-model", - new ApiModelSupport.ApiModelSupportImpl( - ApiModelSupport.SupportStatus.SUPPORTED, Optional.empty()), - false, - "http://testing.com", - new RerankingProvidersConfigImpl.RerankingProviderConfigImpl.ModelConfigImpl - .RequestPropertiesImpl( - /* atMostRetries= */ atMostRetries, - /* initialBackOffMillis= */ 1, - /* readTimeoutMillis= */ 100, - /* maxBackOffMillis= */ 1, - /* jitter= */ 0.0, - /* maxBatchSize= */ 10)); - } - } -} From f87440f2b5c94bd4cc12a073bf1174e7ef5e8b03 Mon Sep 17 00:00:00 2001 From: Eric Hare Date: Fri, 7 Aug 2026 09:30:24 -0700 Subject: [PATCH 4/4] fix: preserve gateway error body for unmapped reranking error codes --- .../reranking/gateway/RerankingEGWClient.java | 11 ++++- .../gateway/RerankingGatewayClientTest.java | 41 +++++++++++++++++++ 2 files changed, 50 insertions(+), 2 deletions(-) diff --git a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java index 301726c35c..0f1b81270c 100644 --- a/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java +++ b/src/main/java/io/stargate/sgv2/jsonapi/service/reranking/gateway/RerankingEGWClient.java @@ -111,8 +111,15 @@ public Uni rerank( // 22-Jan-2026, tatu: This is ugly. But has to be done to work around fragility // of exception mapping - throw SchemaException.Code.valueOf(error.getErrorCode()) - .withPreformattedMessage(error.getErrorBody()); + SchemaException.Code schemaCode; + try { + schemaCode = SchemaException.Code.valueOf(error.getErrorCode()); + } catch (IllegalArgumentException e) { + // Gateway codes outside SchemaException.Code (e.g. RERANKING_PROVIDER_TIMEOUT) + // must not shadow the original error body with an enum-lookup failure. + throw unexpectedServerError.withPreformattedMessage(error.getErrorBody()); + } + throw schemaCode.withPreformattedMessage(error.getErrorBody()); } return new BatchedRerankingResponse( 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 46ae175196..7169c24c24 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 @@ -220,4 +220,45 @@ void handleUnexpectedServerError() { .isEqualTo("SERVER_UNEXPECTED_SERVER_ERROR"); }); } + + @Test + void handleUnmappedErrorCode() { + RerankingService rerankService = mock(RerankingService.class); + final EmbeddingGateway.RerankingResponse gatewayResponse = + EmbeddingGateway.RerankingResponse.newBuilder() + .setError( + EmbeddingGateway.RerankingResponse.ErrorResponse.newBuilder() + .setErrorCode("RERANKING_PROVIDER_TIMEOUT") + .setErrorBody("Reranking provider timed out upstream")) + .build(); + when(rerankService.rerank(any())).thenReturn(Uni.createFrom().item(gatewayResponse)); + + RerankingEGWClient rerankEGWClient = + new RerankingEGWClient( + ModelProvider.NVIDIA, + MODEL_CONFIG, + testConstants.TENANT, + "default", + rerankService, + Map.of(), + TESTING_COMMAND_NAME); + + Throwable result = + rerankEGWClient + .rerank(1, "apple", List.of("orange", "apple"), RERANK_CREDENTIALS) + .subscribe() + .withSubscriber(UniAssertSubscriber.create()) + .awaitFailure() + .getFailure(); + + assertThat(result) + .isInstanceOf(ServerException.class) + .satisfies( + failure -> { + ServerException exception = (ServerException) failure; + assertThat(exception.code) + .isEqualTo(ServerException.Code.UNEXPECTED_SERVER_ERROR.name()); + assertThat(exception.getMessage()).isEqualTo("Reranking provider timed out upstream"); + }); + } }