diff --git a/docs/changelog/142738.yaml b/docs/changelog/142738.yaml new file mode 100644 index 0000000000000..225fe96644765 --- /dev/null +++ b/docs/changelog/142738.yaml @@ -0,0 +1,6 @@ +pr: 142738 +summary: Added service settings update logic for Alibaba Cloud Search provider in the Inference Plugin +area: Inference +type: enhancement +issues: + - 122356 diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchService.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchService.java index db206af40a3d9..5772176e4c0e6 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchService.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchService.java @@ -239,7 +239,7 @@ public AlibabaCloudSearchModel buildModelFromConfigAndSecrets(ModelConfiguration config.getInferenceEntityId(), config.getTaskType(), config.getService(), - ConfigurationParseContext.PERSISTENT + ConfigurationParseContext.REQUEST ).createFromModelConfigurationsAndSecrets(config, secrets); } diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchServiceSettings.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchServiceSettings.java index 6583fe92589ac..a4637dcf74ef7 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchServiceSettings.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchServiceSettings.java @@ -38,25 +38,24 @@ public class AlibabaCloudSearchServiceSettings extends FilteredXContentObject public static final String HOST = "host"; public static final String WORKSPACE_NAME = "workspace"; public static final String HTTP_SCHEMA_NAME = "http_schema"; + private static final Set VALID_SCHEMAS = Set.of("https", "http"); private static final RateLimitSettings DEFAULT_RATE_LIMIT_SETTINGS = new RateLimitSettings(1_000); - public static AlibabaCloudSearchServiceSettings fromMap(Map map, ConfigurationParseContext context) { - ValidationException validationException = new ValidationException(); + public static AlibabaCloudSearchServiceSettings fromMap( + Map map, + ConfigurationParseContext context, + ValidationException validationException + ) { - String modelId = extractRequiredString(map, SERVICE_ID, ModelConfigurations.SERVICE_SETTINGS, validationException); - String host = extractRequiredString(map, HOST, ModelConfigurations.SERVICE_SETTINGS, validationException); + var serviceId = extractRequiredString(map, SERVICE_ID, ModelConfigurations.SERVICE_SETTINGS, validationException); + var host = extractRequiredString(map, HOST, ModelConfigurations.SERVICE_SETTINGS, validationException); var workspaceName = extractRequiredString(map, WORKSPACE_NAME, ModelConfigurations.SERVICE_SETTINGS, validationException); var httpSchema = extractOptionalString(map, HTTP_SCHEMA_NAME, ModelConfigurations.SERVICE_SETTINGS, validationException); - if (httpSchema != null) { - var validSchemas = Set.of("https", "http"); - if (validSchemas.contains(httpSchema) == false) { - validationException.addValidationError("Invalid value for [http_schema]. Must be one of [https, http]"); - } - } + validateHttpSchema(httpSchema, validationException); - RateLimitSettings rateLimitSettings = RateLimitSettings.of( + var rateLimitSettings = RateLimitSettings.of( map, DEFAULT_RATE_LIMIT_SETTINGS, validationException, @@ -64,11 +63,13 @@ public static AlibabaCloudSearchServiceSettings fromMap(Map map, context ); - if (validationException.validationErrors().isEmpty() == false) { - throw validationException; - } + return new AlibabaCloudSearchServiceSettings(serviceId, host, workspaceName, httpSchema, rateLimitSettings); + } - return new AlibabaCloudSearchServiceSettings(modelId, host, workspaceName, httpSchema, rateLimitSettings); + static void validateHttpSchema(String httpSchema, ValidationException validationException) { + if (httpSchema != null && VALID_SCHEMAS.contains(httpSchema) == false) { + validationException.addValidationError("Invalid value for [http_schema]. Must be one of [https, http]"); + } } private final String serviceId; @@ -92,11 +93,11 @@ public AlibabaCloudSearchServiceSettings( } public AlibabaCloudSearchServiceSettings(StreamInput in) throws IOException { - serviceId = in.readString(); - host = in.readString(); - workspaceName = in.readString(); - httpSchema = in.readOptionalString(); - rateLimitSettings = new RateLimitSettings(in); + this.serviceId = in.readString(); + this.host = in.readString(); + this.workspaceName = in.readString(); + this.httpSchema = in.readOptionalString(); + this.rateLimitSettings = new RateLimitSettings(in); } @Override @@ -104,6 +105,36 @@ public String modelId() { return serviceId; } + public AlibabaCloudSearchServiceSettings updateServiceSettings( + Map serviceSettings, + ValidationException validationException + ) { + var extractedHttpSchema = extractOptionalString( + serviceSettings, + HTTP_SCHEMA_NAME, + ModelConfigurations.SERVICE_SETTINGS, + validationException + ); + + validateHttpSchema(extractedHttpSchema, validationException); + + var extractedRateLimitSettings = RateLimitSettings.of( + serviceSettings, + this.rateLimitSettings, + validationException, + AlibabaCloudSearchService.NAME, + ConfigurationParseContext.REQUEST + ); + + return new AlibabaCloudSearchServiceSettings( + this.serviceId, + this.host, + this.workspaceName, + extractedHttpSchema != null ? extractedHttpSchema : this.httpSchema, + extractedRateLimitSettings + ); + } + public String getHost() { return host; } diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/completion/AlibabaCloudSearchCompletionServiceSettings.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/completion/AlibabaCloudSearchCompletionServiceSettings.java index 8c579b0adb04b..12e6cd64a4ab9 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/completion/AlibabaCloudSearchCompletionServiceSettings.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/completion/AlibabaCloudSearchCompletionServiceSettings.java @@ -25,11 +25,11 @@ public class AlibabaCloudSearchCompletionServiceSettings implements ServiceSetti public static final String NAME = "alibabacloud_search_completion_service_settings"; public static AlibabaCloudSearchCompletionServiceSettings fromMap(Map map, ConfigurationParseContext context) { - ValidationException validationException = new ValidationException(); - var commonServiceSettings = AlibabaCloudSearchServiceSettings.fromMap(map, context); - if (validationException.validationErrors().isEmpty() == false) { - throw validationException; - } + var validationException = new ValidationException(); + + var commonServiceSettings = AlibabaCloudSearchServiceSettings.fromMap(map, context, validationException); + + validationException.throwIfValidationErrorsExist(); return new AlibabaCloudSearchCompletionServiceSettings(commonServiceSettings); } @@ -41,7 +41,7 @@ public AlibabaCloudSearchCompletionServiceSettings(AlibabaCloudSearchServiceSett } public AlibabaCloudSearchCompletionServiceSettings(StreamInput in) throws IOException { - commonSettings = new AlibabaCloudSearchServiceSettings(in); + this.commonSettings = new AlibabaCloudSearchServiceSettings(in); } public AlibabaCloudSearchServiceSettings getCommonSettings() { @@ -53,6 +53,17 @@ public String modelId() { return commonSettings.modelId(); } + @Override + public AlibabaCloudSearchCompletionServiceSettings updateServiceSettings(Map serviceSettings) { + var validationException = new ValidationException(); + + var updatedCommonServiceSettings = commonSettings.updateServiceSettings(serviceSettings, validationException); + + validationException.throwIfValidationErrorsExist(); + + return new AlibabaCloudSearchCompletionServiceSettings(updatedCommonServiceSettings); + } + @Override public String getWriteableName() { return NAME; diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/embeddings/AlibabaCloudSearchEmbeddingsServiceSettings.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/embeddings/AlibabaCloudSearchEmbeddingsServiceSettings.java index 363e022b3b06b..372ba1035d7e1 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/embeddings/AlibabaCloudSearchEmbeddingsServiceSettings.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/embeddings/AlibabaCloudSearchEmbeddingsServiceSettings.java @@ -28,25 +28,29 @@ import static org.elasticsearch.xpack.inference.services.ServiceFields.DIMENSIONS; import static org.elasticsearch.xpack.inference.services.ServiceFields.MAX_INPUT_TOKENS; import static org.elasticsearch.xpack.inference.services.ServiceFields.SIMILARITY; +import static org.elasticsearch.xpack.inference.services.ServiceUtils.extractOptionalPositiveInteger; import static org.elasticsearch.xpack.inference.services.ServiceUtils.extractSimilarity; -import static org.elasticsearch.xpack.inference.services.ServiceUtils.removeAsType; public class AlibabaCloudSearchEmbeddingsServiceSettings implements ServiceSettings { public static final String NAME = "alibabacloud_search_embeddings_service_settings"; public static AlibabaCloudSearchEmbeddingsServiceSettings fromMap(Map map, ConfigurationParseContext context) { - ValidationException validationException = new ValidationException(); - var commonServiceSettings = AlibabaCloudSearchServiceSettings.fromMap(map, context); + var validationException = new ValidationException(); - SimilarityMeasure similarity = extractSimilarity(map, ModelConfigurations.SERVICE_SETTINGS, validationException); - Integer dims = removeAsType(map, DIMENSIONS, Integer.class); - Integer maxInputTokens = removeAsType(map, MAX_INPUT_TOKENS, Integer.class); + var commonServiceSettings = AlibabaCloudSearchServiceSettings.fromMap(map, context, validationException); - if (validationException.validationErrors().isEmpty() == false) { - throw validationException; - } + var similarity = extractSimilarity(map, ModelConfigurations.SERVICE_SETTINGS, validationException); + var dimensions = extractOptionalPositiveInteger(map, DIMENSIONS, ModelConfigurations.SERVICE_SETTINGS, validationException); + var maxInputTokens = extractOptionalPositiveInteger( + map, + MAX_INPUT_TOKENS, + ModelConfigurations.SERVICE_SETTINGS, + validationException + ); + + validationException.throwIfValidationErrorsExist(); - return new AlibabaCloudSearchEmbeddingsServiceSettings(commonServiceSettings, similarity, dims, maxInputTokens); + return new AlibabaCloudSearchEmbeddingsServiceSettings(commonServiceSettings, similarity, dimensions, maxInputTokens); } private final AlibabaCloudSearchServiceSettings commonSettings; @@ -67,10 +71,10 @@ public AlibabaCloudSearchEmbeddingsServiceSettings( } public AlibabaCloudSearchEmbeddingsServiceSettings(StreamInput in) throws IOException { - commonSettings = new AlibabaCloudSearchServiceSettings(in); - similarity = in.readOptionalEnum(SimilarityMeasure.class); - dimensions = in.readOptionalVInt(); - maxInputTokens = in.readOptionalVInt(); + this.commonSettings = new AlibabaCloudSearchServiceSettings(in); + this.similarity = in.readOptionalEnum(SimilarityMeasure.class); + this.dimensions = in.readOptionalVInt(); + this.maxInputTokens = in.readOptionalVInt(); } public AlibabaCloudSearchServiceSettings getCommonSettings() { @@ -105,6 +109,28 @@ public String modelId() { return commonSettings.modelId(); } + @Override + public AlibabaCloudSearchEmbeddingsServiceSettings updateServiceSettings(Map serviceSettings) { + var validationException = new ValidationException(); + var commonServiceSettings = commonSettings.updateServiceSettings(serviceSettings, validationException); + + var extractedMaxInputTokens = extractOptionalPositiveInteger( + serviceSettings, + MAX_INPUT_TOKENS, + ModelConfigurations.SERVICE_SETTINGS, + validationException + ); + + validationException.throwIfValidationErrorsExist(); + + return new AlibabaCloudSearchEmbeddingsServiceSettings( + commonServiceSettings, + this.similarity, + this.dimensions, + extractedMaxInputTokens != null ? extractedMaxInputTokens : this.maxInputTokens + ); + } + @Override public String getWriteableName() { return NAME; diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/rerank/AlibabaCloudSearchRerankServiceSettings.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/rerank/AlibabaCloudSearchRerankServiceSettings.java index 556bd0ad3500d..6ef2c7a677f22 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/rerank/AlibabaCloudSearchRerankServiceSettings.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/rerank/AlibabaCloudSearchRerankServiceSettings.java @@ -25,11 +25,11 @@ public class AlibabaCloudSearchRerankServiceSettings implements ServiceSettings public static final String NAME = "alibabacloud_search_rerank_service_settings"; public static AlibabaCloudSearchRerankServiceSettings fromMap(Map map, ConfigurationParseContext context) { - ValidationException validationException = new ValidationException(); - var commonServiceSettings = AlibabaCloudSearchServiceSettings.fromMap(map, context); - if (validationException.validationErrors().isEmpty() == false) { - throw validationException; - } + var validationException = new ValidationException(); + + var commonServiceSettings = AlibabaCloudSearchServiceSettings.fromMap(map, context, validationException); + + validationException.throwIfValidationErrorsExist(); return new AlibabaCloudSearchRerankServiceSettings(commonServiceSettings); } @@ -53,6 +53,17 @@ public String modelId() { return commonSettings.modelId(); } + @Override + public AlibabaCloudSearchRerankServiceSettings updateServiceSettings(Map serviceSettings) { + var validationException = new ValidationException(); + + var updatedCommonServiceSettings = commonSettings.updateServiceSettings(serviceSettings, validationException); + + validationException.throwIfValidationErrorsExist(); + + return new AlibabaCloudSearchRerankServiceSettings(updatedCommonServiceSettings); + } + @Override public String getWriteableName() { return NAME; diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/sparse/AlibabaCloudSearchSparseServiceSettings.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/sparse/AlibabaCloudSearchSparseServiceSettings.java index 963dd17cba327..9b47c41352986 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/sparse/AlibabaCloudSearchSparseServiceSettings.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/sparse/AlibabaCloudSearchSparseServiceSettings.java @@ -25,11 +25,11 @@ public class AlibabaCloudSearchSparseServiceSettings implements ServiceSettings public static final String NAME = "alibabacloud_search_sparse_embeddings_service_settings"; public static AlibabaCloudSearchSparseServiceSettings fromMap(Map map, ConfigurationParseContext context) { - ValidationException validationException = new ValidationException(); - var commonServiceSettings = AlibabaCloudSearchServiceSettings.fromMap(map, context); - if (validationException.validationErrors().isEmpty() == false) { - throw validationException; - } + var validationException = new ValidationException(); + + var commonServiceSettings = AlibabaCloudSearchServiceSettings.fromMap(map, context, validationException); + + validationException.throwIfValidationErrorsExist(); return new AlibabaCloudSearchSparseServiceSettings(commonServiceSettings); } @@ -41,7 +41,7 @@ public AlibabaCloudSearchSparseServiceSettings(AlibabaCloudSearchServiceSettings } public AlibabaCloudSearchSparseServiceSettings(StreamInput in) throws IOException { - commonSettings = new AlibabaCloudSearchServiceSettings(in); + this.commonSettings = new AlibabaCloudSearchServiceSettings(in); } public AlibabaCloudSearchServiceSettings getCommonSettings() { @@ -53,6 +53,17 @@ public String modelId() { return commonSettings.modelId(); } + @Override + public AlibabaCloudSearchSparseServiceSettings updateServiceSettings(Map serviceSettings) { + var validationException = new ValidationException(); + + var updatedCommonServiceSettings = commonSettings.updateServiceSettings(serviceSettings, validationException); + + validationException.throwIfValidationErrorsExist(); + + return new AlibabaCloudSearchSparseServiceSettings(updatedCommonServiceSettings); + } + @Override public String getWriteableName() { return NAME; diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/ai21/Ai21ServiceTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/ai21/Ai21ServiceTests.java index 04f3c02330a48..e5d8d32dafcc6 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/ai21/Ai21ServiceTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/ai21/Ai21ServiceTests.java @@ -241,7 +241,7 @@ public void testParseRequestConfig_CreatesChatCompletionsModel() throws IOExcept service.parseRequestConfig( "id", TaskType.CHAT_COMPLETION, - getRequestConfigMap(getServiceSettingsMap(model), getSecretSettingsMap(secret)), + getRequestConfigMap(getServiceSettingsMap(model, null), getSecretSettingsMap(secret)), modelVerificationListener ); } @@ -272,7 +272,7 @@ public void testParseRequestConfig_ThrowsException_WithoutModelId() throws IOExc service.parseRequestConfig( "id", TaskType.CHAT_COMPLETION, - getRequestConfigMap(getServiceSettingsMap(null), getSecretSettingsMap(secret)), + getRequestConfigMap(getServiceSettingsMap(null, null), getSecretSettingsMap(secret)), modelVerificationListener ); } diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/ai21/completion/Ai21ChatCompletionServiceSettingsTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/ai21/completion/Ai21ChatCompletionServiceSettingsTests.java index 45f33ffca9168..f5e69e60d91a9 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/ai21/completion/Ai21ChatCompletionServiceSettingsTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/ai21/completion/Ai21ChatCompletionServiceSettingsTests.java @@ -12,6 +12,7 @@ import org.elasticsearch.common.ValidationException; import org.elasticsearch.common.io.stream.Writeable; import org.elasticsearch.common.xcontent.XContentHelper; +import org.elasticsearch.core.Nullable; import org.elasticsearch.xcontent.XContentBuilder; import org.elasticsearch.xcontent.XContentFactory; import org.elasticsearch.xcontent.XContentType; @@ -36,44 +37,31 @@ public class Ai21ChatCompletionServiceSettingsTests extends AbstractBWCWireSeria private static final int INITIAL_TEST_RATE_LIMIT = 30; public void testUpdateServiceSettings_AllFields_OnlyMutableFieldsAreUpdated() { - var serviceSettings = new Ai21ChatCompletionServiceSettings(INITIAL_TEST_MODEL_ID, new RateLimitSettings(INITIAL_TEST_RATE_LIMIT)) - .updateServiceSettings( - new HashMap<>( - Map.of( - ServiceFields.MODEL_ID, - TEST_MODEL_ID, - RateLimitSettings.FIELD_NAME, - new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, TEST_RATE_LIMIT)) - ) - ) - ); + var originalServiceSettings = new Ai21ChatCompletionServiceSettings( + INITIAL_TEST_MODEL_ID, + new RateLimitSettings(INITIAL_TEST_RATE_LIMIT) + ); + var updatedServiceSettings = originalServiceSettings.updateServiceSettings(getServiceSettingsMap(TEST_MODEL_ID, TEST_RATE_LIMIT)); assertThat( - serviceSettings, + updatedServiceSettings, is(new Ai21ChatCompletionServiceSettings(INITIAL_TEST_MODEL_ID, new RateLimitSettings(TEST_RATE_LIMIT))) ); } - public void testUpdateServiceSettings_EmptyMap_Success() { - var serviceSettings = new Ai21ChatCompletionServiceSettings(INITIAL_TEST_MODEL_ID, new RateLimitSettings(INITIAL_TEST_RATE_LIMIT)) - .updateServiceSettings(new HashMap<>()); - - assertThat( - serviceSettings, - is(new Ai21ChatCompletionServiceSettings(INITIAL_TEST_MODEL_ID, new RateLimitSettings(INITIAL_TEST_RATE_LIMIT))) + public void testUpdateServiceSettings_EmptyMap_DoesNotChangeSettings() { + var originalServiceSettings = new Ai21ChatCompletionServiceSettings( + INITIAL_TEST_MODEL_ID, + new RateLimitSettings(INITIAL_TEST_RATE_LIMIT) ); + var updatedServiceSettings = originalServiceSettings.updateServiceSettings(new HashMap<>()); + + assertThat(updatedServiceSettings, is(originalServiceSettings)); } public void testFromMap_AllFields_Success() { var serviceSettings = Ai21ChatCompletionServiceSettings.fromMap( - new HashMap<>( - Map.of( - ServiceFields.MODEL_ID, - TEST_MODEL_ID, - RateLimitSettings.FIELD_NAME, - new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, TEST_RATE_LIMIT)) - ) - ), + getServiceSettingsMap(TEST_MODEL_ID, TEST_RATE_LIMIT), ConfigurationParseContext.PERSISTENT ); @@ -84,12 +72,7 @@ public void testFromMap_MissingModelId_ThrowsException() { var thrownException = expectThrows( ValidationException.class, () -> Ai21ChatCompletionServiceSettings.fromMap( - new HashMap<>( - Map.of( - RateLimitSettings.FIELD_NAME, - new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, TEST_RATE_LIMIT)) - ) - ), + getServiceSettingsMap(null, TEST_RATE_LIMIT), ConfigurationParseContext.PERSISTENT ) ); @@ -102,23 +85,16 @@ public void testFromMap_MissingModelId_ThrowsException() { public void testFromMap_MissingRateLimit_Success() { var serviceSettings = Ai21ChatCompletionServiceSettings.fromMap( - new HashMap<>(Map.of(ServiceFields.MODEL_ID, TEST_MODEL_ID)), + getServiceSettingsMap(TEST_MODEL_ID, null), ConfigurationParseContext.PERSISTENT ); - assertThat(serviceSettings, is(new Ai21ChatCompletionServiceSettings(TEST_MODEL_ID, null))); + assertThat(serviceSettings, is(new Ai21ChatCompletionServiceSettings(TEST_MODEL_ID, new RateLimitSettings(200)))); } public void testToXContent_WritesAllValues() throws IOException { var serviceSettings = Ai21ChatCompletionServiceSettings.fromMap( - new HashMap<>( - Map.of( - ServiceFields.MODEL_ID, - TEST_MODEL_ID, - RateLimitSettings.FIELD_NAME, - new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, TEST_RATE_LIMIT)) - ) - ), + getServiceSettingsMap(TEST_MODEL_ID, TEST_RATE_LIMIT), ConfigurationParseContext.PERSISTENT ); @@ -196,10 +172,15 @@ private static Ai21ChatCompletionServiceSettings createRandom() { return new Ai21ChatCompletionServiceSettings(modelId, RateLimitSettingsTests.createRandom()); } - public static Map getServiceSettingsMap(String model) { + public static Map getServiceSettingsMap(@Nullable String modelId, @Nullable Integer rateLimit) { var map = new HashMap(); - map.put(ServiceFields.MODEL_ID, model); + if (modelId != null) { + map.put(ServiceFields.MODEL_ID, modelId); + } + if (rateLimit != null) { + map.put(RateLimitSettings.FIELD_NAME, new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, rateLimit))); + } return map; } diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchServiceSettingsTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchServiceSettingsTests.java index d7965a38c845b..4da789177a5fa 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchServiceSettingsTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchServiceSettingsTests.java @@ -7,97 +7,207 @@ package org.elasticsearch.xpack.inference.services.alibabacloudsearch; +import org.elasticsearch.TransportVersion; import org.elasticsearch.common.Strings; +import org.elasticsearch.common.ValidationException; import org.elasticsearch.common.io.stream.Writeable; -import org.elasticsearch.test.AbstractWireSerializingTestCase; import org.elasticsearch.xcontent.XContentBuilder; import org.elasticsearch.xcontent.XContentFactory; import org.elasticsearch.xcontent.XContentType; +import org.elasticsearch.xpack.core.ml.AbstractBWCWireSerializationTestCase; import org.elasticsearch.xpack.inference.services.settings.RateLimitSettings; import org.elasticsearch.xpack.inference.services.settings.RateLimitSettingsTests; -import org.hamcrest.MatcherAssert; import java.io.IOException; -import java.net.URISyntaxException; import java.util.HashMap; import java.util.Map; +import java.util.Objects; +import static org.hamcrest.Matchers.emptyCollectionOf; import static org.hamcrest.Matchers.is; -public class AlibabaCloudSearchServiceSettingsTests extends AbstractWireSerializingTestCase { +public class AlibabaCloudSearchServiceSettingsTests extends AbstractBWCWireSerializationTestCase { + private static final String TEST_SERVICE_ID = "test-service-id"; + private static final String INITIAL_TEST_SERVICE_ID = "initial-test-service-id"; + private static final String TEST_HOST = "test-host"; + private static final String INITIAL_TEST_HOST = "initial-test-host"; + private static final String TEST_WORKSPACE_NAME = "test-workspace-name"; + private static final String INITIAL_TEST_WORKSPACE_NAME = "initial-test-workspace-name"; + private static final String TEST_HTTP_SCHEMA = "https"; + private static final String INITIAL_TEST_HTTP_SCHEMA = "http"; + private static final int TEST_RATE_LIMIT = 20; + private static final int INITIAL_TEST_RATE_LIMIT = 30; + /** * The created settings can have a url set to null. */ public static AlibabaCloudSearchServiceSettings createRandom() { var model = randomAlphaOfLength(15); - String host = randomAlphaOfLength(15); - String workspaceName = randomAlphaOfLength(10); - String httpSchema = "https"; + var host = randomAlphaOfLength(15); + var workspaceName = randomAlphaOfLength(10); + var httpSchema = randomBoolean() ? "https" : "http"; return new AlibabaCloudSearchServiceSettings(model, host, workspaceName, httpSchema, RateLimitSettingsTests.createRandom()); } - public void testFromMap() throws URISyntaxException { - var model = "model"; - var host = "host"; - var workspaceName = "default"; - var httpSchema = "https"; + public void testUpdateServiceSettings_AllFields_OnlyMutableFieldsAreUpdated() { + var originalServiceSettings = new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + INITIAL_TEST_HTTP_SCHEMA, + new RateLimitSettings(INITIAL_TEST_RATE_LIMIT) + ); + var updatedServiceSettings = originalServiceSettings.updateServiceSettings( + new HashMap<>( + Map.of( + AlibabaCloudSearchServiceSettings.SERVICE_ID, + TEST_SERVICE_ID, + AlibabaCloudSearchServiceSettings.HOST, + TEST_HOST, + AlibabaCloudSearchServiceSettings.WORKSPACE_NAME, + TEST_WORKSPACE_NAME, + AlibabaCloudSearchServiceSettings.HTTP_SCHEMA_NAME, + TEST_HTTP_SCHEMA, + RateLimitSettings.FIELD_NAME, + new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, TEST_RATE_LIMIT)) + ) + ), + new ValidationException() + ); + + assertThat( + updatedServiceSettings, + is( + new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + TEST_HTTP_SCHEMA, + new RateLimitSettings(TEST_RATE_LIMIT) + ) + ) + ); + } + + public void testUpdateServiceSettings_EmptyMap_DoesNotChangeSettings() { + var originalServiceSettings = new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + INITIAL_TEST_HTTP_SCHEMA, + new RateLimitSettings(INITIAL_TEST_RATE_LIMIT) + ); + var updatedServiceSettings = originalServiceSettings.updateServiceSettings(new HashMap<>(), new ValidationException()); + + assertThat(updatedServiceSettings, is(originalServiceSettings)); + } + + public void testFromMap_Success() { var serviceSettings = AlibabaCloudSearchServiceSettings.fromMap( new HashMap<>( Map.of( AlibabaCloudSearchServiceSettings.SERVICE_ID, - model, + TEST_SERVICE_ID, AlibabaCloudSearchServiceSettings.HOST, - host, + TEST_HOST, AlibabaCloudSearchServiceSettings.WORKSPACE_NAME, - workspaceName, + TEST_WORKSPACE_NAME, AlibabaCloudSearchServiceSettings.HTTP_SCHEMA_NAME, - httpSchema + TEST_HTTP_SCHEMA ) ), - null + null, + new ValidationException() ); - MatcherAssert.assertThat(serviceSettings, is(new AlibabaCloudSearchServiceSettings(model, host, workspaceName, httpSchema, null))); + assertThat( + serviceSettings, + is(new AlibabaCloudSearchServiceSettings(TEST_SERVICE_ID, TEST_HOST, TEST_WORKSPACE_NAME, TEST_HTTP_SCHEMA, null)) + ); } public void testFromMap_WithRateLimit() { - var model = "model"; - var host = "host"; - var workspaceName = "default"; - var httpSchema = "https"; var serviceSettings = AlibabaCloudSearchServiceSettings.fromMap( new HashMap<>( Map.of( AlibabaCloudSearchServiceSettings.SERVICE_ID, - model, + TEST_SERVICE_ID, AlibabaCloudSearchServiceSettings.HOST, - host, + TEST_HOST, AlibabaCloudSearchServiceSettings.WORKSPACE_NAME, - workspaceName, + TEST_WORKSPACE_NAME, AlibabaCloudSearchServiceSettings.HTTP_SCHEMA_NAME, - httpSchema, + TEST_HTTP_SCHEMA, RateLimitSettings.FIELD_NAME, - new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, 3)) + new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, TEST_RATE_LIMIT)) ) ), - null + null, + new ValidationException() ); - MatcherAssert.assertThat( + assertThat( serviceSettings, - is(new AlibabaCloudSearchServiceSettings(model, host, workspaceName, httpSchema, new RateLimitSettings(3))) + is( + new AlibabaCloudSearchServiceSettings( + TEST_SERVICE_ID, + TEST_HOST, + TEST_WORKSPACE_NAME, + TEST_HTTP_SCHEMA, + new RateLimitSettings(TEST_RATE_LIMIT) + ) + ) ); } public void testXContent() throws IOException { - var entity = new AlibabaCloudSearchServiceSettings("model_id_name", "host_name", "workspace_name", null, null); + var entity = new AlibabaCloudSearchServiceSettings( + TEST_SERVICE_ID, + TEST_HOST, + TEST_WORKSPACE_NAME, + TEST_HTTP_SCHEMA, + new RateLimitSettings(TEST_RATE_LIMIT) + ); XContentBuilder builder = XContentFactory.contentBuilder(XContentType.JSON); entity.toXContent(builder, null); String xContentResult = Strings.toString(builder); - assertThat(xContentResult, is(""" - {"service_id":"model_id_name","host":"host_name","workspace":"workspace_name","rate_limit":{"requests_per_minute":1000}}""")); + assertThat( + xContentResult, + is( + Strings.format( + """ + {"service_id":"%s","host":"%s","workspace":"%s","http_schema":"%s","rate_limit":{"requests_per_minute":%d}}""", + TEST_SERVICE_ID, + TEST_HOST, + TEST_WORKSPACE_NAME, + TEST_HTTP_SCHEMA, + TEST_RATE_LIMIT + ) + ) + ); + } + + public void testValidateHttpSchema_InvalidSchema_AddsValidationError() { + var validationException = new ValidationException(); + AlibabaCloudSearchServiceSettings.validateHttpSchema("invalid-http-schema", validationException); + assertThat( + validationException.getMessage(), + is("Validation Failed: 1: Invalid value for [http_schema]. Must be one of [https, http];") + ); + } + + public void testValidateHttpSchema_HttpsSchema_Success() { + var validationException = new ValidationException(); + AlibabaCloudSearchServiceSettings.validateHttpSchema("https", validationException); + assertThat(validationException.validationErrors(), is(emptyCollectionOf(String.class))); + } + + public void testValidateHttpSchema_HttpSchema_Success() { + var validationException = new ValidationException(); + AlibabaCloudSearchServiceSettings.validateHttpSchema("http", validationException); + assertThat(validationException.validationErrors(), is(emptyCollectionOf(String.class))); } @Override @@ -112,7 +222,20 @@ protected AlibabaCloudSearchServiceSettings createTestInstance() { @Override protected AlibabaCloudSearchServiceSettings mutateInstance(AlibabaCloudSearchServiceSettings instance) throws IOException { - return null; + var serviceId = instance.modelId(); + var host = instance.getHost(); + var workspaceName = instance.getWorkspaceName(); + var httpSchema = instance.getHttpSchema(); + var rateLimitSettings = instance.rateLimitSettings(); + + switch (between(0, 3)) { + case 0 -> serviceId = randomValueOtherThan(serviceId, () -> randomAlphaOfLength(8)); + case 1 -> host = randomValueOtherThan(host, () -> randomAlphaOfLength(8)); + case 2 -> workspaceName = randomValueOtherThan(workspaceName, () -> randomAlphaOfLength(8)); + case 3 -> httpSchema = Objects.equals(httpSchema, "http") ? "https" : "http"; + default -> throw new AssertionError("Illegal randomisation branch"); + } + return new AlibabaCloudSearchServiceSettings(serviceId, host, workspaceName, httpSchema, rateLimitSettings); } public static Map getServiceSettingsMap(String serviceId, String host, String workspaceName) { @@ -122,4 +245,12 @@ public static Map getServiceSettingsMap(String serviceId, String map.put(AlibabaCloudSearchServiceSettings.WORKSPACE_NAME, workspaceName); return map; } + + @Override + protected AlibabaCloudSearchServiceSettings mutateInstanceForVersion( + AlibabaCloudSearchServiceSettings instance, + TransportVersion version + ) { + return instance; + } } diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchServiceTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchServiceTests.java index ad310505ded21..49089553f258b 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchServiceTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/AlibabaCloudSearchServiceTests.java @@ -706,6 +706,7 @@ private AlibabaCloudSearchModel createEmbeddingsModel( secretSettingsMap, null ) { + @Override public ExecutableAction accept(AlibabaCloudSearchActionVisitor visitor, Map taskSettings) { return (inferenceInputs, timeout, listener) -> { DenseEmbeddingFloatResults results = new DenseEmbeddingFloatResults( @@ -737,6 +738,7 @@ private AlibabaCloudSearchModel createSparseEmbeddingsModel( secretSettingsMap, null ) { + @Override public ExecutableAction accept(AlibabaCloudSearchActionVisitor visitor, Map taskSettings) { return (inferenceInputs, timeout, listener) -> { listener.onResponse(SparseEmbeddingResultsTests.createRandomResults(2, 1)); @@ -821,16 +823,11 @@ public void testBuildModelFromConfigAndSecrets_UnsupportedTaskType() throws IOEx thrownException.getMessage(), is( Strings.format( - """ - Failed to parse stored model [%s] for [%s] service, error: [The [%s] service does not support task type [%s]]. \ - Please delete and add the service again""", - INFERENCE_ENTITY_ID_VALUE, - AlibabaCloudSearchService.NAME, + "The [%s] service does not support task type [%s]", AlibabaCloudSearchService.NAME, TaskType.CHAT_COMPLETION ) ) - ); } } diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/completion/AlibabaCloudSearchCompletionServiceSettingsTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/completion/AlibabaCloudSearchCompletionServiceSettingsTests.java index 167a3f2688292..c9fca5e7e45c4 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/completion/AlibabaCloudSearchCompletionServiceSettingsTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/completion/AlibabaCloudSearchCompletionServiceSettingsTests.java @@ -7,11 +7,12 @@ package org.elasticsearch.xpack.inference.services.alibabacloudsearch.completion; +import org.elasticsearch.TransportVersion; import org.elasticsearch.common.io.stream.Writeable; -import org.elasticsearch.test.AbstractWireSerializingTestCase; +import org.elasticsearch.xpack.core.ml.AbstractBWCWireSerializationTestCase; import org.elasticsearch.xpack.inference.services.alibabacloudsearch.AlibabaCloudSearchServiceSettings; import org.elasticsearch.xpack.inference.services.alibabacloudsearch.AlibabaCloudSearchServiceSettingsTests; -import org.hamcrest.MatcherAssert; +import org.elasticsearch.xpack.inference.services.settings.RateLimitSettings; import java.io.IOException; import java.util.HashMap; @@ -19,39 +20,113 @@ import static org.hamcrest.Matchers.is; -public class AlibabaCloudSearchCompletionServiceSettingsTests extends AbstractWireSerializingTestCase< +public class AlibabaCloudSearchCompletionServiceSettingsTests extends AbstractBWCWireSerializationTestCase< AlibabaCloudSearchCompletionServiceSettings> { + + private static final String TEST_SERVICE_ID = "test-service-id"; + private static final String INITIAL_TEST_SERVICE_ID = "initial-test-service-id"; + private static final String TEST_HOST = "test-host"; + private static final String INITIAL_TEST_HOST = "initial-test-host"; + private static final String TEST_WORKSPACE_NAME = "test-workspace-name"; + private static final String INITIAL_TEST_WORKSPACE_NAME = "initial-test-workspace-name"; + private static final String TEST_HTTP_SCHEMA = "https"; + private static final String INITIAL_TEST_HTTP_SCHEMA = "http"; + private static final int TEST_RATE_LIMIT = 20; + private static final int INITIAL_TEST_RATE_LIMIT = 30; + public static AlibabaCloudSearchCompletionServiceSettings createRandom() { var commonSettings = AlibabaCloudSearchServiceSettingsTests.createRandom(); return new AlibabaCloudSearchCompletionServiceSettings(commonSettings); } - public void testFromMap() { - var model = "model"; - var host = "host"; - var workspaceName = "default"; - var httpSchema = "https"; + public void testUpdateServiceSettings_AllFields_OnlyMutableFieldsAreUpdated() { + var originalServiceSettings = new AlibabaCloudSearchCompletionServiceSettings( + new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + INITIAL_TEST_HTTP_SCHEMA, + new RateLimitSettings(INITIAL_TEST_RATE_LIMIT) + ) + ); + var updatedServiceSettings = originalServiceSettings.updateServiceSettings( + new HashMap<>( + Map.of( + AlibabaCloudSearchServiceSettings.HOST, + TEST_HOST, + AlibabaCloudSearchServiceSettings.SERVICE_ID, + TEST_SERVICE_ID, + AlibabaCloudSearchServiceSettings.WORKSPACE_NAME, + TEST_WORKSPACE_NAME, + AlibabaCloudSearchServiceSettings.HTTP_SCHEMA_NAME, + TEST_HTTP_SCHEMA, + RateLimitSettings.FIELD_NAME, + new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, TEST_RATE_LIMIT)) + ) + ) + ); + + assertThat( + updatedServiceSettings, + is( + new AlibabaCloudSearchCompletionServiceSettings( + new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + TEST_HTTP_SCHEMA, + new RateLimitSettings(TEST_RATE_LIMIT) + ) + ) + ) + ); + } + + public void testUpdateServiceSettings_EmptyMap_DoesNotChangeSettings() { + var originalServiceSettings = new AlibabaCloudSearchCompletionServiceSettings( + new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + INITIAL_TEST_HTTP_SCHEMA, + new RateLimitSettings(INITIAL_TEST_RATE_LIMIT) + ) + ); + var updatedServiceSettings = originalServiceSettings.updateServiceSettings(new HashMap<>()); + + assertThat(updatedServiceSettings, is(originalServiceSettings)); + } + + public void testFromMap_Success() { var serviceSettings = AlibabaCloudSearchCompletionServiceSettings.fromMap( new HashMap<>( Map.of( AlibabaCloudSearchServiceSettings.HOST, - host, + TEST_HOST, AlibabaCloudSearchServiceSettings.SERVICE_ID, - model, + TEST_SERVICE_ID, AlibabaCloudSearchServiceSettings.WORKSPACE_NAME, - workspaceName, + TEST_WORKSPACE_NAME, AlibabaCloudSearchServiceSettings.HTTP_SCHEMA_NAME, - httpSchema + TEST_HTTP_SCHEMA, + RateLimitSettings.FIELD_NAME, + new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, TEST_RATE_LIMIT)) ) ), null ); - MatcherAssert.assertThat( + assertThat( serviceSettings, is( new AlibabaCloudSearchCompletionServiceSettings( - new AlibabaCloudSearchServiceSettings(model, host, workspaceName, httpSchema, null) + new AlibabaCloudSearchServiceSettings( + TEST_SERVICE_ID, + TEST_HOST, + TEST_WORKSPACE_NAME, + TEST_HTTP_SCHEMA, + new RateLimitSettings(TEST_RATE_LIMIT) + ) ) ) ); @@ -70,7 +145,9 @@ protected AlibabaCloudSearchCompletionServiceSettings createTestInstance() { @Override protected AlibabaCloudSearchCompletionServiceSettings mutateInstance(AlibabaCloudSearchCompletionServiceSettings instance) throws IOException { - return createRandom(); + return new AlibabaCloudSearchCompletionServiceSettings( + randomValueOtherThan(instance.getCommonSettings(), AlibabaCloudSearchServiceSettingsTests::createRandom) + ); } public static Map getServiceSettingsMap(String serviceId, String host, String workspaceName) { @@ -80,4 +157,12 @@ public static Map getServiceSettingsMap(String serviceId, String map.put(AlibabaCloudSearchServiceSettings.WORKSPACE_NAME, workspaceName); return map; } + + @Override + protected AlibabaCloudSearchCompletionServiceSettings mutateInstanceForVersion( + AlibabaCloudSearchCompletionServiceSettings instance, + TransportVersion version + ) { + return instance; + } } diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/embeddings/AlibabaCloudSearchEmbeddingsServiceSettingsTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/embeddings/AlibabaCloudSearchEmbeddingsServiceSettingsTests.java index 815e6d0311195..c855693838ae2 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/embeddings/AlibabaCloudSearchEmbeddingsServiceSettingsTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/embeddings/AlibabaCloudSearchEmbeddingsServiceSettingsTests.java @@ -7,13 +7,14 @@ package org.elasticsearch.xpack.inference.services.alibabacloudsearch.embeddings; +import org.elasticsearch.TransportVersion; import org.elasticsearch.common.io.stream.Writeable; import org.elasticsearch.inference.SimilarityMeasure; -import org.elasticsearch.test.AbstractWireSerializingTestCase; +import org.elasticsearch.xpack.core.ml.AbstractBWCWireSerializationTestCase; import org.elasticsearch.xpack.inference.services.ServiceFields; import org.elasticsearch.xpack.inference.services.alibabacloudsearch.AlibabaCloudSearchServiceSettings; import org.elasticsearch.xpack.inference.services.alibabacloudsearch.AlibabaCloudSearchServiceSettingsTests; -import org.hamcrest.MatcherAssert; +import org.elasticsearch.xpack.inference.services.settings.RateLimitSettings; import java.io.IOException; import java.util.HashMap; @@ -21,54 +22,149 @@ import static org.hamcrest.Matchers.is; -public class AlibabaCloudSearchEmbeddingsServiceSettingsTests extends AbstractWireSerializingTestCase< +public class AlibabaCloudSearchEmbeddingsServiceSettingsTests extends AbstractBWCWireSerializationTestCase< AlibabaCloudSearchEmbeddingsServiceSettings> { + + private static final SimilarityMeasure TEST_SIMILARITY_MEASURE = SimilarityMeasure.DOT_PRODUCT; + private static final SimilarityMeasure INITIAL_TEST_SIMILARITY_MEASURE = SimilarityMeasure.COSINE; + private static final int TEST_DIMENSIONS = 1536; + private static final int INITIAL_TEST_DIMENSIONS = 1024; + private static final int TEST_MAX_INPUT_TOKENS = 512; + private static final int INITIAL_TEST_MAX_INPUT_TOKENS = 256; + + private static final String TEST_SERVICE_ID = "test-service-id"; + private static final String INITIAL_TEST_SERVICE_ID = "initial-test-service-id"; + private static final String TEST_HOST = "test-host"; + private static final String INITIAL_TEST_HOST = "initial-test-host"; + private static final String TEST_WORKSPACE_NAME = "test-workspace-name"; + private static final String INITIAL_TEST_WORKSPACE_NAME = "initial-test-workspace-name"; + private static final String TEST_HTTP_SCHEMA = "https"; + private static final String INITIAL_TEST_HTTP_SCHEMA = "http"; + private static final int TEST_RATE_LIMIT = 20; + private static final int INITIAL_TEST_RATE_LIMIT = 30; + public static AlibabaCloudSearchEmbeddingsServiceSettings createRandom() { var commonSettings = AlibabaCloudSearchServiceSettingsTests.createRandom(); - var similarity = SimilarityMeasure.DOT_PRODUCT; - var dims = 1536; - var maxInputTokens = 512; - return new AlibabaCloudSearchEmbeddingsServiceSettings(commonSettings, similarity, dims, maxInputTokens); + return new AlibabaCloudSearchEmbeddingsServiceSettings( + commonSettings, + randomFrom(SimilarityMeasure.values()), + randomInt(TEST_DIMENSIONS), + randomInt(TEST_MAX_INPUT_TOKENS) + ); + } + + public void testUpdateServiceSettings_AllFields_OnlyMutableFieldsAreUpdated() { + var originalServiceSettings = new AlibabaCloudSearchEmbeddingsServiceSettings( + new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + INITIAL_TEST_HTTP_SCHEMA, + new RateLimitSettings(INITIAL_TEST_RATE_LIMIT) + ), + INITIAL_TEST_SIMILARITY_MEASURE, + INITIAL_TEST_DIMENSIONS, + INITIAL_TEST_MAX_INPUT_TOKENS + ); + var updatedServiceSettings = originalServiceSettings.updateServiceSettings( + new HashMap<>( + Map.of( + ServiceFields.SIMILARITY, + TEST_SIMILARITY_MEASURE.toString(), + ServiceFields.DIMENSIONS, + TEST_DIMENSIONS, + ServiceFields.MAX_INPUT_TOKENS, + TEST_MAX_INPUT_TOKENS, + AlibabaCloudSearchServiceSettings.HOST, + TEST_HOST, + AlibabaCloudSearchServiceSettings.SERVICE_ID, + TEST_SERVICE_ID, + AlibabaCloudSearchServiceSettings.WORKSPACE_NAME, + TEST_WORKSPACE_NAME, + AlibabaCloudSearchServiceSettings.HTTP_SCHEMA_NAME, + TEST_HTTP_SCHEMA, + RateLimitSettings.FIELD_NAME, + new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, TEST_RATE_LIMIT)) + ) + ) + ); + + assertThat( + updatedServiceSettings, + is( + new AlibabaCloudSearchEmbeddingsServiceSettings( + new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + TEST_HTTP_SCHEMA, + new RateLimitSettings(TEST_RATE_LIMIT) + ), + INITIAL_TEST_SIMILARITY_MEASURE, + INITIAL_TEST_DIMENSIONS, + TEST_MAX_INPUT_TOKENS + ) + ) + ); + } + + public void testUpdateServiceSettings_EmptyMap_DoesNotChangeSettings() { + var originalServiceSettings = new AlibabaCloudSearchEmbeddingsServiceSettings( + new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + INITIAL_TEST_HTTP_SCHEMA, + new RateLimitSettings(INITIAL_TEST_RATE_LIMIT) + ), + INITIAL_TEST_SIMILARITY_MEASURE, + INITIAL_TEST_DIMENSIONS, + INITIAL_TEST_MAX_INPUT_TOKENS + ); + var updatedServiceSettings = originalServiceSettings.updateServiceSettings(new HashMap<>()); + + assertThat(updatedServiceSettings, is(originalServiceSettings)); } - public void testFromMap() { - var similarity = SimilarityMeasure.DOT_PRODUCT.toString(); - var dims = 1536; - var maxInputTokens = 512; - var model = "model"; - var host = "host"; - var workspaceName = "default"; - var httpSchema = "https"; + public void testFromMap_Success() { var serviceSettings = AlibabaCloudSearchEmbeddingsServiceSettings.fromMap( new HashMap<>( Map.of( ServiceFields.SIMILARITY, - similarity, + TEST_SIMILARITY_MEASURE.toString(), ServiceFields.DIMENSIONS, - dims, + TEST_DIMENSIONS, ServiceFields.MAX_INPUT_TOKENS, - maxInputTokens, + TEST_MAX_INPUT_TOKENS, AlibabaCloudSearchServiceSettings.HOST, - host, + TEST_HOST, AlibabaCloudSearchServiceSettings.SERVICE_ID, - model, + TEST_SERVICE_ID, AlibabaCloudSearchServiceSettings.WORKSPACE_NAME, - workspaceName, + TEST_WORKSPACE_NAME, AlibabaCloudSearchServiceSettings.HTTP_SCHEMA_NAME, - httpSchema + TEST_HTTP_SCHEMA, + RateLimitSettings.FIELD_NAME, + new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, TEST_RATE_LIMIT)) ) ), null ); - MatcherAssert.assertThat( + assertThat( serviceSettings, is( new AlibabaCloudSearchEmbeddingsServiceSettings( - new AlibabaCloudSearchServiceSettings(model, host, workspaceName, httpSchema, null), - SimilarityMeasure.DOT_PRODUCT, - dims, - maxInputTokens + new AlibabaCloudSearchServiceSettings( + TEST_SERVICE_ID, + TEST_HOST, + TEST_WORKSPACE_NAME, + TEST_HTTP_SCHEMA, + new RateLimitSettings(TEST_RATE_LIMIT) + ), + TEST_SIMILARITY_MEASURE, + TEST_DIMENSIONS, + TEST_MAX_INPUT_TOKENS ) ) ); @@ -87,10 +183,34 @@ protected AlibabaCloudSearchEmbeddingsServiceSettings createTestInstance() { @Override protected AlibabaCloudSearchEmbeddingsServiceSettings mutateInstance(AlibabaCloudSearchEmbeddingsServiceSettings instance) throws IOException { - return null; + var commonSettings = instance.getCommonSettings(); + var similarity = instance.similarity(); + var dimensions = instance.dimensions(); + var maxInputTokens = instance.getMaxInputTokens(); + + switch (between(0, 3)) { + case 0 -> commonSettings = randomValueOtherThan( + instance.getCommonSettings(), + AlibabaCloudSearchServiceSettingsTests::createRandom + ); + case 1 -> similarity = randomValueOtherThan(similarity, () -> randomFrom(SimilarityMeasure.values())); + case 2 -> dimensions = randomValueOtherThan(dimensions, () -> randomIntBetween(32, 256)); + case 3 -> maxInputTokens = randomValueOtherThan(maxInputTokens, () -> randomIntBetween(16, 1024)); + default -> throw new AssertionError("Illegal randomisation branch"); + } + return new AlibabaCloudSearchEmbeddingsServiceSettings(commonSettings, similarity, dimensions, maxInputTokens); + } public static Map getServiceSettingsMap(String serviceId, String host, String workspaceName) { return AlibabaCloudSearchServiceSettingsTests.getServiceSettingsMap(serviceId, host, workspaceName); } + + @Override + protected AlibabaCloudSearchEmbeddingsServiceSettings mutateInstanceForVersion( + AlibabaCloudSearchEmbeddingsServiceSettings instance, + TransportVersion version + ) { + return instance; + } } diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/rerank/AlibabaCloudSearchRerankServiceSettingsTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/rerank/AlibabaCloudSearchRerankServiceSettingsTests.java new file mode 100644 index 0000000000000..6a8ca97216587 --- /dev/null +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/rerank/AlibabaCloudSearchRerankServiceSettingsTests.java @@ -0,0 +1,159 @@ +/* + * Copyright Elasticsearch B.V. and/or licensed to Elasticsearch B.V. under one + * or more contributor license agreements. Licensed under the Elastic License + * 2.0; you may not use this file except in compliance with the Elastic License + * 2.0. + */ + +package org.elasticsearch.xpack.inference.services.alibabacloudsearch.rerank; + +import org.elasticsearch.TransportVersion; +import org.elasticsearch.common.io.stream.Writeable; +import org.elasticsearch.xpack.core.ml.AbstractBWCWireSerializationTestCase; +import org.elasticsearch.xpack.inference.services.alibabacloudsearch.AlibabaCloudSearchServiceSettings; +import org.elasticsearch.xpack.inference.services.alibabacloudsearch.AlibabaCloudSearchServiceSettingsTests; +import org.elasticsearch.xpack.inference.services.settings.RateLimitSettings; + +import java.io.IOException; +import java.util.HashMap; +import java.util.Map; + +import static org.hamcrest.Matchers.is; + +public class AlibabaCloudSearchRerankServiceSettingsTests extends AbstractBWCWireSerializationTestCase< + AlibabaCloudSearchRerankServiceSettings> { + + private static final String TEST_SERVICE_ID = "test-service-id"; + private static final String INITIAL_TEST_SERVICE_ID = "initial-test-service-id"; + private static final String TEST_HOST = "test-host"; + private static final String INITIAL_TEST_HOST = "initial-test-host"; + private static final String TEST_WORKSPACE_NAME = "test-workspace-name"; + private static final String INITIAL_TEST_WORKSPACE_NAME = "initial-test-workspace-name"; + private static final String TEST_HTTP_SCHEMA = "https"; + private static final String INITIAL_TEST_HTTP_SCHEMA = "http"; + private static final int TEST_RATE_LIMIT = 20; + private static final int INITIAL_TEST_RATE_LIMIT = 30; + + public static AlibabaCloudSearchRerankServiceSettings createRandom() { + var commonSettings = AlibabaCloudSearchServiceSettingsTests.createRandom(); + return new AlibabaCloudSearchRerankServiceSettings(commonSettings); + } + + public void testUpdateServiceSettings_AllFields_OnlyMutableFieldsAreUpdated() { + var originalServiceSettings = new AlibabaCloudSearchRerankServiceSettings( + new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + INITIAL_TEST_HTTP_SCHEMA, + new RateLimitSettings(INITIAL_TEST_RATE_LIMIT) + ) + ); + var updatedServiceSettings = originalServiceSettings.updateServiceSettings( + new HashMap<>( + Map.of( + AlibabaCloudSearchServiceSettings.HOST, + TEST_HOST, + AlibabaCloudSearchServiceSettings.SERVICE_ID, + TEST_SERVICE_ID, + AlibabaCloudSearchServiceSettings.WORKSPACE_NAME, + TEST_WORKSPACE_NAME, + AlibabaCloudSearchServiceSettings.HTTP_SCHEMA_NAME, + TEST_HTTP_SCHEMA, + RateLimitSettings.FIELD_NAME, + new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, TEST_RATE_LIMIT)) + ) + ) + ); + + assertThat( + updatedServiceSettings, + is( + new AlibabaCloudSearchRerankServiceSettings( + new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + TEST_HTTP_SCHEMA, + new RateLimitSettings(TEST_RATE_LIMIT) + ) + ) + ) + ); + } + + public void testUpdateServiceSettings_EmptyMap_DoesNotChangeSettings() { + var originalServiceSettings = new AlibabaCloudSearchRerankServiceSettings( + new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + INITIAL_TEST_HTTP_SCHEMA, + new RateLimitSettings(INITIAL_TEST_RATE_LIMIT) + ) + ); + var updatedServiceSettings = originalServiceSettings.updateServiceSettings(new HashMap<>()); + + assertThat(updatedServiceSettings, is(originalServiceSettings)); + } + + public void testFromMap_Success() { + var serviceSettings = AlibabaCloudSearchRerankServiceSettings.fromMap( + new HashMap<>( + Map.of( + AlibabaCloudSearchServiceSettings.HOST, + TEST_HOST, + AlibabaCloudSearchServiceSettings.SERVICE_ID, + TEST_SERVICE_ID, + AlibabaCloudSearchServiceSettings.WORKSPACE_NAME, + TEST_WORKSPACE_NAME, + AlibabaCloudSearchServiceSettings.HTTP_SCHEMA_NAME, + TEST_HTTP_SCHEMA, + RateLimitSettings.FIELD_NAME, + new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, TEST_RATE_LIMIT)) + ) + ), + null + ); + + assertThat( + serviceSettings, + is( + new AlibabaCloudSearchRerankServiceSettings( + new AlibabaCloudSearchServiceSettings( + TEST_SERVICE_ID, + TEST_HOST, + TEST_WORKSPACE_NAME, + TEST_HTTP_SCHEMA, + new RateLimitSettings(TEST_RATE_LIMIT) + ) + ) + ) + ); + } + + @Override + protected AlibabaCloudSearchRerankServiceSettings mutateInstanceForVersion( + AlibabaCloudSearchRerankServiceSettings instance, + TransportVersion version + ) { + return instance; + } + + @Override + protected Writeable.Reader instanceReader() { + return AlibabaCloudSearchRerankServiceSettings::new; + } + + @Override + protected AlibabaCloudSearchRerankServiceSettings createTestInstance() { + return createRandom(); + } + + @Override + protected AlibabaCloudSearchRerankServiceSettings mutateInstance(AlibabaCloudSearchRerankServiceSettings instance) throws IOException { + return new AlibabaCloudSearchRerankServiceSettings( + randomValueOtherThan(instance.getCommonSettings(), AlibabaCloudSearchServiceSettingsTests::createRandom) + ); + } +} diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/sparse/AlibabaCloudSearchSparseServiceSettingsTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/sparse/AlibabaCloudSearchSparseServiceSettingsTests.java index 8dc635a52f06f..c260bb9785acb 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/sparse/AlibabaCloudSearchSparseServiceSettingsTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/alibabacloudsearch/sparse/AlibabaCloudSearchSparseServiceSettingsTests.java @@ -7,11 +7,12 @@ package org.elasticsearch.xpack.inference.services.alibabacloudsearch.sparse; +import org.elasticsearch.TransportVersion; import org.elasticsearch.common.io.stream.Writeable; -import org.elasticsearch.test.AbstractWireSerializingTestCase; +import org.elasticsearch.xpack.core.ml.AbstractBWCWireSerializationTestCase; import org.elasticsearch.xpack.inference.services.alibabacloudsearch.AlibabaCloudSearchServiceSettings; import org.elasticsearch.xpack.inference.services.alibabacloudsearch.AlibabaCloudSearchServiceSettingsTests; -import org.hamcrest.MatcherAssert; +import org.elasticsearch.xpack.inference.services.settings.RateLimitSettings; import java.io.IOException; import java.util.HashMap; @@ -19,38 +20,113 @@ import static org.hamcrest.Matchers.is; -public class AlibabaCloudSearchSparseServiceSettingsTests extends AbstractWireSerializingTestCase { +public class AlibabaCloudSearchSparseServiceSettingsTests extends AbstractBWCWireSerializationTestCase< + AlibabaCloudSearchSparseServiceSettings> { + + private static final String TEST_SERVICE_ID = "test-service-id"; + private static final String INITIAL_TEST_SERVICE_ID = "initial-test-service-id"; + private static final String TEST_HOST = "test-host"; + private static final String INITIAL_TEST_HOST = "initial-test-host"; + private static final String TEST_WORKSPACE_NAME = "test-workspace-name"; + private static final String INITIAL_TEST_WORKSPACE_NAME = "initial-test-workspace-name"; + private static final String TEST_HTTP_SCHEMA = "https"; + private static final String INITIAL_TEST_HTTP_SCHEMA = "http"; + private static final int TEST_RATE_LIMIT = 20; + private static final int INITIAL_TEST_RATE_LIMIT = 30; + public static AlibabaCloudSearchSparseServiceSettings createRandom() { var commonSettings = AlibabaCloudSearchServiceSettingsTests.createRandom(); return new AlibabaCloudSearchSparseServiceSettings(commonSettings); } - public void testFromMap() { - var model = "model"; - var host = "host"; - var workspaceName = "default"; - var httpSchema = "https"; + public void testUpdateServiceSettings_AllFields_OnlyMutableFieldsAreUpdated() { + var originalServiceSettings = new AlibabaCloudSearchSparseServiceSettings( + new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + INITIAL_TEST_HTTP_SCHEMA, + new RateLimitSettings(INITIAL_TEST_RATE_LIMIT) + ) + ); + var updatedServiceSettings = originalServiceSettings.updateServiceSettings( + new HashMap<>( + Map.of( + AlibabaCloudSearchServiceSettings.HOST, + TEST_HOST, + AlibabaCloudSearchServiceSettings.SERVICE_ID, + TEST_SERVICE_ID, + AlibabaCloudSearchServiceSettings.WORKSPACE_NAME, + TEST_WORKSPACE_NAME, + AlibabaCloudSearchServiceSettings.HTTP_SCHEMA_NAME, + TEST_HTTP_SCHEMA, + RateLimitSettings.FIELD_NAME, + new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, TEST_RATE_LIMIT)) + ) + ) + ); + + assertThat( + updatedServiceSettings, + is( + new AlibabaCloudSearchSparseServiceSettings( + new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + TEST_HTTP_SCHEMA, + new RateLimitSettings(TEST_RATE_LIMIT) + ) + ) + ) + ); + } + + public void testUpdateServiceSettings_EmptyMap_DoesNotChangeSettings() { + var originalServiceSettings = new AlibabaCloudSearchSparseServiceSettings( + new AlibabaCloudSearchServiceSettings( + INITIAL_TEST_SERVICE_ID, + INITIAL_TEST_HOST, + INITIAL_TEST_WORKSPACE_NAME, + INITIAL_TEST_HTTP_SCHEMA, + new RateLimitSettings(INITIAL_TEST_RATE_LIMIT) + ) + ); + var updatedServiceSettings = originalServiceSettings.updateServiceSettings(new HashMap<>()); + + assertThat(updatedServiceSettings, is(originalServiceSettings)); + } + + public void testFromMap_Success() { var serviceSettings = AlibabaCloudSearchSparseServiceSettings.fromMap( new HashMap<>( Map.of( AlibabaCloudSearchServiceSettings.HOST, - host, + TEST_HOST, AlibabaCloudSearchServiceSettings.SERVICE_ID, - model, + TEST_SERVICE_ID, AlibabaCloudSearchServiceSettings.WORKSPACE_NAME, - workspaceName, + TEST_WORKSPACE_NAME, AlibabaCloudSearchServiceSettings.HTTP_SCHEMA_NAME, - httpSchema + TEST_HTTP_SCHEMA, + RateLimitSettings.FIELD_NAME, + new HashMap<>(Map.of(RateLimitSettings.REQUESTS_PER_MINUTE_FIELD, TEST_RATE_LIMIT)) ) ), null ); - MatcherAssert.assertThat( + assertThat( serviceSettings, is( new AlibabaCloudSearchSparseServiceSettings( - new AlibabaCloudSearchServiceSettings(model, host, workspaceName, httpSchema, null) + new AlibabaCloudSearchServiceSettings( + TEST_SERVICE_ID, + TEST_HOST, + TEST_WORKSPACE_NAME, + TEST_HTTP_SCHEMA, + new RateLimitSettings(TEST_RATE_LIMIT) + ) ) ) ); @@ -68,10 +144,20 @@ protected AlibabaCloudSearchSparseServiceSettings createTestInstance() { @Override protected AlibabaCloudSearchSparseServiceSettings mutateInstance(AlibabaCloudSearchSparseServiceSettings instance) throws IOException { - return null; + return new AlibabaCloudSearchSparseServiceSettings( + randomValueOtherThan(instance.getCommonSettings(), AlibabaCloudSearchServiceSettingsTests::createRandom) + ); } public static Map getServiceSettingsMap(String serviceId, String host, String workspaceName) { return AlibabaCloudSearchServiceSettingsTests.getServiceSettingsMap(serviceId, host, workspaceName); } + + @Override + protected AlibabaCloudSearchSparseServiceSettings mutateInstanceForVersion( + AlibabaCloudSearchSparseServiceSettings instance, + TransportVersion version + ) { + return instance; + } }