This is an automated email from the ASF dual-hosted git repository.
wenjin272 pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/flink-agents.git
The following commit(s) were added to refs/heads/main by this push:
new 4859469c [api][python][java] Track embedding token usage metrics (#870)
4859469c is described below
commit 4859469cc6d734c6c217faadfc219b053cfac6ea
Author: zxs1633079383 <[email protected]>
AuthorDate: Sat Jul 25 00:50:43 2026 -0700
[api][python][java] Track embedding token usage metrics (#870)
---
.../model/BaseEmbeddingModelConnection.java | 9 ++
.../embedding/model/BaseEmbeddingModelSetup.java | 23 ++++
.../api/embedding/model/EmbeddingModelUtils.java | 64 +++++++++++
...beddingModelUtils.java => EmbeddingResult.java} | 35 +++---
...ingModelUtils.java => EmbeddingTokenUsage.java} | 31 ++---
.../python/PythonEmbeddingModelConnection.java | 29 +++++
.../model/python/PythonEmbeddingModelSetup.java | 29 +++++
...BaseEmbeddingModelSetupEmbeddingResultTest.java | 114 ++++++++++++++++++
.../python/PythonEmbeddingModelConnectionTest.java | 37 ++++++
.../python/PythonEmbeddingModelSetupTest.java | 36 ++++++
.../resource/test/EmbeddingCrossLanguageAgent.java | 8 +-
.../bedrock/BedrockEmbeddingModelConnection.java | 113 ++++++++++++++----
.../bedrock/BedrockEmbeddingModelTest.java | 52 +++++++++
.../api/embedding_models/embedding_model.py | 35 +++++-
.../embedding_models/tests/__init__.py} | 23 ----
.../tests/test_embedding_result.py | 127 +++++++++++++++++++++
.../embedding_model_cross_language_agent.py | 6 +-
.../embedding_models/openai_embedding_model.py | 26 ++++-
.../tests/test_openai_embedding_model.py | 32 ++++++
.../tests/test_tongyi_embedding_model.py | 46 ++++++++
.../embedding_models/tongyi_embedding_model.py | 42 ++++++-
.../runtime/java/java_embedding_model.py | 46 ++++++++
python/flink_agents/runtime/python_java_utils.py | 34 +++++-
.../runtime/tests/test_java_embedding_model.py | 70 ++++++++++++
.../runtime/tests/test_python_java_utils.py | 34 +++++-
25 files changed, 1011 insertions(+), 90 deletions(-)
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 4d46e3ee..01210e8f 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 abstract class BaseEmbeddingModelConnection extends
Resource {
* embeddings. The length of each array is determined by the model
itself.
*/
public abstract List<float[]> embed(List<String> texts, Map<String,
Object> parameters);
+
+ public EmbeddingResult<float[]> embedWithUsage(String text, Map<String,
Object> parameters) {
+ return new EmbeddingResult<>(embed(text, parameters), null);
+ }
+
+ public EmbeddingResult<List<float[]>> embedWithUsage(
+ List<String> texts, Map<String, Object> 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 e7c19893..605c189d 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 abstract class BaseEmbeddingModelSetup extends
Resource {
return getConnection().embed(text, params);
}
+ public EmbeddingResult<float[]> embedWithUsage(String text) {
+ return embedWithUsage(text, Collections.emptyMap());
+ }
+
+ public EmbeddingResult<float[]> embedWithUsage(String text, Map<String,
Object> parameters) {
+ Map<String, Object> params = this.getParameters();
+ params.putAll(parameters);
+ BaseEmbeddingModelConnection currentConnection = getConnection();
+ return currentConnection.embedWithUsage(text, params);
+ }
+
/**
* Generate embeddings for multiple texts.
*
@@ -125,4 +136,16 @@ public abstract class BaseEmbeddingModelSetup extends
Resource {
params.putAll(parameters);
return getConnection().embed(texts, params);
}
+
+ public EmbeddingResult<List<float[]>> embedWithUsage(List<String> texts) {
+ return embedWithUsage(texts, Collections.emptyMap());
+ }
+
+ public EmbeddingResult<List<float[]>> embedWithUsage(
+ List<String> texts, Map<String, Object> parameters) {
+ Map<String, Object> 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 c1c71e58..08128eb4 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 class EmbeddingModelUtils {
}
return array;
}
+
+ public static EmbeddingResult<float[]> toSingleEmbeddingResult(Object
result) {
+ Map<?, ?> values = toResultMap(result);
+ return new EmbeddingResult<>(
+ toFloatArray(toEmbeddingList(values.get("embeddings"))),
+ toTokenUsage(values.get("token_usage")));
+ }
+
+ public static EmbeddingResult<List<float[]>> toBatchEmbeddingResult(Object
result) {
+ Map<?, ?> values = toResultMap(result);
+ List<?> rawEmbeddings = toEmbeddingList(values.get("embeddings"));
+ List<float[]> 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/EmbeddingModelUtils.java
b/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingResult.java
similarity index 58%
copy from
api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingModelUtils.java
copy to
api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingResult.java
index c1c71e58..8c23d14d 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/EmbeddingResult.java
@@ -15,23 +15,28 @@
* See the License for the specific language governing permissions and
* limitations under the License.
*/
+
package org.apache.flink.agents.api.embedding.model;
-import java.util.List;
+import javax.annotation.Nullable;
+
+/** Embedding provider result with optional token usage metadata. */
+public class EmbeddingResult<T> {
+ 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;
+ }
-public class EmbeddingModelUtils {
- public static float[] toFloatArray(List list) {
- float[] array = new float[list.size()];
- for (int i = 0; i < list.size(); i++) {
- Object element = list.get(i);
- if (element instanceof Number) {
- array[i] = ((Number) element).floatValue();
- } else {
- throw new IllegalArgumentException(
- "Expected numeric value in embedding result, but got: "
- + element.getClass().getName());
- }
- }
- return array;
+ @Nullable
+ public EmbeddingTokenUsage getTokenUsage() {
+ return tokenUsage;
}
}
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/EmbeddingTokenUsage.java
similarity index 58%
copy from
api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingModelUtils.java
copy to
api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingTokenUsage.java
index c1c71e58..be1abdd9 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/EmbeddingTokenUsage.java
@@ -15,23 +15,24 @@
* See the License for the specific language governing permissions and
* limitations under the License.
*/
+
package org.apache.flink.agents.api.embedding.model;
-import java.util.List;
+/** 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 class EmbeddingModelUtils {
- public static float[] toFloatArray(List list) {
- float[] array = new float[list.size()];
- for (int i = 0; i < list.size(); i++) {
- Object element = list.get(i);
- if (element instanceof Number) {
- array[i] = ((Number) element).floatValue();
- } else {
- throw new IllegalArgumentException(
- "Expected numeric value in embedding result, but got: "
- + element.getClass().getName());
- }
- }
- return array;
+ 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 785ed03a..d9bdb8d4 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 @@ package org.apache.flink.agents.api.embedding.model.python;
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 @@ import static org.apache.flink.util.Preconditions.checkState;
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 class PythonEmbeddingModelConnection extends
BaseEmbeddingModelConnection
+ (results == null ? "null" :
results.getClass().getName()));
}
+ @Override
+ public EmbeddingResult<float[]> embedWithUsage(String text, Map<String,
Object> parameters) {
+ checkState(
+ embeddingModel != null,
+ "EmbeddingModelSetup is not initialized. Cannot perform embed
operation.");
+
+ Map<String, Object> 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<List<float[]>> embedWithUsage(
+ List<String> texts, Map<String, Object> parameters) {
+ checkState(
+ embeddingModel != null,
+ "EmbeddingModelSetup is not initialized. Cannot perform embed
operation.");
+
+ Map<String, Object> 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 c0efd5c3..b6febbc5 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 @@ package org.apache.flink.agents.api.embedding.model.python;
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 @@ import static org.apache.flink.util.Preconditions.checkState;
*/
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 class PythonEmbeddingModelSetup extends
BaseEmbeddingModelSetup
+ (results == null ? "null" :
results.getClass().getName()));
}
+ @Override
+ public EmbeddingResult<float[]> embedWithUsage(String text, Map<String,
Object> parameters) {
+ checkState(
+ embeddingModelSetup != null,
+ "EmbeddingModelSetup is not initialized. Cannot perform embed
operation.");
+
+ Map<String, Object> 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<List<float[]>> embedWithUsage(
+ List<String> texts, Map<String, Object> parameters) {
+ checkState(
+ embeddingModelSetup != null,
+ "EmbeddingModelSetup is not initialized. Cannot perform embed
operation.");
+
+ Map<String, Object> 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<String, Object> 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 00000000..fbbf738d
--- /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<String, Object> getParameters() {
+ return Collections.emptyMap();
+ }
+ };
+ setup.connection =
+ new BaseEmbeddingModelConnection(
+ new ResourceDescriptor("connection",
Collections.emptyMap()),
+ mock(ResourceContext.class)) {
+ @Override
+ public float[] embed(String text, Map<String, Object>
parameters) {
+ return new float[] {0.1f, 0.2f};
+ }
+
+ @Override
+ public java.util.List<float[]> embed(
+ java.util.List<String> texts, Map<String, Object>
parameters) {
+ throw new UnsupportedOperationException();
+ }
+
+ @Override
+ public EmbeddingResult<float[]> embedWithUsage(
+ String text, Map<String, Object> parameters) {
+ return new EmbeddingResult<>(
+ embed(text, parameters), new
EmbeddingTokenUsage(7L, 9L));
+ }
+ };
+
+ EmbeddingResult<float[]> 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<String, Object> getParameters() {
+ return Collections.emptyMap();
+ }
+ };
+ setup.connection =
+ new BaseEmbeddingModelConnection(
+ new ResourceDescriptor("connection",
Collections.emptyMap()),
+ mock(ResourceContext.class)) {
+ @Override
+ public float[] embed(String text, Map<String, Object>
parameters) {
+ return new float[] {0.1f, 0.2f};
+ }
+
+ @Override
+ public java.util.List<float[]> embed(
+ java.util.List<String> texts, Map<String, Object>
parameters) {
+ throw new UnsupportedOperationException();
+ }
+
+ @Override
+ public EmbeddingResult<float[]> embedWithUsage(
+ String text, Map<String, Object> 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 470d5127..a89bf331 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 @@ public class PythonEmbeddingModelConnectionTest {
assertThat(result).hasSize(2);
}
+ @Test
+ void testEmbedWithUsageMultipleTexts() {
+ List<String> 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<List<float[]>> 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<String, Object> values = (Map<String,
Object>) 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 d8071d7c..430ddc4c 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 @@ public class PythonEmbeddingModelSetupTest {
assertThat(result).hasSize(2);
}
+ @Test
+ void testEmbedWithUsageSingleText() {
+ String text = "test text";
+ Map<String, Object> 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<float[]> 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<String, Object> values = (Map<String,
Object>) 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 be5f8d49..14bfbbb4 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.EmbeddingModelConnection;
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 class EmbeddingCrossLanguageAgent extends Agent {
org.apache.flink.agents.api.resource.ResourceType
.EMBEDDING_MODEL);
- float[] embedding = embeddingModel.embed(text);
+ EmbeddingResult<float[]> 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<float[]> embeddings = embeddingModel.embed(List.of(text));
+ EmbeddingResult<List<float[]>> embeddingsResult =
+ embeddingModel.embedWithUsage(List.of(text));
+ List<float[]> 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 a7dc2926..b39ec6b2 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.ObjectMapper;
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<String, Object> parameters) {
+ return embedWithUsage(text, parameters).getEmbeddings();
+ }
+
+ @Override
+ public EmbeddingResult<float[]> embedWithUsage(String text, Map<String,
Object> parameters) {
String model = (String) parameters.getOrDefault("model", defaultModel);
Integer dimensions = (Integer) parameters.get("dimensions");
@@ -138,12 +171,21 @@ public class BedrockEmbeddingModelConnection extends
BaseEmbeddingModelConnectio
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 @@ public class BedrockEmbeddingModelConnection extends
BaseEmbeddingModelConnectio
@Override
public List<float[]> embed(List<String> texts, Map<String, Object>
parameters) {
+ return embedWithUsage(texts, parameters).getEmbeddings();
+ }
+
+ @Override
+ public EmbeddingResult<List<float[]>> embedWithUsage(
+ List<String> texts, Map<String, Object> parameters) {
if (texts.size() <= 1) {
List<float[]> results = new ArrayList<>(texts.size());
+ EmbeddingTokenUsage totalUsage = null;
for (String text : texts) {
- results.add(embed(text, parameters));
+ EmbeddingResult<float[]> result = embedWithUsage(text,
parameters);
+ results.add(result.getEmbeddings());
+ totalUsage = mergeUsage(totalUsage, result.getTokenUsage());
}
- return results;
+ return new EmbeddingResult<>(results, totalUsage);
}
@SuppressWarnings("unchecked")
- CompletableFuture<float[]>[] futures =
+ CompletableFuture<EmbeddingResult<float[]>>[] futures =
texts.stream()
.map(
text ->
CompletableFuture.supplyAsync(
- () -> embed(text, parameters),
embedPool))
+ () -> embedWithUsage(text,
parameters), embedPool))
.toArray(CompletableFuture[]::new);
CompletableFuture.allOf(futures).join();
List<float[]> results = new ArrayList<>(texts.size());
- for (CompletableFuture<float[]> f : futures) {
- results.add(f.join());
+ EmbeddingTokenUsage totalUsage = null;
+ for (CompletableFuture<EmbeddingResult<float[]>> f : futures) {
+ EmbeddingResult<float[]> 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 0891d706..99c8e9b8 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 @@ class BedrockEmbeddingModelTest {
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<List<float[]>> 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 5529c817..c31ab532 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 @@ class BaseEmbeddingModelConnection(Resource, ABC):
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 @@ class BaseEmbeddingModelSetup(Resource, ABC):
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/runtime/tests/test_python_java_utils.py
b/python/flink_agents/api/embedding_models/tests/__init__.py
similarity index 52%
copy from python/flink_agents/runtime/tests/test_python_java_utils.py
copy to python/flink_agents/api/embedding_models/tests/__init__.py
index 7aeeafe9..1b9e7e00 100644
--- a/python/flink_agents/runtime/tests/test_python_java_utils.py
+++ b/python/flink_agents/api/embedding_models/tests/__init__.py
@@ -15,27 +15,4 @@
# See the License for the specific language governing permissions and
# limitations under the License.
#################################################################################
-import json
-from flink_agents.api.decorators import tool
-from flink_agents.api.tools import InjectedArg
-from flink_agents.runtime.python_java_utils import get_python_tool_metadata
-
-
-@tool(injected_args={"tenant_id": InjectedArg.from_config("tenant.id")})
-def decorated_python_tool(order_id: str, tenant_id: str, request_id: str) ->
str:
- """Query order."""
- return f"{tenant_id}:{request_id}:{order_id}"
-
-
-def test_get_python_tool_metadata_merges_callable_injected_args() -> None:
- flat = get_python_tool_metadata(
- __name__, "decorated_python_tool", injected_args=["request_id"]
- )
-
- 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"}
- }
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 00000000..e4219352
--- /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 dbce4891..e57fbc07 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 @@ class EmbeddingModelCrossLanguageAgent(Agent):
)
# 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 @@ class EmbeddingModelCrossLanguageAgent(Agent):
)
# 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 eac4c8f1..92d09b76 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 typing_extensions import override
from flink_agents.api.embedding_models.embedding_model import (
BaseEmbeddingModelConnection,
BaseEmbeddingModelSetup,
+ EmbeddingResult,
+ EmbeddingTokenUsage,
)
DEFAULT_REQUEST_TIMEOUT = 30.0
@@ -115,6 +117,12 @@ class
OpenAIEmbeddingModelConnection(BaseEmbeddingModelConnection):
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 @@ class
OpenAIEmbeddingModelConnection(BaseEmbeddingModelConnection):
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 49907340..9d63c249 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 test_openai_embedding_model() -> None:
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 b60c7559..6262d4bd 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 test_tongyi_embedding_mock(monkeypatch:
pytest.MonkeyPatch) -> None:
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 c63c083d..bea3075b 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 pydantic import Field
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 @@ class
TongyiEmbeddingModelConnection(BaseEmbeddingModelConnection):
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 @@ class
TongyiEmbeddingModelConnection(BaseEmbeddingModelConnection):
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 b2ea4872..230a4944 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 import Any, Dict, Sequence
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 @@ from flink_agents.runtime.java.java_resource_wrapper import (
)
+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 @@ class
JavaEmbeddingModelConnectionImpl(JavaEmbeddingModelConnection):
)
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 @@ class JavaEmbeddingModelSetupImpl(JavaEmbeddingModelSetup):
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 66b9092f..3c4bb60d 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 00000000..2a386668
--- /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()
+
+
[email protected](
+ ("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 7aeeafe9..fa708353 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},
}