diff --git a/server/src/main/java/org/elasticsearch/TransportVersions.java b/server/src/main/java/org/elasticsearch/TransportVersions.java index 6317651beb8df..520bb2640ef8a 100644 --- a/server/src/main/java/org/elasticsearch/TransportVersions.java +++ b/server/src/main/java/org/elasticsearch/TransportVersions.java @@ -165,6 +165,7 @@ static TransportVersion def(int id) { public static final TransportVersion REQUIRE_DATA_STREAM_ADDED = def(8_578_00_0); public static final TransportVersion ML_INFERENCE_COHERE_EMBEDDINGS_ADDED = def(8_579_00_0); public static final TransportVersion DESIRED_NODE_VERSION_OPTIONAL_STRING = def(8_580_00_0); + public static final TransportVersion ML_INFERENCE_REQUEST_INPUT_TYPE_UNSPECIFIED_ADDED = def(8_581_00_0); /* * STOP! READ THIS FIRST! No, really, diff --git a/server/src/main/java/org/elasticsearch/inference/InferenceService.java b/server/src/main/java/org/elasticsearch/inference/InferenceService.java index 235de51d22572..fdeb32de33877 100644 --- a/server/src/main/java/org/elasticsearch/inference/InferenceService.java +++ b/server/src/main/java/org/elasticsearch/inference/InferenceService.java @@ -78,7 +78,13 @@ default void init(Client client) {} * @param taskSettings Settings in the request to override the model's defaults * @param listener Inference result listener */ - void infer(Model model, List input, Map taskSettings, ActionListener listener); + void infer( + Model model, + List input, + Map taskSettings, + InputType inputType, + ActionListener listener + ); /** * Start or prepare the model for use. diff --git a/server/src/main/java/org/elasticsearch/inference/InputType.java b/server/src/main/java/org/elasticsearch/inference/InputType.java index ffc67995c1dda..19f28601409ac 100644 --- a/server/src/main/java/org/elasticsearch/inference/InputType.java +++ b/server/src/main/java/org/elasticsearch/inference/InputType.java @@ -15,9 +15,8 @@ */ public enum InputType { INGEST, - SEARCH; - - public static String NAME = "input_type"; + SEARCH, + UNSPECIFIED; @Override public String toString() { diff --git a/x-pack/plugin/core/src/main/java/org/elasticsearch/xpack/core/inference/action/InferenceAction.java b/x-pack/plugin/core/src/main/java/org/elasticsearch/xpack/core/inference/action/InferenceAction.java index 1fc477927d7b7..2ddba3446d79a 100644 --- a/x-pack/plugin/core/src/main/java/org/elasticsearch/xpack/core/inference/action/InferenceAction.java +++ b/x-pack/plugin/core/src/main/java/org/elasticsearch/xpack/core/inference/action/InferenceAction.java @@ -8,6 +8,7 @@ package org.elasticsearch.xpack.core.inference.action; import org.elasticsearch.ElasticsearchStatusException; +import org.elasticsearch.TransportVersion; import org.elasticsearch.TransportVersions; import org.elasticsearch.action.ActionRequest; import org.elasticsearch.action.ActionRequestValidationException; @@ -60,6 +61,8 @@ public static Request parseRequest(String inferenceEntityId, String taskType, XC Request.Builder builder = PARSER.apply(parser, null); builder.setInferenceEntityId(inferenceEntityId); builder.setTaskType(taskType); + // For rest requests we won't know what the input type is + builder.setInputType(InputType.UNSPECIFIED); return builder.build(); } @@ -96,7 +99,7 @@ public Request(StreamInput in) throws IOException { if (in.getTransportVersion().onOrAfter(TransportVersions.ML_INFERENCE_REQUEST_INPUT_TYPE_ADDED)) { this.inputType = in.readEnum(InputType.class); } else { - this.inputType = InputType.INGEST; + this.inputType = InputType.UNSPECIFIED; } } @@ -146,11 +149,22 @@ public void writeTo(StreamOutput out) throws IOException { out.writeString(input.get(0)); } out.writeGenericMap(taskSettings); + // in version ML_INFERENCE_REQUEST_INPUT_TYPE_ADDED the input type enum was added, so we only want to write the enum if we're + // at that version or later if (out.getTransportVersion().onOrAfter(TransportVersions.ML_INFERENCE_REQUEST_INPUT_TYPE_ADDED)) { - out.writeEnum(inputType); + out.writeEnum(getInputTypeToWrite(out.getTransportVersion())); } } + private InputType getInputTypeToWrite(TransportVersion version) { + // in version ML_INFERENCE_REQUEST_INPUT_TYPE_UNSPECIFIED_ADDED the UNSPECIFIED value was added, so if we're before that + // version other nodes won't know about it, so set it to INGEST instead + if (version.before(TransportVersions.ML_INFERENCE_REQUEST_INPUT_TYPE_UNSPECIFIED_ADDED) && inputType == InputType.UNSPECIFIED) { + return InputType.INGEST; + } + return inputType; + } + @Override public boolean equals(Object o) { if (this == o) return true; @@ -173,6 +187,7 @@ public static class Builder { private TaskType taskType; private String inferenceEntityId; private List input; + private InputType inputType = InputType.UNSPECIFIED; private Map taskSettings = Map.of(); private Builder() {} @@ -197,13 +212,18 @@ public Builder setInput(List input) { return this; } + public Builder setInputType(InputType inputType) { + this.inputType = inputType; + return this; + } + public Builder setTaskSettings(Map taskSettings) { this.taskSettings = taskSettings; return this; } public Request build() { - return new Request(taskType, inferenceEntityId, input, taskSettings, InputType.INGEST); + return new Request(taskType, inferenceEntityId, input, taskSettings, inputType); } } } diff --git a/x-pack/plugin/inference/qa/test-service-plugin/src/main/java/org/elasticsearch/xpack/inference/mock/TestInferenceServiceExtension.java b/x-pack/plugin/inference/qa/test-service-plugin/src/main/java/org/elasticsearch/xpack/inference/mock/TestInferenceServiceExtension.java index eee6f68c20ff7..5ffb4b5df08cc 100644 --- a/x-pack/plugin/inference/qa/test-service-plugin/src/main/java/org/elasticsearch/xpack/inference/mock/TestInferenceServiceExtension.java +++ b/x-pack/plugin/inference/qa/test-service-plugin/src/main/java/org/elasticsearch/xpack/inference/mock/TestInferenceServiceExtension.java @@ -16,6 +16,7 @@ import org.elasticsearch.inference.InferenceService; import org.elasticsearch.inference.InferenceServiceExtension; import org.elasticsearch.inference.InferenceServiceResults; +import org.elasticsearch.inference.InputType; import org.elasticsearch.inference.Model; import org.elasticsearch.inference.ModelConfigurations; import org.elasticsearch.inference.ModelSecrets; @@ -123,11 +124,11 @@ public void infer( Model model, List input, Map taskSettings, + InputType inputType, ActionListener listener ) { switch (model.getConfigurations().getTaskType()) { - case ANY -> listener.onResponse(makeResults(input)); - case SPARSE_EMBEDDING -> listener.onResponse(makeResults(input)); + case ANY, SPARSE_EMBEDDING -> listener.onResponse(makeResults(input)); default -> listener.onFailure( new ElasticsearchStatusException( TaskType.unsupportedTaskTypeErrorMsg(model.getConfigurations().getTaskType(), name()), diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/action/TransportInferenceAction.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/action/TransportInferenceAction.java index b9cc14977b87e..fb3974fc12e8b 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/action/TransportInferenceAction.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/action/TransportInferenceAction.java @@ -92,6 +92,7 @@ private void inferOnService( model, request.getInput(), request.getTaskSettings(), + request.getInputType(), listener.delegateFailureAndWrap((l, inferenceResults) -> l.onResponse(new InferenceAction.Response(inferenceResults))) ); } diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/action/cohere/CohereActionCreator.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/action/cohere/CohereActionCreator.java index 8c9d70f0a7323..0fb5ca9283fae 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/action/cohere/CohereActionCreator.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/action/cohere/CohereActionCreator.java @@ -7,6 +7,7 @@ package org.elasticsearch.xpack.inference.external.action.cohere; +import org.elasticsearch.inference.InputType; import org.elasticsearch.xpack.inference.external.action.ExecutableAction; import org.elasticsearch.xpack.inference.external.http.sender.Sender; import org.elasticsearch.xpack.inference.services.ServiceComponents; @@ -28,8 +29,8 @@ public CohereActionCreator(Sender sender, ServiceComponents serviceComponents) { } @Override - public ExecutableAction create(CohereEmbeddingsModel model, Map taskSettings) { - var overriddenModel = model.overrideWith(taskSettings); + public ExecutableAction create(CohereEmbeddingsModel model, Map taskSettings, InputType inputType) { + var overriddenModel = CohereEmbeddingsModel.of(model, taskSettings, inputType); return new CohereEmbeddingsAction(sender, overriddenModel, serviceComponents); } diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/action/cohere/CohereActionVisitor.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/action/cohere/CohereActionVisitor.java index 1500d48e3c201..cc732e7ab8dc5 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/action/cohere/CohereActionVisitor.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/action/cohere/CohereActionVisitor.java @@ -7,11 +7,12 @@ package org.elasticsearch.xpack.inference.external.action.cohere; +import org.elasticsearch.inference.InputType; import org.elasticsearch.xpack.inference.external.action.ExecutableAction; import org.elasticsearch.xpack.inference.services.cohere.embeddings.CohereEmbeddingsModel; import java.util.Map; public interface CohereActionVisitor { - ExecutableAction create(CohereEmbeddingsModel model, Map taskSettings); + ExecutableAction create(CohereEmbeddingsModel model, Map taskSettings, InputType inputType); } diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/action/openai/OpenAiActionCreator.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/action/openai/OpenAiActionCreator.java index 6c423760d0b35..94583c634fb26 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/action/openai/OpenAiActionCreator.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/action/openai/OpenAiActionCreator.java @@ -29,7 +29,7 @@ public OpenAiActionCreator(Sender sender, ServiceComponents serviceComponents) { @Override public ExecutableAction create(OpenAiEmbeddingsModel model, Map taskSettings) { - var overriddenModel = model.overrideWith(taskSettings); + var overriddenModel = OpenAiEmbeddingsModel.of(model, taskSettings); return new OpenAiEmbeddingsAction(sender, overriddenModel, serviceComponents); } diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/request/cohere/CohereEmbeddingsRequestEntity.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/request/cohere/CohereEmbeddingsRequestEntity.java index a0b5444ee45e4..9e34af5ed6385 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/request/cohere/CohereEmbeddingsRequestEntity.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/external/request/cohere/CohereEmbeddingsRequestEntity.java @@ -20,6 +20,8 @@ import java.util.List; import java.util.Objects; +import static org.elasticsearch.xpack.inference.services.cohere.embeddings.CohereEmbeddingsTaskSettings.invalidInputTypeMessage; + public record CohereEmbeddingsRequestEntity( List input, CohereEmbeddingsTaskSettings taskSettings, @@ -29,14 +31,6 @@ public record CohereEmbeddingsRequestEntity( private static final String SEARCH_DOCUMENT = "search_document"; private static final String SEARCH_QUERY = "search_query"; - /** - * Maps the {@link InputType} to the expected value for cohere for the input_type field in the request using the enum's ordinal. - * The order of these entries is important and needs to match the order in the enum - */ - private static final String[] INPUT_TYPE_MAPPING = { SEARCH_DOCUMENT, SEARCH_QUERY }; - static { - assert INPUT_TYPE_MAPPING.length == InputType.values().length : "input type mapping was incorrectly defined"; - } private static final String TEXTS_FIELD = "texts"; @@ -56,23 +50,31 @@ public XContentBuilder toXContent(XContentBuilder builder, Params params) throws builder.field(CohereServiceSettings.MODEL, model); } - if (taskSettings.inputType() != null) { - builder.field(INPUT_TYPE_FIELD, covertToString(taskSettings.inputType())); + if (taskSettings.getInputType() != null) { + builder.field(INPUT_TYPE_FIELD, covertToString(taskSettings.getInputType())); } if (embeddingType != null) { builder.field(EMBEDDING_TYPES_FIELD, List.of(embeddingType)); } - if (taskSettings.truncation() != null) { - builder.field(CohereServiceFields.TRUNCATE, taskSettings.truncation()); + if (taskSettings.getTruncation() != null) { + builder.field(CohereServiceFields.TRUNCATE, taskSettings.getTruncation()); } builder.endObject(); return builder; } - private static String covertToString(InputType inputType) { - return INPUT_TYPE_MAPPING[inputType.ordinal()]; + // default for testing + static String covertToString(InputType inputType) { + return switch (inputType) { + case INGEST -> SEARCH_DOCUMENT; + case SEARCH -> SEARCH_QUERY; + default -> { + assert false : invalidInputTypeMessage(inputType); + yield null; + } + }; } } diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/SenderService.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/SenderService.java index bb45e8fd684a6..0c40863b37db2 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/SenderService.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/SenderService.java @@ -12,6 +12,7 @@ import org.elasticsearch.core.IOUtils; import org.elasticsearch.inference.InferenceService; import org.elasticsearch.inference.InferenceServiceResults; +import org.elasticsearch.inference.InputType; import org.elasticsearch.inference.Model; import org.elasticsearch.xpack.inference.external.http.sender.HttpRequestSenderFactory; import org.elasticsearch.xpack.inference.external.http.sender.Sender; @@ -41,16 +42,23 @@ protected ServiceComponents getServiceComponents() { } @Override - public void infer(Model model, List input, Map taskSettings, ActionListener listener) { + public void infer( + Model model, + List input, + Map taskSettings, + InputType inputType, + ActionListener listener + ) { init(); - doInfer(model, input, taskSettings, listener); + doInfer(model, input, taskSettings, inputType, listener); } protected abstract void doInfer( Model model, List input, Map taskSettings, + InputType inputType, ActionListener listener ); diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/ServiceUtils.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/ServiceUtils.java index c218a0ff12c22..7637bd9740670 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/ServiceUtils.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/ServiceUtils.java @@ -11,10 +11,10 @@ import org.elasticsearch.action.ActionListener; import org.elasticsearch.common.ValidationException; import org.elasticsearch.common.settings.SecureString; -import org.elasticsearch.core.CheckedFunction; import org.elasticsearch.core.Nullable; import org.elasticsearch.core.Strings; import org.elasticsearch.inference.InferenceService; +import org.elasticsearch.inference.InputType; import org.elasticsearch.inference.Model; import org.elasticsearch.inference.TaskType; import org.elasticsearch.rest.RestStatus; @@ -24,7 +24,7 @@ import java.net.URI; import java.net.URISyntaxException; -import java.util.Arrays; +import java.util.EnumSet; import java.util.List; import java.util.Locale; import java.util.Map; @@ -110,7 +110,7 @@ public static String mustBeNonEmptyString(String settingName, String scope) { return Strings.format("[%s] Invalid value empty string. [%s] must be a non-empty string", scope, settingName); } - public static String invalidValue(String settingName, String scope, String invalidType, String... requiredTypes) { + public static String invalidValue(String settingName, String scope, String invalidType, String[] requiredTypes) { return Strings.format( "[%s] Invalid value [%s] received. [%s] must be one of [%s]", scope, @@ -221,12 +221,12 @@ public static String extractOptionalString( return optionalField; } - public static T extractOptionalEnum( + public static > E extractOptionalEnum( Map map, String settingName, String scope, - CheckedFunction converter, - T[] validTypes, + EnumConstructor constructor, + EnumSet validValues, ValidationException validationException ) { var enumString = extractOptionalString(map, settingName, scope, validationException); @@ -234,16 +234,34 @@ public static T extractOptionalEnum( return null; } - var validTypesAsStrings = Arrays.stream(validTypes).map(type -> type.toString().toLowerCase(Locale.ROOT)).toArray(String[]::new); + var validValuesAsStrings = validValues.stream().map(value -> value.toString().toLowerCase(Locale.ROOT)).toArray(String[]::new); try { - return converter.apply(enumString); + var createdEnum = constructor.apply(enumString); + validateEnumValue(createdEnum, validValues); + + return createdEnum; } catch (IllegalArgumentException e) { - validationException.addValidationError(invalidValue(settingName, scope, enumString, validTypesAsStrings)); + validationException.addValidationError(invalidValue(settingName, scope, enumString, validValuesAsStrings)); } return null; } + private static > void validateEnumValue(E enumValue, EnumSet validValues) { + if (validValues.contains(enumValue) == false) { + throw new IllegalArgumentException(Strings.format("Enum value [%s] is not one of the acceptable values", enumValue.toString())); + } + } + + /** + * Functional interface for creating an enum from a string. + * @param + */ + @FunctionalInterface + public interface EnumConstructor> { + E apply(String name) throws IllegalArgumentException; + } + public static String parsePersistedConfigErrorMsg(String inferenceEntityId, String serviceName) { return format( "Failed to parse stored model [%s] for [%s] service, please delete and add the service again", @@ -272,7 +290,7 @@ public static ElasticsearchStatusException createInvalidModelException(Model mod public static void getEmbeddingSize(Model model, InferenceService service, ActionListener listener) { assert model.getTaskType() == TaskType.TEXT_EMBEDDING; - service.infer(model, List.of(TEST_EMBEDDING_INPUT), Map.of(), listener.delegateFailureAndWrap((delegate, r) -> { + service.infer(model, List.of(TEST_EMBEDDING_INPUT), Map.of(), InputType.INGEST, listener.delegateFailureAndWrap((delegate, r) -> { if (r instanceof TextEmbedding embeddingResults) { try { delegate.onResponse(embeddingResults.getFirstEmbeddingSize()); diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/CohereModel.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/CohereModel.java index 1b4843e441248..81a27e1e536f3 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/CohereModel.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/CohereModel.java @@ -7,6 +7,7 @@ package org.elasticsearch.xpack.inference.services.cohere; +import org.elasticsearch.inference.InputType; import org.elasticsearch.inference.Model; import org.elasticsearch.inference.ModelConfigurations; import org.elasticsearch.inference.ModelSecrets; @@ -30,5 +31,5 @@ protected CohereModel(CohereModel model, ServiceSettings serviceSettings) { super(model, serviceSettings); } - public abstract ExecutableAction accept(CohereActionVisitor creator, Map taskSettings); + public abstract ExecutableAction accept(CohereActionVisitor creator, Map taskSettings, InputType inputType); } diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/CohereService.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/CohereService.java index 8783f12852ec8..3f608c977f686 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/CohereService.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/CohereService.java @@ -14,6 +14,7 @@ import org.elasticsearch.action.ActionListener; import org.elasticsearch.core.Nullable; import org.elasticsearch.inference.InferenceServiceResults; +import org.elasticsearch.inference.InputType; import org.elasticsearch.inference.Model; import org.elasticsearch.inference.ModelConfigurations; import org.elasticsearch.inference.ModelSecrets; @@ -123,6 +124,7 @@ public void doInfer( Model model, List input, Map taskSettings, + InputType inputType, ActionListener listener ) { if (model instanceof CohereModel == false) { @@ -133,7 +135,7 @@ public void doInfer( CohereModel cohereModel = (CohereModel) model; var actionCreator = new CohereActionCreator(getSender(), getServiceComponents()); - var action = cohereModel.accept(actionCreator, taskSettings); + var action = cohereModel.accept(actionCreator, taskSettings, inputType); action.execute(input, listener); } @@ -174,6 +176,6 @@ private CohereEmbeddingsModel updateModelWithEmbeddingDetails(CohereEmbeddingsMo @Override public TransportVersion getMinimalSupportedVersion() { - return TransportVersions.ML_INFERENCE_COHERE_EMBEDDINGS_ADDED; + return TransportVersions.ML_INFERENCE_REQUEST_INPUT_TYPE_UNSPECIFIED_ADDED; } } diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsModel.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsModel.java index c92700e87cd96..a3afdc306b217 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsModel.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsModel.java @@ -8,6 +8,7 @@ package org.elasticsearch.xpack.inference.services.cohere.embeddings; import org.elasticsearch.core.Nullable; +import org.elasticsearch.inference.InputType; import org.elasticsearch.inference.ModelConfigurations; import org.elasticsearch.inference.ModelSecrets; import org.elasticsearch.inference.TaskType; @@ -19,6 +20,11 @@ import java.util.Map; public class CohereEmbeddingsModel extends CohereModel { + public static CohereEmbeddingsModel of(CohereEmbeddingsModel model, Map taskSettings, InputType inputType) { + var requestTaskSettings = CohereEmbeddingsTaskSettings.fromMap(taskSettings); + return new CohereEmbeddingsModel(model, CohereEmbeddingsTaskSettings.of(model.getTaskSettings(), requestTaskSettings, inputType)); + } + public CohereEmbeddingsModel( String modelId, TaskType taskType, @@ -73,16 +79,7 @@ public DefaultSecretSettings getSecretSettings() { } @Override - public ExecutableAction accept(CohereActionVisitor visitor, Map taskSettings) { - return visitor.create(this, taskSettings); - } - - public CohereEmbeddingsModel overrideWith(Map taskSettings) { - if (taskSettings == null || taskSettings.isEmpty()) { - return this; - } - - var requestTaskSettings = CohereEmbeddingsTaskSettings.fromMap(taskSettings); - return new CohereEmbeddingsModel(this, getTaskSettings().overrideWith(requestTaskSettings)); + public ExecutableAction accept(CohereActionVisitor visitor, Map taskSettings, InputType inputType) { + return visitor.create(this, taskSettings, inputType); } } diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsServiceSettings.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsServiceSettings.java index 5327bcbcf22dd..916e7fadcc8fb 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsServiceSettings.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsServiceSettings.java @@ -19,6 +19,7 @@ import org.elasticsearch.xpack.inference.services.cohere.CohereServiceSettings; import java.io.IOException; +import java.util.EnumSet; import java.util.Map; import java.util.Objects; @@ -37,7 +38,7 @@ public static CohereEmbeddingsServiceSettings fromMap(Map map) { EMBEDDING_TYPE, ModelConfigurations.SERVICE_SETTINGS, CohereEmbeddingType::fromString, - CohereEmbeddingType.values(), + EnumSet.allOf(CohereEmbeddingType.class), validationException ); diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsTaskSettings.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsTaskSettings.java index 858efdb0d1ace..b294350580a2e 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsTaskSettings.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsTaskSettings.java @@ -9,6 +9,7 @@ import org.elasticsearch.TransportVersion; import org.elasticsearch.TransportVersions; +import org.elasticsearch.common.Strings; import org.elasticsearch.common.ValidationException; import org.elasticsearch.common.io.stream.StreamInput; import org.elasticsearch.common.io.stream.StreamOutput; @@ -20,7 +21,9 @@ import org.elasticsearch.xpack.inference.services.cohere.CohereTruncation; import java.io.IOException; +import java.util.EnumSet; import java.util.Map; +import java.util.Objects; import static org.elasticsearch.xpack.inference.services.ServiceUtils.extractOptionalEnum; import static org.elasticsearch.xpack.inference.services.cohere.CohereServiceFields.TRUNCATE; @@ -31,18 +34,16 @@ *

* See api docs for details. *

- * - * @param inputType Specifies the type of input you're giving to the model - * @param truncation Specifies how the API will handle inputs longer than the maximum token length */ -public record CohereEmbeddingsTaskSettings(@Nullable InputType inputType, @Nullable CohereTruncation truncation) implements TaskSettings { +public class CohereEmbeddingsTaskSettings implements TaskSettings { public static final String NAME = "cohere_embeddings_task_settings"; public static final CohereEmbeddingsTaskSettings EMPTY_SETTINGS = new CohereEmbeddingsTaskSettings(null, null); static final String INPUT_TYPE = "input_type"; + private static final EnumSet VALID_REQUEST_VALUES2 = EnumSet.of(InputType.INGEST, InputType.SEARCH); public static CohereEmbeddingsTaskSettings fromMap(Map map) { - if (map.isEmpty()) { + if (map == null || map.isEmpty()) { return EMPTY_SETTINGS; } @@ -53,7 +54,7 @@ public static CohereEmbeddingsTaskSettings fromMap(Map map) { INPUT_TYPE, ModelConfigurations.TASK_SETTINGS, InputType::fromString, - InputType.values(), + VALID_REQUEST_VALUES2, validationException ); CohereTruncation truncation = extractOptionalEnum( @@ -61,7 +62,7 @@ public static CohereEmbeddingsTaskSettings fromMap(Map map) { TRUNCATE, ModelConfigurations.TASK_SETTINGS, CohereTruncation::fromString, - CohereTruncation.values(), + EnumSet.allOf(CohereTruncation.class), validationException ); @@ -72,10 +73,73 @@ public static CohereEmbeddingsTaskSettings fromMap(Map map) { return new CohereEmbeddingsTaskSettings(inputType, truncation); } + /** + * Creates a new {@link CohereEmbeddingsTaskSettings} by preferring non-null fields from the provided parameters. + * For the input type, preference is given to requestInputType if it is not null and not UNSPECIFIED. + * Then preference is given to the requestTaskSettings and finally to originalSettings even if the value is null. + * + * Similarly, for the truncation field preference is given to requestTaskSettings if it is not null and then to + * originalSettings. + * @param originalSettings the settings stored as part of the inference entity configuration + * @param requestTaskSettings the settings passed in within the task_settings field of the request + * @param requestInputType the input type passed in the request parameters + * @return a constructed {@link CohereEmbeddingsTaskSettings} + */ + public static CohereEmbeddingsTaskSettings of( + CohereEmbeddingsTaskSettings originalSettings, + CohereEmbeddingsTaskSettings requestTaskSettings, + InputType requestInputType + ) { + var inputTypeToUse = getValidInputType(originalSettings, requestTaskSettings, requestInputType); + var truncationToUse = getValidTruncation(originalSettings, requestTaskSettings); + + return new CohereEmbeddingsTaskSettings(inputTypeToUse, truncationToUse); + } + + private static InputType getValidInputType( + CohereEmbeddingsTaskSettings originalSettings, + CohereEmbeddingsTaskSettings requestTaskSettings, + InputType requestInputType + ) { + InputType inputTypeToUse = originalSettings.inputType; + + if (VALID_REQUEST_VALUES2.contains(requestInputType)) { + inputTypeToUse = requestInputType; + } else if (requestTaskSettings.inputType != null) { + inputTypeToUse = requestTaskSettings.inputType; + } + + return inputTypeToUse; + } + + private static CohereTruncation getValidTruncation( + CohereEmbeddingsTaskSettings originalSettings, + CohereEmbeddingsTaskSettings requestTaskSettings + ) { + return requestTaskSettings.getTruncation() == null ? originalSettings.truncation : requestTaskSettings.getTruncation(); + } + + private final InputType inputType; + private final CohereTruncation truncation; + public CohereEmbeddingsTaskSettings(StreamInput in) throws IOException { this(in.readOptionalEnum(InputType.class), in.readOptionalEnum(CohereTruncation.class)); } + public CohereEmbeddingsTaskSettings(@Nullable InputType inputType, @Nullable CohereTruncation truncation) { + validateInputType(inputType); + this.inputType = inputType; + this.truncation = truncation; + } + + private static void validateInputType(InputType inputType) { + if (inputType == null) { + return; + } + + assert VALID_REQUEST_VALUES2.contains(inputType) : invalidInputTypeMessage(inputType); + } + @Override public XContentBuilder toXContent(XContentBuilder builder, Params params) throws IOException { builder.startObject(); @@ -90,6 +154,14 @@ public XContentBuilder toXContent(XContentBuilder builder, Params params) throws return builder; } + public InputType getInputType() { + return inputType; + } + + public CohereTruncation getTruncation() { + return truncation; + } + @Override public String getWriteableName() { return NAME; @@ -106,10 +178,20 @@ public void writeTo(StreamOutput out) throws IOException { out.writeOptionalEnum(truncation); } - public CohereEmbeddingsTaskSettings overrideWith(CohereEmbeddingsTaskSettings requestTaskSettings) { - var inputTypeToUse = requestTaskSettings.inputType() == null ? inputType : requestTaskSettings.inputType(); - var truncationToUse = requestTaskSettings.truncation() == null ? truncation : requestTaskSettings.truncation(); + @Override + public boolean equals(Object o) { + if (this == o) return true; + if (o == null || getClass() != o.getClass()) return false; + CohereEmbeddingsTaskSettings that = (CohereEmbeddingsTaskSettings) o; + return Objects.equals(inputType, that.inputType) && Objects.equals(truncation, that.truncation); + } - return new CohereEmbeddingsTaskSettings(inputTypeToUse, truncationToUse); + @Override + public int hashCode() { + return Objects.hash(inputType, truncation); + } + + public static String invalidInputTypeMessage(InputType inputType) { + return Strings.format("received invalid input type value [%s]", inputType.toString()); } } diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/elser/ElserMlNodeService.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/elser/ElserMlNodeService.java index 12bdcd3f20614..1d0bd123c69f3 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/elser/ElserMlNodeService.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/elser/ElserMlNodeService.java @@ -19,6 +19,7 @@ import org.elasticsearch.inference.InferenceService; import org.elasticsearch.inference.InferenceServiceExtension; import org.elasticsearch.inference.InferenceServiceResults; +import org.elasticsearch.inference.InputType; import org.elasticsearch.inference.Model; import org.elasticsearch.inference.ModelConfigurations; import org.elasticsearch.inference.TaskType; @@ -210,7 +211,13 @@ public void stop(String inferenceEntityId, ActionListener listener) { } @Override - public void infer(Model model, List input, Map taskSettings, ActionListener listener) { + public void infer( + Model model, + List input, + Map taskSettings, + InputType inputType, + ActionListener listener + ) { // No task settings to override with requestTaskSettings if (TaskType.SPARSE_EMBEDDING.isAnyOrSame(model.getConfigurations().getTaskType()) == false) { diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/huggingface/HuggingFaceBaseService.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/huggingface/HuggingFaceBaseService.java index ef93cdd57b756..dcaa760868c49 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/huggingface/HuggingFaceBaseService.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/huggingface/HuggingFaceBaseService.java @@ -10,6 +10,7 @@ import org.apache.lucene.util.SetOnce; import org.elasticsearch.action.ActionListener; import org.elasticsearch.inference.InferenceServiceResults; +import org.elasticsearch.inference.InputType; import org.elasticsearch.inference.Model; import org.elasticsearch.inference.ModelConfigurations; import org.elasticsearch.inference.ModelSecrets; @@ -96,6 +97,7 @@ public void doInfer( Model model, List input, Map taskSettings, + InputType inputType, ActionListener listener ) { if (model instanceof HuggingFaceModel == false) { diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/openai/OpenAiService.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/openai/OpenAiService.java index 9b5283ef4f803..594d7cf2cf31c 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/openai/OpenAiService.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/openai/OpenAiService.java @@ -14,6 +14,7 @@ import org.elasticsearch.action.ActionListener; import org.elasticsearch.core.Nullable; import org.elasticsearch.inference.InferenceServiceResults; +import org.elasticsearch.inference.InputType; import org.elasticsearch.inference.Model; import org.elasticsearch.inference.ModelConfigurations; import org.elasticsearch.inference.ModelSecrets; @@ -136,6 +137,7 @@ public void doInfer( Model model, List input, Map taskSettings, + InputType inputType, ActionListener listener ) { if (model instanceof OpenAiModel == false) { diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsModel.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsModel.java index 98b0161665d8e..74d97099bbb76 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsModel.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsModel.java @@ -21,6 +21,15 @@ public class OpenAiEmbeddingsModel extends OpenAiModel { + public static OpenAiEmbeddingsModel of(OpenAiEmbeddingsModel model, Map taskSettings) { + if (taskSettings == null || taskSettings.isEmpty()) { + return model; + } + + var requestTaskSettings = OpenAiEmbeddingsRequestTaskSettings.fromMap(taskSettings); + return new OpenAiEmbeddingsModel(model, OpenAiEmbeddingsTaskSettings.of(model.getTaskSettings(), requestTaskSettings)); + } + public OpenAiEmbeddingsModel( String inferenceEntityId, TaskType taskType, @@ -78,13 +87,4 @@ public DefaultSecretSettings getSecretSettings() { public ExecutableAction accept(OpenAiActionVisitor creator, Map taskSettings) { return creator.create(this, taskSettings); } - - public OpenAiEmbeddingsModel overrideWith(Map taskSettings) { - if (taskSettings == null || taskSettings.isEmpty()) { - return this; - } - - var requestTaskSettings = OpenAiEmbeddingsRequestTaskSettings.fromMap(taskSettings); - return new OpenAiEmbeddingsModel(this, getTaskSettings().overrideWith(requestTaskSettings)); - } } diff --git a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsTaskSettings.java b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsTaskSettings.java index 45a9ce1cabbc3..c6f3179a4f088 100644 --- a/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsTaskSettings.java +++ b/x-pack/plugin/inference/src/main/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsTaskSettings.java @@ -50,6 +50,23 @@ public static OpenAiEmbeddingsTaskSettings fromMap(Map map) { return new OpenAiEmbeddingsTaskSettings(model, user); } + /** + * Creates a new {@link OpenAiEmbeddingsTaskSettings} object by overriding the values in originalSettings with the ones + * passed in via requestSettings if the fields are not null. + * @param originalSettings the original task settings from the inference entity configuration from storage + * @param requestSettings the task settings from the request + * @return a new {@link OpenAiEmbeddingsTaskSettings} + */ + public static OpenAiEmbeddingsTaskSettings of( + OpenAiEmbeddingsTaskSettings originalSettings, + OpenAiEmbeddingsRequestTaskSettings requestSettings + ) { + var modelToUse = requestSettings.model() == null ? originalSettings.model : requestSettings.model(); + var userToUse = requestSettings.user() == null ? originalSettings.user : requestSettings.user(); + + return new OpenAiEmbeddingsTaskSettings(modelToUse, userToUse); + } + public OpenAiEmbeddingsTaskSettings { Objects.requireNonNull(model); } @@ -84,11 +101,4 @@ public void writeTo(StreamOutput out) throws IOException { out.writeString(model); out.writeOptionalString(user); } - - public OpenAiEmbeddingsTaskSettings overrideWith(OpenAiEmbeddingsRequestTaskSettings requestSettings) { - var modelToUse = requestSettings.model() == null ? model : requestSettings.model(); - var userToUse = requestSettings.user() == null ? user : requestSettings.user(); - - return new OpenAiEmbeddingsTaskSettings(modelToUse, userToUse); - } } diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/InputTypeTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/InputTypeTests.java new file mode 100644 index 0000000000000..088f93507d35f --- /dev/null +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/InputTypeTests.java @@ -0,0 +1,21 @@ +/* + * 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; + +import org.elasticsearch.inference.InputType; +import org.elasticsearch.test.ESTestCase; + +public class InputTypeTests extends ESTestCase { + public static InputType randomWithoutUnspecified() { + return randomFrom(InputType.INGEST, InputType.SEARCH); + } + + public static InputType[] valuesWithoutUnspecified() { + return new InputType[] { InputType.INGEST, InputType.SEARCH }; + } +} diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/action/InferenceActionRequestTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/action/InferenceActionRequestTests.java index 4f7ae9436418f..396af55ce5616 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/action/InferenceActionRequestTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/action/InferenceActionRequestTests.java @@ -7,22 +7,26 @@ package org.elasticsearch.xpack.inference.action; +import org.elasticsearch.TransportVersion; +import org.elasticsearch.TransportVersions; import org.elasticsearch.common.io.stream.Writeable; import org.elasticsearch.core.Tuple; import org.elasticsearch.inference.InputType; import org.elasticsearch.inference.TaskType; -import org.elasticsearch.test.AbstractWireSerializingTestCase; import org.elasticsearch.xcontent.json.JsonXContent; import org.elasticsearch.xpack.core.inference.action.InferenceAction; +import org.elasticsearch.xpack.core.ml.AbstractBWCWireSerializationTestCase; import java.io.IOException; import java.util.ArrayList; import java.util.HashMap; +import java.util.List; +import java.util.Map; import static org.hamcrest.Matchers.is; import static org.hamcrest.collection.IsIterableContainingInOrder.contains; -public class InferenceActionRequestTests extends AbstractWireSerializingTestCase { +public class InferenceActionRequestTests extends AbstractBWCWireSerializationTestCase { @Override protected Writeable.Reader instanceReader() { @@ -70,7 +74,7 @@ public void testParseRequest_DefaultsInputTypeToIngest() throws IOException { """; try (var parser = createParser(JsonXContent.jsonXContent, singleInputRequest)) { var request = InferenceAction.Request.parseRequest("model_id", "sparse_embedding", parser); - assertThat(request.getInputType(), is(InputType.INGEST)); + assertThat(request.getInputType(), is(InputType.UNSPECIFIED)); } } @@ -135,4 +139,76 @@ protected InferenceAction.Request mutateInstance(InferenceAction.Request instanc default -> throw new UnsupportedOperationException(); }; } + + @Override + protected InferenceAction.Request mutateInstanceForVersion(InferenceAction.Request instance, TransportVersion version) { + if (version.before(TransportVersions.INFERENCE_MULTIPLE_INPUTS)) { + return new InferenceAction.Request( + instance.getTaskType(), + instance.getInferenceEntityId(), + instance.getInput().subList(0, 1), + instance.getTaskSettings(), + InputType.UNSPECIFIED + ); + } else if (version.before(TransportVersions.ML_INFERENCE_REQUEST_INPUT_TYPE_ADDED)) { + return new InferenceAction.Request( + instance.getTaskType(), + instance.getInferenceEntityId(), + instance.getInput(), + instance.getTaskSettings(), + InputType.UNSPECIFIED + ); + } else if (version.before(TransportVersions.ML_INFERENCE_REQUEST_INPUT_TYPE_UNSPECIFIED_ADDED) + && instance.getInputType() == InputType.UNSPECIFIED) { + return new InferenceAction.Request( + instance.getTaskType(), + instance.getInferenceEntityId(), + instance.getInput(), + instance.getTaskSettings(), + InputType.INGEST + ); + } + + return instance; + } + + public void testWriteTo_WhenVersionIsOnAfterUnspecifiedAdded() throws IOException { + assertBwcSerialization( + new InferenceAction.Request(TaskType.TEXT_EMBEDDING, "model", List.of(), Map.of(), InputType.UNSPECIFIED), + TransportVersions.ML_INFERENCE_REQUEST_INPUT_TYPE_UNSPECIFIED_ADDED + ); + } + + public void testWriteTo_WhenVersionIsBeforeUnspecifiedAdded_ButAfterInputTypeAdded_ShouldSetToIngest() throws IOException { + assertBwcSerialization( + new InferenceAction.Request(TaskType.TEXT_EMBEDDING, "model", List.of(), Map.of(), InputType.UNSPECIFIED), + TransportVersions.ML_INFERENCE_REQUEST_INPUT_TYPE_ADDED + ); + } + + public void testWriteTo_WhenVersionIsBeforeUnspecifiedAdded_ButAfterInputTypeAdded_ShouldSetToIngest_ManualCheck() throws IOException { + var instance = new InferenceAction.Request(TaskType.TEXT_EMBEDDING, "model", List.of(), Map.of(), InputType.UNSPECIFIED); + + InferenceAction.Request deserializedInstance = copyWriteable( + instance, + getNamedWriteableRegistry(), + instanceReader(), + TransportVersions.ML_INFERENCE_REQUEST_INPUT_TYPE_ADDED + ); + + assertThat(deserializedInstance.getInputType(), is(InputType.INGEST)); + } + + public void testWriteTo_WhenVersionIsBeforeInputTypeAdded_ShouldSetInputTypeToUnspecified() throws IOException { + var instance = new InferenceAction.Request(TaskType.TEXT_EMBEDDING, "model", List.of(), Map.of(), InputType.INGEST); + + InferenceAction.Request deserializedInstance = copyWriteable( + instance, + getNamedWriteableRegistry(), + instanceReader(), + TransportVersions.HOT_THREADS_AS_BYTES + ); + + assertThat(deserializedInstance.getInputType(), is(InputType.UNSPECIFIED)); + } } diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/external/action/cohere/CohereActionCreatorTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/external/action/cohere/CohereActionCreatorTests.java index 67a95265f093d..e7cfc784db117 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/external/action/cohere/CohereActionCreatorTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/external/action/cohere/CohereActionCreatorTests.java @@ -110,7 +110,7 @@ public void testCreate_CohereEmbeddingsModel() throws IOException { ); var actionCreator = new CohereActionCreator(sender, createWithEmptySettings(threadPool)); var overriddenTaskSettings = CohereEmbeddingsTaskSettingsTests.getTaskSettingsMap(InputType.SEARCH, CohereTruncation.END); - var action = actionCreator.create(model, overriddenTaskSettings); + var action = actionCreator.create(model, overriddenTaskSettings, InputType.UNSPECIFIED); PlainActionFuture listener = new PlainActionFuture<>(); action.execute(List.of("abc"), listener); diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/external/request/cohere/CohereEmbeddingsRequestEntityTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/external/request/cohere/CohereEmbeddingsRequestEntityTests.java index 8ef9ea4b0316b..2d3ff25222ab9 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/external/request/cohere/CohereEmbeddingsRequestEntityTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/external/request/cohere/CohereEmbeddingsRequestEntityTests.java @@ -66,4 +66,9 @@ public void testXContent_WritesNoOptionalFields_WhenTheyAreNotDefined() throws I MatcherAssert.assertThat(xContentResult, is(""" {"texts":["abc"]}""")); } + + public void testConvertToString_ThrowsAssertionFailure_WhenInputTypeIsUnspecified() { + var thrownException = expectThrows(AssertionError.class, () -> CohereEmbeddingsRequestEntity.covertToString(InputType.UNSPECIFIED)); + MatcherAssert.assertThat(thrownException.getMessage(), is("received invalid input type value [unspecified]")); + } } diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/SenderServiceTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/SenderServiceTests.java index 31d7667fa6665..8b596aa5cf0c8 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/SenderServiceTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/SenderServiceTests.java @@ -13,6 +13,7 @@ import org.elasticsearch.action.support.PlainActionFuture; import org.elasticsearch.core.TimeValue; import org.elasticsearch.inference.InferenceServiceResults; +import org.elasticsearch.inference.InputType; import org.elasticsearch.inference.Model; import org.elasticsearch.inference.TaskType; import org.elasticsearch.test.ESTestCase; @@ -105,6 +106,7 @@ protected void doInfer( Model model, List input, Map taskSettings, + InputType inputType, ActionListener listener ) { diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/ServiceUtilsTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/ServiceUtilsTests.java index b935c5a8c64b3..689c9f9b08a2b 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/ServiceUtilsTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/ServiceUtilsTests.java @@ -23,6 +23,7 @@ import org.elasticsearch.xpack.inference.results.TextEmbeddingByteResultsTests; import org.elasticsearch.xpack.inference.results.TextEmbeddingResultsTests; +import java.util.EnumSet; import java.util.HashMap; import java.util.List; import java.util.Map; @@ -261,7 +262,7 @@ public void testExtractOptionalString_AddsException_WhenFieldIsEmpty() { public void testExtractOptionalEnum_ReturnsNull_WhenFieldDoesNotExist() { var validation = new ValidationException(); Map map = modifiableMap(Map.of("key", "value")); - var createdEnum = extractOptionalEnum(map, "abc", "scope", InputType::fromString, InputType.values(), validation); + var createdEnum = extractOptionalEnum(map, "abc", "scope", InputType::fromString, EnumSet.allOf(InputType.class), validation); assertNull(createdEnum); assertTrue(validation.validationErrors().isEmpty()); @@ -271,7 +272,14 @@ public void testExtractOptionalEnum_ReturnsNull_WhenFieldDoesNotExist() { public void testExtractOptionalEnum_ReturnsNullAndAddsException_WhenAnInvalidValueExists() { var validation = new ValidationException(); Map map = modifiableMap(Map.of("key", "invalid_value")); - var createdEnum = extractOptionalEnum(map, "key", "scope", InputType::fromString, InputType.values(), validation); + var createdEnum = extractOptionalEnum( + map, + "key", + "scope", + InputType::fromString, + EnumSet.of(InputType.INGEST, InputType.SEARCH), + validation + ); assertNull(createdEnum); assertFalse(validation.validationErrors().isEmpty()); @@ -282,6 +290,27 @@ public void testExtractOptionalEnum_ReturnsNullAndAddsException_WhenAnInvalidVal ); } + public void testExtractOptionalEnum_ReturnsNullAndAddsException_WhenValueIsNotPartOfTheAcceptableValues() { + var validation = new ValidationException(); + Map map = modifiableMap(Map.of("key", InputType.UNSPECIFIED.toString())); + var createdEnum = extractOptionalEnum(map, "key", "scope", InputType::fromString, EnumSet.of(InputType.INGEST), validation); + + assertNull(createdEnum); + assertFalse(validation.validationErrors().isEmpty()); + assertTrue(map.isEmpty()); + assertThat(validation.validationErrors().get(0), is("[scope] Invalid value [unspecified] received. [key] must be one of [ingest]")); + } + + public void testExtractOptionalEnum_ReturnsIngest_WhenValueIsAcceptable() { + var validation = new ValidationException(); + Map map = modifiableMap(Map.of("key", InputType.INGEST.toString())); + var createdEnum = extractOptionalEnum(map, "key", "scope", InputType::fromString, EnumSet.of(InputType.INGEST), validation); + + assertThat(createdEnum, is(InputType.INGEST)); + assertTrue(validation.validationErrors().isEmpty()); + assertTrue(map.isEmpty()); + } + public void testGetEmbeddingSize_ReturnsError_WhenTextEmbeddingResults_IsEmpty() { var service = mock(InferenceService.class); @@ -290,11 +319,11 @@ public void testGetEmbeddingSize_ReturnsError_WhenTextEmbeddingResults_IsEmpty() doAnswer(invocation -> { @SuppressWarnings("unchecked") - ActionListener listener = (ActionListener) invocation.getArguments()[3]; + ActionListener listener = (ActionListener) invocation.getArguments()[4]; listener.onResponse(new TextEmbeddingResults(List.of())); return Void.TYPE; - }).when(service).infer(any(), any(), any(), any()); + }).when(service).infer(any(), any(), any(), any(), any()); PlainActionFuture listener = new PlainActionFuture<>(); getEmbeddingSize(model, service, listener); @@ -313,11 +342,11 @@ public void testGetEmbeddingSize_ReturnsError_WhenTextEmbeddingByteResults_IsEmp doAnswer(invocation -> { @SuppressWarnings("unchecked") - ActionListener listener = (ActionListener) invocation.getArguments()[3]; + ActionListener listener = (ActionListener) invocation.getArguments()[4]; listener.onResponse(new TextEmbeddingByteResults(List.of())); return Void.TYPE; - }).when(service).infer(any(), any(), any(), any()); + }).when(service).infer(any(), any(), any(), any(), any()); PlainActionFuture listener = new PlainActionFuture<>(); getEmbeddingSize(model, service, listener); @@ -338,11 +367,11 @@ public void testGetEmbeddingSize_ReturnsSize_ForTextEmbeddingResults() { doAnswer(invocation -> { @SuppressWarnings("unchecked") - ActionListener listener = (ActionListener) invocation.getArguments()[3]; + ActionListener listener = (ActionListener) invocation.getArguments()[4]; listener.onResponse(textEmbedding); return Void.TYPE; - }).when(service).infer(any(), any(), any(), any()); + }).when(service).infer(any(), any(), any(), any(), any()); PlainActionFuture listener = new PlainActionFuture<>(); getEmbeddingSize(model, service, listener); @@ -362,11 +391,11 @@ public void testGetEmbeddingSize_ReturnsSize_ForTextEmbeddingByteResults() { doAnswer(invocation -> { @SuppressWarnings("unchecked") - ActionListener listener = (ActionListener) invocation.getArguments()[3]; + ActionListener listener = (ActionListener) invocation.getArguments()[4]; listener.onResponse(textEmbedding); return Void.TYPE; - }).when(service).infer(any(), any(), any(), any()); + }).when(service).infer(any(), any(), any(), any(), any()); PlainActionFuture listener = new PlainActionFuture<>(); getEmbeddingSize(model, service, listener); diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/cohere/CohereServiceTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/cohere/CohereServiceTests.java index 0250e08a48452..7daad207f9068 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/cohere/CohereServiceTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/cohere/CohereServiceTests.java @@ -34,6 +34,7 @@ import org.elasticsearch.xpack.inference.services.cohere.embeddings.CohereEmbeddingsModelTests; import org.elasticsearch.xpack.inference.services.cohere.embeddings.CohereEmbeddingsServiceSettingsTests; import org.elasticsearch.xpack.inference.services.cohere.embeddings.CohereEmbeddingsTaskSettings; +import org.elasticsearch.xpack.inference.services.cohere.embeddings.CohereEmbeddingsTaskSettingsTests; import org.hamcrest.MatcherAssert; import org.hamcrest.Matchers; import org.junit.After; @@ -686,7 +687,7 @@ public void testInfer_ThrowsErrorWhenModelIsNotCohereModel() throws IOException try (var service = new CohereService(new SetOnce<>(factory), new SetOnce<>(createWithEmptySettings(threadPool)))) { PlainActionFuture listener = new PlainActionFuture<>(); - service.infer(mockModel, List.of(""), new HashMap<>(), listener); + service.infer(mockModel, List.of(""), new HashMap<>(), InputType.INGEST, listener); var thrownException = expectThrows(ElasticsearchStatusException.class, () -> listener.actionGet(TIMEOUT)); MatcherAssert.assertThat( @@ -745,7 +746,7 @@ public void testInfer_SendsRequest() throws IOException { null ); PlainActionFuture listener = new PlainActionFuture<>(); - service.infer(model, List.of("abc"), new HashMap<>(), listener); + service.infer(model, List.of("abc"), new HashMap<>(), InputType.INGEST, listener); var result = listener.actionGet(TIMEOUT); @@ -848,7 +849,7 @@ public void testInfer_UnauthorisedResponse() throws IOException { null ); PlainActionFuture listener = new PlainActionFuture<>(); - service.infer(model, List.of("abc"), new HashMap<>(), listener); + service.infer(model, List.of("abc"), new HashMap<>(), InputType.INGEST, listener); var error = expectThrows(ElasticsearchException.class, () -> listener.actionGet(TIMEOUT)); MatcherAssert.assertThat(error.getMessage(), containsString("Received an authentication error status code for request")); @@ -857,6 +858,193 @@ public void testInfer_UnauthorisedResponse() throws IOException { } } + public void testInfer_SetsInputTypeToIngest_FromInferParameter_WhenTaskSettingsAreEmpty() throws IOException { + var senderFactory = new HttpRequestSenderFactory(threadPool, clientManager, mockClusterServiceEmpty(), Settings.EMPTY); + + try (var service = new CohereService(new SetOnce<>(senderFactory), new SetOnce<>(createWithEmptySettings(threadPool)))) { + + String responseJson = """ + { + "id": "de37399c-5df6-47cb-bc57-e3c5680c977b", + "texts": [ + "hello" + ], + "embeddings": { + "float": [ + [ + 0.123, + -0.123 + ] + ] + }, + "meta": { + "api_version": { + "version": "1" + }, + "billed_units": { + "input_tokens": 1 + } + }, + "response_type": "embeddings_by_type" + } + """; + webServer.enqueue(new MockResponse().setResponseCode(200).setBody(responseJson)); + + var model = CohereEmbeddingsModelTests.createModel( + getUrl(webServer), + "secret", + CohereEmbeddingsTaskSettings.EMPTY_SETTINGS, + 1024, + 1024, + "model", + null + ); + PlainActionFuture listener = new PlainActionFuture<>(); + service.infer(model, List.of("abc"), new HashMap<>(), InputType.INGEST, listener); + + var result = listener.actionGet(TIMEOUT); + + MatcherAssert.assertThat(result.asMap(), Matchers.is(buildExpectation(List.of(List.of(0.123F, -0.123F))))); + MatcherAssert.assertThat(webServer.requests(), hasSize(1)); + assertNull(webServer.requests().get(0).getUri().getQuery()); + MatcherAssert.assertThat( + webServer.requests().get(0).getHeader(HttpHeaders.CONTENT_TYPE), + equalTo(XContentType.JSON.mediaType()) + ); + MatcherAssert.assertThat(webServer.requests().get(0).getHeader(HttpHeaders.AUTHORIZATION), equalTo("Bearer secret")); + + var requestMap = entityAsMap(webServer.requests().get(0).getBody()); + MatcherAssert.assertThat(requestMap, is(Map.of("texts", List.of("abc"), "model", "model", "input_type", "search_document"))); + } + } + + public void testInfer_SetsInputTypeToIngestFromInferParameter_WhenModelSettingIsNull_AndRequestTaskSettingsIsSearch() + throws IOException { + var senderFactory = new HttpRequestSenderFactory(threadPool, clientManager, mockClusterServiceEmpty(), Settings.EMPTY); + + try (var service = new CohereService(new SetOnce<>(senderFactory), new SetOnce<>(createWithEmptySettings(threadPool)))) { + + String responseJson = """ + { + "id": "de37399c-5df6-47cb-bc57-e3c5680c977b", + "texts": [ + "hello" + ], + "embeddings": { + "float": [ + [ + 0.123, + -0.123 + ] + ] + }, + "meta": { + "api_version": { + "version": "1" + }, + "billed_units": { + "input_tokens": 1 + } + }, + "response_type": "embeddings_by_type" + } + """; + webServer.enqueue(new MockResponse().setResponseCode(200).setBody(responseJson)); + + var model = CohereEmbeddingsModelTests.createModel( + getUrl(webServer), + "secret", + new CohereEmbeddingsTaskSettings(null, null), + 1024, + 1024, + "model", + null + ); + PlainActionFuture listener = new PlainActionFuture<>(); + service.infer( + model, + List.of("abc"), + CohereEmbeddingsTaskSettingsTests.getTaskSettingsMap(InputType.SEARCH, null), + InputType.INGEST, + listener + ); + + var result = listener.actionGet(TIMEOUT); + + MatcherAssert.assertThat(result.asMap(), Matchers.is(buildExpectation(List.of(List.of(0.123F, -0.123F))))); + MatcherAssert.assertThat(webServer.requests(), hasSize(1)); + assertNull(webServer.requests().get(0).getUri().getQuery()); + MatcherAssert.assertThat( + webServer.requests().get(0).getHeader(HttpHeaders.CONTENT_TYPE), + equalTo(XContentType.JSON.mediaType()) + ); + MatcherAssert.assertThat(webServer.requests().get(0).getHeader(HttpHeaders.AUTHORIZATION), equalTo("Bearer secret")); + + var requestMap = entityAsMap(webServer.requests().get(0).getBody()); + MatcherAssert.assertThat(requestMap, is(Map.of("texts", List.of("abc"), "model", "model", "input_type", "search_document"))); + } + } + + public void testInfer_DoesNotSetInputType_WhenNotPresentInTaskSettings_AndUnspecifiedIsPassedInRequest() throws IOException { + var senderFactory = new HttpRequestSenderFactory(threadPool, clientManager, mockClusterServiceEmpty(), Settings.EMPTY); + + try (var service = new CohereService(new SetOnce<>(senderFactory), new SetOnce<>(createWithEmptySettings(threadPool)))) { + + String responseJson = """ + { + "id": "de37399c-5df6-47cb-bc57-e3c5680c977b", + "texts": [ + "hello" + ], + "embeddings": { + "float": [ + [ + 0.123, + -0.123 + ] + ] + }, + "meta": { + "api_version": { + "version": "1" + }, + "billed_units": { + "input_tokens": 1 + } + }, + "response_type": "embeddings_by_type" + } + """; + webServer.enqueue(new MockResponse().setResponseCode(200).setBody(responseJson)); + + var model = CohereEmbeddingsModelTests.createModel( + getUrl(webServer), + "secret", + new CohereEmbeddingsTaskSettings(null, null), + 1024, + 1024, + "model", + null + ); + PlainActionFuture listener = new PlainActionFuture<>(); + service.infer(model, List.of("abc"), new HashMap<>(), InputType.UNSPECIFIED, listener); + + var result = listener.actionGet(TIMEOUT); + + MatcherAssert.assertThat(result.asMap(), Matchers.is(buildExpectation(List.of(List.of(0.123F, -0.123F))))); + MatcherAssert.assertThat(webServer.requests(), hasSize(1)); + assertNull(webServer.requests().get(0).getUri().getQuery()); + MatcherAssert.assertThat( + webServer.requests().get(0).getHeader(HttpHeaders.CONTENT_TYPE), + equalTo(XContentType.JSON.mediaType()) + ); + MatcherAssert.assertThat(webServer.requests().get(0).getHeader(HttpHeaders.AUTHORIZATION), equalTo("Bearer secret")); + + var requestMap = entityAsMap(webServer.requests().get(0).getBody()); + MatcherAssert.assertThat(requestMap, is(Map.of("texts", List.of("abc"), "model", "model"))); + } + } + private Map getRequestConfigMap( Map serviceSettings, Map taskSettings, diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsModelTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsModelTests.java index 1961d6b168d54..5570731dbe8d9 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsModelTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsModelTests.java @@ -21,12 +21,36 @@ import static org.elasticsearch.xpack.inference.services.cohere.embeddings.CohereEmbeddingsTaskSettingsTests.getTaskSettingsMap; import static org.hamcrest.Matchers.is; -import static org.hamcrest.Matchers.sameInstance; public class CohereEmbeddingsModelTests extends ESTestCase { - public void testOverrideWith_OverridesInputType_WithSearch() { + public void testOverrideWith_DoesNotOverrideAndModelRemainsEqual_WhenSettingsAreEmpty_AndInputTypeIsInvalid() { + var model = createModel("url", "api_key", null, null, null); + + var overriddenModel = CohereEmbeddingsModel.of(model, Map.of(), InputType.UNSPECIFIED); + MatcherAssert.assertThat(overriddenModel, is(model)); + } + + public void testOverrideWith_DoesNotOverrideAndModelRemainsEqual_WhenSettingsAreNull_AndInputTypeIsInvalid() { + var model = createModel("url", "api_key", null, null, null); + + var overriddenModel = CohereEmbeddingsModel.of(model, null, InputType.UNSPECIFIED); + MatcherAssert.assertThat(overriddenModel, is(model)); + } + + public void testOverrideWith_SetsInputTypeToIngest_WhenTheFieldIsNullInModelTaskSettings_AndNullInRequestTaskSettings() { var model = createModel( + "url", + "api_key", + new CohereEmbeddingsTaskSettings(null, null), + null, + null, + "model", + CohereEmbeddingType.FLOAT + ); + + var overriddenModel = CohereEmbeddingsModel.of(model, getTaskSettingsMap(null, null), InputType.INGEST); + var expectedModel = createModel( "url", "api_key", new CohereEmbeddingsTaskSettings(InputType.INGEST, null), @@ -35,8 +59,21 @@ public void testOverrideWith_OverridesInputType_WithSearch() { "model", CohereEmbeddingType.FLOAT ); + MatcherAssert.assertThat(overriddenModel, is(expectedModel)); + } - var overriddenModel = model.overrideWith(getTaskSettingsMap(InputType.SEARCH, null)); + public void testOverrideWith_SetsInputType_FromRequest_IfValid_OverridingStoredTaskSettings() { + var model = createModel( + "url", + "api_key", + new CohereEmbeddingsTaskSettings(InputType.INGEST, null), + null, + null, + "model", + CohereEmbeddingType.FLOAT + ); + + var overriddenModel = CohereEmbeddingsModel.of(model, getTaskSettingsMap(null, null), InputType.SEARCH); var expectedModel = createModel( "url", "api_key", @@ -49,18 +86,100 @@ public void testOverrideWith_OverridesInputType_WithSearch() { MatcherAssert.assertThat(overriddenModel, is(expectedModel)); } - public void testOverrideWith_DoesNotOverride_WhenSettingsAreEmpty() { - var model = createModel("url", "api_key", null, null, null); + public void testOverrideWith_SetsInputType_FromRequest_IfValid_OverridingRequestTaskSettings() { + var model = createModel( + "url", + "api_key", + new CohereEmbeddingsTaskSettings(null, null), + null, + null, + "model", + CohereEmbeddingType.FLOAT + ); - var overriddenModel = model.overrideWith(Map.of()); - MatcherAssert.assertThat(overriddenModel, sameInstance(model)); + var overriddenModel = CohereEmbeddingsModel.of(model, getTaskSettingsMap(InputType.INGEST, null), InputType.SEARCH); + var expectedModel = createModel( + "url", + "api_key", + new CohereEmbeddingsTaskSettings(InputType.SEARCH, null), + null, + null, + "model", + CohereEmbeddingType.FLOAT + ); + MatcherAssert.assertThat(overriddenModel, is(expectedModel)); } - public void testOverrideWith_DoesNotOverride_WhenSettingsAreNull() { - var model = createModel("url", "api_key", null, null, null); + public void testOverrideWith_OverridesInputType_WithRequestTaskSettingsSearch_WhenRequestInputTypeIsInvalid() { + var model = createModel( + "url", + "api_key", + new CohereEmbeddingsTaskSettings(InputType.INGEST, null), + null, + null, + "model", + CohereEmbeddingType.FLOAT + ); + + var overriddenModel = CohereEmbeddingsModel.of(model, getTaskSettingsMap(InputType.SEARCH, null), InputType.UNSPECIFIED); + var expectedModel = createModel( + "url", + "api_key", + new CohereEmbeddingsTaskSettings(InputType.SEARCH, null), + null, + null, + "model", + CohereEmbeddingType.FLOAT + ); + MatcherAssert.assertThat(overriddenModel, is(expectedModel)); + } + + public void testOverrideWith_DoesNotSetInputType_FromRequest_IfInputTypeIsInvalid() { + var model = createModel( + "url", + "api_key", + new CohereEmbeddingsTaskSettings(null, null), + null, + null, + "model", + CohereEmbeddingType.FLOAT + ); - var overriddenModel = model.overrideWith(null); - MatcherAssert.assertThat(overriddenModel, sameInstance(model)); + var overriddenModel = CohereEmbeddingsModel.of(model, getTaskSettingsMap(null, null), InputType.UNSPECIFIED); + var expectedModel = createModel( + "url", + "api_key", + new CohereEmbeddingsTaskSettings(null, null), + null, + null, + "model", + CohereEmbeddingType.FLOAT + ); + MatcherAssert.assertThat(overriddenModel, is(expectedModel)); + } + + public void testOverrideWith_DoesNotSetInputType_WhenRequestTaskSettingsIsNull_AndRequestInputTypeIsInvalid() { + var model = createModel( + "url", + "api_key", + new CohereEmbeddingsTaskSettings(InputType.INGEST, null), + null, + null, + "model", + CohereEmbeddingType.FLOAT + ); + + var overriddenModel = CohereEmbeddingsModel.of(model, getTaskSettingsMap(null, null), InputType.UNSPECIFIED); + var expectedModel = createModel( + "url", + "api_key", + new CohereEmbeddingsTaskSettings(InputType.INGEST, null), + null, + null, + "model", + CohereEmbeddingType.FLOAT + ); + MatcherAssert.assertThat(overriddenModel, is(expectedModel)); } public static CohereEmbeddingsModel createModel( diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsTaskSettingsTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsTaskSettingsTests.java index 164d3998f138f..77e3280d18f93 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsTaskSettingsTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/cohere/embeddings/CohereEmbeddingsTaskSettingsTests.java @@ -13,20 +13,21 @@ import org.elasticsearch.inference.InputType; import org.elasticsearch.test.AbstractWireSerializingTestCase; import org.elasticsearch.xpack.inference.services.cohere.CohereServiceFields; -import org.elasticsearch.xpack.inference.services.cohere.CohereServiceSettings; import org.elasticsearch.xpack.inference.services.cohere.CohereTruncation; +import org.hamcrest.CoreMatchers; import org.hamcrest.MatcherAssert; import java.io.IOException; import java.util.HashMap; import java.util.Map; +import static org.elasticsearch.xpack.inference.InputTypeTests.randomWithoutUnspecified; import static org.hamcrest.Matchers.is; public class CohereEmbeddingsTaskSettingsTests extends AbstractWireSerializingTestCase { public static CohereEmbeddingsTaskSettings createRandom() { - var inputType = randomBoolean() ? randomFrom(InputType.values()) : null; + var inputType = randomBoolean() ? randomWithoutUnspecified() : null; var truncation = randomBoolean() ? randomFrom(CohereTruncation.values()) : null; return new CohereEmbeddingsTaskSettings(inputType, truncation); @@ -39,6 +40,10 @@ public void testFromMap_CreatesEmptySettings_WhenAllFieldsAreNull() { ); } + public void testFromMap_CreatesEmptySettings_WhenMapIsNull() { + MatcherAssert.assertThat(CohereEmbeddingsTaskSettings.fromMap(null), is(new CohereEmbeddingsTaskSettings(null, null))); + } + public void testFromMap_CreatesSettings_WhenAllFieldsOfSettingsArePresent() { MatcherAssert.assertThat( CohereEmbeddingsTaskSettings.fromMap( @@ -67,26 +72,55 @@ public void testFromMap_ReturnsFailure_WhenInputTypeIsInvalid() { ); } - public void testOverrideWith_KeepsOriginalValuesWhenOverridesAreNull() { - var taskSettings = CohereEmbeddingsTaskSettings.fromMap( - new HashMap<>(Map.of(CohereServiceSettings.MODEL, "model", CohereServiceFields.TRUNCATE, CohereTruncation.END.toString())) + public void testFromMap_ReturnsFailure_WhenInputTypeIsUnspecified() { + var exception = expectThrows( + ValidationException.class, + () -> CohereEmbeddingsTaskSettings.fromMap( + new HashMap<>(Map.of(CohereEmbeddingsTaskSettings.INPUT_TYPE, InputType.UNSPECIFIED.toString())) + ) + ); + + MatcherAssert.assertThat( + exception.getMessage(), + is("Validation Failed: 1: [task_settings] Invalid value [unspecified] received. [input_type] must be one of [ingest, search];") ); + } + + public void testXContent_ThrowsAssertionFailure_WhenInputTypeIsUnspecified() { + var thrownException = expectThrows(AssertionError.class, () -> new CohereEmbeddingsTaskSettings(InputType.UNSPECIFIED, null)); + MatcherAssert.assertThat(thrownException.getMessage(), CoreMatchers.is("received invalid input type value [unspecified]")); + } - var overriddenTaskSettings = taskSettings.overrideWith(CohereEmbeddingsTaskSettings.EMPTY_SETTINGS); + public void testOf_KeepsOriginalValuesWhenRequestSettingsAreNull_AndRequestInputTypeIsInvalid() { + var taskSettings = new CohereEmbeddingsTaskSettings(InputType.INGEST, CohereTruncation.NONE); + var overriddenTaskSettings = CohereEmbeddingsTaskSettings.of( + taskSettings, + CohereEmbeddingsTaskSettings.EMPTY_SETTINGS, + InputType.UNSPECIFIED + ); MatcherAssert.assertThat(overriddenTaskSettings, is(taskSettings)); } - public void testOverrideWith_UsesOverriddenSettings() { - var taskSettings = CohereEmbeddingsTaskSettings.fromMap( - new HashMap<>(Map.of(CohereServiceFields.TRUNCATE, CohereTruncation.END.toString())) + public void testOf_UsesRequestTaskSettings() { + var taskSettings = new CohereEmbeddingsTaskSettings(null, CohereTruncation.NONE); + var overriddenTaskSettings = CohereEmbeddingsTaskSettings.of( + taskSettings, + new CohereEmbeddingsTaskSettings(InputType.INGEST, CohereTruncation.END), + InputType.UNSPECIFIED ); - var requestTaskSettings = CohereEmbeddingsTaskSettings.fromMap( - new HashMap<>(Map.of(CohereServiceFields.TRUNCATE, CohereTruncation.START.toString())) + MatcherAssert.assertThat(overriddenTaskSettings, is(new CohereEmbeddingsTaskSettings(InputType.INGEST, CohereTruncation.END))); + } + + public void testOf_UsesRequestTaskSettings_AndRequestInputType() { + var taskSettings = new CohereEmbeddingsTaskSettings(InputType.SEARCH, CohereTruncation.NONE); + var overriddenTaskSettings = CohereEmbeddingsTaskSettings.of( + taskSettings, + new CohereEmbeddingsTaskSettings(null, CohereTruncation.END), + InputType.INGEST ); - var overriddenTaskSettings = taskSettings.overrideWith(requestTaskSettings); - MatcherAssert.assertThat(overriddenTaskSettings, is(new CohereEmbeddingsTaskSettings(null, CohereTruncation.START))); + MatcherAssert.assertThat(overriddenTaskSettings, is(new CohereEmbeddingsTaskSettings(InputType.INGEST, CohereTruncation.END))); } @Override diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/huggingface/HuggingFaceBaseServiceTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/huggingface/HuggingFaceBaseServiceTests.java index e9fb835016b4f..dcf8b3a900a22 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/huggingface/HuggingFaceBaseServiceTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/huggingface/HuggingFaceBaseServiceTests.java @@ -13,6 +13,7 @@ import org.elasticsearch.action.support.PlainActionFuture; import org.elasticsearch.core.TimeValue; import org.elasticsearch.inference.InferenceServiceResults; +import org.elasticsearch.inference.InputType; import org.elasticsearch.inference.TaskType; import org.elasticsearch.test.ESTestCase; import org.elasticsearch.threadpool.ThreadPool; @@ -64,7 +65,7 @@ public void testInfer_ThrowsErrorWhenModelIsNotHuggingFaceModel() throws IOExcep try (var service = new TestService(new SetOnce<>(factory), new SetOnce<>(createWithEmptySettings(threadPool)))) { PlainActionFuture listener = new PlainActionFuture<>(); - service.infer(mockModel, List.of(""), new HashMap<>(), listener); + service.infer(mockModel, List.of(""), new HashMap<>(), InputType.INGEST, listener); var thrownException = expectThrows(ElasticsearchStatusException.class, () -> listener.actionGet(TIMEOUT)); assertThat( diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/huggingface/HuggingFaceServiceTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/huggingface/HuggingFaceServiceTests.java index a76cce41b4fe4..36a4d144d8c5c 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/huggingface/HuggingFaceServiceTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/huggingface/HuggingFaceServiceTests.java @@ -15,6 +15,7 @@ import org.elasticsearch.core.Nullable; import org.elasticsearch.core.TimeValue; import org.elasticsearch.inference.InferenceServiceResults; +import org.elasticsearch.inference.InputType; import org.elasticsearch.inference.Model; import org.elasticsearch.inference.ModelConfigurations; import org.elasticsearch.inference.ModelSecrets; @@ -492,7 +493,7 @@ public void testInfer_SendsEmbeddingsRequest() throws IOException { var model = HuggingFaceEmbeddingsModelTests.createModel(getUrl(webServer), "secret"); PlainActionFuture listener = new PlainActionFuture<>(); - service.infer(model, List.of("abc"), new HashMap<>(), listener); + service.infer(model, List.of("abc"), new HashMap<>(), InputType.INGEST, listener); var result = listener.actionGet(TIMEOUT); @@ -527,7 +528,7 @@ public void testInfer_SendsElserRequest() throws IOException { var model = HuggingFaceElserModelTests.createModel(getUrl(webServer), "secret"); PlainActionFuture listener = new PlainActionFuture<>(); - service.infer(model, List.of("abc"), new HashMap<>(), listener); + service.infer(model, List.of("abc"), new HashMap<>(), InputType.INGEST, listener); var result = listener.actionGet(TIMEOUT); diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/openai/OpenAiServiceTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/openai/OpenAiServiceTests.java index 394286ee5287b..2659715771686 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/openai/OpenAiServiceTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/openai/OpenAiServiceTests.java @@ -15,6 +15,7 @@ import org.elasticsearch.common.settings.Settings; import org.elasticsearch.core.TimeValue; import org.elasticsearch.inference.InferenceServiceResults; +import org.elasticsearch.inference.InputType; import org.elasticsearch.inference.Model; import org.elasticsearch.inference.ModelConfigurations; import org.elasticsearch.inference.ModelSecrets; @@ -667,7 +668,7 @@ public void testInfer_ThrowsErrorWhenModelIsNotOpenAiModel() throws IOException try (var service = new OpenAiService(new SetOnce<>(factory), new SetOnce<>(createWithEmptySettings(threadPool)))) { PlainActionFuture listener = new PlainActionFuture<>(); - service.infer(mockModel, List.of(""), new HashMap<>(), listener); + service.infer(mockModel, List.of(""), new HashMap<>(), InputType.INGEST, listener); var thrownException = expectThrows(ElasticsearchStatusException.class, () -> listener.actionGet(TIMEOUT)); assertThat( @@ -713,7 +714,7 @@ public void testInfer_SendsRequest() throws IOException { var model = OpenAiEmbeddingsModelTests.createModel(getUrl(webServer), "org", "secret", "model", "user"); PlainActionFuture listener = new PlainActionFuture<>(); - service.infer(model, List.of("abc"), new HashMap<>(), listener); + service.infer(model, List.of("abc"), new HashMap<>(), InputType.INGEST, listener); var result = listener.actionGet(TIMEOUT); @@ -787,7 +788,7 @@ public void testInfer_UnauthorisedResponse() throws IOException { var model = OpenAiEmbeddingsModelTests.createModel(getUrl(webServer), "org", "secret", "model", "user"); PlainActionFuture listener = new PlainActionFuture<>(); - service.infer(model, List.of("abc"), new HashMap<>(), listener); + service.infer(model, List.of("abc"), new HashMap<>(), InputType.INGEST, listener); var error = expectThrows(ElasticsearchException.class, () -> listener.actionGet(TIMEOUT)); assertThat(error.getMessage(), containsString("Received an authentication error status code for request")); diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsModelTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsModelTests.java index 10e856ec8a27e..e2144132af6c1 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsModelTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsModelTests.java @@ -27,7 +27,7 @@ public void testOverrideWith_OverridesUser() { var model = createModel("url", "org", "api_key", "model_name", null); var requestTaskSettingsMap = getRequestTaskSettingsMap(null, "user_override"); - var overriddenModel = model.overrideWith(requestTaskSettingsMap); + var overriddenModel = OpenAiEmbeddingsModel.of(model, requestTaskSettingsMap); assertThat(overriddenModel, is(createModel("url", "org", "api_key", "model_name", "user_override"))); } @@ -37,14 +37,14 @@ public void testOverrideWith_EmptyMap() { var requestTaskSettingsMap = Map.of(); - var overriddenModel = model.overrideWith(requestTaskSettingsMap); + var overriddenModel = OpenAiEmbeddingsModel.of(model, requestTaskSettingsMap); assertThat(overriddenModel, sameInstance(model)); } public void testOverrideWith_NullMap() { var model = createModel("url", "org", "api_key", "model_name", null); - var overriddenModel = model.overrideWith(null); + var overriddenModel = OpenAiEmbeddingsModel.of(model, null); assertThat(overriddenModel, sameInstance(model)); } diff --git a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsTaskSettingsTests.java b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsTaskSettingsTests.java index f297eb622c421..103fab071098e 100644 --- a/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsTaskSettingsTests.java +++ b/x-pack/plugin/inference/src/test/java/org/elasticsearch/xpack/inference/services/openai/embeddings/OpenAiEmbeddingsTaskSettingsTests.java @@ -72,7 +72,7 @@ public void testOverrideWith_KeepsOriginalValuesWithOverridesAreNull() { new HashMap<>(Map.of(OpenAiEmbeddingsTaskSettings.MODEL, "model", OpenAiEmbeddingsTaskSettings.USER, "user")) ); - var overriddenTaskSettings = taskSettings.overrideWith(OpenAiEmbeddingsRequestTaskSettings.EMPTY_SETTINGS); + var overriddenTaskSettings = OpenAiEmbeddingsTaskSettings.of(taskSettings, OpenAiEmbeddingsRequestTaskSettings.EMPTY_SETTINGS); MatcherAssert.assertThat(overriddenTaskSettings, is(taskSettings)); } @@ -85,7 +85,7 @@ public void testOverrideWith_UsesOverriddenSettings() { new HashMap<>(Map.of(OpenAiEmbeddingsTaskSettings.MODEL, "model2", OpenAiEmbeddingsTaskSettings.USER, "user2")) ); - var overriddenTaskSettings = taskSettings.overrideWith(requestTaskSettings); + var overriddenTaskSettings = OpenAiEmbeddingsTaskSettings.of(taskSettings, requestTaskSettings); MatcherAssert.assertThat(overriddenTaskSettings, is(new OpenAiEmbeddingsTaskSettings("model2", "user2"))); } @@ -98,7 +98,7 @@ public void testOverrideWith_UsesOnlyNonNullModelSetting() { new HashMap<>(Map.of(OpenAiEmbeddingsTaskSettings.MODEL, "model2")) ); - var overriddenTaskSettings = taskSettings.overrideWith(requestTaskSettings); + var overriddenTaskSettings = OpenAiEmbeddingsTaskSettings.of(taskSettings, requestTaskSettings); MatcherAssert.assertThat(overriddenTaskSettings, is(new OpenAiEmbeddingsTaskSettings("model2", "user"))); }