diff --git a/foundation-models/sap-rpt/src/main/java/com/sap/ai/sdk/foundationmodels/rpt/RptClient.java b/foundation-models/sap-rpt/src/main/java/com/sap/ai/sdk/foundationmodels/rpt/RptClient.java index e968aff49..8bca6da74 100644 --- a/foundation-models/sap-rpt/src/main/java/com/sap/ai/sdk/foundationmodels/rpt/RptClient.java +++ b/foundation-models/sap-rpt/src/main/java/com/sap/ai/sdk/foundationmodels/rpt/RptClient.java @@ -9,12 +9,15 @@ import com.sap.ai.sdk.core.JacksonConfiguration; import com.sap.ai.sdk.foundationmodels.rpt.generated.client.DefaultApi; import com.sap.ai.sdk.foundationmodels.rpt.generated.model.PredictRequestPayload; +import com.sap.ai.sdk.foundationmodels.rpt.generated.model.PredictRequestPayloadOneOf; +import com.sap.ai.sdk.foundationmodels.rpt.generated.model.PredictRequestPayloadOneOf1; import com.sap.ai.sdk.foundationmodels.rpt.generated.model.PredictResponsePayload; import com.sap.ai.sdk.foundationmodels.rpt.generated.model.PredictionConfig; import com.sap.cloud.sdk.cloudplatform.connectivity.Destination; import com.sap.cloud.sdk.services.openapi.apache.apiclient.ApiClient; import java.io.File; import java.util.Map; +import java.util.Set; import javax.annotation.Nonnull; import lombok.AccessLevel; import lombok.RequiredArgsConstructor; @@ -28,6 +31,15 @@ public class RptClient { @Nonnull private final DefaultApi api; @Nonnull private final DefaultApi apiWithGzipEncoding; + private final boolean usesOldModel; + + private static final Set PRE1_6MODELS = + Set.of( + "sap-rpt-1-large", + "sap-rpt-1-small", + "sap-rpt-1.1-preview", + "sap-rpt-1.5", + "sap-rpt-1.5-large"); /** * Creates a new RptClient for the specified foundation model. @@ -39,8 +51,9 @@ public class RptClient { @Nonnull public static RptClient forModel(@Nonnull final RptModel foundationModel) throws DeploymentResolutionException { + final var usesOldModel = PRE1_6MODELS.contains(foundationModel.name()); final var destination = new AiCoreService().getInferenceDestination().forModel(foundationModel); - return forDestination(destination); + return forDestination(destination, usesOldModel); } /** @@ -49,10 +62,12 @@ public static RptClient forModel(@Nonnull final RptModel foundationModel) * @param destination The destination to use. * @return A new instance of RptClient. */ - static RptClient forDestination(@Nonnull final Destination destination) { + static RptClient forDestination( + @Nonnull final Destination destination, final boolean usesOldModel) { final var apiClient = ApiClient.create(destination).withObjectMapper(getDefaultObjectMapper()); final var api = new DefaultApi(apiClient); - return new RptClient(api, api.withDefaultHeaders(Map.of("Content-Encoding", "gzip"))); + return new RptClient( + api, api.withDefaultHeaders(Map.of("Content-Encoding", "gzip")), usesOldModel); } /** @@ -73,13 +88,31 @@ static RptClient forDestination(@Nonnull final Destination destination) { * * @param requestBody The prediction request * @return prediction response from the RPT model + * @apiNote When used with a pre-1.6 model, the {@code contextMode} field of the embedded {@link + * com.sap.ai.sdk.foundationmodels.rpt.generated.model.PredictionConfig} is set to {@code + * null} on the passed-in object as a side effect. */ @Beta @Nonnull public PredictResponsePayload tableCompletion(@Nonnull final PredictRequestPayload requestBody) { + // contextMode has to be null for models < 1.6 + if (usesOldModel) { + configFrom(requestBody).setContextMode(null); + } return apiWithGzipEncoding.predict(requestBody); } + @Nonnull + private static PredictionConfig configFrom(@Nonnull final PredictRequestPayload requestBody) { + if (requestBody instanceof PredictRequestPayloadOneOf rb) { + return rb.getPredictionConfig(); + } else if (requestBody instanceof PredictRequestPayloadOneOf1 rb1) { + return rb1.getPredictionConfig(); + } + throw new IllegalArgumentException( + "Unsupported PredictRequestPayload type: " + requestBody.getClass().getName()); + } + /** * Make in-context predictions for specified target columns based on provided table data Parquet * file. @@ -100,12 +133,19 @@ public PredictResponsePayload tableCompletion(@Nonnull final PredictRequestPaylo * @param parquetFile Parquet file * @param predictionConfig The prediction configuration * @return prediction response from the RPT model + * @apiNote When used with a pre-1.6 model, the {@code contextMode} field of the passed-in {@link + * com.sap.ai.sdk.foundationmodels.rpt.generated.model.PredictionConfig} is set to {@code + * null} as a side effect. * @since 1.16.0 */ @Beta @Nonnull public PredictResponsePayload tableCompletion( @Nonnull final File parquetFile, @Nonnull final PredictionConfig predictionConfig) { + // contextMode has to be null for models < 1.6 + if (usesOldModel) { + predictionConfig.setContextMode(null); + } try { final var config = JacksonConfiguration.getDefaultObjectMapper().writeValueAsString(predictionConfig); diff --git a/foundation-models/sap-rpt/src/test/java/com/sap/ai/sdk/foundationmodels/rpt/RptClientTest.java b/foundation-models/sap-rpt/src/test/java/com/sap/ai/sdk/foundationmodels/rpt/RptClientTest.java index 0047144c2..4f0dc7be0 100644 --- a/foundation-models/sap-rpt/src/test/java/com/sap/ai/sdk/foundationmodels/rpt/RptClientTest.java +++ b/foundation-models/sap-rpt/src/test/java/com/sap/ai/sdk/foundationmodels/rpt/RptClientTest.java @@ -1,5 +1,9 @@ package com.sap.ai.sdk.foundationmodels.rpt; +import static com.github.tomakehurst.wiremock.client.WireMock.aResponse; +import static com.github.tomakehurst.wiremock.client.WireMock.post; +import static com.github.tomakehurst.wiremock.client.WireMock.stubFor; +import static com.github.tomakehurst.wiremock.client.WireMock.urlEqualTo; import static com.sap.ai.sdk.foundationmodels.rpt.generated.model.ColumnType.STRING; import static com.sap.ai.sdk.foundationmodels.rpt.generated.model.TargetColumnConfig.TaskTypeEnum.CLASSIFICATION; import static org.assertj.core.api.Assertions.assertThat; @@ -60,7 +64,7 @@ class RptClientTest { void setup(final WireMockRuntimeInfo server) { final DefaultHttpDestination destination = DefaultHttpDestination.builder(server.getHttpBaseUrl()).build(); - client = RptClient.forDestination(destination); + client = RptClient.forDestination(destination, false); ApacheHttpClient5Accessor.setHttpClientCache(ApacheHttpClient5Cache.DISABLED); } @@ -302,4 +306,83 @@ void testTableCompletionWithParquetThrowsIllegalArgumentException() { .hasMessageContaining("Failed to serialize PredictionConfig"); } } + + @Test + void testOldModelThrowsOnUnknownPayloadType() { + final var oldModelClient = + RptClient.forDestination(DefaultHttpDestination.builder("http://localhost").build(), true); + final var unknownPayload = mock(PredictRequestPayload.class); + + assertThatThrownBy(() -> oldModelClient.tableCompletion(unknownPayload)) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("Unsupported PredictRequestPayload type"); + } + + @Test + void testOldModelStripsContextModeFromRowWiseRequest(final WireMockRuntimeInfo server) { + stubFor(post(urlEqualTo("/predict")).willReturn(aResponse().withStatus(200).withBody("{}"))); + final var oldModelClient = + RptClient.forDestination( + DefaultHttpDestination.builder(server.getHttpBaseUrl()).build(), true); + + val config = + PredictionConfig.create() + .targetColumns(TARGET_COLUMN) + .contextMode(PredictionConfig.ContextModeEnum.DEEP); + val request = + PredictRequestPayloadOneOf.create() + .predictionConfig(config) + .rows(List.of()) + .indexColumn("ID") + .dataSchema(DATA_SCHEMA) + .parseDataTypes(true); + + assertThat(config.getContextMode()).isEqualTo(PredictionConfig.ContextModeEnum.DEEP); + oldModelClient.tableCompletion(request); + assertThat(config.getContextMode()).isNull(); + } + + @Test + void testOldModelStripsContextModeFromColumnWiseRequest(final WireMockRuntimeInfo server) { + stubFor(post(urlEqualTo("/predict")).willReturn(aResponse().withStatus(200).withBody("{}"))); + final var oldModelClient = + RptClient.forDestination( + DefaultHttpDestination.builder(server.getHttpBaseUrl()).build(), true); + + val config = + PredictionConfig.create() + .targetColumns(TARGET_COLUMN) + .contextMode(PredictionConfig.ContextModeEnum.DEFAULT); + val request = + PredictRequestPayloadOneOf1.create() + .predictionConfig(config) + .columns(Map.of()) + .indexColumn("ID") + .dataSchema(DATA_SCHEMA) + .parseDataTypes(true); + + assertThat(config.getContextMode()).isEqualTo(PredictionConfig.ContextModeEnum.DEFAULT); + oldModelClient.tableCompletion(request); + assertThat(config.getContextMode()).isNull(); + } + + @Test + void testOldModelStripsContextModeFromParquetRequest(final WireMockRuntimeInfo server) { + stubFor( + post(urlEqualTo("/predict_parquet")) + .willReturn(aResponse().withStatus(200).withBody("{}"))); + final var oldModelClient = + RptClient.forDestination( + DefaultHttpDestination.builder(server.getHttpBaseUrl()).build(), true); + + val parquetFile = Path.of("src/test/resources/rpt/test-data.parquet").toFile(); + val predictionConfig = + PredictionConfig.create() + .targetColumns(TARGET_COLUMN) + .contextMode(PredictionConfig.ContextModeEnum.DEEP); + + assertThat(predictionConfig.getContextMode()).isEqualTo(PredictionConfig.ContextModeEnum.DEEP); + oldModelClient.tableCompletion(parquetFile, predictionConfig); + assertThat(predictionConfig.getContextMode()).isNull(); + } }