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 6168ea55c..03603951e 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 @@ -48,7 +48,7 @@ import org.springframework.ai.ollama.api.OllamaApi.Message.Role; import org.springframework.ai.ollama.api.OllamaApi.Message.ToolCall; import org.springframework.ai.ollama.api.OllamaApi.Message.ToolCallFunction; import org.springframework.ai.ollama.api.OllamaOptions; -import org.springframework.ai.ollama.metadata.OllamaUsage; +import org.springframework.ai.ollama.metadata.OllamaChatUsage; import org.springframework.util.Assert; import org.springframework.util.CollectionUtils; import org.springframework.util.StringUtils; @@ -177,7 +177,7 @@ public class OllamaChatModel extends AbstractToolCallSupport implements ChatMode public static ChatResponseMetadata from(OllamaApi.ChatResponse response) { Assert.notNull(response, "OllamaApi.ChatResponse must not be null"); return ChatResponseMetadata.builder() - .withUsage(OllamaUsage.from(response)) + .withUsage(OllamaChatUsage.from(response)) .withModel(response.model()) .withKeyValue("created-at", response.createdAt()) .withKeyValue("eval-duration", response.evalDuration()) 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 3c43b9154..0984034d2 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 @@ -33,6 +33,7 @@ import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.ollama.api.OllamaApi; import org.springframework.ai.ollama.api.OllamaApi.EmbeddingsResponse; import org.springframework.ai.ollama.api.OllamaOptions; +import org.springframework.ai.ollama.metadata.OllamaEmbeddingUsage; import org.springframework.util.Assert; import org.springframework.util.StringUtils; @@ -125,7 +126,7 @@ public class OllamaEmbeddingModel extends AbstractEmbeddingModel { .toList(); EmbeddingResponseMetadata embeddingResponseMetadata = new EmbeddingResponseMetadata(response.model(), - new EmptyUsage()); + OllamaEmbeddingUsage.from(response)); EmbeddingResponse embeddingResponse = new EmbeddingResponse(embeddings, embeddingResponseMetadata); diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaApi.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaApi.java index 565c44e1d..b20e206bc 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaApi.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaApi.java @@ -767,7 +767,11 @@ public class OllamaApi { @JsonInclude(Include.NON_NULL) public record EmbeddingsResponse( @JsonProperty("model") String model, - @JsonProperty("embeddings") List embeddings) { + @JsonProperty("embeddings") List embeddings, + @JsonProperty("total_duration") Long totalDuration, + @JsonProperty("load_duration") Long loadDuration, + @JsonProperty("prompt_eval_count") Integer promptEvalCount) { + } /** diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaUsage.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaChatUsage.java similarity index 88% rename from models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaUsage.java rename to models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaChatUsage.java index a437557d5..e1c1bfac8 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaUsage.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaChatUsage.java @@ -26,18 +26,18 @@ import org.springframework.util.Assert; * @see Usage * @author Fu Cheng */ -public class OllamaUsage implements Usage { +public class OllamaChatUsage implements Usage { protected static final String AI_USAGE_STRING = "{ promptTokens: %1$d, generationTokens: %2$d, totalTokens: %3$d }"; - public static OllamaUsage from(OllamaApi.ChatResponse response) { + public static OllamaChatUsage from(OllamaApi.ChatResponse response) { Assert.notNull(response, "OllamaApi.ChatResponse must not be null"); - return new OllamaUsage(response); + return new OllamaChatUsage(response); } private final OllamaApi.ChatResponse response; - public OllamaUsage(OllamaApi.ChatResponse response) { + public OllamaChatUsage(OllamaApi.ChatResponse response) { this.response = response; } diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaEmbeddingUsage.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaEmbeddingUsage.java new file mode 100644 index 000000000..61ea60b33 --- /dev/null +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaEmbeddingUsage.java @@ -0,0 +1,60 @@ +/* + * Copyright 2023 - 2024 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. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.ai.ollama.metadata; + +import java.util.Optional; + +import org.springframework.ai.chat.metadata.Usage; +import org.springframework.ai.ollama.api.OllamaApi.EmbeddingsResponse; +import org.springframework.util.Assert; + +/** + * {@link Usage} implementation for {@literal Ollama} embeddings. + * + * @see Usage + * @author Christian Tzolov + */ +public class OllamaEmbeddingUsage implements Usage { + + protected static final String AI_USAGE_STRING = "{ promptTokens: %1$d, generationTokens: %2$d, totalTokens: %3$d }"; + + public static OllamaEmbeddingUsage from(EmbeddingsResponse response) { + Assert.notNull(response, "OllamaApi.EmbeddingsResponse must not be null"); + return new OllamaEmbeddingUsage(response); + } + + private Long promptTokens; + + public OllamaEmbeddingUsage(EmbeddingsResponse response) { + this.promptTokens = Optional.ofNullable(response.promptEvalCount()).map(Integer::longValue).orElse(0L); + } + + @Override + public Long getPromptTokens() { + return this.promptTokens; + } + + @Override + public Long getGenerationTokens() { + return 0L; + } + + @Override + public String toString() { + return AI_USAGE_STRING.formatted(getPromptTokens(), getGenerationTokens(), getTotalTokens()); + } + +} diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelIT.java index ebede2211..d311fe1af 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelIT.java @@ -15,6 +15,11 @@ */ package org.springframework.ai.ollama; +import static org.assertj.core.api.Assertions.assertThat; + +import java.io.IOException; +import java.util.List; + import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; import org.junit.jupiter.api.BeforeAll; @@ -24,7 +29,6 @@ import org.springframework.ai.embedding.EmbeddingRequest; import org.springframework.ai.embedding.EmbeddingResponse; import org.springframework.ai.ollama.api.OllamaApi; import org.springframework.ai.ollama.api.OllamaApiIT; -import org.springframework.ai.ollama.api.OllamaModel; import org.springframework.ai.ollama.api.OllamaOptions; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.SpringBootConfiguration; @@ -32,24 +36,15 @@ import org.springframework.boot.test.context.SpringBootTest; import org.springframework.context.annotation.Bean; import org.testcontainers.junit.jupiter.Testcontainers; -import java.io.IOException; -import java.util.List; - -import static org.assertj.core.api.Assertions.assertThat; - @SpringBootTest @DisabledIf("isDisabled") @Testcontainers class OllamaEmbeddingModelIT extends BaseOllamaIT { - private static final String MODEL = OllamaModel.MISTRAL.getName(); + private static final String MODEL = "mxbai-embed-large"; private static final Log logger = LogFactory.getLog(OllamaApiIT.class); - // @Container - // static OllamaContainer ollamaContainer = new - // OllamaContainer(OllamaImage.DEFAULT_IMAGE); - static String baseUrl = "http://localhost:11434"; @BeforeAll @@ -75,8 +70,10 @@ class OllamaEmbeddingModelIT extends BaseOllamaIT { assertThat(embeddingResponse.getResults().get(1).getIndex()).isEqualTo(1); assertThat(embeddingResponse.getResults().get(1).getOutput()).isNotEmpty(); assertThat(embeddingResponse.getMetadata().getModel()).isEqualTo(MODEL); + assertThat(embeddingResponse.getMetadata().getUsage().getPromptTokens()).isEqualTo(4); + assertThat(embeddingResponse.getMetadata().getUsage().getTotalTokens()).isEqualTo(4); - assertThat(embeddingModel.dimensions()).isEqualTo(4096); + assertThat(embeddingModel.dimensions()).isEqualTo(1024); } @SpringBootConfiguration diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelTests.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelTests.java index 192c99864..6b6569b8f 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelTests.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelTests.java @@ -54,25 +54,25 @@ public class OllamaEmbeddingModelTests { public void options() { when(ollamaApi.embed(embeddingsRequestCaptor.capture())) - .thenReturn( - new EmbeddingsResponse("RESPONSE_MODEL_NAME", List.of(new float[]{1f, 2f, 3f}, new float[]{4f, 5f, 6f}))) + .thenReturn(new EmbeddingsResponse("RESPONSE_MODEL_NAME", + List.of(new float[] { 1f, 2f, 3f }, new float[] { 4f, 5f, 6f }), 0L, 0L, 0)) .thenReturn(new EmbeddingsResponse("RESPONSE_MODEL_NAME2", - List.of(new float[]{7f, 8f, 9f}, new float[]{10f, 11f, 12f}))); + List.of(new float[] { 7f, 8f, 9f }, new float[] { 10f, 11f, 12f }), 0L, 0L, 0)); // Tests default options var defaultOptions = OllamaOptions.builder().withModel("DEFAULT_MODEL").build(); var embeddingModel = new OllamaEmbeddingModel(ollamaApi, defaultOptions); - EmbeddingResponse response = embeddingModel - .call(new EmbeddingRequest(List.of("Input1", "Input2", "Input3"), EmbeddingOptionsBuilder.builder().build())); + EmbeddingResponse response = embeddingModel.call( + new EmbeddingRequest(List.of("Input1", "Input2", "Input3"), EmbeddingOptionsBuilder.builder().build())); assertThat(response.getResults()).hasSize(2); assertThat(response.getResults().get(0).getIndex()).isEqualTo(0); - assertThat(response.getResults().get(0).getOutput()).isEqualTo(new float[]{1f, 2f, 3f}); + assertThat(response.getResults().get(0).getOutput()).isEqualTo(new float[] { 1f, 2f, 3f }); assertThat(response.getResults().get(0).getMetadata()).isEqualTo(EmbeddingResultMetadata.EMPTY); assertThat(response.getResults().get(1).getIndex()).isEqualTo(1); - assertThat(response.getResults().get(1).getOutput()).isEqualTo(new float[]{4f, 5f, 6f}); + assertThat(response.getResults().get(1).getOutput()).isEqualTo(new float[] { 4f, 5f, 6f }); assertThat(response.getResults().get(1).getMetadata()).isEqualTo(EmbeddingResultMetadata.EMPTY); assertThat(response.getMetadata().getModel()).isEqualTo("RESPONSE_MODEL_NAME"); @@ -94,10 +94,10 @@ public class OllamaEmbeddingModelTests { assertThat(response.getResults()).hasSize(2); assertThat(response.getResults().get(0).getIndex()).isEqualTo(0); - assertThat(response.getResults().get(0).getOutput()).isEqualTo(new float[]{7f, 8f, 9f}); + assertThat(response.getResults().get(0).getOutput()).isEqualTo(new float[] { 7f, 8f, 9f }); assertThat(response.getResults().get(0).getMetadata()).isEqualTo(EmbeddingResultMetadata.EMPTY); assertThat(response.getResults().get(1).getIndex()).isEqualTo(1); - assertThat(response.getResults().get(1).getOutput()).isEqualTo(new float[]{10f, 11f, 12f}); + assertThat(response.getResults().get(1).getOutput()).isEqualTo(new float[] { 10f, 11f, 12f }); assertThat(response.getResults().get(1).getMetadata()).isEqualTo(EmbeddingResultMetadata.EMPTY); assertThat(response.getMetadata().getModel()).isEqualTo("RESPONSE_MODEL_NAME2"); diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/OllamaApiIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/OllamaApiIT.java index af050bee4..427ae8f91 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/OllamaApiIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/OllamaApiIT.java @@ -146,6 +146,11 @@ public class OllamaApiIT extends BaseOllamaIT { assertThat(response).isNotNull(); assertThat(response.embeddings()).hasSize(1); assertThat(response.embeddings().get(0)).hasSize(3200); + assertThat(response.model()).isEqualTo(MODEL); + assertThat(response.promptEvalCount()).isEqualTo(5); + assertThat(response.loadDuration()).isGreaterThan(1); + assertThat(response.totalDuration()).isGreaterThan(1); + } } \ No newline at end of file