diff --git a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java index 72e327ed6..f36a0eb11 100644 --- a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java +++ b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java @@ -177,7 +177,6 @@ public class AnthropicChatModel implements ChatModel { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) .provider(AnthropicApi.PROVIDER_NAME) - .requestOptions(prompt.getOptions()) .build(); ChatResponse response = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION @@ -240,7 +239,6 @@ public class AnthropicChatModel implements ChatModel { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) .provider(AnthropicApi.PROVIDER_NAME) - .requestOptions(prompt.getOptions()) .build(); Observation observation = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION.observation( diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java index e85f9e033..bed72e982 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java @@ -245,7 +245,6 @@ public class AzureOpenAiChatModel implements ChatModel { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) .provider(AiProvider.AZURE_OPENAI.value()) - .requestOptions(prompt.getOptions() != null ? prompt.getOptions() : this.defaultOptions) .build(); ChatResponse response = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION @@ -300,7 +299,6 @@ public class AzureOpenAiChatModel implements ChatModel { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) .provider(AiProvider.AZURE_OPENAI.value()) - .requestOptions(prompt.getOptions() != null ? prompt.getOptions() : this.defaultOptions) .build(); Observation observation = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION.observation( diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModel.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModel.java index a5f5b3357..c63ed598b 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModel.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModel.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2024 the original author or authors. + * Copyright 2023-2025 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -51,6 +51,7 @@ import org.springframework.util.CollectionUtils; * @author Mark Pollack * @author Christian Tzolov * @author Thomas Vitale + * @author Soby Chacko * @since 1.0.0 */ public class AzureOpenAiEmbeddingModel extends AbstractEmbeddingModel { @@ -124,12 +125,15 @@ public class AzureOpenAiEmbeddingModel extends AbstractEmbeddingModel { .from(this.defaultOptions) .merge(embeddingRequest.getOptions()) .build(); - EmbeddingsOptions azureOptions = options.toAzureOptions(embeddingRequest.getInstructions()); + + EmbeddingRequest embeddingRequestWithMergedOptions = new EmbeddingRequest(embeddingRequest.getInstructions(), + options); + + EmbeddingsOptions azureOptions = options.toAzureOptions(embeddingRequestWithMergedOptions.getInstructions()); var observationContext = EmbeddingModelObservationContext.builder() - .embeddingRequest(embeddingRequest) + .embeddingRequest(embeddingRequestWithMergedOptions) .provider(AiProvider.AZURE_OPENAI.value()) - .requestOptions(options) .build(); return EmbeddingModelObservationDocumentation.EMBEDDING_MODEL_OPERATION diff --git a/models/spring-ai-bedrock-converse/src/main/java/org/springframework/ai/bedrock/converse/BedrockProxyChatModel.java b/models/spring-ai-bedrock-converse/src/main/java/org/springframework/ai/bedrock/converse/BedrockProxyChatModel.java index 1cb4b2e54..0f4d136a1 100644 --- a/models/spring-ai-bedrock-converse/src/main/java/org/springframework/ai/bedrock/converse/BedrockProxyChatModel.java +++ b/models/spring-ai-bedrock-converse/src/main/java/org/springframework/ai/bedrock/converse/BedrockProxyChatModel.java @@ -220,7 +220,6 @@ public class BedrockProxyChatModel implements ChatModel { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) .provider(AiProvider.BEDROCK_CONVERSE.value()) - .requestOptions(prompt.getOptions()) .build(); ChatResponse chatResponse = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION @@ -647,7 +646,6 @@ public class BedrockProxyChatModel implements ChatModel { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) .provider(AiProvider.BEDROCK_CONVERSE.value()) - .requestOptions(prompt.getOptions()) .build(); Observation observation = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION.observation( diff --git a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java index 740666389..e5a774cac 100644 --- a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java +++ b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java @@ -241,7 +241,6 @@ public class MiniMaxChatModel implements ChatModel { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(requestPrompt) .provider(MiniMaxApiConstants.PROVIDER_NAME) - .requestOptions(requestPrompt.getOptions()) .build(); ChatResponse response = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION @@ -334,7 +333,6 @@ public class MiniMaxChatModel implements ChatModel { final ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(requestPrompt) .provider(MiniMaxApiConstants.PROVIDER_NAME) - .requestOptions(requestPrompt.getOptions()) .build(); Observation observation = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION.observation( diff --git a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxEmbeddingModel.java b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxEmbeddingModel.java index fec3b0c31..9a5983785 100644 --- a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxEmbeddingModel.java +++ b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxEmbeddingModel.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2024 the original author or authors. + * Copyright 2023-2025 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -43,13 +43,15 @@ import org.springframework.ai.retry.RetryUtils; import org.springframework.lang.Nullable; import org.springframework.retry.support.RetryTemplate; import org.springframework.util.Assert; +import org.springframework.util.StringUtils; /** * MiniMax Embedding Model implementation. * * @author Geng Rong * @author Thomas Vitale - * @since 1.0.0 M1 + * @author Soby Chacko + * @since 1.0.0 */ public class MiniMaxEmbeddingModel extends AbstractEmbeddingModel { @@ -149,14 +151,15 @@ public class MiniMaxEmbeddingModel extends AbstractEmbeddingModel { @Override public EmbeddingResponse call(EmbeddingRequest request) { - MiniMaxEmbeddingOptions requestOptions = mergeOptions(request.getOptions(), this.defaultOptions); + + EmbeddingRequest embeddingRequest = buildEmbeddingRequest(request); + MiniMaxApi.EmbeddingRequest apiRequest = new MiniMaxApi.EmbeddingRequest(request.getInstructions(), - requestOptions.getModel()); + embeddingRequest.getOptions().getModel()); var observationContext = EmbeddingModelObservationContext.builder() .embeddingRequest(request) .provider(MiniMaxApiConstants.PROVIDER_NAME) - .requestOptions(requestOptions) .build(); return EmbeddingModelObservationDocumentation.EMBEDDING_MODEL_OPERATION @@ -188,26 +191,24 @@ public class MiniMaxEmbeddingModel extends AbstractEmbeddingModel { return new DefaultUsage(0, 0, apiEmbeddingList.totalTokens()); } - /** - * Merge runtime and default {@link EmbeddingOptions} to compute the final options to - * use in the request. - */ - private MiniMaxEmbeddingOptions mergeOptions(@Nullable EmbeddingOptions runtimeOptions, - MiniMaxEmbeddingOptions defaultOptions) { - var runtimeOptionsForProvider = ModelOptionsUtils.copyToTarget(runtimeOptions, EmbeddingOptions.class, + EmbeddingRequest buildEmbeddingRequest(EmbeddingRequest embeddingRequest) { + // Process runtime options + MiniMaxEmbeddingOptions runtimeOptions = null; + if (embeddingRequest.getOptions() != null) { + runtimeOptions = ModelOptionsUtils.copyToTarget(embeddingRequest.getOptions(), EmbeddingOptions.class, + MiniMaxEmbeddingOptions.class); + } + + // Define request options by merging runtime options and default options + MiniMaxEmbeddingOptions requestOptions = ModelOptionsUtils.merge(runtimeOptions, this.defaultOptions, MiniMaxEmbeddingOptions.class); - var optionBuilder = MiniMaxEmbeddingOptions.builder(); - if (runtimeOptionsForProvider != null && runtimeOptionsForProvider.getModel() != null) { - optionBuilder.model(runtimeOptionsForProvider.getModel()); + // Validate request options + if (!StringUtils.hasText(requestOptions.getModel())) { + throw new IllegalArgumentException("model cannot be null or empty"); } - else if (defaultOptions.getModel() != null) { - optionBuilder.model(defaultOptions.getModel()); - } - else { - optionBuilder.model(MiniMaxApi.DEFAULT_EMBEDDING_MODEL); - } - return optionBuilder.build(); + + return new EmbeddingRequest(embeddingRequest.getInstructions(), requestOptions); } public void setObservationConvention(EmbeddingModelObservationConvention observationConvention) { diff --git a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java index 5e860f33a..9f165e27c 100644 --- a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java +++ b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2024 the original author or authors. + * Copyright 2023-2025 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -29,6 +29,7 @@ import reactor.core.publisher.Flux; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.document.MetadataMode; +import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.minimax.MiniMaxChatModel; import org.springframework.ai.minimax.MiniMaxChatOptions; import org.springframework.ai.minimax.MiniMaxEmbeddingModel; @@ -57,6 +58,7 @@ import static org.mockito.BDDMockito.given; /** * @author Geng Rong + * @author Soby Chacko */ @SuppressWarnings("unchecked") @ExtendWith(MockitoExtension.class) @@ -150,8 +152,9 @@ public class MiniMaxRetryTests { .willThrow(new TransientAiException("Transient Error 2")) .willReturn(ResponseEntity.of(Optional.of(expectedEmbeddings))); + EmbeddingOptions options = MiniMaxEmbeddingOptions.builder().model("model").build(); var result = this.embeddingModel - .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), null)); + .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), options)); assertThat(result).isNotNull(); assertThat(result.getResult().getOutput()).isEqualTo(new float[] { 9.9f, 8.8f }); @@ -163,8 +166,9 @@ public class MiniMaxRetryTests { public void miniMaxEmbeddingNonTransientError() { given(this.miniMaxApi.embeddings(isA(EmbeddingRequest.class))) .willThrow(new RuntimeException("Non Transient Error")); + EmbeddingOptions options = MiniMaxEmbeddingOptions.builder().model("model").build(); assertThrows(RuntimeException.class, () -> this.embeddingModel - .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), null))); + .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), options))); } private class TestRetryListener implements RetryListener { diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java index 7b8a3ee91..738e2640d 100644 --- a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java +++ b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java @@ -187,7 +187,6 @@ public class MistralAiChatModel implements ChatModel { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) .provider(MistralAiApi.PROVIDER_NAME) - .requestOptions(prompt.getOptions()) .build(); ChatResponse response = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION @@ -260,7 +259,6 @@ public class MistralAiChatModel implements ChatModel { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) .provider(MistralAiApi.PROVIDER_NAME) - .requestOptions(prompt.getOptions()) .build(); Observation observation = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION.observation( diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiEmbeddingModel.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiEmbeddingModel.java index 0908f57c5..347128f0e 100644 --- a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiEmbeddingModel.java +++ b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiEmbeddingModel.java @@ -116,9 +116,8 @@ public class MistralAiEmbeddingModel extends AbstractEmbeddingModel { var apiRequest = createRequest(embeddingRequest); var observationContext = EmbeddingModelObservationContext.builder() - .embeddingRequest(request) + .embeddingRequest(embeddingRequest) .provider(MistralAiApi.PROVIDER_NAME) - .requestOptions(embeddingRequest.getOptions()) .build(); return EmbeddingModelObservationDocumentation.EMBEDDING_MODEL_OPERATION diff --git a/models/spring-ai-oci-genai/src/main/java/org/springframework/ai/oci/OCIEmbeddingModel.java b/models/spring-ai-oci-genai/src/main/java/org/springframework/ai/oci/OCIEmbeddingModel.java index 81b02107c..8ed9299a7 100644 --- a/models/spring-ai-oci-genai/src/main/java/org/springframework/ai/oci/OCIEmbeddingModel.java +++ b/models/spring-ai-oci-genai/src/main/java/org/springframework/ai/oci/OCIEmbeddingModel.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2024 the original author or authors. + * Copyright 2023-2025 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -43,6 +43,7 @@ import org.springframework.ai.embedding.observation.EmbeddingModelObservationDoc import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.observation.conventions.AiProvider; import org.springframework.util.Assert; +import org.springframework.util.StringUtils; /** * {@link org.springframework.ai.embedding.EmbeddingModel} implementation that uses the @@ -83,13 +84,15 @@ public class OCIEmbeddingModel extends AbstractEmbeddingModel { @Override public EmbeddingResponse call(EmbeddingRequest request) { Assert.notEmpty(request.getInstructions(), "At least one text is required!"); - OCIEmbeddingOptions runtimeOptions = mergeOptions(request.getOptions(), this.options); - List embedTextRequests = createRequests(request.getInstructions(), runtimeOptions); + + EmbeddingRequest embeddingRequest = buildEmbeddingRequest(request); + + List embedTextRequests = createRequests(embeddingRequest.getInstructions(), + (OCIEmbeddingOptions) embeddingRequest.getOptions()); EmbeddingModelObservationContext context = EmbeddingModelObservationContext.builder() - .embeddingRequest(request) + .embeddingRequest(embeddingRequest) .provider(AiProvider.OCI_GENAI.value()) - .requestOptions(runtimeOptions) .build(); return EmbeddingModelObservationDocumentation.EMBEDDING_MODEL_OPERATION @@ -158,6 +161,26 @@ public class OCIEmbeddingModel extends AbstractEmbeddingModel { return defaultOptions; } + EmbeddingRequest buildEmbeddingRequest(EmbeddingRequest embeddingRequest) { + // Process runtime options + OCIEmbeddingOptions runtimeOptions = null; + if (embeddingRequest.getOptions() != null) { + runtimeOptions = ModelOptionsUtils.copyToTarget(embeddingRequest.getOptions(), EmbeddingOptions.class, + OCIEmbeddingOptions.class); + } + + // Define request options by merging runtime options and default options + OCIEmbeddingOptions requestOptions = ModelOptionsUtils.merge(runtimeOptions, this.options, + OCIEmbeddingOptions.class); + + // Validate request options + if (!StringUtils.hasText(requestOptions.getModel())) { + throw new IllegalArgumentException("model cannot be null or empty"); + } + + return new EmbeddingRequest(embeddingRequest.getInstructions(), requestOptions); + } + private float[] toFloats(List embedding) { float[] floats = new float[embedding.size()]; for (int i = 0; i < embedding.size(); i++) { diff --git a/models/spring-ai-oci-genai/src/main/java/org/springframework/ai/oci/cohere/OCICohereChatModel.java b/models/spring-ai-oci-genai/src/main/java/org/springframework/ai/oci/cohere/OCICohereChatModel.java index d3a23b5d9..462a9773a 100644 --- a/models/spring-ai-oci-genai/src/main/java/org/springframework/ai/oci/cohere/OCICohereChatModel.java +++ b/models/spring-ai-oci-genai/src/main/java/org/springframework/ai/oci/cohere/OCICohereChatModel.java @@ -53,6 +53,7 @@ import org.springframework.ai.chat.observation.DefaultChatModelObservationConven import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.ModelOptionsUtils; +import org.springframework.ai.model.tool.ToolCallingChatOptions; import org.springframework.ai.observation.conventions.AiProvider; import org.springframework.ai.oci.ServingModeHelper; import org.springframework.util.Assert; @@ -104,10 +105,10 @@ public class OCICohereChatModel implements ChatModel { @Override public ChatResponse call(Prompt prompt) { + Prompt requestPrompt = this.buildRequestPrompt(prompt); ChatModelObservationContext observationContext = ChatModelObservationContext.builder() - .prompt(prompt) + .prompt(requestPrompt) .provider(AiProvider.OCI_GENAI.value()) - .requestOptions(prompt.getOptions() != null ? prompt.getOptions() : this.defaultOptions) .build(); return ChatModelObservationDocumentation.CHAT_MODEL_OPERATION @@ -120,6 +121,21 @@ public class OCICohereChatModel implements ChatModel { }); } + Prompt buildRequestPrompt(Prompt prompt) { + // Process runtime options + OCICohereChatOptions runtimeOptions = null; + if (prompt.getOptions() != null) { + runtimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, + OCICohereChatOptions.class); + } + + // Define request options by merging runtime options and default options + OCICohereChatOptions requestOptions = ModelOptionsUtils.merge(runtimeOptions, this.defaultOptions, + OCICohereChatOptions.class); + + return new Prompt(prompt.getInstructions(), requestOptions); + } + @Override public ChatOptions getDefaultOptions() { return OCICohereChatOptions.fromOptions(this.defaultOptions); diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java index bc886108f..051dabde9 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java @@ -225,7 +225,6 @@ public class OllamaChatModel implements ChatModel { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) .provider(OllamaApi.PROVIDER_NAME) - .requestOptions(prompt.getOptions()) .build(); ChatResponse response = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION @@ -296,7 +295,6 @@ public class OllamaChatModel implements ChatModel { final ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) .provider(OllamaApi.PROVIDER_NAME) - .requestOptions(prompt.getOptions()) .build(); Observation observation = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION.observation( diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaEmbeddingModel.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaEmbeddingModel.java index da0408782..4a5710c9a 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaEmbeddingModel.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaEmbeddingModel.java @@ -113,7 +113,6 @@ public class OllamaEmbeddingModel extends AbstractEmbeddingModel { var observationContext = EmbeddingModelObservationContext.builder() .embeddingRequest(request) .provider(OllamaApi.PROVIDER_NAME) - .requestOptions(embeddingRequest.getOptions()) .build(); return EmbeddingModelObservationDocumentation.EMBEDDING_MODEL_OPERATION diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java index 87636a18a..5fbb8a283 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java @@ -186,7 +186,6 @@ public class OpenAiChatModel implements ChatModel { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) .provider(OpenAiApiConstants.PROVIDER_NAME) - .requestOptions(prompt.getOptions()) .build(); ChatResponse response = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION @@ -289,7 +288,6 @@ public class OpenAiChatModel implements ChatModel { final ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) .provider(OpenAiApiConstants.PROVIDER_NAME) - .requestOptions(prompt.getOptions()) .build(); Observation observation = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION.observation( diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingModel.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingModel.java index 2ac56916f..47c06ac5a 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingModel.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingModel.java @@ -156,9 +156,8 @@ public class OpenAiEmbeddingModel extends AbstractEmbeddingModel { OpenAiApi.EmbeddingRequest> apiRequest = createRequest(embeddingRequest); var observationContext = EmbeddingModelObservationContext.builder() - .embeddingRequest(request) + .embeddingRequest(embeddingRequest) .provider(OpenAiApiConstants.PROVIDER_NAME) - .requestOptions(embeddingRequest.getOptions()) .build(); return EmbeddingModelObservationDocumentation.EMBEDDING_MODEL_OPERATION diff --git a/models/spring-ai-transformers/src/main/java/org/springframework/ai/transformers/TransformersEmbeddingModel.java b/models/spring-ai-transformers/src/main/java/org/springframework/ai/transformers/TransformersEmbeddingModel.java index 5fd6f4a8b..a7324cad7 100644 --- a/models/spring-ai-transformers/src/main/java/org/springframework/ai/transformers/TransformersEmbeddingModel.java +++ b/models/spring-ai-transformers/src/main/java/org/springframework/ai/transformers/TransformersEmbeddingModel.java @@ -289,7 +289,6 @@ public class TransformersEmbeddingModel extends AbstractEmbeddingModel implement var observationContext = EmbeddingModelObservationContext.builder() .embeddingRequest(request) .provider(AiProvider.ONNX.value()) - .requestOptions(request.getOptions()) .build(); return EmbeddingModelObservationDocumentation.EMBEDDING_MODEL_OPERATION diff --git a/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingModel.java b/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingModel.java index 836137095..4bef9d114 100644 --- a/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingModel.java +++ b/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingModel.java @@ -35,6 +35,7 @@ import org.springframework.ai.chat.metadata.Usage; import org.springframework.ai.document.Document; import org.springframework.ai.embedding.AbstractEmbeddingModel; import org.springframework.ai.embedding.Embedding; +import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingRequest; import org.springframework.ai.embedding.EmbeddingResponse; import org.springframework.ai.embedding.EmbeddingResponseMetadata; @@ -59,6 +60,7 @@ import org.springframework.util.StringUtils; * @author Christian Tzolov * @author Mark Pollack * @author Rodrigo Malara + * @author Soby Chacko * @since 1.0.0 */ public class VertexAiTextEmbeddingModel extends AbstractEmbeddingModel { @@ -117,12 +119,11 @@ public class VertexAiTextEmbeddingModel extends AbstractEmbeddingModel { @Override public EmbeddingResponse call(EmbeddingRequest request) { - final VertexAiTextEmbeddingOptions finalOptions = mergedOptions(request); + EmbeddingRequest embeddingRequest = buildEmbeddingRequest(request); var observationContext = EmbeddingModelObservationContext.builder() - .embeddingRequest(request) + .embeddingRequest(embeddingRequest) .provider(AiProvider.VERTEX_AI.value()) - .requestOptions(finalOptions) .build(); return EmbeddingModelObservationDocumentation.EMBEDDING_MODEL_OPERATION @@ -131,10 +132,11 @@ public class VertexAiTextEmbeddingModel extends AbstractEmbeddingModel { .observe(() -> { try (PredictionServiceClient client = createPredictionServiceClient()) { - EndpointName endpointName = this.connectionDetails.getEndpointName(finalOptions.getModel()); + EmbeddingOptions options = embeddingRequest.getOptions(); + EndpointName endpointName = this.connectionDetails.getEndpointName(options.getModel()); PredictRequest.Builder predictRequestBuilder = getPredictRequestBuilder(request, endpointName, - finalOptions); + (VertexAiTextEmbeddingOptions) options); PredictResponse embeddingResponse = this.retryTemplate .execute(context -> getPredictResponse(client, predictRequestBuilder)); @@ -155,7 +157,7 @@ public class VertexAiTextEmbeddingModel extends AbstractEmbeddingModel { embeddingList.add(new Embedding(vectorValues, index++)); } EmbeddingResponse response = new EmbeddingResponse(embeddingList, - generateResponseMetadata(finalOptions.getModel(), totalTokenCount)); + generateResponseMetadata(options.getModel(), totalTokenCount)); observationContext.setResponse(response); @@ -164,17 +166,24 @@ public class VertexAiTextEmbeddingModel extends AbstractEmbeddingModel { }); } - private VertexAiTextEmbeddingOptions mergedOptions(EmbeddingRequest request) { - - VertexAiTextEmbeddingOptions mergedOptions = this.defaultOptions; - - if (request.getOptions() != null) { - var defaultOptionsCopy = VertexAiTextEmbeddingOptions.builder().from(this.defaultOptions).build(); - mergedOptions = ModelOptionsUtils.merge(request.getOptions(), defaultOptionsCopy, + EmbeddingRequest buildEmbeddingRequest(EmbeddingRequest embeddingRequest) { + // Process runtime options + VertexAiTextEmbeddingOptions runtimeOptions = null; + if (embeddingRequest.getOptions() != null) { + runtimeOptions = ModelOptionsUtils.copyToTarget(embeddingRequest.getOptions(), EmbeddingOptions.class, VertexAiTextEmbeddingOptions.class); } - return mergedOptions; + // Define request options by merging runtime options and default options + VertexAiTextEmbeddingOptions requestOptions = ModelOptionsUtils.merge(runtimeOptions, this.defaultOptions, + VertexAiTextEmbeddingOptions.class); + + // Validate request options + if (!StringUtils.hasText(requestOptions.getModel())) { + throw new IllegalArgumentException("model cannot be null or empty"); + } + + return new EmbeddingRequest(embeddingRequest.getInstructions(), requestOptions); } protected PredictRequest.Builder getPredictRequestBuilder(EmbeddingRequest request, EndpointName endpointName, diff --git a/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingRetryTests.java b/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingRetryTests.java index 3430791d5..088c87bc7 100644 --- a/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingRetryTests.java +++ b/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingRetryTests.java @@ -30,6 +30,7 @@ import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingRequest; import org.springframework.ai.embedding.EmbeddingResponse; import org.springframework.ai.retry.RetryUtils; @@ -116,7 +117,8 @@ public class VertexAiTextEmbeddingRetryTests { .willThrow(new TransientAiException("Transient Error 2")) .willReturn(mockResponse); - EmbeddingResponse result = this.embeddingModel.call(new EmbeddingRequest(List.of("text1", "text2"), null)); + EmbeddingOptions options = VertexAiTextEmbeddingOptions.builder().model("model").build(); + EmbeddingResponse result = this.embeddingModel.call(new EmbeddingRequest(List.of("text1", "text2"), options)); assertThat(result).isNotNull(); assertThat(result.getResults()).hasSize(1); @@ -132,8 +134,9 @@ public class VertexAiTextEmbeddingRetryTests { // Setup the mock PredictionServiceClient to throw a non-transient error given(this.mockPredictionServiceClient.predict(any())).willThrow(new RuntimeException("Non Transient Error")); + EmbeddingOptions options = VertexAiTextEmbeddingOptions.builder().model("model").build(); // Assert that a RuntimeException is thrown and not retried - assertThatThrownBy(() -> this.embeddingModel.call(new EmbeddingRequest(List.of("text1", "text2"), null))) + assertThatThrownBy(() -> this.embeddingModel.call(new EmbeddingRequest(List.of("text1", "text2"), options))) .isInstanceOf(RuntimeException.class); // Verify that predict was called only once (no retries for non-transient errors) diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java index 0ecd83848..36fcfd649 100644 --- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java +++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java @@ -366,7 +366,6 @@ public class VertexAiGeminiChatModel implements ChatModel, DisposableBean { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) .provider(VertexAiGeminiConstants.PROVIDER_NAME) - .requestOptions(prompt.getOptions()) .build(); ChatResponse response = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION @@ -478,7 +477,6 @@ public class VertexAiGeminiChatModel implements ChatModel, DisposableBean { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(prompt) .provider(VertexAiGeminiConstants.PROVIDER_NAME) - .requestOptions(prompt.getOptions()) .build(); Observation observation = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION.observation( diff --git a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatModel.java b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatModel.java index 57a9527ca..408666fdc 100644 --- a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatModel.java +++ b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatModel.java @@ -242,7 +242,6 @@ public class ZhiPuAiChatModel implements ChatModel { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(requestPrompt) .provider(ZhiPuApiConstants.PROVIDER_NAME) - .requestOptions(prompt.getOptions()) .build(); ChatResponse response = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION @@ -319,7 +318,6 @@ public class ZhiPuAiChatModel implements ChatModel { final ChatModelObservationContext observationContext = ChatModelObservationContext.builder() .prompt(requestPrompt) .provider(ZhiPuApiConstants.PROVIDER_NAME) - .requestOptions(buildRequestOptions(request)) .build(); Observation observation = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION.observation( diff --git a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiEmbeddingModel.java b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiEmbeddingModel.java index f310f65cc..f8c3f6205 100644 --- a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiEmbeddingModel.java +++ b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiEmbeddingModel.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2024 the original author or authors. + * Copyright 2023-2025 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -41,15 +41,16 @@ import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.retry.RetryUtils; import org.springframework.ai.zhipuai.api.ZhiPuAiApi; import org.springframework.ai.zhipuai.api.ZhiPuApiConstants; -import org.springframework.lang.Nullable; import org.springframework.retry.support.RetryTemplate; import org.springframework.util.Assert; +import org.springframework.util.StringUtils; /** * ZhiPuAI Embedding Model implementation. * * @author Geng Rong - * @since 1.0.0 M1 + * @author Soby Chacko + * @since 1.0.0 */ public class ZhiPuAiEmbeddingModel extends AbstractEmbeddingModel { @@ -153,12 +154,12 @@ public class ZhiPuAiEmbeddingModel extends AbstractEmbeddingModel { logger.warn( "ZhiPu Embedding does not support batch embedding. Will make multiple API calls to embed(Document)"); } - ZhiPuAiEmbeddingOptions requestOptions = mergeOptions(request.getOptions(), this.defaultOptions); + + EmbeddingRequest embeddingRequest = buildEmbeddingRequest(request); var observationContext = EmbeddingModelObservationContext.builder() - .embeddingRequest(request) + .embeddingRequest(embeddingRequest) .provider(ZhiPuApiConstants.PROVIDER_NAME) - .requestOptions(requestOptions) .build(); return EmbeddingModelObservationDocumentation.EMBEDDING_MODEL_OPERATION @@ -170,7 +171,7 @@ public class ZhiPuAiEmbeddingModel extends AbstractEmbeddingModel { var totalUsage = new ZhiPuAiApi.Usage(0, 0, 0); for (String inputContent : request.getInstructions()) { - var apiRequest = createEmbeddingRequest(inputContent, requestOptions); + var apiRequest = createEmbeddingRequest(inputContent, embeddingRequest.getOptions()); ZhiPuAiApi.EmbeddingList response = this.retryTemplate .execute(ctx -> this.zhiPuAiApi.embeddings(apiRequest).getBody()); @@ -210,24 +211,24 @@ public class ZhiPuAiEmbeddingModel extends AbstractEmbeddingModel { return new DefaultUsage(usage.promptTokens(), usage.completionTokens(), usage.totalTokens(), usage); } - /** - * Merge runtime and default {@link EmbeddingOptions} to compute the final options to - * use in the request. - */ - private ZhiPuAiEmbeddingOptions mergeOptions(@Nullable EmbeddingOptions runtimeOptions, - ZhiPuAiEmbeddingOptions defaultOptions) { - var runtimeOptionsForProvider = ModelOptionsUtils.copyToTarget(runtimeOptions, EmbeddingOptions.class, - ZhiPuAiEmbeddingOptions.class); - - if (runtimeOptionsForProvider == null) { - return defaultOptions; + EmbeddingRequest buildEmbeddingRequest(EmbeddingRequest embeddingRequest) { + // Process runtime options + ZhiPuAiEmbeddingOptions runtimeOptions = null; + if (embeddingRequest.getOptions() != null) { + runtimeOptions = ModelOptionsUtils.copyToTarget(embeddingRequest.getOptions(), EmbeddingOptions.class, + ZhiPuAiEmbeddingOptions.class); } - return ZhiPuAiEmbeddingOptions.builder() - .model(ModelOptionsUtils.mergeOption(runtimeOptionsForProvider.getModel(), defaultOptions.getModel())) - .dimensions(ModelOptionsUtils.mergeOption(runtimeOptionsForProvider.getDimensions(), - defaultOptions.getDimensions())) - .build(); + // Define request options by merging runtime options and default options + ZhiPuAiEmbeddingOptions requestOptions = ModelOptionsUtils.merge(runtimeOptions, this.defaultOptions, + ZhiPuAiEmbeddingOptions.class); + + // Validate request options + if (!StringUtils.hasText(requestOptions.getModel())) { + throw new IllegalArgumentException("model cannot be null or empty"); + } + + return new EmbeddingRequest(embeddingRequest.getInstructions(), requestOptions); } private ZhiPuAiApi.EmbeddingRequest createEmbeddingRequest(String text, EmbeddingOptions requestOptions) { diff --git a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/api/ZhiPuAiRetryTests.java b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/api/ZhiPuAiRetryTests.java index 3ef3225e6..b78db1620 100644 --- a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/api/ZhiPuAiRetryTests.java +++ b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/api/ZhiPuAiRetryTests.java @@ -29,6 +29,7 @@ import reactor.core.publisher.Flux; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.document.MetadataMode; +import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.image.ImageMessage; import org.springframework.ai.image.ImagePrompt; import org.springframework.ai.retry.RetryUtils; @@ -164,9 +165,9 @@ public class ZhiPuAiRetryTests { .willThrow(new TransientAiException("Transient Error 1")) .willThrow(new TransientAiException("Transient Error 2")) .willReturn(ResponseEntity.of(Optional.of(expectedEmbeddings))); - + EmbeddingOptions options = ZhiPuAiEmbeddingOptions.builder().model("model").build(); var result = this.embeddingModel - .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), null)); + .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), options)); assertThat(result).isNotNull(); assertThat(result.getResult().getOutput()).isEqualTo(new float[] { 9.9f, 8.8f }); @@ -178,8 +179,9 @@ public class ZhiPuAiRetryTests { public void zhiPuAiEmbeddingNonTransientError() { given(this.zhiPuAiApi.embeddings(isA(EmbeddingRequest.class))) .willThrow(new RuntimeException("Non Transient Error")); + EmbeddingOptions options = ZhiPuAiEmbeddingOptions.builder().model("model").build(); assertThrows(RuntimeException.class, () -> this.embeddingModel - .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), null))); + .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), options))); } @Test diff --git a/spring-ai-model/src/main/java/org/springframework/ai/chat/observation/ChatModelObservationContext.java b/spring-ai-model/src/main/java/org/springframework/ai/chat/observation/ChatModelObservationContext.java index 64689201a..819edec41 100644 --- a/spring-ai-model/src/main/java/org/springframework/ai/chat/observation/ChatModelObservationContext.java +++ b/spring-ai-model/src/main/java/org/springframework/ai/chat/observation/ChatModelObservationContext.java @@ -32,35 +32,21 @@ import org.springframework.util.Assert; */ public class ChatModelObservationContext extends ModelObservationContext { - private final ChatOptions requestOptions; - - ChatModelObservationContext(Prompt prompt, String provider, ChatOptions requestOptions) { + ChatModelObservationContext(Prompt prompt, String provider) { super(prompt, AiOperationMetadata.builder().operationType(AiOperationType.CHAT.value()).provider(provider).build()); - Assert.notNull(requestOptions, "requestOptions cannot be null"); - this.requestOptions = requestOptions; } public static Builder builder() { return new Builder(); } - /** - * @deprecated Use {@link #getRequest().getOptions()} instead. - */ - @Deprecated(forRemoval = true) - public ChatOptions getRequestOptions() { - return this.requestOptions; - } - public static final class Builder { private Prompt prompt; private String provider; - private ChatOptions requestOptions; - private Builder() { } @@ -74,18 +60,8 @@ public class ChatModelObservationContext extends ModelObservationContext stopSequencesJoiner.add("\"" + value + "\"")); + options.getStopSequences().forEach(value -> stopSequencesJoiner.add("\"" + value + "\"")); KeyValue.of(ChatModelObservationDocumentation.HighCardinalityKeyNames.REQUEST_STOP_SEQUENCES, - context.getRequestOptions().getStopSequences(), Objects::nonNull); + options.getStopSequences(), Objects::nonNull); return keyValues.and( ChatModelObservationDocumentation.HighCardinalityKeyNames.REQUEST_STOP_SEQUENCES.asString(), stopSequencesJoiner.toString()); @@ -153,26 +158,29 @@ public class DefaultChatModelObservationConvention implements ChatModelObservati } protected KeyValues requestTemperature(KeyValues keyValues, ChatModelObservationContext context) { - if (context.getRequestOptions().getTemperature() != null) { + ChatOptions options = context.getRequest().getOptions(); + if (options.getTemperature() != null) { return keyValues.and( ChatModelObservationDocumentation.HighCardinalityKeyNames.REQUEST_TEMPERATURE.asString(), - String.valueOf(context.getRequestOptions().getTemperature())); + String.valueOf(options.getTemperature())); } return keyValues; } protected KeyValues requestTopK(KeyValues keyValues, ChatModelObservationContext context) { - if (context.getRequestOptions().getTopK() != null) { + ChatOptions options = context.getRequest().getOptions(); + if (options.getTopK() != null) { return keyValues.and(ChatModelObservationDocumentation.HighCardinalityKeyNames.REQUEST_TOP_K.asString(), - String.valueOf(context.getRequestOptions().getTopK())); + String.valueOf(options.getTopK())); } return keyValues; } protected KeyValues requestTopP(KeyValues keyValues, ChatModelObservationContext context) { - if (context.getRequestOptions().getTopP() != null) { + ChatOptions options = context.getRequest().getOptions(); + if (options.getTopP() != null) { return keyValues.and(ChatModelObservationDocumentation.HighCardinalityKeyNames.REQUEST_TOP_P.asString(), - String.valueOf(context.getRequestOptions().getTopP())); + String.valueOf(options.getTopP())); } return keyValues; } diff --git a/spring-ai-model/src/main/java/org/springframework/ai/embedding/observation/DefaultEmbeddingModelObservationConvention.java b/spring-ai-model/src/main/java/org/springframework/ai/embedding/observation/DefaultEmbeddingModelObservationConvention.java index 6949f0e00..97d5146a6 100644 --- a/spring-ai-model/src/main/java/org/springframework/ai/embedding/observation/DefaultEmbeddingModelObservationConvention.java +++ b/spring-ai-model/src/main/java/org/springframework/ai/embedding/observation/DefaultEmbeddingModelObservationConvention.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2024 the original author or authors. + * Copyright 2023-2025 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -25,6 +25,7 @@ import org.springframework.util.StringUtils; * Default conventions to populate observations for embedding model operations. * * @author Thomas Vitale + * @author Soby Chacko * @since 1.0.0 */ public class DefaultEmbeddingModelObservationConvention implements EmbeddingModelObservationConvention { @@ -44,9 +45,9 @@ public class DefaultEmbeddingModelObservationConvention implements EmbeddingMode @Override public String getContextualName(EmbeddingModelObservationContext context) { - if (StringUtils.hasText(context.getRequestOptions().getModel())) { + if (StringUtils.hasText(context.getRequest().getOptions().getModel())) { return "%s %s".formatted(context.getOperationMetadata().operationType(), - context.getRequestOptions().getModel()); + context.getRequest().getOptions().getModel()); } return context.getOperationMetadata().operationType(); } @@ -68,9 +69,9 @@ public class DefaultEmbeddingModelObservationConvention implements EmbeddingMode } protected KeyValue requestModel(EmbeddingModelObservationContext context) { - if (StringUtils.hasText(context.getRequestOptions().getModel())) { + if (StringUtils.hasText(context.getRequest().getOptions().getModel())) { return KeyValue.of(EmbeddingModelObservationDocumentation.LowCardinalityKeyNames.REQUEST_MODEL, - context.getRequestOptions().getModel()); + context.getRequest().getOptions().getModel()); } return REQUEST_MODEL_NONE; } @@ -98,10 +99,10 @@ public class DefaultEmbeddingModelObservationConvention implements EmbeddingMode // Request protected KeyValues requestEmbeddingDimension(KeyValues keyValues, EmbeddingModelObservationContext context) { - if (context.getRequestOptions().getDimensions() != null) { + if (context.getRequest().getOptions().getDimensions() != null) { return keyValues .and(EmbeddingModelObservationDocumentation.HighCardinalityKeyNames.REQUEST_EMBEDDING_DIMENSIONS - .asString(), String.valueOf(context.getRequestOptions().getDimensions())); + .asString(), String.valueOf(context.getRequest().getOptions().getDimensions())); } return keyValues; } diff --git a/spring-ai-model/src/main/java/org/springframework/ai/embedding/observation/EmbeddingModelObservationContext.java b/spring-ai-model/src/main/java/org/springframework/ai/embedding/observation/EmbeddingModelObservationContext.java index 07bc7b0ed..bc35c7253 100644 --- a/spring-ai-model/src/main/java/org/springframework/ai/embedding/observation/EmbeddingModelObservationContext.java +++ b/spring-ai-model/src/main/java/org/springframework/ai/embedding/observation/EmbeddingModelObservationContext.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2024 the original author or authors. + * Copyright 2023-2025 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -22,49 +22,34 @@ import org.springframework.ai.embedding.EmbeddingResponse; import org.springframework.ai.model.observation.ModelObservationContext; import org.springframework.ai.observation.AiOperationMetadata; import org.springframework.ai.observation.conventions.AiOperationType; -import org.springframework.util.Assert; /** * Context used to store metadata for embedding model exchanges. * * @author Thomas Vitale + * @author Soby Chacko * @since 1.0.0 */ public class EmbeddingModelObservationContext extends ModelObservationContext { - private final EmbeddingOptions requestOptions; - - EmbeddingModelObservationContext(EmbeddingRequest embeddingRequest, String provider, - EmbeddingOptions requestOptions) { + EmbeddingModelObservationContext(EmbeddingRequest embeddingRequest, String provider) { super(embeddingRequest, AiOperationMetadata.builder() .operationType(AiOperationType.EMBEDDING.value()) .provider(provider) .build()); - Assert.notNull(requestOptions, "requestOptions cannot be null"); - this.requestOptions = requestOptions; } public static Builder builder() { return new Builder(); } - /** - * @deprecated Use {@link #getRequest().getOptions()} instead. - */ - @Deprecated(forRemoval = true) - public EmbeddingOptions getRequestOptions() { - return this.requestOptions; - } - public static final class Builder { private EmbeddingRequest embeddingRequest; private String provider; - private EmbeddingOptions requestOptions; - private Builder() { } @@ -78,18 +63,8 @@ public class EmbeddingModelObservationContext extends ModelObservationContext ChatModelObservationContext.builder() - .prompt(generatePrompt()) - .provider("superprovider") - .requestOptions(null) - .build()).isInstanceOf(IllegalArgumentException.class) - .hasMessageContaining("requestOptions cannot be null"); - } - - private Prompt generatePrompt() { - return new Prompt("hello"); + private Prompt generatePrompt(ChatOptions chatOptions) { + return new Prompt("hello", chatOptions); } } diff --git a/spring-ai-model/src/test/java/org/springframework/ai/chat/observation/ChatModelPromptContentObservationFilterTests.java b/spring-ai-model/src/test/java/org/springframework/ai/chat/observation/ChatModelPromptContentObservationFilterTests.java index c05dd3ef9..638f1320b 100644 --- a/spring-ai-model/src/test/java/org/springframework/ai/chat/observation/ChatModelPromptContentObservationFilterTests.java +++ b/spring-ai-model/src/test/java/org/springframework/ai/chat/observation/ChatModelPromptContentObservationFilterTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2024 the original author or authors. + * Copyright 2023-2025 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -50,9 +50,8 @@ class ChatModelPromptContentObservationFilterTests { @Test void whenEmptyPromptThenReturnOriginalContext() { var expectedContext = ChatModelObservationContext.builder() - .prompt(new Prompt(List.of())) + .prompt(new Prompt(List.of(), ChatOptions.builder().model("mistral").build())) .provider("superprovider") - .requestOptions(ChatOptions.builder().model("mistral").build()) .build(); var actualContext = this.observationFilter.map(expectedContext); @@ -62,9 +61,8 @@ class ChatModelPromptContentObservationFilterTests { @Test void whenPromptWithTextThenAugmentContext() { var originalContext = ChatModelObservationContext.builder() - .prompt(new Prompt("supercalifragilisticexpialidocious")) + .prompt(new Prompt("supercalifragilisticexpialidocious", ChatOptions.builder().model("mistral").build())) .provider("superprovider") - .requestOptions(ChatOptions.builder().model("mistral").build()) .build(); var augmentedContext = this.observationFilter.map(originalContext); @@ -75,10 +73,11 @@ class ChatModelPromptContentObservationFilterTests { @Test void whenPromptWithMessagesThenAugmentContext() { var originalContext = ChatModelObservationContext.builder() - .prompt(new Prompt(List.of(new SystemMessage("you're a chimney sweep"), - new UserMessage("supercalifragilisticexpialidocious")))) + .prompt(new Prompt( + List.of(new SystemMessage("you're a chimney sweep"), + new UserMessage("supercalifragilisticexpialidocious")), + ChatOptions.builder().model("mistral").build())) .provider("superprovider") - .requestOptions(ChatOptions.builder().model("mistral").build()) .build(); var augmentedContext = this.observationFilter.map(originalContext); diff --git a/spring-ai-model/src/test/java/org/springframework/ai/chat/observation/ChatModelPromptContentObservationHandlerTests.java b/spring-ai-model/src/test/java/org/springframework/ai/chat/observation/ChatModelPromptContentObservationHandlerTests.java index ab90a8551..c76b52470 100644 --- a/spring-ai-model/src/test/java/org/springframework/ai/chat/observation/ChatModelPromptContentObservationHandlerTests.java +++ b/spring-ai-model/src/test/java/org/springframework/ai/chat/observation/ChatModelPromptContentObservationHandlerTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2024 the original author or authors. + * Copyright 2023-2025 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -42,9 +42,9 @@ class ChatModelPromptContentObservationHandlerTests { @Test void whenPromptWithTextThenSpanEvent() { var observationContext = ChatModelObservationContext.builder() - .prompt(new Prompt("supercalifragilisticexpialidocious")) + .prompt(new Prompt("supercalifragilisticexpialidocious", + ChatOptions.builder().model("spoonful-of-sugar").build())) .provider("mary-poppins") - .requestOptions(ChatOptions.builder().model("spoonful-of-sugar").build()) .build(); var sdkTracer = SdkTracerProvider.builder().build().get("test"); var otelTracer = new OtelTracer(sdkTracer, new OtelCurrentTraceContext(), null); diff --git a/spring-ai-model/src/test/java/org/springframework/ai/chat/observation/DefaultChatModelObservationConventionTests.java b/spring-ai-model/src/test/java/org/springframework/ai/chat/observation/DefaultChatModelObservationConventionTests.java index c70e8b8a0..5629a1de4 100644 --- a/spring-ai-model/src/test/java/org/springframework/ai/chat/observation/DefaultChatModelObservationConventionTests.java +++ b/spring-ai-model/src/test/java/org/springframework/ai/chat/observation/DefaultChatModelObservationConventionTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2024 the original author or authors. + * Copyright 2023-2025 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -55,9 +55,8 @@ class DefaultChatModelObservationConventionTests { @Test void contextualNameWhenModelIsDefined() { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() - .prompt(generatePrompt()) + .prompt(generatePrompt(ChatOptions.builder().model("mistral").build())) .provider("superprovider") - .requestOptions(ChatOptions.builder().model("mistral").build()) .build(); assertThat(this.observationConvention.getContextualName(observationContext)).isEqualTo("chat mistral"); } @@ -65,9 +64,8 @@ class DefaultChatModelObservationConventionTests { @Test void contextualNameWhenModelIsNotDefined() { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() - .prompt(generatePrompt()) + .prompt(generatePrompt(ChatOptions.builder().build())) .provider("superprovider") - .requestOptions(ChatOptions.builder().build()) .build(); assertThat(this.observationConvention.getContextualName(observationContext)).isEqualTo("chat"); } @@ -75,9 +73,8 @@ class DefaultChatModelObservationConventionTests { @Test void supportsOnlyChatModelObservationContext() { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() - .prompt(generatePrompt()) + .prompt(generatePrompt(ChatOptions.builder().model("mistral").build())) .provider("superprovider") - .requestOptions(ChatOptions.builder().model("mistral").build()) .build(); assertThat(this.observationConvention.supportsContext(observationContext)).isTrue(); assertThat(this.observationConvention.supportsContext(new Observation.Context())).isFalse(); @@ -86,9 +83,8 @@ class DefaultChatModelObservationConventionTests { @Test void shouldHaveLowCardinalityKeyValuesWhenDefined() { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() - .prompt(generatePrompt()) + .prompt(generatePrompt(ChatOptions.builder().model("mistral").build())) .provider("superprovider") - .requestOptions(ChatOptions.builder().model("mistral").build()) .build(); assertThat(this.observationConvention.getLowCardinalityKeyValues(observationContext)).contains( KeyValue.of(LowCardinalityKeyNames.AI_OPERATION_TYPE.asString(), "chat"), @@ -99,9 +95,7 @@ class DefaultChatModelObservationConventionTests { @Test void shouldHaveKeyValuesWhenDefinedAndResponse() { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() - .prompt(generatePrompt()) - .provider("superprovider") - .requestOptions(ChatOptions.builder() + .prompt(generatePrompt(ChatOptions.builder() .model("mistral") .frequencyPenalty(0.8) .maxTokens(200) @@ -110,7 +104,8 @@ class DefaultChatModelObservationConventionTests { .temperature(0.5) .topK(1) .topP(0.9) - .build()) + .build())) + .provider("superprovider") .build(); observationContext.setResponse(new ChatResponse( List.of(new Generation(new AssistantMessage("response"), @@ -136,9 +131,8 @@ class DefaultChatModelObservationConventionTests { @Test void shouldNotHaveKeyValuesWhenMissing() { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() - .prompt(generatePrompt()) + .prompt(generatePrompt(ChatOptions.builder().build())) .provider("superprovider") - .requestOptions(ChatOptions.builder().build()) .build(); assertThat(this.observationConvention.getLowCardinalityKeyValues(observationContext)) .contains(KeyValue.of(LowCardinalityKeyNames.REQUEST_MODEL.asString(), KeyValue.NONE_VALUE)) @@ -162,9 +156,8 @@ class DefaultChatModelObservationConventionTests { @Test void shouldNotHaveKeyValuesWhenEmptyValues() { ChatModelObservationContext observationContext = ChatModelObservationContext.builder() - .prompt(generatePrompt()) + .prompt(generatePrompt(ChatOptions.builder().stopSequences(List.of()).build())) .provider("superprovider") - .requestOptions(ChatOptions.builder().stopSequences(List.of()).build()) .build(); observationContext.setResponse(new ChatResponse( List.of(new Generation(new AssistantMessage("response"), @@ -178,8 +171,8 @@ class DefaultChatModelObservationConventionTests { HighCardinalityKeyNames.RESPONSE_ID.asString()); } - private Prompt generatePrompt() { - return new Prompt("Who let the dogs out?"); + private Prompt generatePrompt(ChatOptions chatOptions) { + return new Prompt("Who let the dogs out?", chatOptions); } static class TestUsage implements Usage { diff --git a/spring-ai-model/src/test/java/org/springframework/ai/embedding/observation/DefaultEmbeddingModelObservationConventionTests.java b/spring-ai-model/src/test/java/org/springframework/ai/embedding/observation/DefaultEmbeddingModelObservationConventionTests.java index 7ed12ac16..ba5c6467d 100644 --- a/spring-ai-model/src/test/java/org/springframework/ai/embedding/observation/DefaultEmbeddingModelObservationConventionTests.java +++ b/spring-ai-model/src/test/java/org/springframework/ai/embedding/observation/DefaultEmbeddingModelObservationConventionTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2024 the original author or authors. + * Copyright 2023-2025 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -22,9 +22,11 @@ import java.util.Map; import io.micrometer.common.KeyValue; import io.micrometer.observation.Observation; +import org.jetbrains.annotations.NotNull; import org.junit.jupiter.api.Test; import org.springframework.ai.chat.metadata.Usage; +import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingOptionsBuilder; import org.springframework.ai.embedding.EmbeddingRequest; import org.springframework.ai.embedding.EmbeddingResponse; @@ -52,9 +54,8 @@ class DefaultEmbeddingModelObservationConventionTests { @Test void contextualNameWhenModelIsDefined() { EmbeddingModelObservationContext observationContext = EmbeddingModelObservationContext.builder() - .embeddingRequest(generateEmbeddingRequest()) + .embeddingRequest(generateEmbeddingRequest(EmbeddingOptionsBuilder.builder().withModel("mistral").build())) .provider("superprovider") - .requestOptions(EmbeddingOptionsBuilder.builder().withModel("mistral").build()) .build(); assertThat(this.observationConvention.getContextualName(observationContext)).isEqualTo("embedding mistral"); } @@ -62,9 +63,8 @@ class DefaultEmbeddingModelObservationConventionTests { @Test void contextualNameWhenModelIsNotDefined() { EmbeddingModelObservationContext observationContext = EmbeddingModelObservationContext.builder() - .embeddingRequest(generateEmbeddingRequest()) + .embeddingRequest(generateEmbeddingRequest(EmbeddingOptionsBuilder.builder().build())) .provider("superprovider") - .requestOptions(EmbeddingOptionsBuilder.builder().build()) .build(); assertThat(this.observationConvention.getContextualName(observationContext)).isEqualTo("embedding"); } @@ -72,9 +72,9 @@ class DefaultEmbeddingModelObservationConventionTests { @Test void supportsOnlyEmbeddingModelObservationContext() { EmbeddingModelObservationContext observationContext = EmbeddingModelObservationContext.builder() - .embeddingRequest(generateEmbeddingRequest()) + .embeddingRequest( + generateEmbeddingRequest(EmbeddingOptionsBuilder.builder().withModel("supermodel").build())) .provider("superprovider") - .requestOptions(EmbeddingOptionsBuilder.builder().withModel("supermodel").build()) .build(); assertThat(this.observationConvention.supportsContext(observationContext)).isTrue(); assertThat(this.observationConvention.supportsContext(new Observation.Context())).isFalse(); @@ -83,9 +83,8 @@ class DefaultEmbeddingModelObservationConventionTests { @Test void shouldHaveLowCardinalityKeyValuesWhenDefined() { EmbeddingModelObservationContext observationContext = EmbeddingModelObservationContext.builder() - .embeddingRequest(generateEmbeddingRequest()) + .embeddingRequest(generateEmbeddingRequest(EmbeddingOptionsBuilder.builder().withModel("mistral").build())) .provider("superprovider") - .requestOptions(EmbeddingOptionsBuilder.builder().withModel("mistral").build()) .build(); assertThat(this.observationConvention.getLowCardinalityKeyValues(observationContext)).contains( KeyValue.of(LowCardinalityKeyNames.AI_OPERATION_TYPE.asString(), "embedding"), @@ -96,9 +95,9 @@ class DefaultEmbeddingModelObservationConventionTests { @Test void shouldHaveLowCardinalityKeyValuesWhenDefinedAndResponse() { EmbeddingModelObservationContext observationContext = EmbeddingModelObservationContext.builder() - .embeddingRequest(generateEmbeddingRequest()) + .embeddingRequest(generateEmbeddingRequest( + EmbeddingOptionsBuilder.builder().withModel("mistral").withDimensions(1492).build())) .provider("superprovider") - .requestOptions(EmbeddingOptionsBuilder.builder().withModel("mistral").withDimensions(1492).build()) .build(); observationContext.setResponse(new EmbeddingResponse(List.of(), new EmbeddingResponseMetadata("mistral-42", new TestUsage(), Map.of()))); @@ -113,9 +112,8 @@ class DefaultEmbeddingModelObservationConventionTests { @Test void shouldNotHaveKeyValuesWhenMissing() { EmbeddingModelObservationContext observationContext = EmbeddingModelObservationContext.builder() - .embeddingRequest(generateEmbeddingRequest()) + .embeddingRequest(generateEmbeddingRequest(EmbeddingOptionsBuilder.builder().build())) .provider("superprovider") - .requestOptions(EmbeddingOptionsBuilder.builder().build()) .build(); assertThat(this.observationConvention.getLowCardinalityKeyValues(observationContext)) .contains(KeyValue.of(LowCardinalityKeyNames.REQUEST_MODEL.asString(), KeyValue.NONE_VALUE)) @@ -128,8 +126,8 @@ class DefaultEmbeddingModelObservationConventionTests { HighCardinalityKeyNames.USAGE_TOTAL_TOKENS.asString()); } - private EmbeddingRequest generateEmbeddingRequest() { - return new EmbeddingRequest(List.of(), EmbeddingOptionsBuilder.builder().build()); + private EmbeddingRequest generateEmbeddingRequest(EmbeddingOptions embeddingOptions) { + return new EmbeddingRequest(List.of(), embeddingOptions); } static class TestUsage implements Usage { diff --git a/spring-ai-model/src/test/java/org/springframework/ai/embedding/observation/EmbeddingModelMeterObservationHandlerTests.java b/spring-ai-model/src/test/java/org/springframework/ai/embedding/observation/EmbeddingModelMeterObservationHandlerTests.java index b7b1c65f3..dada880a1 100644 --- a/spring-ai-model/src/test/java/org/springframework/ai/embedding/observation/EmbeddingModelMeterObservationHandlerTests.java +++ b/spring-ai-model/src/test/java/org/springframework/ai/embedding/observation/EmbeddingModelMeterObservationHandlerTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2024 the original author or authors. + * Copyright 2023-2025 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -28,6 +28,7 @@ import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.springframework.ai.chat.metadata.Usage; +import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingOptionsBuilder; import org.springframework.ai.embedding.EmbeddingRequest; import org.springframework.ai.embedding.EmbeddingResponse; @@ -92,14 +93,13 @@ class EmbeddingModelMeterObservationHandlerTests { private EmbeddingModelObservationContext generateObservationContext() { return EmbeddingModelObservationContext.builder() - .embeddingRequest(generateEmbeddingRequest()) + .embeddingRequest(generateEmbeddingRequest(EmbeddingOptionsBuilder.builder().withModel("mistral").build())) .provider("superprovider") - .requestOptions(EmbeddingOptionsBuilder.builder().withModel("mistral").build()) .build(); } - private EmbeddingRequest generateEmbeddingRequest() { - return new EmbeddingRequest(List.of(), EmbeddingOptionsBuilder.builder().build()); + private EmbeddingRequest generateEmbeddingRequest(EmbeddingOptions embeddingOptions) { + return new EmbeddingRequest(List.of(), embeddingOptions); } static class TestUsage implements Usage { diff --git a/spring-ai-model/src/test/java/org/springframework/ai/embedding/observation/EmbeddingModelObservationContextTests.java b/spring-ai-model/src/test/java/org/springframework/ai/embedding/observation/EmbeddingModelObservationContextTests.java index 0678fe26a..780e881a4 100644 --- a/spring-ai-model/src/test/java/org/springframework/ai/embedding/observation/EmbeddingModelObservationContextTests.java +++ b/spring-ai-model/src/test/java/org/springframework/ai/embedding/observation/EmbeddingModelObservationContextTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2024 the original author or authors. + * Copyright 2023-2025 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -20,6 +20,7 @@ import java.util.List; import org.junit.jupiter.api.Test; +import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingOptionsBuilder; import org.springframework.ai.embedding.EmbeddingRequest; @@ -36,26 +37,16 @@ class EmbeddingModelObservationContextTests { @Test void whenMandatoryRequestOptionsThenReturn() { var observationContext = EmbeddingModelObservationContext.builder() - .embeddingRequest(generateEmbeddingRequest()) + .embeddingRequest( + generateEmbeddingRequest(EmbeddingOptionsBuilder.builder().withModel("supermodel").build())) .provider("superprovider") - .requestOptions(EmbeddingOptionsBuilder.builder().withModel("supermodel").build()) .build(); assertThat(observationContext).isNotNull(); } - @Test - void whenRequestOptionsIsNullThenThrow() { - assertThatThrownBy(() -> EmbeddingModelObservationContext.builder() - .embeddingRequest(generateEmbeddingRequest()) - .provider("superprovider") - .requestOptions(null) - .build()).isInstanceOf(IllegalArgumentException.class) - .hasMessageContaining("requestOptions cannot be null"); - } - - private EmbeddingRequest generateEmbeddingRequest() { - return new EmbeddingRequest(List.of(), EmbeddingOptionsBuilder.builder().build()); + private EmbeddingRequest generateEmbeddingRequest(EmbeddingOptions embeddingOptions) { + return new EmbeddingRequest(List.of(), embeddingOptions); } }