diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelConnection.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelConnection.java index 4d46e3eed..01210e8f5 100644 --- a/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelConnection.java +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelConnection.java @@ -80,4 +80,13 @@ public ResourceType getResourceType() { * embeddings. The length of each array is determined by the model itself. */ public abstract List embed(List texts, Map parameters); + + public EmbeddingResult embedWithUsage(String text, Map parameters) { + return new EmbeddingResult<>(embed(text, parameters), null); + } + + public EmbeddingResult> embedWithUsage( + List texts, Map parameters) { + return new EmbeddingResult<>(embed(texts, parameters), null); + } } diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetup.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetup.java index e7c19893d..605c189d5 100644 --- a/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetup.java +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetup.java @@ -109,6 +109,17 @@ public float[] embed(String text, Map parameters) { return getConnection().embed(text, params); } + public EmbeddingResult embedWithUsage(String text) { + return embedWithUsage(text, Collections.emptyMap()); + } + + public EmbeddingResult embedWithUsage(String text, Map parameters) { + Map params = this.getParameters(); + params.putAll(parameters); + BaseEmbeddingModelConnection currentConnection = getConnection(); + return currentConnection.embedWithUsage(text, params); + } + /** * Generate embeddings for multiple texts. * @@ -125,4 +136,16 @@ public List embed(List texts, Map parameters) { params.putAll(parameters); return getConnection().embed(texts, params); } + + public EmbeddingResult> embedWithUsage(List texts) { + return embedWithUsage(texts, Collections.emptyMap()); + } + + public EmbeddingResult> embedWithUsage( + List texts, Map parameters) { + Map params = this.getParameters(); + params.putAll(parameters); + BaseEmbeddingModelConnection currentConnection = getConnection(); + return currentConnection.embedWithUsage(texts, params); + } } diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingModelUtils.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingModelUtils.java index c1c71e589..08128eb45 100644 --- a/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingModelUtils.java +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingModelUtils.java @@ -17,7 +17,9 @@ */ package org.apache.flink.agents.api.embedding.model; +import java.util.ArrayList; import java.util.List; +import java.util.Map; public class EmbeddingModelUtils { public static float[] toFloatArray(List list) { @@ -34,4 +36,66 @@ public static float[] toFloatArray(List list) { } return array; } + + public static EmbeddingResult toSingleEmbeddingResult(Object result) { + Map values = toResultMap(result); + return new EmbeddingResult<>( + toFloatArray(toEmbeddingList(values.get("embeddings"))), + toTokenUsage(values.get("token_usage"))); + } + + public static EmbeddingResult> toBatchEmbeddingResult(Object result) { + Map values = toResultMap(result); + List rawEmbeddings = toEmbeddingList(values.get("embeddings")); + List embeddings = new ArrayList<>(); + for (Object embedding : rawEmbeddings) { + embeddings.add(toFloatArray(toEmbeddingList(embedding))); + } + return new EmbeddingResult<>(embeddings, toTokenUsage(values.get("token_usage"))); + } + + private static Map toResultMap(Object result) { + if (result instanceof Map) { + return (Map) result; + } + throw new IllegalArgumentException( + "Expected Map from Python embed_with_usage method, but got: " + + (result == null ? "null" : result.getClass().getName())); + } + + private static List toEmbeddingList(Object embeddings) { + if (embeddings instanceof List) { + return (List) embeddings; + } + throw new IllegalArgumentException( + "Expected List value in Python embedding result, but got: " + + (embeddings == null ? "null" : embeddings.getClass().getName())); + } + + private static EmbeddingTokenUsage toTokenUsage(Object tokenUsage) { + if (tokenUsage == null) { + return null; + } + if (!(tokenUsage instanceof Map)) { + throw new IllegalArgumentException( + "Expected Map token_usage in Python embedding result, but got: " + + tokenUsage.getClass().getName()); + } + + Map usage = (Map) tokenUsage; + return new EmbeddingTokenUsage( + toLong(usage.get("prompt_tokens"), "prompt_tokens"), + toLong(usage.get("total_tokens"), "total_tokens")); + } + + private static long toLong(Object value, String fieldName) { + if (value instanceof Number) { + return ((Number) value).longValue(); + } + throw new IllegalArgumentException( + "Expected numeric " + + fieldName + + " in Python embedding token usage, but got: " + + (value == null ? "null" : value.getClass().getName())); + } } diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingResult.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingResult.java new file mode 100644 index 000000000..8c23d14d7 --- /dev/null +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingResult.java @@ -0,0 +1,42 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you 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 + * + * http://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.apache.flink.agents.api.embedding.model; + +import javax.annotation.Nullable; + +/** Embedding provider result with optional token usage metadata. */ +public class EmbeddingResult { + private final T embeddings; + + @Nullable private final EmbeddingTokenUsage tokenUsage; + + public EmbeddingResult(T embeddings, @Nullable EmbeddingTokenUsage tokenUsage) { + this.embeddings = embeddings; + this.tokenUsage = tokenUsage; + } + + public T getEmbeddings() { + return embeddings; + } + + @Nullable + public EmbeddingTokenUsage getTokenUsage() { + return tokenUsage; + } +} diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingTokenUsage.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingTokenUsage.java new file mode 100644 index 000000000..be1abdd98 --- /dev/null +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingTokenUsage.java @@ -0,0 +1,38 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you 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 + * + * http://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.apache.flink.agents.api.embedding.model; + +/** Token usage reported by an embedding provider. */ +public class EmbeddingTokenUsage { + private final long promptTokens; + private final long totalTokens; + + public EmbeddingTokenUsage(long promptTokens, long totalTokens) { + this.promptTokens = promptTokens; + this.totalTokens = totalTokens; + } + + public long getPromptTokens() { + return promptTokens; + } + + public long getTotalTokens() { + return totalTokens; + } +} diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnection.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnection.java index 785ed03a8..d9bdb8d41 100644 --- a/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnection.java +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnection.java @@ -19,6 +19,7 @@ import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelConnection; import org.apache.flink.agents.api.embedding.model.EmbeddingModelUtils; +import org.apache.flink.agents.api.embedding.model.EmbeddingResult; import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; @@ -42,6 +43,9 @@ public class PythonEmbeddingModelConnection extends BaseEmbeddingModelConnection implements PythonResourceWrapper { + private static final String CALL_EMBED_WITH_USAGE = + "python_java_utils.call_embedding_with_usage"; + private final PyObject embeddingModel; private final PythonResourceAdapter adapter; @@ -119,6 +123,31 @@ public List embed(List texts, Map parameters) { + (results == null ? "null" : results.getClass().getName())); } + @Override + public EmbeddingResult embedWithUsage(String text, Map parameters) { + checkState( + embeddingModel != null, + "EmbeddingModelSetup is not initialized. Cannot perform embed operation."); + + Map kwargs = new HashMap<>(parameters); + kwargs.put("text", text); + Object result = adapter.invoke(CALL_EMBED_WITH_USAGE, embeddingModel, kwargs); + return EmbeddingModelUtils.toSingleEmbeddingResult(result); + } + + @Override + public EmbeddingResult> embedWithUsage( + List texts, Map parameters) { + checkState( + embeddingModel != null, + "EmbeddingModelSetup is not initialized. Cannot perform embed operation."); + + Map kwargs = new HashMap<>(parameters); + kwargs.put("text", texts); + Object result = adapter.invoke(CALL_EMBED_WITH_USAGE, embeddingModel, kwargs); + return EmbeddingModelUtils.toBatchEmbeddingResult(result); + } + @Override public Object getPythonResource() { return embeddingModel; diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetup.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetup.java index c0efd5c35..b6febbc58 100644 --- a/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetup.java +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetup.java @@ -19,6 +19,7 @@ import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelSetup; import org.apache.flink.agents.api.embedding.model.EmbeddingModelUtils; +import org.apache.flink.agents.api.embedding.model.EmbeddingResult; import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; @@ -42,6 +43,9 @@ */ public class PythonEmbeddingModelSetup extends BaseEmbeddingModelSetup implements PythonResourceWrapper { + private static final String CALL_EMBED_WITH_USAGE = + "python_java_utils.call_embedding_with_usage"; + private final PyObject embeddingModelSetup; private final PythonResourceAdapter adapter; @@ -124,6 +128,31 @@ public List embed(List texts, Map parameters) { + (results == null ? "null" : results.getClass().getName())); } + @Override + public EmbeddingResult embedWithUsage(String text, Map parameters) { + checkState( + embeddingModelSetup != null, + "EmbeddingModelSetup is not initialized. Cannot perform embed operation."); + + Map kwargs = new HashMap<>(parameters); + kwargs.put("text", text); + Object result = adapter.invoke(CALL_EMBED_WITH_USAGE, embeddingModelSetup, kwargs); + return EmbeddingModelUtils.toSingleEmbeddingResult(result); + } + + @Override + public EmbeddingResult> embedWithUsage( + List texts, Map parameters) { + checkState( + embeddingModelSetup != null, + "EmbeddingModelSetup is not initialized. Cannot perform embed operation."); + + Map kwargs = new HashMap<>(parameters); + kwargs.put("text", texts); + Object result = adapter.invoke(CALL_EMBED_WITH_USAGE, embeddingModelSetup, kwargs); + return EmbeddingModelUtils.toBatchEmbeddingResult(result); + } + @Override public Map getParameters() { return Map.of(); diff --git a/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupEmbeddingResultTest.java b/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupEmbeddingResultTest.java new file mode 100644 index 000000000..fbbf738da --- /dev/null +++ b/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupEmbeddingResultTest.java @@ -0,0 +1,114 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you 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 + * + * http://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.apache.flink.agents.api.embedding.model; + +import org.apache.flink.agents.api.resource.ResourceContext; +import org.apache.flink.agents.api.resource.ResourceDescriptor; +import org.junit.jupiter.api.Test; + +import java.util.Collections; +import java.util.Map; + +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.mockito.Mockito.mock; + +/** Test cases for embedding results returned by model setups. */ +class BaseEmbeddingModelSetupEmbeddingResultTest { + + @Test + void testEmbedWithUsageDelegatesProviderUsage() { + BaseEmbeddingModelSetup setup = + new BaseEmbeddingModelSetup( + new ResourceDescriptor( + "test", Map.of("connection", "connection", "model", "model")), + mock(ResourceContext.class)) { + @Override + public Map getParameters() { + return Collections.emptyMap(); + } + }; + setup.connection = + new BaseEmbeddingModelConnection( + new ResourceDescriptor("connection", Collections.emptyMap()), + mock(ResourceContext.class)) { + @Override + public float[] embed(String text, Map parameters) { + return new float[] {0.1f, 0.2f}; + } + + @Override + public java.util.List embed( + java.util.List texts, Map parameters) { + throw new UnsupportedOperationException(); + } + + @Override + public EmbeddingResult embedWithUsage( + String text, Map parameters) { + return new EmbeddingResult<>( + embed(text, parameters), new EmbeddingTokenUsage(7L, 9L)); + } + }; + + EmbeddingResult result = setup.embedWithUsage("hello"); + + assertArrayEquals(new float[] {0.1f, 0.2f}, result.getEmbeddings()); + assertEquals(7L, result.getTokenUsage().getPromptTokens()); + assertEquals(9L, result.getTokenUsage().getTotalTokens()); + } + + @Test + void testEmbedDelegatesToExistingConnectionMethod() { + BaseEmbeddingModelSetup setup = + new BaseEmbeddingModelSetup( + new ResourceDescriptor( + "test", Map.of("connection", "connection", "model", "model")), + mock(ResourceContext.class)) { + @Override + public Map getParameters() { + return Collections.emptyMap(); + } + }; + setup.connection = + new BaseEmbeddingModelConnection( + new ResourceDescriptor("connection", Collections.emptyMap()), + mock(ResourceContext.class)) { + @Override + public float[] embed(String text, Map parameters) { + return new float[] {0.1f, 0.2f}; + } + + @Override + public java.util.List embed( + java.util.List texts, Map parameters) { + throw new UnsupportedOperationException(); + } + + @Override + public EmbeddingResult embedWithUsage( + String text, Map parameters) { + return new EmbeddingResult<>( + new float[] {0.3f, 0.4f}, new EmbeddingTokenUsage(7L, 9L)); + } + }; + + assertArrayEquals(new float[] {0.1f, 0.2f}, setup.embed("hello")); + } +} diff --git a/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnectionTest.java b/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnectionTest.java index 470d51270..a89bf3313 100644 --- a/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnectionTest.java +++ b/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnectionTest.java @@ -17,6 +17,7 @@ */ package org.apache.flink.agents.api.embedding.model.python; +import org.apache.flink.agents.api.embedding.model.EmbeddingResult; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.apache.flink.agents.api.resource.python.PythonResourceAdapter; @@ -136,6 +137,42 @@ void testEmbedSingleTextWithEmptyParameters() { assertThat(result).hasSize(2); } + @Test + void testEmbedWithUsageMultipleTexts() { + List texts = List.of("first", "second"); + when(mockAdapter.invoke( + eq("python_java_utils.call_embedding_with_usage"), + eq(mockEmbeddingModel), + any(Map.class))) + .thenReturn( + Map.of( + "embeddings", + List.of(List.of(0.1, 0.2), List.of(0.3, 0.4)), + "token_usage", + Map.of("prompt_tokens", 7, "total_tokens", 9))); + + EmbeddingResult> result = + pythonEmbeddingModelConnection.embedWithUsage(texts, Map.of("batch_size", 2)); + + assertThat(result.getEmbeddings()).hasSize(2); + assertThat(result.getEmbeddings().get(0)).containsExactly(0.1f, 0.2f); + assertThat(result.getEmbeddings().get(1)).containsExactly(0.3f, 0.4f); + assertThat(result.getTokenUsage()).isNotNull(); + assertThat(result.getTokenUsage().getPromptTokens()).isEqualTo(7L); + assertThat(result.getTokenUsage().getTotalTokens()).isEqualTo(9L); + verify(mockAdapter) + .invoke( + eq("python_java_utils.call_embedding_with_usage"), + eq(mockEmbeddingModel), + argThat( + kwargs -> { + Map values = (Map) kwargs; + assertThat(values).containsEntry("text", texts); + assertThat(values).containsEntry("batch_size", 2); + return true; + })); + } + @Test void testEmbedSingleTextWithNullEmbeddingModelThrowsException() { PythonEmbeddingModelConnection connectionWithNullModel = diff --git a/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetupTest.java b/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetupTest.java index d8071d7cd..430ddc4c3 100644 --- a/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetupTest.java +++ b/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetupTest.java @@ -17,6 +17,7 @@ */ package org.apache.flink.agents.api.embedding.model.python; +import org.apache.flink.agents.api.embedding.model.EmbeddingResult; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.apache.flink.agents.api.resource.python.PythonResourceAdapter; @@ -143,6 +144,41 @@ void testEmbedSingleTextWithEmptyParameters() { assertThat(result).hasSize(2); } + @Test + void testEmbedWithUsageSingleText() { + String text = "test text"; + Map parameters = Map.of("model", "test-model"); + when(mockAdapter.invoke( + eq("python_java_utils.call_embedding_with_usage"), + eq(mockEmbeddingModelSetup), + any(Map.class))) + .thenReturn( + Map.of( + "embeddings", + List.of(0.1, 0.2), + "token_usage", + Map.of("prompt_tokens", 7, "total_tokens", 9))); + + EmbeddingResult result = + pythonEmbeddingModelSetup.embedWithUsage(text, parameters); + + assertThat(result.getEmbeddings()).containsExactly(0.1f, 0.2f); + assertThat(result.getTokenUsage()).isNotNull(); + assertThat(result.getTokenUsage().getPromptTokens()).isEqualTo(7L); + assertThat(result.getTokenUsage().getTotalTokens()).isEqualTo(9L); + verify(mockAdapter) + .invoke( + eq("python_java_utils.call_embedding_with_usage"), + eq(mockEmbeddingModelSetup), + argThat( + kwargs -> { + Map values = (Map) kwargs; + assertThat(values).containsEntry("text", text); + assertThat(values).containsEntry("model", "test-model"); + return true; + })); + } + @Test void testEmbedSingleTextWithNullEmbeddingModelSetupThrowsException() { PythonEmbeddingModelSetup setupWithNullModel = diff --git a/e2e-test/flink-agents-end-to-end-tests-resource-cross-language/src/test/java/org/apache/flink/agents/resource/test/EmbeddingCrossLanguageAgent.java b/e2e-test/flink-agents-end-to-end-tests-resource-cross-language/src/test/java/org/apache/flink/agents/resource/test/EmbeddingCrossLanguageAgent.java index be5f8d494..14bfbbb4d 100644 --- a/e2e-test/flink-agents-end-to-end-tests-resource-cross-language/src/test/java/org/apache/flink/agents/resource/test/EmbeddingCrossLanguageAgent.java +++ b/e2e-test/flink-agents-end-to-end-tests-resource-cross-language/src/test/java/org/apache/flink/agents/resource/test/EmbeddingCrossLanguageAgent.java @@ -29,6 +29,7 @@ import org.apache.flink.agents.api.annotation.EmbeddingModelSetup; import org.apache.flink.agents.api.context.RunnerContext; import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelSetup; +import org.apache.flink.agents.api.embedding.model.EmbeddingResult; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.apache.flink.agents.api.resource.ResourceName; @@ -105,11 +106,14 @@ public static void testEmbeddingGeneration(Event event, RunnerContext ctx) throw org.apache.flink.agents.api.resource.ResourceType .EMBEDDING_MODEL); - float[] embedding = embeddingModel.embed(text); + EmbeddingResult embeddingResult = embeddingModel.embedWithUsage(text); + float[] embedding = embeddingResult.getEmbeddings(); System.out.printf("[TEST] Generated embedding with dimension: %d%n", embedding.length); validateEmbeddingResult(id, text, embedding); - List embeddings = embeddingModel.embed(List.of(text)); + EmbeddingResult> embeddingsResult = + embeddingModel.embedWithUsage(List.of(text)); + List embeddings = embeddingsResult.getEmbeddings(); validateEmbeddingResults(id, List.of(text), embeddings); // Create a minimal test result to avoid serialization issues diff --git a/integrations/embedding-models/bedrock/src/main/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelConnection.java b/integrations/embedding-models/bedrock/src/main/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelConnection.java index a7dc2926c..b39ec6b20 100644 --- a/integrations/embedding-models/bedrock/src/main/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelConnection.java +++ b/integrations/embedding-models/bedrock/src/main/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelConnection.java @@ -23,6 +23,8 @@ import com.fasterxml.jackson.databind.node.ObjectNode; import org.apache.flink.agents.api.RetryExecutor; import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelConnection; +import org.apache.flink.agents.api.embedding.model.EmbeddingResult; +import org.apache.flink.agents.api.embedding.model.EmbeddingTokenUsage; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import software.amazon.awssdk.auth.credentials.DefaultCredentialsProvider; @@ -80,36 +82,67 @@ public class BedrockEmbeddingModelConnection extends BaseEmbeddingModelConnectio public BedrockEmbeddingModelConnection( ResourceDescriptor descriptor, ResourceContext resourceContext) { super(descriptor, resourceContext); + this.client = createClient(resolveRegion(descriptor)); + this.defaultModel = resolveDefaultModel(descriptor); + this.embedPool = Executors.newFixedThreadPool(resolveEmbedConcurrency(descriptor)); + this.retryExecutor = createRetryExecutor(descriptor); + } + + BedrockEmbeddingModelConnection( + ResourceDescriptor descriptor, + ResourceContext resourceContext, + BedrockRuntimeClient client, + ExecutorService embedPool, + RetryExecutor retryExecutor, + String defaultModel) { + super(descriptor, resourceContext); + this.client = client; + this.defaultModel = defaultModel; + this.embedPool = embedPool; + this.retryExecutor = retryExecutor; + } + + private static BedrockRuntimeClient createClient(String region) { + return BedrockRuntimeClient.builder() + .region(Region.of(region)) + .credentialsProvider(DefaultCredentialsProvider.create()) + .build(); + } + private static String resolveRegion(ResourceDescriptor descriptor) { String region = descriptor.getArgument("region"); if (region == null || region.isBlank()) { region = "us-east-1"; } + return region; + } - this.client = - BedrockRuntimeClient.builder() - .region(Region.of(region)) - .credentialsProvider(DefaultCredentialsProvider.create()) - .build(); - + private static String resolveDefaultModel(ResourceDescriptor descriptor) { String model = descriptor.getArgument("model"); - this.defaultModel = (model != null && !model.isBlank()) ? model : DEFAULT_MODEL; + return (model != null && !model.isBlank()) ? model : DEFAULT_MODEL; + } + private static int resolveEmbedConcurrency(ResourceDescriptor descriptor) { Integer concurrency = descriptor.getArgument("embed_concurrency"); - int threads = concurrency != null ? concurrency : 4; - this.embedPool = Executors.newFixedThreadPool(threads); + return concurrency != null ? concurrency : 4; + } + private static RetryExecutor createRetryExecutor(ResourceDescriptor descriptor) { Integer retries = descriptor.getArgument("max_retries"); - this.retryExecutor = - RetryExecutor.builder() - .maxRetries(retries != null ? retries : 5) - .initialBackoffMs(200) - .retryablePredicate(BedrockEmbeddingModelConnection::isRetryable) - .build(); + return RetryExecutor.builder() + .maxRetries(retries != null ? retries : 5) + .initialBackoffMs(200) + .retryablePredicate(BedrockEmbeddingModelConnection::isRetryable) + .build(); } @Override public float[] embed(String text, Map parameters) { + return embedWithUsage(text, parameters).getEmbeddings(); + } + + @Override + public EmbeddingResult embedWithUsage(String text, Map parameters) { String model = (String) parameters.getOrDefault("model", defaultModel); Integer dimensions = (Integer) parameters.get("dimensions"); @@ -138,12 +171,21 @@ public float[] embed(String text, Map parameters) { for (int i = 0; i < embeddingNode.size(); i++) { embedding[i] = (float) embeddingNode.get(i).asDouble(); } - return embedding; + return new EmbeddingResult<>(embedding, extractTokenUsage(result)); } catch (Exception e) { throw new RuntimeException("Failed to parse Bedrock embedding response.", e); } } + private static EmbeddingTokenUsage extractTokenUsage(JsonNode result) { + JsonNode inputTokenCount = result.get("inputTextTokenCount"); + if (inputTokenCount == null || !inputTokenCount.isNumber()) { + return null; + } + long tokens = inputTokenCount.asLong(); + return new EmbeddingTokenUsage(tokens, tokens); + } + private static boolean isRetryable(Exception e) { String msg = e.toString(); return msg.contains("ThrottlingException") @@ -155,27 +197,52 @@ private static boolean isRetryable(Exception e) { @Override public List embed(List texts, Map parameters) { + return embedWithUsage(texts, parameters).getEmbeddings(); + } + + @Override + public EmbeddingResult> embedWithUsage( + List texts, Map parameters) { if (texts.size() <= 1) { List results = new ArrayList<>(texts.size()); + EmbeddingTokenUsage totalUsage = null; for (String text : texts) { - results.add(embed(text, parameters)); + EmbeddingResult result = embedWithUsage(text, parameters); + results.add(result.getEmbeddings()); + totalUsage = mergeUsage(totalUsage, result.getTokenUsage()); } - return results; + return new EmbeddingResult<>(results, totalUsage); } @SuppressWarnings("unchecked") - CompletableFuture[] futures = + CompletableFuture>[] futures = texts.stream() .map( text -> CompletableFuture.supplyAsync( - () -> embed(text, parameters), embedPool)) + () -> embedWithUsage(text, parameters), embedPool)) .toArray(CompletableFuture[]::new); CompletableFuture.allOf(futures).join(); List results = new ArrayList<>(texts.size()); - for (CompletableFuture f : futures) { - results.add(f.join()); + EmbeddingTokenUsage totalUsage = null; + for (CompletableFuture> f : futures) { + EmbeddingResult result = f.join(); + results.add(result.getEmbeddings()); + totalUsage = mergeUsage(totalUsage, result.getTokenUsage()); + } + return new EmbeddingResult<>(results, totalUsage); + } + + private static EmbeddingTokenUsage mergeUsage( + EmbeddingTokenUsage left, EmbeddingTokenUsage right) { + if (left == null) { + return right; + } + if (right == null) { + return left; } - return results; + return new EmbeddingTokenUsage( + left.getPromptTokens() + right.getPromptTokens(), + left.getTotalTokens() + right.getTotalTokens()); } @Override diff --git a/integrations/embedding-models/bedrock/src/test/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelTest.java b/integrations/embedding-models/bedrock/src/test/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelTest.java index 0891d7067..99c8e9b81 100644 --- a/integrations/embedding-models/bedrock/src/test/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelTest.java +++ b/integrations/embedding-models/bedrock/src/test/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelTest.java @@ -18,17 +18,31 @@ package org.apache.flink.agents.integrations.embeddingmodels.bedrock; +import org.apache.flink.agents.api.RetryExecutor; import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelConnection; import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelSetup; +import org.apache.flink.agents.api.embedding.model.EmbeddingResult; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; +import software.amazon.awssdk.core.SdkBytes; +import software.amazon.awssdk.services.bedrockruntime.BedrockRuntimeClient; +import software.amazon.awssdk.services.bedrockruntime.model.InvokeModelRequest; +import software.amazon.awssdk.services.bedrockruntime.model.InvokeModelResponse; +import java.util.List; import java.util.Map; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; import static org.assertj.core.api.Assertions.assertThat; import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; /** Tests for {@link BedrockEmbeddingModelConnection} and {@link BedrockEmbeddingModelSetup}. */ class BedrockEmbeddingModelTest { @@ -94,4 +108,42 @@ void testSetupParametersNoDimensions() { assertThat(setup.getParameters()).doesNotContainKey("dimensions"); } + + @Test + @DisplayName("Batch embedding aggregates token usage from worker results") + void testBatchEmbeddingAggregatesTokenUsage() throws Exception { + BedrockRuntimeClient client = mock(BedrockRuntimeClient.class); + when(client.invokeModel(any(InvokeModelRequest.class))) + .thenReturn( + InvokeModelResponse.builder() + .body( + SdkBytes.fromUtf8String( + "{\"embedding\":[0.1,0.2],\"inputTextTokenCount\":5}")) + .build()); + + ExecutorService embedPool = Executors.newFixedThreadPool(2); + BedrockEmbeddingModelConnection conn = + new BedrockEmbeddingModelConnection( + connDescriptor(null), + NOOP, + client, + embedPool, + RetryExecutor.builder().maxRetries(0).build(), + "mock-model"); + + try { + EmbeddingResult> result = + conn.embedWithUsage(List.of("first", "second"), Map.of("model", "mock-model")); + + assertThat(result.getEmbeddings()).hasSize(2); + assertThat(result.getEmbeddings().get(0)).containsExactly(0.1f, 0.2f); + assertThat(result.getEmbeddings().get(1)).containsExactly(0.1f, 0.2f); + assertThat(result.getTokenUsage()).isNotNull(); + assertThat(result.getTokenUsage().getPromptTokens()).isEqualTo(10L); + assertThat(result.getTokenUsage().getTotalTokens()).isEqualTo(10L); + verify(client, times(2)).invokeModel(any(InvokeModelRequest.class)); + } finally { + conn.close(); + } + } } diff --git a/python/flink_agents/api/embedding_models/embedding_model.py b/python/flink_agents/api/embedding_models/embedding_model.py index 5529c817e..c31ab5324 100644 --- a/python/flink_agents/api/embedding_models/embedding_model.py +++ b/python/flink_agents/api/embedding_models/embedding_model.py @@ -16,13 +16,32 @@ # limitations under the License. ################################################################################# from abc import ABC, abstractmethod -from typing import Any, Dict, Sequence, cast +from dataclasses import dataclass +from typing import Any, Dict, Generic, Sequence, TypeVar, cast from pydantic import Field from typing_extensions import override from flink_agents.api.resource import Resource, ResourceType +EmbeddingValue = TypeVar("EmbeddingValue", list[float], list[list[float]]) + + +@dataclass(frozen=True) +class EmbeddingTokenUsage: + """Token usage reported by an embedding provider.""" + + prompt_tokens: int = 0 + total_tokens: int = 0 + + +@dataclass(frozen=True) +class EmbeddingResult(Generic[EmbeddingValue]): + """Embedding provider result with optional token usage metadata.""" + + embeddings: EmbeddingValue + token_usage: EmbeddingTokenUsage | None = None + class BaseEmbeddingModelConnection(Resource, ABC): """Base abstract class for text embedding model connection. @@ -62,6 +81,12 @@ def embed( The dimension of the vector depends on the specific embedding model used. """ + def embed_with_usage( + self, text: str | Sequence[str], **kwargs: Any + ) -> EmbeddingResult[list[float] | list[list[float]]]: + """Generate embeddings and return provider token usage when available.""" + return EmbeddingResult(embeddings=self.embed(text, **kwargs)) + class BaseEmbeddingModelSetup(Resource, ABC): """Base abstract class for text embedding model setup. @@ -123,3 +148,11 @@ def embed( merged_kwargs = self.model_kwargs.copy() merged_kwargs.update(kwargs) return self._get_connection().embed(text, **merged_kwargs) + + def embed_with_usage( + self, text: str | Sequence[str], **kwargs: Any + ) -> EmbeddingResult[list[float] | list[list[float]]]: + """Generate embeddings and return provider token usage when available.""" + merged_kwargs = self.model_kwargs.copy() + merged_kwargs.update(kwargs) + return self._get_connection().embed_with_usage(text, **merged_kwargs) diff --git a/python/flink_agents/api/embedding_models/tests/__init__.py b/python/flink_agents/api/embedding_models/tests/__init__.py new file mode 100644 index 000000000..1b9e7e00a --- /dev/null +++ b/python/flink_agents/api/embedding_models/tests/__init__.py @@ -0,0 +1,18 @@ +################################################################################ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you 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 +# +# http://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. +################################################################################# + diff --git a/python/flink_agents/api/embedding_models/tests/test_embedding_result.py b/python/flink_agents/api/embedding_models/tests/test_embedding_result.py new file mode 100644 index 000000000..e42193520 --- /dev/null +++ b/python/flink_agents/api/embedding_models/tests/test_embedding_result.py @@ -0,0 +1,127 @@ +################################################################################ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you 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 +# +# http://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. +################################################################################# +from typing import Any, Dict, Sequence +from unittest.mock import MagicMock + +from flink_agents.api.embedding_models.embedding_model import ( + BaseEmbeddingModelConnection, + BaseEmbeddingModelSetup, + EmbeddingResult, + EmbeddingTokenUsage, +) +from flink_agents.api.resource import Resource, ResourceType +from flink_agents.api.resource_context import ResourceContext + + +class FakeEmbeddingModelConnection(BaseEmbeddingModelConnection): + def embed( + self, text: str | Sequence[str], **kwargs: Any + ) -> list[float] | list[list[float]]: + if isinstance(text, str): + return [0.1, 0.2] + return [[0.1, 0.2] for _ in text] + + def embed_with_usage( + self, text: str | Sequence[str], **kwargs: Any + ) -> EmbeddingResult[list[float] | list[list[float]]]: + return EmbeddingResult( + embeddings=self.embed(text, **kwargs), + token_usage=EmbeddingTokenUsage(prompt_tokens=7, total_tokens=9), + ) + + +class FakeEmbeddingModelConnectionWithoutUsage(BaseEmbeddingModelConnection): + def embed( + self, text: str | Sequence[str], **kwargs: Any + ) -> list[float] | list[list[float]]: + if isinstance(text, str): + return [0.1, 0.2] + return [[0.1, 0.2] for _ in text] + + +class DispatchAwareEmbeddingModelConnection(BaseEmbeddingModelConnection): + def embed( + self, text: str | Sequence[str], **kwargs: Any + ) -> list[float] | list[list[float]]: + if isinstance(text, str): + return [0.1, 0.2] + return [[0.1, 0.2] for _ in text] + + def embed_with_usage( + self, text: str | Sequence[str], **kwargs: Any + ) -> EmbeddingResult[list[float] | list[list[float]]]: + return EmbeddingResult( + embeddings=[0.3, 0.4] + if isinstance(text, str) + else [[0.3, 0.4] for _ in text], + token_usage=EmbeddingTokenUsage(prompt_tokens=7, total_tokens=9), + ) + + +class FakeEmbeddingModelSetup(BaseEmbeddingModelSetup): + @property + def model_kwargs(self) -> Dict[str, Any]: + return {} + + +def _make_setup(connection: BaseEmbeddingModelConnection) -> FakeEmbeddingModelSetup: + def get_resource(name: str, resource_type: ResourceType) -> Resource: + assert name == "mock-connection" + assert resource_type == ResourceType.EMBEDDING_MODEL_CONNECTION + return connection + + ctx = MagicMock(spec=ResourceContext) + ctx.get_resource = get_resource + setup = FakeEmbeddingModelSetup( + name="embedding", + connection="mock-connection", + model="mock-model", + resource_context=ctx, + ) + setup.open() + return setup + + +def test_embedding_result_returns_provider_usage() -> None: + setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) + + result = setup.embed_with_usage("hello") + + assert result.embeddings == [0.1, 0.2] + assert result.token_usage == EmbeddingTokenUsage(prompt_tokens=7, total_tokens=9) + + +def test_embedding_result_defaults_to_no_usage() -> None: + setup = _make_setup(FakeEmbeddingModelConnectionWithoutUsage(name="connection")) + + result = setup.embed_with_usage("hello") + + assert result.embeddings == [0.1, 0.2] + assert result.token_usage is None + + +def test_embed_preserves_existing_embedding_only_api() -> None: + setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) + + assert setup.embed("hello") == [0.1, 0.2] + + +def test_embed_preserves_connection_embed_dispatch() -> None: + setup = _make_setup(DispatchAwareEmbeddingModelConnection(name="connection")) + + assert setup.embed("hello") == [0.1, 0.2] diff --git a/python/flink_agents/e2e_tests/e2e_tests_resource_cross_language/embedding_model_cross_language_agent.py b/python/flink_agents/e2e_tests/e2e_tests_resource_cross_language/embedding_model_cross_language_agent.py index dbce4891b..e57fbc07f 100644 --- a/python/flink_agents/e2e_tests/e2e_tests_resource_cross_language/embedding_model_cross_language_agent.py +++ b/python/flink_agents/e2e_tests/e2e_tests_resource_cross_language/embedding_model_cross_language_agent.py @@ -77,7 +77,8 @@ def process_input(event: Event, ctx: RunnerContext) -> None: ) # Test single text embedding - embedding = embeddingModel.embed(input_text) + embedding_result = embeddingModel.embed_with_usage(input_text) + embedding = embedding_result.embeddings print(f"[TEST] Generated embedding with dimension: {len(embedding)}") # Validate single embedding result @@ -98,7 +99,8 @@ def process_input(event: Event, ctx: RunnerContext) -> None: ) # Test batch embedding - embeddings = embeddingModel.embed([input_text]) + embeddings_result = embeddingModel.embed_with_usage([input_text]) + embeddings = embeddings_result.embeddings print(f"[TEST] Generated batch embeddings: count={len(embeddings)}") # Validate batch embedding results diff --git a/python/flink_agents/integrations/embedding_models/openai_embedding_model.py b/python/flink_agents/integrations/embedding_models/openai_embedding_model.py index eac4c8f10..92d09b765 100644 --- a/python/flink_agents/integrations/embedding_models/openai_embedding_model.py +++ b/python/flink_agents/integrations/embedding_models/openai_embedding_model.py @@ -24,6 +24,8 @@ from flink_agents.api.embedding_models.embedding_model import ( BaseEmbeddingModelConnection, BaseEmbeddingModelSetup, + EmbeddingResult, + EmbeddingTokenUsage, ) DEFAULT_REQUEST_TIMEOUT = 30.0 @@ -115,6 +117,12 @@ def embed( self, text: str | Sequence[str], **kwargs: Any ) -> list[float] | list[list[float]]: """Generate embedding vector for a single text query.""" + return self.embed_with_usage(text, **kwargs).embeddings + + def embed_with_usage( + self, text: str | Sequence[str], **kwargs: Any + ) -> EmbeddingResult[list[float] | list[list[float]]]: + """Generate embeddings and return OpenAI token usage when available.""" # Extract OpenAI specific parameters model = kwargs.pop("model") encoding_format = kwargs.pop("encoding_format", None) @@ -132,8 +140,24 @@ def embed( user=user if user is not None else NOT_GIVEN, ) + usage = getattr(response, "usage", None) + token_usage = None + if usage is not None: + prompt_tokens = getattr(usage, "prompt_tokens", None) + total_tokens = getattr(usage, "total_tokens", None) + if prompt_tokens is not None or total_tokens is not None: + token_usage = EmbeddingTokenUsage( + prompt_tokens=int(prompt_tokens or 0), + total_tokens=int( + total_tokens if total_tokens is not None else prompt_tokens + ), + ) + embeddings = [list(embedding.embedding) for embedding in response.data] - return embeddings[0] if isinstance(text, str) else embeddings + return EmbeddingResult( + embeddings=embeddings[0] if isinstance(text, str) else embeddings, + token_usage=token_usage, + ) @override def close(self) -> None: diff --git a/python/flink_agents/integrations/embedding_models/tests/test_openai_embedding_model.py b/python/flink_agents/integrations/embedding_models/tests/test_openai_embedding_model.py index 49907340f..9d63c2494 100644 --- a/python/flink_agents/integrations/embedding_models/tests/test_openai_embedding_model.py +++ b/python/flink_agents/integrations/embedding_models/tests/test_openai_embedding_model.py @@ -16,6 +16,7 @@ # limitations under the License. ################################################################################ import os +from types import SimpleNamespace from unittest.mock import MagicMock import pytest @@ -57,3 +58,34 @@ def get_resource(name: str, type: ResourceType) -> Resource: assert isinstance(response, list) assert len(response) > 0 assert all(isinstance(x, float) for x in response) # + + +def test_openai_embedding_model_returns_token_usage() -> None: + """Test OpenAI embedding usage is returned with the embedding result.""" + connection = OpenAIEmbeddingModelConnection(name="openai", api_key="fake-key") + mock_client = MagicMock() + mock_client.embeddings.create.return_value = SimpleNamespace( + data=[SimpleNamespace(embedding=[0.1, 0.2, 0.3])], + usage=SimpleNamespace(prompt_tokens=5, total_tokens=5), + ) + connection._OpenAIEmbeddingModelConnection__client = mock_client + + def get_resource(name: str, type: ResourceType) -> Resource: + if type == ResourceType.EMBEDDING_MODEL_CONNECTION: + return connection + else: + msg = f"Unknown resource type: {type}" + raise ValueError(msg) + + mock_ctx = MagicMock(spec=ResourceContext) + mock_ctx.get_resource = get_resource + embedding_model = OpenAIEmbeddingModelSetup( + name="openai", model=test_model, connection="openai", resource_context=mock_ctx + ) + embedding_model.open() + + result = embedding_model.embed_with_usage("Hello, Flink Agent!") + assert result.embeddings == [0.1, 0.2, 0.3] + assert result.token_usage is not None + assert result.token_usage.prompt_tokens == 5 + assert result.token_usage.total_tokens == 5 diff --git a/python/flink_agents/integrations/embedding_models/tests/test_tongyi_embedding_model.py b/python/flink_agents/integrations/embedding_models/tests/test_tongyi_embedding_model.py index b60c75596..6262d4bd6 100644 --- a/python/flink_agents/integrations/embedding_models/tests/test_tongyi_embedding_model.py +++ b/python/flink_agents/integrations/embedding_models/tests/test_tongyi_embedding_model.py @@ -155,6 +155,52 @@ def get_resource(name: str, type: ResourceType) -> Resource: assert len(response) == 5 +def test_tongyi_embedding_returns_token_usage( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Test DashScope embedding usage is recorded as model token metrics.""" + mock_embedding = [0.1, 0.2, 0.3] + mocked_response = SimpleNamespace( + status_code=HTTPStatus.OK, + output={ + "embeddings": [{"embedding": mock_embedding}], + }, + usage={"total_tokens": 6}, + message="Success", + ) + mock_call = MagicMock(return_value=mocked_response) + monkeypatch.setattr( + "flink_agents.integrations.embedding_models.tongyi_embedding_model.dashscope.TextEmbedding.call", + mock_call, + ) + + connection = TongyiEmbeddingModelConnection( + name="tongyi", + api_key="fake-key", + ) + + def get_resource(name: str, type: ResourceType) -> Resource: + if type == ResourceType.EMBEDDING_MODEL_CONNECTION: + return connection + else: + msg = f"Unknown resource type: {type}" + raise ValueError(msg) + + embedding_model = TongyiEmbeddingModelSetup( + name="tongyi", + model=test_model, + connection="tongyi", + resource_context=_make_ctx(get_resource), + ) + embedding_model.open() + + result = embedding_model.embed_with_usage("Test text") + assert result.embeddings == mock_embedding + assert result.token_usage is not None + assert result.token_usage.prompt_tokens == 6 + assert result.token_usage.total_tokens == 6 + + def test_tongyi_embedding_batch_mock(monkeypatch: pytest.MonkeyPatch) -> None: """Test batch embedding functionality with mocked DashScope API.""" mock_embeddings = [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]] diff --git a/python/flink_agents/integrations/embedding_models/tongyi_embedding_model.py b/python/flink_agents/integrations/embedding_models/tongyi_embedding_model.py index c63c083d4..bea3075bd 100644 --- a/python/flink_agents/integrations/embedding_models/tongyi_embedding_model.py +++ b/python/flink_agents/integrations/embedding_models/tongyi_embedding_model.py @@ -25,12 +25,27 @@ from flink_agents.api.embedding_models.embedding_model import ( BaseEmbeddingModelConnection, BaseEmbeddingModelSetup, + EmbeddingResult, + EmbeddingTokenUsage, ) DEFAULT_REQUEST_TIMEOUT = 30.0 DEFAULT_MODEL = "text-embedding-v4" +def _get_usage_value(obj: Any, *names: str) -> int | None: + """Read a token usage value from dict-like or object-like provider responses.""" + if obj is None: + return None + for name in names: + if isinstance(obj, dict) and name in obj: + return int(obj[name]) + value = getattr(obj, name, None) + if value is not None: + return int(value) + return None + + class TongyiEmbeddingModelConnection(BaseEmbeddingModelConnection): """Tongyi Embedding Model Connection which manages connection to DashScope API. @@ -78,6 +93,12 @@ def embed( self, text: str | Sequence[str], **kwargs: Any ) -> list[float] | list[list[float]]: """Generate embedding vector for text input.""" + return self.embed_with_usage(text, **kwargs).embeddings + + def embed_with_usage( + self, text: str | Sequence[str], **kwargs: Any + ) -> EmbeddingResult[list[float] | list[list[float]]]: + """Generate embeddings and return DashScope token usage when available.""" model = kwargs.pop("model", DEFAULT_MODEL) text_type = kwargs.pop("text_type", None) dimension = kwargs.pop("dimension", None) @@ -103,8 +124,27 @@ def embed( msg = f"DashScope TextEmbedding call failed: {response.message}" raise RuntimeError(msg) + usage = getattr(response, "usage", None) + if usage is None and isinstance(response.output, dict): + usage = response.output.get("usage") + prompt_tokens = _get_usage_value(usage, "input_tokens", "prompt_tokens") + total_tokens = _get_usage_value(usage, "total_tokens") + if prompt_tokens is None: + prompt_tokens = total_tokens + token_usage = None + if prompt_tokens is not None or total_tokens is not None: + token_usage = EmbeddingTokenUsage( + prompt_tokens=int(prompt_tokens or 0), + total_tokens=int( + total_tokens if total_tokens is not None else prompt_tokens + ), + ) + embeddings = [e["embedding"] for e in response.output["embeddings"]] - return embeddings[0] if isinstance(text, str) else embeddings + return EmbeddingResult( + embeddings=embeddings[0] if isinstance(text, str) else embeddings, + token_usage=token_usage, + ) class TongyiEmbeddingModelSetup(BaseEmbeddingModelSetup): diff --git a/python/flink_agents/runtime/java/java_embedding_model.py b/python/flink_agents/runtime/java/java_embedding_model.py index b2ea48724..230a49441 100644 --- a/python/flink_agents/runtime/java/java_embedding_model.py +++ b/python/flink_agents/runtime/java/java_embedding_model.py @@ -19,6 +19,10 @@ from typing_extensions import override +from flink_agents.api.embedding_models.embedding_model import ( + EmbeddingResult, + EmbeddingTokenUsage, +) from flink_agents.api.embedding_models.java_embedding_model import ( JavaEmbeddingModelConnection, JavaEmbeddingModelSetup, @@ -28,6 +32,28 @@ ) +def _from_java_embedding_result( + j_result: Any, text: str | Sequence[str] +) -> EmbeddingResult[list[float] | list[list[float]]]: + """Convert a Java EmbeddingResult into the Python result contract.""" + j_embeddings = j_result.getEmbeddings() + embeddings = ( + list(j_embeddings) + if isinstance(text, str) + else [list(embedding) for embedding in j_embeddings] + ) + j_usage = j_result.getTokenUsage() + token_usage = ( + None + if j_usage is None + else EmbeddingTokenUsage( + prompt_tokens=j_usage.getPromptTokens(), + total_tokens=j_usage.getTotalTokens(), + ) + ) + return EmbeddingResult(embeddings=embeddings, token_usage=token_usage) + + class JavaEmbeddingModelConnectionImpl(JavaEmbeddingModelConnection): """Java-based implementation of EmbeddingModelConnection that wraps a Java embedding model object. @@ -72,6 +98,16 @@ def embed( ) return list(result) if isinstance(text, str) else [list(emb) for emb in result] + @override + def embed_with_usage( + self, text: str | Sequence[str], **kwargs: Any + ) -> EmbeddingResult[list[float] | list[list[float]]]: + """Generate embeddings through Java and preserve provider token usage.""" + result = self._j_resource.embedWithUsage( + text if isinstance(text, str) else list(text), kwargs + ) + return _from_java_embedding_result(result, text) + class JavaEmbeddingModelSetupImpl(JavaEmbeddingModelSetup): """Java-based implementation of EmbeddingModelSetup that wraps a Java embedding @@ -134,3 +170,13 @@ def embed( text if isinstance(text, str) else list(text), kwargs ) return list(result) if isinstance(text, str) else [list(emb) for emb in result] + + @override + def embed_with_usage( + self, text: str | Sequence[str], **kwargs: Any + ) -> EmbeddingResult[list[float] | list[list[float]]]: + """Generate embeddings through Java and preserve provider token usage.""" + result = self._j_resource.embedWithUsage( + text if isinstance(text, str) else list(text), kwargs + ) + return _from_java_embedding_result(result, text) diff --git a/python/flink_agents/runtime/python_java_utils.py b/python/flink_agents/runtime/python_java_utils.py index 66b9092f7..3c4bb60dc 100644 --- a/python/flink_agents/runtime/python_java_utils.py +++ b/python/flink_agents/runtime/python_java_utils.py @@ -154,7 +154,9 @@ def get_python_tool_metadata( descriptor = PythonFunction(module=module, qualname=qual_name) callable_ = descriptor.as_callable() name = callable_.__name__ - description = (parse(callable_.__doc__).description or "") if callable_.__doc__ else "" + description = ( + (parse(callable_.__doc__).description or "") if callable_.__doc__ else "" + ) callable_injected_args = normalize_injected_args( getattr(callable_, "_injected_args", None) ) @@ -177,9 +179,7 @@ def _dump_injected_args(injected_args: Dict[str, Any]) -> str: ) -def invoke_python_tool( - module: str, qual_name: str, kwargs: Dict[str, Any] -) -> Any: +def invoke_python_tool(module: str, qual_name: str, kwargs: Dict[str, Any]) -> Any: """Invoke a Python callable as a tool, passing the provided keyword arguments. Used by the Java-side ``PythonResourceAdapter.invokePythonTool`` so a Java host can @@ -226,6 +226,32 @@ def from_java_resource(type_name: str, kwargs: Dict[str, Any]) -> Resource: return cls(**kwargs) +def embedding_result_to_java(result: Any) -> Dict[str, Any]: + """Convert an embedding result into Pemja-safe Python primitives.""" + usage = result.token_usage + return { + "embeddings": result.embeddings, + "token_usage": ( + None + if usage is None + else { + "prompt_tokens": usage.prompt_tokens, + "total_tokens": usage.total_tokens, + } + ), + } + + +def call_embedding_with_usage( + embedding_model: Any, kwargs: Dict[str, Any] +) -> Dict[str, Any]: + """Call ``embed_with_usage`` and return Java-safe primitives. + + This avoids PyObject attribute access in the Java caller. + """ + return embedding_result_to_java(embedding_model.embed_with_usage(**kwargs)) + + def normalize_tool_call_id(tool_call: Dict[str, Any]) -> Dict[str, Any]: """Normalize tool call by converting the ID field to string format while preserving all other fields. diff --git a/python/flink_agents/runtime/tests/test_java_embedding_model.py b/python/flink_agents/runtime/tests/test_java_embedding_model.py new file mode 100644 index 000000000..2a3866687 --- /dev/null +++ b/python/flink_agents/runtime/tests/test_java_embedding_model.py @@ -0,0 +1,70 @@ +################################################################################ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you 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 +# +# http://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. +################################################################################# +from unittest.mock import MagicMock + +import pytest + +from flink_agents.runtime.java.java_embedding_model import ( + JavaEmbeddingModelConnectionImpl, + JavaEmbeddingModelSetupImpl, +) + + +class _JavaTokenUsage: + def getPromptTokens(self) -> int: + return 7 + + def getTotalTokens(self) -> int: + return 9 + + +class _JavaEmbeddingResult: + def getEmbeddings(self) -> list[list[float]]: + return [[0.1, 0.2], [0.3, 0.4]] + + def getTokenUsage(self) -> _JavaTokenUsage: + return _JavaTokenUsage() + + +@pytest.mark.parametrize( + ("wrapper_class", "kwargs"), + [ + (JavaEmbeddingModelConnectionImpl, {}), + ( + JavaEmbeddingModelSetupImpl, + {"connection": "connection", "model": "test-model"}, + ), + ], +) +def test_java_embedding_wrappers_preserve_usage( + wrapper_class: type[JavaEmbeddingModelConnectionImpl | JavaEmbeddingModelSetupImpl], + kwargs: dict[str, str], +) -> None: + j_resource = MagicMock() + j_resource.embedWithUsage.return_value = _JavaEmbeddingResult() + wrapper = wrapper_class(j_resource, MagicMock(), **kwargs) + + result = wrapper.embed_with_usage(["first", "second"], batch_size=2) + + assert result.embeddings == [[0.1, 0.2], [0.3, 0.4]] + assert result.token_usage is not None + assert result.token_usage.prompt_tokens == 7 + assert result.token_usage.total_tokens == 9 + j_resource.embedWithUsage.assert_called_once_with( + ["first", "second"], {"batch_size": 2} + ) diff --git a/python/flink_agents/runtime/tests/test_python_java_utils.py b/python/flink_agents/runtime/tests/test_python_java_utils.py index 7aeeafe9c..fa708353b 100644 --- a/python/flink_agents/runtime/tests/test_python_java_utils.py +++ b/python/flink_agents/runtime/tests/test_python_java_utils.py @@ -18,8 +18,15 @@ import json from flink_agents.api.decorators import tool +from flink_agents.api.embedding_models.embedding_model import ( + EmbeddingResult, + EmbeddingTokenUsage, +) from flink_agents.api.tools import InjectedArg -from flink_agents.runtime.python_java_utils import get_python_tool_metadata +from flink_agents.runtime.python_java_utils import ( + call_embedding_with_usage, + get_python_tool_metadata, +) @tool(injected_args={"tenant_id": InjectedArg.from_config("tenant.id")}) @@ -36,6 +43,27 @@ def test_get_python_tool_metadata_merges_callable_injected_args() -> None: schema = json.loads(flat["inputSchema"]) assert set(schema["properties"]) == {"order_id"} injected_args = json.loads(flat["injectedArgs"]) - assert injected_args == { - "tenant_id": {"source": "config", "key": "tenant.id"} + assert injected_args == {"tenant_id": {"source": "config", "key": "tenant.id"}} + + +class _UsageAwareEmbeddingModel: + def embed_with_usage( + self, text: str, **kwargs: object + ) -> EmbeddingResult[list[float]]: + assert text == "hello" + assert kwargs == {"model": "test-model"} + return EmbeddingResult( + embeddings=[0.1, 0.2], + token_usage=EmbeddingTokenUsage(prompt_tokens=7, total_tokens=9), + ) + + +def test_call_embedding_with_usage_returns_pemja_safe_primitives() -> None: + result = call_embedding_with_usage( + _UsageAwareEmbeddingModel(), {"text": "hello", "model": "test-model"} + ) + + assert result == { + "embeddings": [0.1, 0.2], + "token_usage": {"prompt_tokens": 7, "total_tokens": 9}, }