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},
     }

Reply via email to