From 81a7d3d14d827f0c7a8d49c4c5483c94a8cc9888 Mon Sep 17 00:00:00 2001 From: zlc <1633079383@qq.com> Date: Sat, 4 Jul 2026 14:35:49 +0800 Subject: [PATCH 1/5] [api][python][java] Track embedding token usage metrics --- .../model/BaseEmbeddingModelConnection.java | 9 + .../model/BaseEmbeddingModelSetup.java | 25 +- .../api/embedding/model/EmbeddingResult.java | 42 +++ .../embedding/model/EmbeddingTokenUsage.java | 38 +++ ...seEmbeddingModelSetupTokenMetricsTest.java | 264 ++++++++++++++++++ .../BedrockEmbeddingModelConnection.java | 113 ++++++-- .../bedrock/BedrockEmbeddingModelTest.java | 52 ++++ .../api/embedding_models/embedding_model.py | 43 ++- .../api/embedding_models/tests/__init__.py | 18 ++ .../tests/test_token_metrics.py | 198 +++++++++++++ .../openai_embedding_model.py | 26 +- .../tests/test_openai_embedding_model.py | 43 +++ .../tests/test_tongyi_embedding_model.py | 57 ++++ .../tongyi_embedding_model.py | 40 ++- 14 files changed, 939 insertions(+), 29 deletions(-) create mode 100644 api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingResult.java create mode 100644 api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingTokenUsage.java create mode 100644 api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupTokenMetricsTest.java create mode 100644 python/flink_agents/api/embedding_models/tests/__init__.py create mode 100644 python/flink_agents/api/embedding_models/tests/test_token_metrics.py diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelConnection.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelConnection.java index 4d46e3eed..01210e8f5 100644 --- a/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelConnection.java +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelConnection.java @@ -80,4 +80,13 @@ public ResourceType getResourceType() { * embeddings. The length of each array is determined by the model itself. */ public abstract List embed(List texts, Map parameters); + + public EmbeddingResult embedWithUsage(String text, Map parameters) { + return new EmbeddingResult<>(embed(text, parameters), null); + } + + public EmbeddingResult> embedWithUsage( + List texts, Map parameters) { + return new EmbeddingResult<>(embed(texts, parameters), null); + } } diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetup.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetup.java index e7c19893d..1ae3b7e86 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 @@ -18,6 +18,7 @@ package org.apache.flink.agents.api.embedding.model; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.resource.Resource; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; @@ -106,7 +107,10 @@ public float[] embed(String text) { public float[] embed(String text, Map parameters) { Map params = this.getParameters(); params.putAll(parameters); - return getConnection().embed(text, params); + BaseEmbeddingModelConnection currentConnection = getConnection(); + EmbeddingResult result = currentConnection.embedWithUsage(text, params); + recordTokenMetrics(result.getTokenUsage()); + return result.getEmbeddings(); } /** @@ -123,6 +127,23 @@ public List embed(List texts) { public List embed(List texts, Map parameters) { Map params = this.getParameters(); params.putAll(parameters); - return getConnection().embed(texts, params); + BaseEmbeddingModelConnection currentConnection = getConnection(); + EmbeddingResult> result = currentConnection.embedWithUsage(texts, params); + recordTokenMetrics(result.getTokenUsage()); + return result.getEmbeddings(); + } + + private void recordTokenMetrics(EmbeddingTokenUsage usage) { + if (usage == null) { + return; + } + FlinkAgentsMetricGroup metricGroup = getMetricGroup(); + if (metricGroup == null) { + return; + } + + FlinkAgentsMetricGroup modelGroup = metricGroup.getSubGroup("model", model); + modelGroup.getCounter("promptTokens").inc(usage.getPromptTokens()); + modelGroup.getCounter("totalTokens").inc(usage.getTotalTokens()); } } diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingResult.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingResult.java new file mode 100644 index 000000000..8c23d14d7 --- /dev/null +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingResult.java @@ -0,0 +1,42 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.flink.agents.api.embedding.model; + +import javax.annotation.Nullable; + +/** Embedding provider result with optional token usage metadata. */ +public class EmbeddingResult { + private final T embeddings; + + @Nullable private final EmbeddingTokenUsage tokenUsage; + + public EmbeddingResult(T embeddings, @Nullable EmbeddingTokenUsage tokenUsage) { + this.embeddings = embeddings; + this.tokenUsage = tokenUsage; + } + + public T getEmbeddings() { + return embeddings; + } + + @Nullable + public EmbeddingTokenUsage getTokenUsage() { + return tokenUsage; + } +} diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingTokenUsage.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingTokenUsage.java new file mode 100644 index 000000000..be1abdd98 --- /dev/null +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingTokenUsage.java @@ -0,0 +1,38 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.flink.agents.api.embedding.model; + +/** Token usage reported by an embedding provider. */ +public class EmbeddingTokenUsage { + private final long promptTokens; + private final long totalTokens; + + public EmbeddingTokenUsage(long promptTokens, long totalTokens) { + this.promptTokens = promptTokens; + this.totalTokens = totalTokens; + } + + public long getPromptTokens() { + return promptTokens; + } + + public long getTotalTokens() { + return totalTokens; + } +} diff --git a/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupTokenMetricsTest.java b/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupTokenMetricsTest.java new file mode 100644 index 000000000..354937d85 --- /dev/null +++ b/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupTokenMetricsTest.java @@ -0,0 +1,264 @@ +/* + * 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.metrics.FlinkAgentsMetricGroup; +import org.apache.flink.agents.api.metrics.UpdatableGauge; +import org.apache.flink.agents.api.resource.ResourceContext; +import org.apache.flink.agents.api.resource.ResourceDescriptor; +import org.apache.flink.metrics.Counter; +import org.apache.flink.metrics.Histogram; +import org.apache.flink.metrics.Meter; +import org.apache.flink.metrics.SimpleCounter; +import org.junit.jupiter.api.Test; + +import java.util.ArrayList; +import java.util.Collections; +import java.util.HashMap; +import java.util.List; +import java.util.Map; + +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verifyNoInteractions; + +/** Test cases for embedding model token metrics. */ +class BaseEmbeddingModelSetupTokenMetricsTest { + + private static class TestEmbeddingModelSetup extends BaseEmbeddingModelSetup { + + TestEmbeddingModelSetup(BaseEmbeddingModelConnection connection) { + super( + new ResourceDescriptor( + TestEmbeddingModelSetup.class.getName(), + Map.of("connection", "mock-connection", "model", "mock-model")), + mock(ResourceContext.class)); + this.connection = connection; + } + + @Override + public Map getParameters() { + return new HashMap<>(); + } + } + + private static class TestEmbeddingModelConnection extends BaseEmbeddingModelConnection { + + TestEmbeddingModelConnection() { + super( + new ResourceDescriptor( + TestEmbeddingModelConnection.class.getName(), Collections.emptyMap()), + mock(ResourceContext.class)); + } + + @Override + public float[] embed(String text, Map parameters) { + return new float[] {0.1f, 0.2f}; + } + + @Override + public EmbeddingResult embedWithUsage( + String text, Map parameters) { + return new EmbeddingResult<>(embed(text, parameters), new EmbeddingTokenUsage(7L, 9L)); + } + + @Override + public List embed(List texts, Map parameters) { + List embeddings = new ArrayList<>(); + for (String ignored : texts) { + embeddings.add(new float[] {0.1f, 0.2f}); + } + return embeddings; + } + + @Override + public EmbeddingResult> embedWithUsage( + List texts, Map parameters) { + return new EmbeddingResult<>( + embed(texts, parameters), new EmbeddingTokenUsage(11L, 13L)); + } + } + + private static class TestEmbeddingModelConnectionWithoutUsage + extends BaseEmbeddingModelConnection { + + TestEmbeddingModelConnectionWithoutUsage() { + super( + new ResourceDescriptor( + TestEmbeddingModelConnectionWithoutUsage.class.getName(), + Collections.emptyMap()), + mock(ResourceContext.class)); + } + + @Override + public float[] embed(String text, Map parameters) { + return new float[] {0.1f, 0.2f}; + } + + @Override + public List embed(List texts, Map parameters) { + List embeddings = new ArrayList<>(); + for (String ignored : texts) { + embeddings.add(new float[] {0.1f, 0.2f}); + } + return embeddings; + } + } + + private static class ThrowThenReportUsageConnection extends BaseEmbeddingModelConnection { + private int calls; + + ThrowThenReportUsageConnection() { + super( + new ResourceDescriptor( + ThrowThenReportUsageConnection.class.getName(), Collections.emptyMap()), + mock(ResourceContext.class)); + } + + @Override + public float[] embed(String text, Map parameters) { + return new float[] {0.1f, 0.2f}; + } + + @Override + public EmbeddingResult embedWithUsage( + String text, Map parameters) { + calls++; + if (calls == 1) { + throw new RuntimeException("provider failure"); + } + return new EmbeddingResult<>(embed(text, parameters), new EmbeddingTokenUsage(3L, 4L)); + } + + @Override + public List embed(List texts, Map parameters) { + List embeddings = new ArrayList<>(); + for (String ignored : texts) { + embeddings.add(new float[] {0.1f, 0.2f}); + } + return embeddings; + } + } + + @Test + void testEmbeddingTokenMetricsAreRecordedWhenUsageIsReported() { + TestEmbeddingModelSetup setup = + new TestEmbeddingModelSetup(new TestEmbeddingModelConnection()); + TestMetricGroup metricGroup = new TestMetricGroup(); + setup.setMetricGroup(metricGroup); + + assertArrayEquals(new float[] {0.1f, 0.2f}, setup.embed("hello")); + + TestMetricGroup modelGroup = + (TestMetricGroup) metricGroup.getSubGroup("model", "mock-model"); + assertEquals(7L, modelGroup.counters.get("promptTokens").getCount()); + assertEquals(9L, modelGroup.counters.get("totalTokens").getCount()); + } + + @Test + void testEmbeddingTokenMetricsAreNoopWhenUsageIsAbsent() { + TestEmbeddingModelSetup setup = + new TestEmbeddingModelSetup(new TestEmbeddingModelConnectionWithoutUsage()); + FlinkAgentsMetricGroup metricGroup = mock(FlinkAgentsMetricGroup.class); + setup.setMetricGroup(metricGroup); + + setup.embed("hello"); + + verifyNoInteractions(metricGroup); + } + + @Test + void testEmbeddingTokenMetricsAccumulateAcrossRequests() { + TestEmbeddingModelSetup setup = + new TestEmbeddingModelSetup(new TestEmbeddingModelConnection()); + TestMetricGroup metricGroup = new TestMetricGroup(); + setup.setMetricGroup(metricGroup); + + setup.embed("hello"); + setup.embed(List.of("hello", "flink")); + + TestMetricGroup modelGroup = + (TestMetricGroup) metricGroup.getSubGroup("model", "mock-model"); + assertEquals(18L, modelGroup.counters.get("promptTokens").getCount()); + assertEquals(22L, modelGroup.counters.get("totalTokens").getCount()); + } + + @Test + void testEmbeddingTokenMetricsDoNotLeakAfterProviderFailure() { + TestEmbeddingModelSetup setup = + new TestEmbeddingModelSetup(new ThrowThenReportUsageConnection()); + TestMetricGroup metricGroup = new TestMetricGroup(); + setup.setMetricGroup(metricGroup); + + assertThrows(RuntimeException.class, () -> setup.embed("first")); + assertArrayEquals(new float[] {0.1f, 0.2f}, setup.embed("second")); + + TestMetricGroup modelGroup = + (TestMetricGroup) metricGroup.getSubGroup("model", "mock-model"); + assertEquals(3L, modelGroup.counters.get("promptTokens").getCount()); + assertEquals(4L, modelGroup.counters.get("totalTokens").getCount()); + } + + private static class TestMetricGroup implements FlinkAgentsMetricGroup { + final Map subGroups = new HashMap<>(); + final Map counters = new HashMap<>(); + + @Override + public FlinkAgentsMetricGroup getSubGroup(String name) { + return subGroups.computeIfAbsent(name, ignored -> new TestMetricGroup()); + } + + @Override + public FlinkAgentsMetricGroup getSubGroup(String key, String value) { + return subGroups.computeIfAbsent(key + "=" + value, ignored -> new TestMetricGroup()); + } + + @Override + public UpdatableGauge getGauge(String name) { + return null; + } + + @Override + public Counter getCounter(String name) { + return counters.computeIfAbsent(name, ignored -> new SimpleCounter()); + } + + @Override + public Meter getMeter(String name) { + return null; + } + + @Override + public Meter getMeter(String name, Counter counter) { + return null; + } + + @Override + public Histogram getHistogram(String name) { + return null; + } + + @Override + public Histogram getHistogram(String name, int windowSize) { + return null; + } + } +} diff --git a/integrations/embedding-models/bedrock/src/main/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelConnection.java b/integrations/embedding-models/bedrock/src/main/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelConnection.java index a7dc2926c..b39ec6b20 100644 --- a/integrations/embedding-models/bedrock/src/main/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelConnection.java +++ b/integrations/embedding-models/bedrock/src/main/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelConnection.java @@ -23,6 +23,8 @@ import com.fasterxml.jackson.databind.node.ObjectNode; import org.apache.flink.agents.api.RetryExecutor; import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelConnection; +import org.apache.flink.agents.api.embedding.model.EmbeddingResult; +import org.apache.flink.agents.api.embedding.model.EmbeddingTokenUsage; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import software.amazon.awssdk.auth.credentials.DefaultCredentialsProvider; @@ -80,36 +82,67 @@ public class BedrockEmbeddingModelConnection extends BaseEmbeddingModelConnectio public BedrockEmbeddingModelConnection( ResourceDescriptor descriptor, ResourceContext resourceContext) { super(descriptor, resourceContext); + this.client = createClient(resolveRegion(descriptor)); + this.defaultModel = resolveDefaultModel(descriptor); + this.embedPool = Executors.newFixedThreadPool(resolveEmbedConcurrency(descriptor)); + this.retryExecutor = createRetryExecutor(descriptor); + } + + BedrockEmbeddingModelConnection( + ResourceDescriptor descriptor, + ResourceContext resourceContext, + BedrockRuntimeClient client, + ExecutorService embedPool, + RetryExecutor retryExecutor, + String defaultModel) { + super(descriptor, resourceContext); + this.client = client; + this.defaultModel = defaultModel; + this.embedPool = embedPool; + this.retryExecutor = retryExecutor; + } + + private static BedrockRuntimeClient createClient(String region) { + return BedrockRuntimeClient.builder() + .region(Region.of(region)) + .credentialsProvider(DefaultCredentialsProvider.create()) + .build(); + } + private static String resolveRegion(ResourceDescriptor descriptor) { String region = descriptor.getArgument("region"); if (region == null || region.isBlank()) { region = "us-east-1"; } + return region; + } - this.client = - BedrockRuntimeClient.builder() - .region(Region.of(region)) - .credentialsProvider(DefaultCredentialsProvider.create()) - .build(); - + private static String resolveDefaultModel(ResourceDescriptor descriptor) { String model = descriptor.getArgument("model"); - this.defaultModel = (model != null && !model.isBlank()) ? model : DEFAULT_MODEL; + return (model != null && !model.isBlank()) ? model : DEFAULT_MODEL; + } + private static int resolveEmbedConcurrency(ResourceDescriptor descriptor) { Integer concurrency = descriptor.getArgument("embed_concurrency"); - int threads = concurrency != null ? concurrency : 4; - this.embedPool = Executors.newFixedThreadPool(threads); + return concurrency != null ? concurrency : 4; + } + private static RetryExecutor createRetryExecutor(ResourceDescriptor descriptor) { Integer retries = descriptor.getArgument("max_retries"); - this.retryExecutor = - RetryExecutor.builder() - .maxRetries(retries != null ? retries : 5) - .initialBackoffMs(200) - .retryablePredicate(BedrockEmbeddingModelConnection::isRetryable) - .build(); + return RetryExecutor.builder() + .maxRetries(retries != null ? retries : 5) + .initialBackoffMs(200) + .retryablePredicate(BedrockEmbeddingModelConnection::isRetryable) + .build(); } @Override public float[] embed(String text, Map parameters) { + return embedWithUsage(text, parameters).getEmbeddings(); + } + + @Override + public EmbeddingResult embedWithUsage(String text, Map parameters) { String model = (String) parameters.getOrDefault("model", defaultModel); Integer dimensions = (Integer) parameters.get("dimensions"); @@ -138,12 +171,21 @@ public float[] embed(String text, Map parameters) { for (int i = 0; i < embeddingNode.size(); i++) { embedding[i] = (float) embeddingNode.get(i).asDouble(); } - return embedding; + return new EmbeddingResult<>(embedding, extractTokenUsage(result)); } catch (Exception e) { throw new RuntimeException("Failed to parse Bedrock embedding response.", e); } } + private static EmbeddingTokenUsage extractTokenUsage(JsonNode result) { + JsonNode inputTokenCount = result.get("inputTextTokenCount"); + if (inputTokenCount == null || !inputTokenCount.isNumber()) { + return null; + } + long tokens = inputTokenCount.asLong(); + return new EmbeddingTokenUsage(tokens, tokens); + } + private static boolean isRetryable(Exception e) { String msg = e.toString(); return msg.contains("ThrottlingException") @@ -155,27 +197,52 @@ private static boolean isRetryable(Exception e) { @Override public List embed(List texts, Map parameters) { + return embedWithUsage(texts, parameters).getEmbeddings(); + } + + @Override + public EmbeddingResult> embedWithUsage( + List texts, Map parameters) { if (texts.size() <= 1) { List results = new ArrayList<>(texts.size()); + EmbeddingTokenUsage totalUsage = null; for (String text : texts) { - results.add(embed(text, parameters)); + EmbeddingResult result = embedWithUsage(text, parameters); + results.add(result.getEmbeddings()); + totalUsage = mergeUsage(totalUsage, result.getTokenUsage()); } - return results; + return new EmbeddingResult<>(results, totalUsage); } @SuppressWarnings("unchecked") - CompletableFuture[] futures = + CompletableFuture>[] futures = texts.stream() .map( text -> CompletableFuture.supplyAsync( - () -> embed(text, parameters), embedPool)) + () -> embedWithUsage(text, parameters), embedPool)) .toArray(CompletableFuture[]::new); CompletableFuture.allOf(futures).join(); List results = new ArrayList<>(texts.size()); - for (CompletableFuture f : futures) { - results.add(f.join()); + EmbeddingTokenUsage totalUsage = null; + for (CompletableFuture> f : futures) { + EmbeddingResult result = f.join(); + results.add(result.getEmbeddings()); + totalUsage = mergeUsage(totalUsage, result.getTokenUsage()); + } + return new EmbeddingResult<>(results, totalUsage); + } + + private static EmbeddingTokenUsage mergeUsage( + EmbeddingTokenUsage left, EmbeddingTokenUsage right) { + if (left == null) { + return right; + } + if (right == null) { + return left; } - return results; + return new EmbeddingTokenUsage( + left.getPromptTokens() + right.getPromptTokens(), + left.getTotalTokens() + right.getTotalTokens()); } @Override diff --git a/integrations/embedding-models/bedrock/src/test/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelTest.java b/integrations/embedding-models/bedrock/src/test/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelTest.java index 0891d7067..99c8e9b81 100644 --- a/integrations/embedding-models/bedrock/src/test/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelTest.java +++ b/integrations/embedding-models/bedrock/src/test/java/org/apache/flink/agents/integrations/embeddingmodels/bedrock/BedrockEmbeddingModelTest.java @@ -18,17 +18,31 @@ package org.apache.flink.agents.integrations.embeddingmodels.bedrock; +import org.apache.flink.agents.api.RetryExecutor; import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelConnection; import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelSetup; +import org.apache.flink.agents.api.embedding.model.EmbeddingResult; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; +import software.amazon.awssdk.core.SdkBytes; +import software.amazon.awssdk.services.bedrockruntime.BedrockRuntimeClient; +import software.amazon.awssdk.services.bedrockruntime.model.InvokeModelRequest; +import software.amazon.awssdk.services.bedrockruntime.model.InvokeModelResponse; +import java.util.List; import java.util.Map; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; import static org.assertj.core.api.Assertions.assertThat; import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; /** Tests for {@link BedrockEmbeddingModelConnection} and {@link BedrockEmbeddingModelSetup}. */ class BedrockEmbeddingModelTest { @@ -94,4 +108,42 @@ void testSetupParametersNoDimensions() { assertThat(setup.getParameters()).doesNotContainKey("dimensions"); } + + @Test + @DisplayName("Batch embedding aggregates token usage from worker results") + void testBatchEmbeddingAggregatesTokenUsage() throws Exception { + BedrockRuntimeClient client = mock(BedrockRuntimeClient.class); + when(client.invokeModel(any(InvokeModelRequest.class))) + .thenReturn( + InvokeModelResponse.builder() + .body( + SdkBytes.fromUtf8String( + "{\"embedding\":[0.1,0.2],\"inputTextTokenCount\":5}")) + .build()); + + ExecutorService embedPool = Executors.newFixedThreadPool(2); + BedrockEmbeddingModelConnection conn = + new BedrockEmbeddingModelConnection( + connDescriptor(null), + NOOP, + client, + embedPool, + RetryExecutor.builder().maxRetries(0).build(), + "mock-model"); + + try { + EmbeddingResult> result = + conn.embedWithUsage(List.of("first", "second"), Map.of("model", "mock-model")); + + assertThat(result.getEmbeddings()).hasSize(2); + assertThat(result.getEmbeddings().get(0)).containsExactly(0.1f, 0.2f); + assertThat(result.getEmbeddings().get(1)).containsExactly(0.1f, 0.2f); + assertThat(result.getTokenUsage()).isNotNull(); + assertThat(result.getTokenUsage().getPromptTokens()).isEqualTo(10L); + assertThat(result.getTokenUsage().getTotalTokens()).isEqualTo(10L); + verify(client, times(2)).invokeModel(any(InvokeModelRequest.class)); + } finally { + conn.close(); + } + } } diff --git a/python/flink_agents/api/embedding_models/embedding_model.py b/python/flink_agents/api/embedding_models/embedding_model.py index 5529c817e..ecdee5ce5 100644 --- a/python/flink_agents/api/embedding_models/embedding_model.py +++ b/python/flink_agents/api/embedding_models/embedding_model.py @@ -16,13 +16,32 @@ # limitations under the License. ################################################################################# from abc import ABC, abstractmethod -from typing import Any, Dict, Sequence, cast +from dataclasses import dataclass +from typing import Any, Dict, Generic, Sequence, TypeVar, cast from pydantic import Field from typing_extensions import override from flink_agents.api.resource import Resource, ResourceType +EmbeddingValue = TypeVar("EmbeddingValue", list[float], list[list[float]]) + + +@dataclass(frozen=True) +class EmbeddingTokenUsage: + """Token usage reported by an embedding provider.""" + + prompt_tokens: int = 0 + total_tokens: int = 0 + + +@dataclass(frozen=True) +class EmbeddingResult(Generic[EmbeddingValue]): + """Embedding provider result with optional token usage metadata.""" + + embeddings: EmbeddingValue + token_usage: EmbeddingTokenUsage | None = None + class BaseEmbeddingModelConnection(Resource, ABC): """Base abstract class for text embedding model connection. @@ -62,6 +81,12 @@ def embed( The dimension of the vector depends on the specific embedding model used. """ + def embed_with_usage( + self, text: str | Sequence[str], **kwargs: Any + ) -> EmbeddingResult[list[float] | list[list[float]]]: + """Generate embeddings and return provider token usage when available.""" + return EmbeddingResult(embeddings=self.embed(text, **kwargs)) + class BaseEmbeddingModelSetup(Resource, ABC): """Base abstract class for text embedding model setup. @@ -122,4 +147,18 @@ def embed( """ merged_kwargs = self.model_kwargs.copy() merged_kwargs.update(kwargs) - return self._get_connection().embed(text, **merged_kwargs) + result = self._get_connection().embed_with_usage(text, **merged_kwargs) + self._record_token_metrics(result.token_usage) + return result.embeddings + + def _record_token_metrics(self, usage: EmbeddingTokenUsage | None) -> None: + """Record embedding token metrics under the current model metric group.""" + if usage is None: + return + metric_group = self.metric_group + if metric_group is None: + return + + model_group = metric_group.get_sub_group("model", self.model) + model_group.get_counter("promptTokens").inc(usage.prompt_tokens) + model_group.get_counter("totalTokens").inc(usage.total_tokens) diff --git a/python/flink_agents/api/embedding_models/tests/__init__.py b/python/flink_agents/api/embedding_models/tests/__init__.py new file mode 100644 index 000000000..1b9e7e00a --- /dev/null +++ b/python/flink_agents/api/embedding_models/tests/__init__.py @@ -0,0 +1,18 @@ +################################################################################ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +################################################################################# + diff --git a/python/flink_agents/api/embedding_models/tests/test_token_metrics.py b/python/flink_agents/api/embedding_models/tests/test_token_metrics.py new file mode 100644 index 000000000..c71850c81 --- /dev/null +++ b/python/flink_agents/api/embedding_models/tests/test_token_metrics.py @@ -0,0 +1,198 @@ +################################################################################ +# 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 + +import pytest + +from flink_agents.api.embedding_models.embedding_model import ( + BaseEmbeddingModelConnection, + BaseEmbeddingModelSetup, + EmbeddingResult, + EmbeddingTokenUsage, +) +from flink_agents.api.metric_group import Counter, MetricGroup +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 ThrowThenReportUsageConnection(BaseEmbeddingModelConnection): + def __init__(self, **kwargs: Any) -> None: + super().__init__(**kwargs) + self._calls = 0 + + 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]]]: + self._calls += 1 + if self._calls == 1: + msg = "provider failure" + raise RuntimeError(msg) + return EmbeddingResult( + embeddings=self.embed(text, **kwargs), + token_usage=EmbeddingTokenUsage(prompt_tokens=3, total_tokens=4), + ) + + +class FakeEmbeddingModelSetup(BaseEmbeddingModelSetup): + @property + def model_kwargs(self) -> Dict[str, Any]: + return {} + + +class _MockCounter(Counter): + def __init__(self) -> None: + self._count = 0 + + def inc(self, n: int = 1) -> None: + self._count += n + + def dec(self, n: int = 1) -> None: + self._count -= n + + def get_count(self) -> int: + return self._count + + +class _MockMetricGroup(MetricGroup): + def __init__(self) -> None: + self._sub_groups: dict[str, _MockMetricGroup] = {} + self._counters: dict[str, _MockCounter] = {} + + def get_sub_group(self, name: str, value: str | None = None) -> "_MockMetricGroup": + key = f"{name}={value}" if value is not None else name + if key not in self._sub_groups: + self._sub_groups[key] = _MockMetricGroup() + return self._sub_groups[key] + + def get_counter(self, name: str) -> _MockCounter: + if name not in self._counters: + self._counters[name] = _MockCounter() + return self._counters[name] + + def get_meter(self, name: str) -> Any: + return MagicMock() + + def get_gauge(self, name: str) -> Any: + return MagicMock() + + def get_histogram(self, name: str, window_size: int = 100) -> Any: + return MagicMock() + + +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_token_metrics_are_recorded_when_usage_is_reported() -> None: + setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) + metric_group = _MockMetricGroup() + setup.set_metric_group(metric_group) + + assert setup.embed("hello") == [0.1, 0.2] + + model_group = metric_group.get_sub_group("model", "mock-model") + assert model_group.get_counter("promptTokens").get_count() == 7 + assert model_group.get_counter("totalTokens").get_count() == 9 + + +def test_embedding_token_metrics_are_noop_when_usage_is_absent() -> None: + setup = _make_setup(FakeEmbeddingModelConnectionWithoutUsage(name="connection")) + metric_group = _MockMetricGroup() + setup.set_metric_group(metric_group) + + assert setup.embed("hello") == [0.1, 0.2] + + model_group = metric_group.get_sub_group("model", "mock-model") + assert "promptTokens" not in model_group._counters + assert "totalTokens" not in model_group._counters + + +def test_embedding_token_metrics_accumulate_across_requests() -> None: + setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) + metric_group = _MockMetricGroup() + setup.set_metric_group(metric_group) + + setup.embed("hello") + setup.embed(["hello", "flink"]) + + model_group = metric_group.get_sub_group("model", "mock-model") + assert model_group.get_counter("promptTokens").get_count() == 14 + assert model_group.get_counter("totalTokens").get_count() == 18 + + +def test_embedding_token_metrics_do_not_leak_after_provider_failure() -> None: + setup = _make_setup(ThrowThenReportUsageConnection(name="connection")) + metric_group = _MockMetricGroup() + setup.set_metric_group(metric_group) + + with pytest.raises(RuntimeError, match="provider failure"): + setup.embed("first") + + assert setup.embed("second") == [0.1, 0.2] + + model_group = metric_group.get_sub_group("model", "mock-model") + assert model_group.get_counter("promptTokens").get_count() == 3 + assert model_group.get_counter("totalTokens").get_count() == 4 diff --git a/python/flink_agents/integrations/embedding_models/openai_embedding_model.py b/python/flink_agents/integrations/embedding_models/openai_embedding_model.py index eac4c8f10..92d09b765 100644 --- a/python/flink_agents/integrations/embedding_models/openai_embedding_model.py +++ b/python/flink_agents/integrations/embedding_models/openai_embedding_model.py @@ -24,6 +24,8 @@ from flink_agents.api.embedding_models.embedding_model import ( BaseEmbeddingModelConnection, BaseEmbeddingModelSetup, + EmbeddingResult, + EmbeddingTokenUsage, ) DEFAULT_REQUEST_TIMEOUT = 30.0 @@ -115,6 +117,12 @@ def embed( self, text: str | Sequence[str], **kwargs: Any ) -> list[float] | list[list[float]]: """Generate embedding vector for a single text query.""" + return self.embed_with_usage(text, **kwargs).embeddings + + def embed_with_usage( + self, text: str | Sequence[str], **kwargs: Any + ) -> EmbeddingResult[list[float] | list[list[float]]]: + """Generate embeddings and return OpenAI token usage when available.""" # Extract OpenAI specific parameters model = kwargs.pop("model") encoding_format = kwargs.pop("encoding_format", None) @@ -132,8 +140,24 @@ def embed( user=user if user is not None else NOT_GIVEN, ) + usage = getattr(response, "usage", None) + token_usage = None + if usage is not None: + prompt_tokens = getattr(usage, "prompt_tokens", None) + total_tokens = getattr(usage, "total_tokens", None) + if prompt_tokens is not None or total_tokens is not None: + token_usage = EmbeddingTokenUsage( + prompt_tokens=int(prompt_tokens or 0), + total_tokens=int( + total_tokens if total_tokens is not None else prompt_tokens + ), + ) + embeddings = [list(embedding.embedding) for embedding in response.data] - return embeddings[0] if isinstance(text, str) else embeddings + return EmbeddingResult( + embeddings=embeddings[0] if isinstance(text, str) else embeddings, + token_usage=token_usage, + ) @override def close(self) -> None: diff --git a/python/flink_agents/integrations/embedding_models/tests/test_openai_embedding_model.py b/python/flink_agents/integrations/embedding_models/tests/test_openai_embedding_model.py index 49907340f..76add0228 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,45 @@ def get_resource(name: str, type: ResourceType) -> Resource: assert isinstance(response, list) assert len(response) > 0 assert all(isinstance(x, float) for x in response) # + + +def test_openai_embedding_model_records_token_metrics() -> None: + """Test OpenAI embedding usage is recorded as model token metrics.""" + 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 + ) + metric_group = MagicMock() + model_group = MagicMock() + prompt_counter = MagicMock() + total_counter = MagicMock() + metric_group.get_sub_group.return_value = model_group + model_group.get_counter.side_effect = { + "promptTokens": prompt_counter, + "totalTokens": total_counter, + }.__getitem__ + + embedding_model.open() + embedding_model.set_metric_group(metric_group) + + assert embedding_model.embed("Hello, Flink Agent!") == [0.1, 0.2, 0.3] + + metric_group.get_sub_group.assert_called_once_with("model", test_model) + prompt_counter.inc.assert_called_once_with(5) + total_counter.inc.assert_called_once_with(5) diff --git a/python/flink_agents/integrations/embedding_models/tests/test_tongyi_embedding_model.py b/python/flink_agents/integrations/embedding_models/tests/test_tongyi_embedding_model.py index b60c75596..faa40bd73 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,63 @@ def get_resource(name: str, type: ResourceType) -> Resource: assert len(response) == 5 +def test_tongyi_embedding_records_token_metrics( + 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": {"input_tokens": 6, "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), + ) + metric_group = MagicMock() + model_group = MagicMock() + prompt_counter = MagicMock() + total_counter = MagicMock() + metric_group.get_sub_group.return_value = model_group + model_group.get_counter.side_effect = { + "promptTokens": prompt_counter, + "totalTokens": total_counter, + }.__getitem__ + + embedding_model.open() + embedding_model.set_metric_group(metric_group) + + assert embedding_model.embed("Test text") == mock_embedding + + metric_group.get_sub_group.assert_called_once_with("model", test_model) + prompt_counter.inc.assert_called_once_with(6) + total_counter.inc.assert_called_once_with(6) + + def test_tongyi_embedding_batch_mock(monkeypatch: pytest.MonkeyPatch) -> None: """Test batch embedding functionality with mocked DashScope API.""" mock_embeddings = [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]] diff --git a/python/flink_agents/integrations/embedding_models/tongyi_embedding_model.py b/python/flink_agents/integrations/embedding_models/tongyi_embedding_model.py index c63c083d4..bbcc981e4 100644 --- a/python/flink_agents/integrations/embedding_models/tongyi_embedding_model.py +++ b/python/flink_agents/integrations/embedding_models/tongyi_embedding_model.py @@ -25,12 +25,27 @@ from flink_agents.api.embedding_models.embedding_model import ( BaseEmbeddingModelConnection, BaseEmbeddingModelSetup, + EmbeddingResult, + EmbeddingTokenUsage, ) DEFAULT_REQUEST_TIMEOUT = 30.0 DEFAULT_MODEL = "text-embedding-v4" +def _get_usage_value(obj: Any, *names: str) -> int | None: + """Read a token usage value from dict-like or object-like provider responses.""" + if obj is None: + return None + for name in names: + if isinstance(obj, dict) and name in obj: + return int(obj[name]) + value = getattr(obj, name, None) + if value is not None: + return int(value) + return None + + class TongyiEmbeddingModelConnection(BaseEmbeddingModelConnection): """Tongyi Embedding Model Connection which manages connection to DashScope API. @@ -78,6 +93,12 @@ def embed( self, text: str | Sequence[str], **kwargs: Any ) -> list[float] | list[list[float]]: """Generate embedding vector for text input.""" + return self.embed_with_usage(text, **kwargs).embeddings + + def embed_with_usage( + self, text: str | Sequence[str], **kwargs: Any + ) -> EmbeddingResult[list[float] | list[list[float]]]: + """Generate embeddings and return DashScope token usage when available.""" model = kwargs.pop("model", DEFAULT_MODEL) text_type = kwargs.pop("text_type", None) dimension = kwargs.pop("dimension", None) @@ -103,8 +124,25 @@ def embed( msg = f"DashScope TextEmbedding call failed: {response.message}" raise RuntimeError(msg) + usage = getattr(response, "usage", None) + if usage is None and isinstance(response.output, dict): + usage = response.output.get("usage") + prompt_tokens = _get_usage_value(usage, "input_tokens", "prompt_tokens") + total_tokens = _get_usage_value(usage, "total_tokens") + 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): From 84f73e1e8a514a0352470caffcbec1a2f39968bf Mon Sep 17 00:00:00 2001 From: zlc <1633079383@qq.com> Date: Wed, 8 Jul 2026 14:21:18 +0800 Subject: [PATCH 2/5] fix: record embedding metrics at request boundary --- .../model/BaseEmbeddingModelSetup.java | 55 +++++++++++---- .../api/vectorstores/BaseVectorStore.java | 15 +++- ...seEmbeddingModelSetupTokenMetricsTest.java | 69 ++++++++++++++++--- .../api/embedding_models/embedding_model.py | 22 +++--- .../tests/test_token_metrics.py | 68 +++++++++++++++--- .../api/vector_stores/vector_store.py | 15 +++- .../tests/test_openai_embedding_model.py | 5 +- .../tests/test_tongyi_embedding_model.py | 5 +- 8 files changed, 204 insertions(+), 50 deletions(-) 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 1ae3b7e86..a3c8a061a 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 @@ -105,12 +105,27 @@ public float[] embed(String text) { } public float[] embed(String text, Map parameters) { + return embedWithUsage(text, parameters).getEmbeddings(); + } + + public float[] embed( + String text, + Map parameters, + @Nullable FlinkAgentsMetricGroup requestMetricGroup) { + EmbeddingResult result = embedWithUsage(text, parameters); + recordTokenMetrics(requestMetricGroup, result.getTokenUsage()); + return result.getEmbeddings(); + } + + public EmbeddingResult embedWithUsage(String text) { + return embedWithUsage(text, Collections.emptyMap()); + } + + public EmbeddingResult embedWithUsage(String text, Map parameters) { Map params = this.getParameters(); params.putAll(parameters); BaseEmbeddingModelConnection currentConnection = getConnection(); - EmbeddingResult result = currentConnection.embedWithUsage(text, params); - recordTokenMetrics(result.getTokenUsage()); - return result.getEmbeddings(); + return currentConnection.embedWithUsage(text, params); } /** @@ -125,24 +140,38 @@ public List embed(List texts) { } public List embed(List texts, Map parameters) { + return embedWithUsage(texts, parameters).getEmbeddings(); + } + + public List embed( + List texts, + Map parameters, + @Nullable FlinkAgentsMetricGroup requestMetricGroup) { + EmbeddingResult> result = embedWithUsage(texts, parameters); + recordTokenMetrics(requestMetricGroup, result.getTokenUsage()); + return result.getEmbeddings(); + } + + public EmbeddingResult> embedWithUsage(List texts) { + return embedWithUsage(texts, Collections.emptyMap()); + } + + public EmbeddingResult> embedWithUsage( + List texts, Map parameters) { Map params = this.getParameters(); params.putAll(parameters); BaseEmbeddingModelConnection currentConnection = getConnection(); - EmbeddingResult> result = currentConnection.embedWithUsage(texts, params); - recordTokenMetrics(result.getTokenUsage()); - return result.getEmbeddings(); + return currentConnection.embedWithUsage(texts, params); } - private void recordTokenMetrics(EmbeddingTokenUsage usage) { - if (usage == null) { - return; - } - FlinkAgentsMetricGroup metricGroup = getMetricGroup(); - if (metricGroup == null) { + public void recordTokenMetrics( + @Nullable FlinkAgentsMetricGroup requestMetricGroup, + @Nullable EmbeddingTokenUsage usage) { + if (requestMetricGroup == null || usage == null) { return; } - FlinkAgentsMetricGroup modelGroup = metricGroup.getSubGroup("model", model); + FlinkAgentsMetricGroup modelGroup = requestMetricGroup.getSubGroup("model", model); modelGroup.getCounter("promptTokens").inc(usage.getPromptTokens()); modelGroup.getCounter("totalTokens").inc(usage.getTotalTokens()); } diff --git a/api/src/main/java/org/apache/flink/agents/api/vectorstores/BaseVectorStore.java b/api/src/main/java/org/apache/flink/agents/api/vectorstores/BaseVectorStore.java index d0fc4ce72..aad119f3f 100644 --- a/api/src/main/java/org/apache/flink/agents/api/vectorstores/BaseVectorStore.java +++ b/api/src/main/java/org/apache/flink/agents/api/vectorstores/BaseVectorStore.java @@ -19,6 +19,7 @@ package org.apache.flink.agents.api.vectorstores; import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelSetup; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.resource.Resource; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; @@ -27,6 +28,7 @@ import javax.annotation.Nullable; import java.io.IOException; +import java.util.Collections; import java.util.List; import java.util.Map; @@ -176,7 +178,10 @@ public void update( * @return VectorStoreQueryResult containing the retrieved documents */ public VectorStoreQueryResult query(VectorStoreQuery query) { - final float[] queryEmbedding = getEmbeddingModel().embed(query.getQueryText()); + final FlinkAgentsMetricGroup requestMetricGroup = getMetricGroup(); + final float[] queryEmbedding = + getEmbeddingModel() + .embed(query.getQueryText(), Collections.emptyMap(), requestMetricGroup); final Map storeKwargs = this.getStoreKwargs(); storeKwargs.putAll(query.getExtraArgs()); @@ -332,9 +337,15 @@ public abstract void updateEmbedding( /** Auto-embed any documents whose {@code embedding} field is {@code null}. */ protected void ensureEmbeddings(List documents) { + final FlinkAgentsMetricGroup requestMetricGroup = getMetricGroup(); for (Document doc : documents) { if (doc.getEmbedding() == null) { - doc.setEmbedding(getEmbeddingModel().embed(doc.getContent())); + doc.setEmbedding( + getEmbeddingModel() + .embed( + doc.getContent(), + Collections.emptyMap(), + requestMetricGroup)); } } } diff --git a/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupTokenMetricsTest.java b/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupTokenMetricsTest.java index 354937d85..3e4e6ad9e 100644 --- a/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupTokenMetricsTest.java +++ b/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupTokenMetricsTest.java @@ -163,9 +163,10 @@ void testEmbeddingTokenMetricsAreRecordedWhenUsageIsReported() { TestEmbeddingModelSetup setup = new TestEmbeddingModelSetup(new TestEmbeddingModelConnection()); TestMetricGroup metricGroup = new TestMetricGroup(); - setup.setMetricGroup(metricGroup); - assertArrayEquals(new float[] {0.1f, 0.2f}, setup.embed("hello")); + assertArrayEquals( + new float[] {0.1f, 0.2f}, + setup.embed("hello", Collections.emptyMap(), metricGroup)); TestMetricGroup modelGroup = (TestMetricGroup) metricGroup.getSubGroup("model", "mock-model"); @@ -178,22 +179,55 @@ void testEmbeddingTokenMetricsAreNoopWhenUsageIsAbsent() { TestEmbeddingModelSetup setup = new TestEmbeddingModelSetup(new TestEmbeddingModelConnectionWithoutUsage()); FlinkAgentsMetricGroup metricGroup = mock(FlinkAgentsMetricGroup.class); - setup.setMetricGroup(metricGroup); - setup.embed("hello"); + setup.embed("hello", Collections.emptyMap(), metricGroup); verifyNoInteractions(metricGroup); } + @Test + void testEmbeddingTokenMetricsAreNoopWhenMetricGroupIsAbsent() { + TestEmbeddingModelSetup setup = + new TestEmbeddingModelSetup(new TestEmbeddingModelConnection()); + FlinkAgentsMetricGroup boundMetricGroup = mock(FlinkAgentsMetricGroup.class); + setup.setMetricGroup(boundMetricGroup); + + setup.embed("hello", Collections.emptyMap(), null); + + verifyNoInteractions(boundMetricGroup); + } + + @Test + void testEmbeddingTokenMetricsUseRequestScopedMetricGroup() { + TestEmbeddingModelSetup setup = + new TestEmbeddingModelSetup(new TestEmbeddingModelConnection()); + TestMetricGroup actionA = new TestMetricGroup(); + TestMetricGroup actionB = new TestMetricGroup(); + + setup.setMetricGroup(actionB); + + assertArrayEquals( + new float[] {0.1f, 0.2f}, setup.embed("hello", Collections.emptyMap(), actionA)); + + TestMetricGroup actionAModelGroup = + (TestMetricGroup) actionA.getSubGroup("model", "mock-model"); + assertEquals(7L, actionAModelGroup.counters.get("promptTokens").getCount()); + assertEquals(9L, actionAModelGroup.counters.get("totalTokens").getCount()); + + TestMetricGroup actionBModelGroup = + (TestMetricGroup) actionB.getSubGroup("model", "mock-model"); + assertEquals(0L, actionBModelGroup.getCounter("promptTokens").getCount()); + assertEquals(0L, actionBModelGroup.getCounter("totalTokens").getCount()); + } + @Test void testEmbeddingTokenMetricsAccumulateAcrossRequests() { TestEmbeddingModelSetup setup = new TestEmbeddingModelSetup(new TestEmbeddingModelConnection()); TestMetricGroup metricGroup = new TestMetricGroup(); - setup.setMetricGroup(metricGroup); - setup.embed("hello"); - setup.embed(List.of("hello", "flink")); + setup.embed("hello", Collections.emptyMap(), metricGroup); + setup.embed(List.of("hello", "flink"), Collections.emptyMap(), metricGroup); TestMetricGroup modelGroup = (TestMetricGroup) metricGroup.getSubGroup("model", "mock-model"); @@ -206,10 +240,13 @@ void testEmbeddingTokenMetricsDoNotLeakAfterProviderFailure() { TestEmbeddingModelSetup setup = new TestEmbeddingModelSetup(new ThrowThenReportUsageConnection()); TestMetricGroup metricGroup = new TestMetricGroup(); - setup.setMetricGroup(metricGroup); - assertThrows(RuntimeException.class, () -> setup.embed("first")); - assertArrayEquals(new float[] {0.1f, 0.2f}, setup.embed("second")); + assertThrows( + RuntimeException.class, + () -> setup.embed("first", Collections.emptyMap(), metricGroup)); + assertArrayEquals( + new float[] {0.1f, 0.2f}, + setup.embed("second", Collections.emptyMap(), metricGroup)); TestMetricGroup modelGroup = (TestMetricGroup) metricGroup.getSubGroup("model", "mock-model"); @@ -217,6 +254,18 @@ void testEmbeddingTokenMetricsDoNotLeakAfterProviderFailure() { assertEquals(4L, modelGroup.counters.get("totalTokens").getCount()); } + @Test + void testEmbeddingWithoutExplicitMetricGroupDoesNotReadBoundMetricGroup() { + TestEmbeddingModelSetup setup = + new TestEmbeddingModelSetup(new TestEmbeddingModelConnection()); + FlinkAgentsMetricGroup boundMetricGroup = mock(FlinkAgentsMetricGroup.class); + setup.setMetricGroup(boundMetricGroup); + + assertArrayEquals(new float[] {0.1f, 0.2f}, setup.embed("hello")); + + verifyNoInteractions(boundMetricGroup); + } + private static class TestMetricGroup implements FlinkAgentsMetricGroup { final Map subGroups = new HashMap<>(); final Map counters = new HashMap<>(); diff --git a/python/flink_agents/api/embedding_models/embedding_model.py b/python/flink_agents/api/embedding_models/embedding_model.py index ecdee5ce5..d9943c270 100644 --- a/python/flink_agents/api/embedding_models/embedding_model.py +++ b/python/flink_agents/api/embedding_models/embedding_model.py @@ -22,6 +22,7 @@ from pydantic import Field from typing_extensions import override +from flink_agents.api.metric_group import MetricGroup from flink_agents.api.resource import Resource, ResourceType EmbeddingValue = TypeVar("EmbeddingValue", list[float], list[list[float]]) @@ -145,18 +146,21 @@ def embed( A list of floating-point numbers representing the embedding vector. The dimension of the vector depends on the specific embedding model used. """ + 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 provider token usage when available.""" merged_kwargs = self.model_kwargs.copy() merged_kwargs.update(kwargs) - result = self._get_connection().embed_with_usage(text, **merged_kwargs) - self._record_token_metrics(result.token_usage) - return result.embeddings + return self._get_connection().embed_with_usage(text, **merged_kwargs) - def _record_token_metrics(self, usage: EmbeddingTokenUsage | None) -> None: - """Record embedding token metrics under the current model metric group.""" - if usage is None: - return - metric_group = self.metric_group - if metric_group is None: + def record_token_metrics( + self, metric_group: MetricGroup | None, usage: EmbeddingTokenUsage | None + ) -> None: + """Record embedding token metrics under the request-scoped metric group.""" + if metric_group is None or usage is None: return model_group = metric_group.get_sub_group("model", self.model) diff --git a/python/flink_agents/api/embedding_models/tests/test_token_metrics.py b/python/flink_agents/api/embedding_models/tests/test_token_metrics.py index c71850c81..7a50db4e8 100644 --- a/python/flink_agents/api/embedding_models/tests/test_token_metrics.py +++ b/python/flink_agents/api/embedding_models/tests/test_token_metrics.py @@ -149,9 +149,10 @@ def get_resource(name: str, resource_type: ResourceType) -> Resource: def test_embedding_token_metrics_are_recorded_when_usage_is_reported() -> None: setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) metric_group = _MockMetricGroup() - setup.set_metric_group(metric_group) - assert setup.embed("hello") == [0.1, 0.2] + result = setup.embed_with_usage("hello") + setup.record_token_metrics(metric_group, result.token_usage) + assert result.embeddings == [0.1, 0.2] model_group = metric_group.get_sub_group("model", "mock-model") assert model_group.get_counter("promptTokens").get_count() == 7 @@ -161,22 +162,56 @@ def test_embedding_token_metrics_are_recorded_when_usage_is_reported() -> None: def test_embedding_token_metrics_are_noop_when_usage_is_absent() -> None: setup = _make_setup(FakeEmbeddingModelConnectionWithoutUsage(name="connection")) metric_group = _MockMetricGroup() - setup.set_metric_group(metric_group) - assert setup.embed("hello") == [0.1, 0.2] + result = setup.embed_with_usage("hello") + setup.record_token_metrics(metric_group, result.token_usage) + assert result.embeddings == [0.1, 0.2] model_group = metric_group.get_sub_group("model", "mock-model") assert "promptTokens" not in model_group._counters assert "totalTokens" not in model_group._counters +def test_embedding_token_metrics_are_noop_when_metric_group_is_absent() -> None: + setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) + bound_metric_group = _MockMetricGroup() + setup.set_metric_group(bound_metric_group) + + result = setup.embed_with_usage("hello") + setup.record_token_metrics(None, result.token_usage) + + model_group = bound_metric_group.get_sub_group("model", "mock-model") + assert model_group.get_counter("promptTokens").get_count() == 0 + assert model_group.get_counter("totalTokens").get_count() == 0 + + +def test_embedding_token_metrics_use_request_scoped_metric_group() -> None: + setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) + action_a_metric_group = _MockMetricGroup() + action_b_metric_group = _MockMetricGroup() + setup.set_metric_group(action_b_metric_group) + + result = setup.embed_with_usage("hello") + setup.record_token_metrics(action_a_metric_group, result.token_usage) + assert result.embeddings == [0.1, 0.2] + + action_a_model_group = action_a_metric_group.get_sub_group("model", "mock-model") + assert action_a_model_group.get_counter("promptTokens").get_count() == 7 + assert action_a_model_group.get_counter("totalTokens").get_count() == 9 + + action_b_model_group = action_b_metric_group.get_sub_group("model", "mock-model") + assert action_b_model_group.get_counter("promptTokens").get_count() == 0 + assert action_b_model_group.get_counter("totalTokens").get_count() == 0 + + def test_embedding_token_metrics_accumulate_across_requests() -> None: setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) metric_group = _MockMetricGroup() - setup.set_metric_group(metric_group) - setup.embed("hello") - setup.embed(["hello", "flink"]) + first = setup.embed_with_usage("hello") + setup.record_token_metrics(metric_group, first.token_usage) + second = setup.embed_with_usage(["hello", "flink"]) + setup.record_token_metrics(metric_group, second.token_usage) model_group = metric_group.get_sub_group("model", "mock-model") assert model_group.get_counter("promptTokens").get_count() == 14 @@ -186,13 +221,26 @@ def test_embedding_token_metrics_accumulate_across_requests() -> None: def test_embedding_token_metrics_do_not_leak_after_provider_failure() -> None: setup = _make_setup(ThrowThenReportUsageConnection(name="connection")) metric_group = _MockMetricGroup() - setup.set_metric_group(metric_group) with pytest.raises(RuntimeError, match="provider failure"): - setup.embed("first") + setup.embed_with_usage("first") - assert setup.embed("second") == [0.1, 0.2] + result = setup.embed_with_usage("second") + setup.record_token_metrics(metric_group, result.token_usage) + assert result.embeddings == [0.1, 0.2] model_group = metric_group.get_sub_group("model", "mock-model") assert model_group.get_counter("promptTokens").get_count() == 3 assert model_group.get_counter("totalTokens").get_count() == 4 + + +def test_embedding_without_explicit_metric_group_does_not_read_bound_group() -> None: + setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) + bound_metric_group = _MockMetricGroup() + setup.set_metric_group(bound_metric_group) + + assert setup.embed("hello") == [0.1, 0.2] + + model_group = bound_metric_group.get_sub_group("model", "mock-model") + assert model_group.get_counter("promptTokens").get_count() == 0 + assert model_group.get_counter("totalTokens").get_count() == 0 diff --git a/python/flink_agents/api/vector_stores/vector_store.py b/python/flink_agents/api/vector_stores/vector_store.py index 2cf04c49a..464ab024b 100644 --- a/python/flink_agents/api/vector_stores/vector_store.py +++ b/python/flink_agents/api/vector_stores/vector_store.py @@ -287,7 +287,12 @@ def query(self, query: VectorStoreQuery) -> VectorStoreQueryResult: VectorStoreQueryResult containing the retrieved documents """ # Generate embedding from the query text - query_embedding = self._get_embedding_model().embed(query.query_text) + embedding_model = self._get_embedding_model() + embedding_result = embedding_model.embed_with_usage(query.query_text) + embedding_model.record_token_metrics( + self.metric_group, embedding_result.token_usage + ) + query_embedding = embedding_result.embeddings # Merge setup kwargs with query-specific args merged_kwargs = self.store_kwargs.copy() @@ -344,9 +349,15 @@ def update( def _ensure_embeddings(self, documents: List[Document]) -> None: """Auto-embed any documents whose ``embedding`` field is ``None``.""" + embedding_model = self._get_embedding_model() + request_metric_group = self.metric_group for doc in documents: if doc.embedding is None: - doc.embedding = self._get_embedding_model().embed(doc.content) + result = embedding_model.embed_with_usage(doc.content) + embedding_model.record_token_metrics( + request_metric_group, result.token_usage + ) + doc.embedding = result.embeddings @abstractmethod def get( 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 76add0228..0f335c6fa 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 @@ -93,9 +93,10 @@ def get_resource(name: str, type: ResourceType) -> Resource: }.__getitem__ embedding_model.open() - embedding_model.set_metric_group(metric_group) - assert embedding_model.embed("Hello, Flink Agent!") == [0.1, 0.2, 0.3] + result = embedding_model.embed_with_usage("Hello, Flink Agent!") + embedding_model.record_token_metrics(metric_group, result.token_usage) + assert result.embeddings == [0.1, 0.2, 0.3] metric_group.get_sub_group.assert_called_once_with("model", test_model) prompt_counter.inc.assert_called_once_with(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 faa40bd73..efb0e79f3 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 @@ -203,9 +203,10 @@ def get_resource(name: str, type: ResourceType) -> Resource: }.__getitem__ embedding_model.open() - embedding_model.set_metric_group(metric_group) - assert embedding_model.embed("Test text") == mock_embedding + result = embedding_model.embed_with_usage("Test text") + embedding_model.record_token_metrics(metric_group, result.token_usage) + assert result.embeddings == mock_embedding metric_group.get_sub_group.assert_called_once_with("model", test_model) prompt_counter.inc.assert_called_once_with(6) From 03fbe7aa9617d4f33f1b7da27f05ee8d95293ba1 Mon Sep 17 00:00:00 2001 From: zlc <1633079383@qq.com> Date: Fri, 10 Jul 2026 17:52:43 +0800 Subject: [PATCH 3/5] =?UTF-8?q?refactor(embedding):=20=E6=94=B6=E7=AA=84?= =?UTF-8?q?=20token=20usage=20=E5=8F=98=E6=9B=B4=E5=88=B0=E7=BB=93?= =?UTF-8?q?=E6=9E=9C=E8=BF=94=E5=9B=9E=20API?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit # 影响范围 - Java/Python `BaseEmbeddingModelConnection` 与 `BaseEmbeddingModelSetup` 保留 `embedWithUsage`/`embed_with_usage` 的结果携带链路;普通 `embed` API 的返回类型不变。 - `BaseVectorStore.query`、自动补 embedding 以及 Python 对应实现恢复仅调用 `embed`,不再进入指标组或计数器。 - OpenAI、Tongyi、Bedrock 的 provider usage 提取及批量聚合能力保持不变;测试改为断言 usage 结果本身。 # 改动影响面 - 主链路:embedding provider -> `EmbeddingResult` -> 调用方显式消费 usage;本提交不再把 usage 写入 `MetricGroup`。 - 向量库扩展点:恢复 `BaseEmbeddingModelSetup.embed` 覆盖语义,并且仅在确有缺失 embedding 时延迟解析模型,避免 Mem0 无模型路径报错。 - 不受影响范围:#860/#861 的 resource metric-group 传播和 action-scoped metric 语义未改动;未引入新的 metric abstraction。 # 功能改进/开发/新增 - 类型: refactor - 触发条件: 普通资源 API 可并发调用,无法安全保证 `MetricGroup`/`SimpleCounter` 的线程模型;Python vector-store 还绕过了 `embed` 扩展点。 - 根因类别: 指标写入边界不明确、公共 API 扩展点回归、惰性资源解析边界遗漏。 - 行为变化: 有;PR 仅暴露 provider usage 返回 API,不再自动记录 embedding token metrics。 # 验证 - `mvn --batch-mode --no-transfer-progress -pl api -Dtest=BaseEmbeddingModelSetupEmbeddingResultTest test` PASS - `mvn --batch-mode --no-transfer-progress -pl integrations/embedding-models/bedrock -am -Dtest=BedrockEmbeddingModelTest -Dsurefire.failIfNoSpecifiedTests=false test` PASS - `uv run --python 3.12 --extra test pytest ...vector_stores... -q` PASS (49 passed, 4 skipped) - `./tools/lint.sh --check` PASS - `uv run --python 3.12 --extra lint ruff format --check ... && uv run --python 3.12 --extra lint ruff check ...` PASS - `git diff --check` PASS --- .../model/BaseEmbeddingModelSetup.java | 30 -- .../api/vectorstores/BaseVectorStore.java | 15 +- ...mbeddingModelSetupEmbeddingResultTest.java | 76 +++++ ...seEmbeddingModelSetupTokenMetricsTest.java | 313 ------------------ .../api/embedding_models/embedding_model.py | 12 - .../tests/test_embedding_result.py | 102 ++++++ .../tests/test_token_metrics.py | 246 -------------- .../api/vector_stores/vector_store.py | 15 +- .../tests/test_openai_embedding_model.py | 22 +- .../tests/test_tongyi_embedding_model.py | 20 +- 10 files changed, 191 insertions(+), 660 deletions(-) create mode 100644 api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupEmbeddingResultTest.java delete mode 100644 api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupTokenMetricsTest.java create mode 100644 python/flink_agents/api/embedding_models/tests/test_embedding_result.py delete mode 100644 python/flink_agents/api/embedding_models/tests/test_token_metrics.py 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 a3c8a061a..51d8b6ead 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 @@ -18,7 +18,6 @@ package org.apache.flink.agents.api.embedding.model; -import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.resource.Resource; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; @@ -108,15 +107,6 @@ public float[] embed(String text, Map parameters) { return embedWithUsage(text, parameters).getEmbeddings(); } - public float[] embed( - String text, - Map parameters, - @Nullable FlinkAgentsMetricGroup requestMetricGroup) { - EmbeddingResult result = embedWithUsage(text, parameters); - recordTokenMetrics(requestMetricGroup, result.getTokenUsage()); - return result.getEmbeddings(); - } - public EmbeddingResult embedWithUsage(String text) { return embedWithUsage(text, Collections.emptyMap()); } @@ -143,15 +133,6 @@ public List embed(List texts, Map parameters) { return embedWithUsage(texts, parameters).getEmbeddings(); } - public List embed( - List texts, - Map parameters, - @Nullable FlinkAgentsMetricGroup requestMetricGroup) { - EmbeddingResult> result = embedWithUsage(texts, parameters); - recordTokenMetrics(requestMetricGroup, result.getTokenUsage()); - return result.getEmbeddings(); - } - public EmbeddingResult> embedWithUsage(List texts) { return embedWithUsage(texts, Collections.emptyMap()); } @@ -164,15 +145,4 @@ public EmbeddingResult> embedWithUsage( return currentConnection.embedWithUsage(texts, params); } - public void recordTokenMetrics( - @Nullable FlinkAgentsMetricGroup requestMetricGroup, - @Nullable EmbeddingTokenUsage usage) { - if (requestMetricGroup == null || usage == null) { - return; - } - - FlinkAgentsMetricGroup modelGroup = requestMetricGroup.getSubGroup("model", model); - modelGroup.getCounter("promptTokens").inc(usage.getPromptTokens()); - modelGroup.getCounter("totalTokens").inc(usage.getTotalTokens()); - } } diff --git a/api/src/main/java/org/apache/flink/agents/api/vectorstores/BaseVectorStore.java b/api/src/main/java/org/apache/flink/agents/api/vectorstores/BaseVectorStore.java index aad119f3f..d0fc4ce72 100644 --- a/api/src/main/java/org/apache/flink/agents/api/vectorstores/BaseVectorStore.java +++ b/api/src/main/java/org/apache/flink/agents/api/vectorstores/BaseVectorStore.java @@ -19,7 +19,6 @@ package org.apache.flink.agents.api.vectorstores; import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelSetup; -import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.resource.Resource; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; @@ -28,7 +27,6 @@ import javax.annotation.Nullable; import java.io.IOException; -import java.util.Collections; import java.util.List; import java.util.Map; @@ -178,10 +176,7 @@ public void update( * @return VectorStoreQueryResult containing the retrieved documents */ public VectorStoreQueryResult query(VectorStoreQuery query) { - final FlinkAgentsMetricGroup requestMetricGroup = getMetricGroup(); - final float[] queryEmbedding = - getEmbeddingModel() - .embed(query.getQueryText(), Collections.emptyMap(), requestMetricGroup); + final float[] queryEmbedding = getEmbeddingModel().embed(query.getQueryText()); final Map storeKwargs = this.getStoreKwargs(); storeKwargs.putAll(query.getExtraArgs()); @@ -337,15 +332,9 @@ public abstract void updateEmbedding( /** Auto-embed any documents whose {@code embedding} field is {@code null}. */ protected void ensureEmbeddings(List documents) { - final FlinkAgentsMetricGroup requestMetricGroup = getMetricGroup(); for (Document doc : documents) { if (doc.getEmbedding() == null) { - doc.setEmbedding( - getEmbeddingModel() - .embed( - doc.getContent(), - Collections.emptyMap(), - requestMetricGroup)); + doc.setEmbedding(getEmbeddingModel().embed(doc.getContent())); } } } diff --git a/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupEmbeddingResultTest.java b/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupEmbeddingResultTest.java new file mode 100644 index 000000000..a0d93cc65 --- /dev/null +++ b/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupEmbeddingResultTest.java @@ -0,0 +1,76 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.flink.agents.api.embedding.model; + +import org.apache.flink.agents.api.resource.ResourceContext; +import org.apache.flink.agents.api.resource.ResourceDescriptor; +import org.junit.jupiter.api.Test; + +import java.util.Collections; +import java.util.Map; + +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.mockito.Mockito.mock; + +/** Test cases for embedding results returned by model setups. */ +class BaseEmbeddingModelSetupEmbeddingResultTest { + + @Test + void testEmbedWithUsageDelegatesProviderUsage() { + BaseEmbeddingModelSetup setup = + new BaseEmbeddingModelSetup( + new ResourceDescriptor( + "test", Map.of("connection", "connection", "model", "model")), + mock(ResourceContext.class)) { + @Override + public Map getParameters() { + return Collections.emptyMap(); + } + }; + setup.connection = + new BaseEmbeddingModelConnection( + new ResourceDescriptor("connection", Collections.emptyMap()), + mock(ResourceContext.class)) { + @Override + public float[] embed(String text, Map parameters) { + return new float[] {0.1f, 0.2f}; + } + + @Override + public java.util.List embed( + java.util.List texts, Map parameters) { + throw new UnsupportedOperationException(); + } + + @Override + public EmbeddingResult embedWithUsage( + String text, Map parameters) { + return new EmbeddingResult<>( + embed(text, parameters), new EmbeddingTokenUsage(7L, 9L)); + } + }; + + EmbeddingResult result = setup.embedWithUsage("hello"); + + assertArrayEquals(new float[] {0.1f, 0.2f}, result.getEmbeddings()); + assertEquals(7L, result.getTokenUsage().getPromptTokens()); + assertEquals(9L, result.getTokenUsage().getTotalTokens()); + } +} diff --git a/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupTokenMetricsTest.java b/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupTokenMetricsTest.java deleted file mode 100644 index 3e4e6ad9e..000000000 --- a/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupTokenMetricsTest.java +++ /dev/null @@ -1,313 +0,0 @@ -/* - * 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.metrics.FlinkAgentsMetricGroup; -import org.apache.flink.agents.api.metrics.UpdatableGauge; -import org.apache.flink.agents.api.resource.ResourceContext; -import org.apache.flink.agents.api.resource.ResourceDescriptor; -import org.apache.flink.metrics.Counter; -import org.apache.flink.metrics.Histogram; -import org.apache.flink.metrics.Meter; -import org.apache.flink.metrics.SimpleCounter; -import org.junit.jupiter.api.Test; - -import java.util.ArrayList; -import java.util.Collections; -import java.util.HashMap; -import java.util.List; -import java.util.Map; - -import static org.junit.jupiter.api.Assertions.assertArrayEquals; -import static org.junit.jupiter.api.Assertions.assertEquals; -import static org.junit.jupiter.api.Assertions.assertThrows; -import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.verifyNoInteractions; - -/** Test cases for embedding model token metrics. */ -class BaseEmbeddingModelSetupTokenMetricsTest { - - private static class TestEmbeddingModelSetup extends BaseEmbeddingModelSetup { - - TestEmbeddingModelSetup(BaseEmbeddingModelConnection connection) { - super( - new ResourceDescriptor( - TestEmbeddingModelSetup.class.getName(), - Map.of("connection", "mock-connection", "model", "mock-model")), - mock(ResourceContext.class)); - this.connection = connection; - } - - @Override - public Map getParameters() { - return new HashMap<>(); - } - } - - private static class TestEmbeddingModelConnection extends BaseEmbeddingModelConnection { - - TestEmbeddingModelConnection() { - super( - new ResourceDescriptor( - TestEmbeddingModelConnection.class.getName(), Collections.emptyMap()), - mock(ResourceContext.class)); - } - - @Override - public float[] embed(String text, Map parameters) { - return new float[] {0.1f, 0.2f}; - } - - @Override - public EmbeddingResult embedWithUsage( - String text, Map parameters) { - return new EmbeddingResult<>(embed(text, parameters), new EmbeddingTokenUsage(7L, 9L)); - } - - @Override - public List embed(List texts, Map parameters) { - List embeddings = new ArrayList<>(); - for (String ignored : texts) { - embeddings.add(new float[] {0.1f, 0.2f}); - } - return embeddings; - } - - @Override - public EmbeddingResult> embedWithUsage( - List texts, Map parameters) { - return new EmbeddingResult<>( - embed(texts, parameters), new EmbeddingTokenUsage(11L, 13L)); - } - } - - private static class TestEmbeddingModelConnectionWithoutUsage - extends BaseEmbeddingModelConnection { - - TestEmbeddingModelConnectionWithoutUsage() { - super( - new ResourceDescriptor( - TestEmbeddingModelConnectionWithoutUsage.class.getName(), - Collections.emptyMap()), - mock(ResourceContext.class)); - } - - @Override - public float[] embed(String text, Map parameters) { - return new float[] {0.1f, 0.2f}; - } - - @Override - public List embed(List texts, Map parameters) { - List embeddings = new ArrayList<>(); - for (String ignored : texts) { - embeddings.add(new float[] {0.1f, 0.2f}); - } - return embeddings; - } - } - - private static class ThrowThenReportUsageConnection extends BaseEmbeddingModelConnection { - private int calls; - - ThrowThenReportUsageConnection() { - super( - new ResourceDescriptor( - ThrowThenReportUsageConnection.class.getName(), Collections.emptyMap()), - mock(ResourceContext.class)); - } - - @Override - public float[] embed(String text, Map parameters) { - return new float[] {0.1f, 0.2f}; - } - - @Override - public EmbeddingResult embedWithUsage( - String text, Map parameters) { - calls++; - if (calls == 1) { - throw new RuntimeException("provider failure"); - } - return new EmbeddingResult<>(embed(text, parameters), new EmbeddingTokenUsage(3L, 4L)); - } - - @Override - public List embed(List texts, Map parameters) { - List embeddings = new ArrayList<>(); - for (String ignored : texts) { - embeddings.add(new float[] {0.1f, 0.2f}); - } - return embeddings; - } - } - - @Test - void testEmbeddingTokenMetricsAreRecordedWhenUsageIsReported() { - TestEmbeddingModelSetup setup = - new TestEmbeddingModelSetup(new TestEmbeddingModelConnection()); - TestMetricGroup metricGroup = new TestMetricGroup(); - - assertArrayEquals( - new float[] {0.1f, 0.2f}, - setup.embed("hello", Collections.emptyMap(), metricGroup)); - - TestMetricGroup modelGroup = - (TestMetricGroup) metricGroup.getSubGroup("model", "mock-model"); - assertEquals(7L, modelGroup.counters.get("promptTokens").getCount()); - assertEquals(9L, modelGroup.counters.get("totalTokens").getCount()); - } - - @Test - void testEmbeddingTokenMetricsAreNoopWhenUsageIsAbsent() { - TestEmbeddingModelSetup setup = - new TestEmbeddingModelSetup(new TestEmbeddingModelConnectionWithoutUsage()); - FlinkAgentsMetricGroup metricGroup = mock(FlinkAgentsMetricGroup.class); - - setup.embed("hello", Collections.emptyMap(), metricGroup); - - verifyNoInteractions(metricGroup); - } - - @Test - void testEmbeddingTokenMetricsAreNoopWhenMetricGroupIsAbsent() { - TestEmbeddingModelSetup setup = - new TestEmbeddingModelSetup(new TestEmbeddingModelConnection()); - FlinkAgentsMetricGroup boundMetricGroup = mock(FlinkAgentsMetricGroup.class); - setup.setMetricGroup(boundMetricGroup); - - setup.embed("hello", Collections.emptyMap(), null); - - verifyNoInteractions(boundMetricGroup); - } - - @Test - void testEmbeddingTokenMetricsUseRequestScopedMetricGroup() { - TestEmbeddingModelSetup setup = - new TestEmbeddingModelSetup(new TestEmbeddingModelConnection()); - TestMetricGroup actionA = new TestMetricGroup(); - TestMetricGroup actionB = new TestMetricGroup(); - - setup.setMetricGroup(actionB); - - assertArrayEquals( - new float[] {0.1f, 0.2f}, setup.embed("hello", Collections.emptyMap(), actionA)); - - TestMetricGroup actionAModelGroup = - (TestMetricGroup) actionA.getSubGroup("model", "mock-model"); - assertEquals(7L, actionAModelGroup.counters.get("promptTokens").getCount()); - assertEquals(9L, actionAModelGroup.counters.get("totalTokens").getCount()); - - TestMetricGroup actionBModelGroup = - (TestMetricGroup) actionB.getSubGroup("model", "mock-model"); - assertEquals(0L, actionBModelGroup.getCounter("promptTokens").getCount()); - assertEquals(0L, actionBModelGroup.getCounter("totalTokens").getCount()); - } - - @Test - void testEmbeddingTokenMetricsAccumulateAcrossRequests() { - TestEmbeddingModelSetup setup = - new TestEmbeddingModelSetup(new TestEmbeddingModelConnection()); - TestMetricGroup metricGroup = new TestMetricGroup(); - - setup.embed("hello", Collections.emptyMap(), metricGroup); - setup.embed(List.of("hello", "flink"), Collections.emptyMap(), metricGroup); - - TestMetricGroup modelGroup = - (TestMetricGroup) metricGroup.getSubGroup("model", "mock-model"); - assertEquals(18L, modelGroup.counters.get("promptTokens").getCount()); - assertEquals(22L, modelGroup.counters.get("totalTokens").getCount()); - } - - @Test - void testEmbeddingTokenMetricsDoNotLeakAfterProviderFailure() { - TestEmbeddingModelSetup setup = - new TestEmbeddingModelSetup(new ThrowThenReportUsageConnection()); - TestMetricGroup metricGroup = new TestMetricGroup(); - - assertThrows( - RuntimeException.class, - () -> setup.embed("first", Collections.emptyMap(), metricGroup)); - assertArrayEquals( - new float[] {0.1f, 0.2f}, - setup.embed("second", Collections.emptyMap(), metricGroup)); - - TestMetricGroup modelGroup = - (TestMetricGroup) metricGroup.getSubGroup("model", "mock-model"); - assertEquals(3L, modelGroup.counters.get("promptTokens").getCount()); - assertEquals(4L, modelGroup.counters.get("totalTokens").getCount()); - } - - @Test - void testEmbeddingWithoutExplicitMetricGroupDoesNotReadBoundMetricGroup() { - TestEmbeddingModelSetup setup = - new TestEmbeddingModelSetup(new TestEmbeddingModelConnection()); - FlinkAgentsMetricGroup boundMetricGroup = mock(FlinkAgentsMetricGroup.class); - setup.setMetricGroup(boundMetricGroup); - - assertArrayEquals(new float[] {0.1f, 0.2f}, setup.embed("hello")); - - verifyNoInteractions(boundMetricGroup); - } - - private static class TestMetricGroup implements FlinkAgentsMetricGroup { - final Map subGroups = new HashMap<>(); - final Map counters = new HashMap<>(); - - @Override - public FlinkAgentsMetricGroup getSubGroup(String name) { - return subGroups.computeIfAbsent(name, ignored -> new TestMetricGroup()); - } - - @Override - public FlinkAgentsMetricGroup getSubGroup(String key, String value) { - return subGroups.computeIfAbsent(key + "=" + value, ignored -> new TestMetricGroup()); - } - - @Override - public UpdatableGauge getGauge(String name) { - return null; - } - - @Override - public Counter getCounter(String name) { - return counters.computeIfAbsent(name, ignored -> new SimpleCounter()); - } - - @Override - public Meter getMeter(String name) { - return null; - } - - @Override - public Meter getMeter(String name, Counter counter) { - return null; - } - - @Override - public Histogram getHistogram(String name) { - return null; - } - - @Override - public Histogram getHistogram(String name, int windowSize) { - return null; - } - } -} diff --git a/python/flink_agents/api/embedding_models/embedding_model.py b/python/flink_agents/api/embedding_models/embedding_model.py index d9943c270..ce96c0095 100644 --- a/python/flink_agents/api/embedding_models/embedding_model.py +++ b/python/flink_agents/api/embedding_models/embedding_model.py @@ -22,7 +22,6 @@ from pydantic import Field from typing_extensions import override -from flink_agents.api.metric_group import MetricGroup from flink_agents.api.resource import Resource, ResourceType EmbeddingValue = TypeVar("EmbeddingValue", list[float], list[list[float]]) @@ -155,14 +154,3 @@ def embed_with_usage( merged_kwargs = self.model_kwargs.copy() merged_kwargs.update(kwargs) return self._get_connection().embed_with_usage(text, **merged_kwargs) - - def record_token_metrics( - self, metric_group: MetricGroup | None, usage: EmbeddingTokenUsage | None - ) -> None: - """Record embedding token metrics under the request-scoped metric group.""" - if metric_group is None or usage is None: - return - - model_group = metric_group.get_sub_group("model", self.model) - model_group.get_counter("promptTokens").inc(usage.prompt_tokens) - model_group.get_counter("totalTokens").inc(usage.total_tokens) diff --git a/python/flink_agents/api/embedding_models/tests/test_embedding_result.py b/python/flink_agents/api/embedding_models/tests/test_embedding_result.py new file mode 100644 index 000000000..88534dfd3 --- /dev/null +++ b/python/flink_agents/api/embedding_models/tests/test_embedding_result.py @@ -0,0 +1,102 @@ +################################################################################ +# 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 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] diff --git a/python/flink_agents/api/embedding_models/tests/test_token_metrics.py b/python/flink_agents/api/embedding_models/tests/test_token_metrics.py deleted file mode 100644 index 7a50db4e8..000000000 --- a/python/flink_agents/api/embedding_models/tests/test_token_metrics.py +++ /dev/null @@ -1,246 +0,0 @@ -################################################################################ -# 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 - -import pytest - -from flink_agents.api.embedding_models.embedding_model import ( - BaseEmbeddingModelConnection, - BaseEmbeddingModelSetup, - EmbeddingResult, - EmbeddingTokenUsage, -) -from flink_agents.api.metric_group import Counter, MetricGroup -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 ThrowThenReportUsageConnection(BaseEmbeddingModelConnection): - def __init__(self, **kwargs: Any) -> None: - super().__init__(**kwargs) - self._calls = 0 - - 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]]]: - self._calls += 1 - if self._calls == 1: - msg = "provider failure" - raise RuntimeError(msg) - return EmbeddingResult( - embeddings=self.embed(text, **kwargs), - token_usage=EmbeddingTokenUsage(prompt_tokens=3, total_tokens=4), - ) - - -class FakeEmbeddingModelSetup(BaseEmbeddingModelSetup): - @property - def model_kwargs(self) -> Dict[str, Any]: - return {} - - -class _MockCounter(Counter): - def __init__(self) -> None: - self._count = 0 - - def inc(self, n: int = 1) -> None: - self._count += n - - def dec(self, n: int = 1) -> None: - self._count -= n - - def get_count(self) -> int: - return self._count - - -class _MockMetricGroup(MetricGroup): - def __init__(self) -> None: - self._sub_groups: dict[str, _MockMetricGroup] = {} - self._counters: dict[str, _MockCounter] = {} - - def get_sub_group(self, name: str, value: str | None = None) -> "_MockMetricGroup": - key = f"{name}={value}" if value is not None else name - if key not in self._sub_groups: - self._sub_groups[key] = _MockMetricGroup() - return self._sub_groups[key] - - def get_counter(self, name: str) -> _MockCounter: - if name not in self._counters: - self._counters[name] = _MockCounter() - return self._counters[name] - - def get_meter(self, name: str) -> Any: - return MagicMock() - - def get_gauge(self, name: str) -> Any: - return MagicMock() - - def get_histogram(self, name: str, window_size: int = 100) -> Any: - return MagicMock() - - -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_token_metrics_are_recorded_when_usage_is_reported() -> None: - setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) - metric_group = _MockMetricGroup() - - result = setup.embed_with_usage("hello") - setup.record_token_metrics(metric_group, result.token_usage) - assert result.embeddings == [0.1, 0.2] - - model_group = metric_group.get_sub_group("model", "mock-model") - assert model_group.get_counter("promptTokens").get_count() == 7 - assert model_group.get_counter("totalTokens").get_count() == 9 - - -def test_embedding_token_metrics_are_noop_when_usage_is_absent() -> None: - setup = _make_setup(FakeEmbeddingModelConnectionWithoutUsage(name="connection")) - metric_group = _MockMetricGroup() - - result = setup.embed_with_usage("hello") - setup.record_token_metrics(metric_group, result.token_usage) - assert result.embeddings == [0.1, 0.2] - - model_group = metric_group.get_sub_group("model", "mock-model") - assert "promptTokens" not in model_group._counters - assert "totalTokens" not in model_group._counters - - -def test_embedding_token_metrics_are_noop_when_metric_group_is_absent() -> None: - setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) - bound_metric_group = _MockMetricGroup() - setup.set_metric_group(bound_metric_group) - - result = setup.embed_with_usage("hello") - setup.record_token_metrics(None, result.token_usage) - - model_group = bound_metric_group.get_sub_group("model", "mock-model") - assert model_group.get_counter("promptTokens").get_count() == 0 - assert model_group.get_counter("totalTokens").get_count() == 0 - - -def test_embedding_token_metrics_use_request_scoped_metric_group() -> None: - setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) - action_a_metric_group = _MockMetricGroup() - action_b_metric_group = _MockMetricGroup() - setup.set_metric_group(action_b_metric_group) - - result = setup.embed_with_usage("hello") - setup.record_token_metrics(action_a_metric_group, result.token_usage) - assert result.embeddings == [0.1, 0.2] - - action_a_model_group = action_a_metric_group.get_sub_group("model", "mock-model") - assert action_a_model_group.get_counter("promptTokens").get_count() == 7 - assert action_a_model_group.get_counter("totalTokens").get_count() == 9 - - action_b_model_group = action_b_metric_group.get_sub_group("model", "mock-model") - assert action_b_model_group.get_counter("promptTokens").get_count() == 0 - assert action_b_model_group.get_counter("totalTokens").get_count() == 0 - - -def test_embedding_token_metrics_accumulate_across_requests() -> None: - setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) - metric_group = _MockMetricGroup() - - first = setup.embed_with_usage("hello") - setup.record_token_metrics(metric_group, first.token_usage) - second = setup.embed_with_usage(["hello", "flink"]) - setup.record_token_metrics(metric_group, second.token_usage) - - model_group = metric_group.get_sub_group("model", "mock-model") - assert model_group.get_counter("promptTokens").get_count() == 14 - assert model_group.get_counter("totalTokens").get_count() == 18 - - -def test_embedding_token_metrics_do_not_leak_after_provider_failure() -> None: - setup = _make_setup(ThrowThenReportUsageConnection(name="connection")) - metric_group = _MockMetricGroup() - - with pytest.raises(RuntimeError, match="provider failure"): - setup.embed_with_usage("first") - - result = setup.embed_with_usage("second") - setup.record_token_metrics(metric_group, result.token_usage) - assert result.embeddings == [0.1, 0.2] - - model_group = metric_group.get_sub_group("model", "mock-model") - assert model_group.get_counter("promptTokens").get_count() == 3 - assert model_group.get_counter("totalTokens").get_count() == 4 - - -def test_embedding_without_explicit_metric_group_does_not_read_bound_group() -> None: - setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) - bound_metric_group = _MockMetricGroup() - setup.set_metric_group(bound_metric_group) - - assert setup.embed("hello") == [0.1, 0.2] - - model_group = bound_metric_group.get_sub_group("model", "mock-model") - assert model_group.get_counter("promptTokens").get_count() == 0 - assert model_group.get_counter("totalTokens").get_count() == 0 diff --git a/python/flink_agents/api/vector_stores/vector_store.py b/python/flink_agents/api/vector_stores/vector_store.py index 464ab024b..2cf04c49a 100644 --- a/python/flink_agents/api/vector_stores/vector_store.py +++ b/python/flink_agents/api/vector_stores/vector_store.py @@ -287,12 +287,7 @@ def query(self, query: VectorStoreQuery) -> VectorStoreQueryResult: VectorStoreQueryResult containing the retrieved documents """ # Generate embedding from the query text - embedding_model = self._get_embedding_model() - embedding_result = embedding_model.embed_with_usage(query.query_text) - embedding_model.record_token_metrics( - self.metric_group, embedding_result.token_usage - ) - query_embedding = embedding_result.embeddings + query_embedding = self._get_embedding_model().embed(query.query_text) # Merge setup kwargs with query-specific args merged_kwargs = self.store_kwargs.copy() @@ -349,15 +344,9 @@ def update( def _ensure_embeddings(self, documents: List[Document]) -> None: """Auto-embed any documents whose ``embedding`` field is ``None``.""" - embedding_model = self._get_embedding_model() - request_metric_group = self.metric_group for doc in documents: if doc.embedding is None: - result = embedding_model.embed_with_usage(doc.content) - embedding_model.record_token_metrics( - request_metric_group, result.token_usage - ) - doc.embedding = result.embeddings + doc.embedding = self._get_embedding_model().embed(doc.content) @abstractmethod def get( 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 0f335c6fa..9d63c2494 100644 --- a/python/flink_agents/integrations/embedding_models/tests/test_openai_embedding_model.py +++ b/python/flink_agents/integrations/embedding_models/tests/test_openai_embedding_model.py @@ -60,8 +60,8 @@ def get_resource(name: str, type: ResourceType) -> Resource: assert all(isinstance(x, float) for x in response) # -def test_openai_embedding_model_records_token_metrics() -> None: - """Test OpenAI embedding usage is recorded as model token metrics.""" +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( @@ -82,22 +82,10 @@ def get_resource(name: str, type: ResourceType) -> Resource: embedding_model = OpenAIEmbeddingModelSetup( name="openai", model=test_model, connection="openai", resource_context=mock_ctx ) - metric_group = MagicMock() - model_group = MagicMock() - prompt_counter = MagicMock() - total_counter = MagicMock() - metric_group.get_sub_group.return_value = model_group - model_group.get_counter.side_effect = { - "promptTokens": prompt_counter, - "totalTokens": total_counter, - }.__getitem__ - embedding_model.open() result = embedding_model.embed_with_usage("Hello, Flink Agent!") - embedding_model.record_token_metrics(metric_group, result.token_usage) assert result.embeddings == [0.1, 0.2, 0.3] - - metric_group.get_sub_group.assert_called_once_with("model", test_model) - prompt_counter.inc.assert_called_once_with(5) - total_counter.inc.assert_called_once_with(5) + 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 efb0e79f3..98508a9b2 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,7 +155,7 @@ def get_resource(name: str, type: ResourceType) -> Resource: assert len(response) == 5 -def test_tongyi_embedding_records_token_metrics( +def test_tongyi_embedding_returns_token_usage( monkeypatch: pytest.MonkeyPatch, ) -> None: """Test DashScope embedding usage is recorded as model token metrics.""" @@ -192,25 +192,13 @@ def get_resource(name: str, type: ResourceType) -> Resource: connection="tongyi", resource_context=_make_ctx(get_resource), ) - metric_group = MagicMock() - model_group = MagicMock() - prompt_counter = MagicMock() - total_counter = MagicMock() - metric_group.get_sub_group.return_value = model_group - model_group.get_counter.side_effect = { - "promptTokens": prompt_counter, - "totalTokens": total_counter, - }.__getitem__ - embedding_model.open() result = embedding_model.embed_with_usage("Test text") - embedding_model.record_token_metrics(metric_group, result.token_usage) assert result.embeddings == mock_embedding - - metric_group.get_sub_group.assert_called_once_with("model", test_model) - prompt_counter.inc.assert_called_once_with(6) - total_counter.inc.assert_called_once_with(6) + 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: From 2fa0ae04d6502a711049a9fd70bcb48ba0a61d7d Mon Sep 17 00:00:00 2001 From: zlc <1633079383@qq.com> Date: Mon, 13 Jul 2026 14:30:20 +0800 Subject: [PATCH 4/5] =?UTF-8?q?fix(embedding):=20=E4=BF=AE=E5=A4=8D=20Base?= =?UTF-8?q?EmbeddingModelSetup=20=E6=A0=BC=E5=BC=8F=E6=A3=80=E6=9F=A5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit # 影响范围 - 直接修改 api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetup.java。 - 影响 BaseEmbeddingModelSetup.embedWithUsage(List, Map) 末尾格式区域;不改变方法签名、返回值或调用链。 # 改动影响面 - 主链路:embedding setup -> BaseEmbeddingModelConnection.embedWithUsage(...) 保持不变,仅移除 Spotless 标记的类尾多余空行。 - 数据/协议变化:无字段、API、指标、事件或响应结构变化。 - 不受影响范围:EmbeddingResult、provider usage 提取、vector-store embed 调用路径、#860/#861 metric group 传播语义均不变。 # 功能改进/开发/新增 - 类型: fix - 触发条件: PR #870 的 Code Style Check 在 BaseEmbeddingModelSetup.java 上报 Spotless format violation。 - 根因类别: 格式边界遗漏。 - 行为变化: 无;仅修复 CI 格式检查。 # 验证 - ./tools/lint.sh -c PASS - mvn --batch-mode --no-transfer-progress -pl api spotless:check PASS(当前配置显示 Spotless check skipped,但 Maven 目标成功) --- .../agents/api/embedding/model/BaseEmbeddingModelSetup.java | 1 - 1 file changed, 1 deletion(-) 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 51d8b6ead..5d829dff2 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 @@ -144,5 +144,4 @@ public EmbeddingResult> embedWithUsage( BaseEmbeddingModelConnection currentConnection = getConnection(); return currentConnection.embedWithUsage(texts, params); } - } From 0dc947f65e31a0e25ec615649fc8b4877408fbfe Mon Sep 17 00:00:00 2001 From: zlc <1633079383@qq.com> Date: Mon, 20 Jul 2026 19:15:10 -0700 Subject: [PATCH 5/5] =?UTF-8?q?fix(embedding):=20=E4=BF=AE=E5=A4=8D?= =?UTF-8?q?=E7=94=A8=E9=87=8F=20API=20=E8=B7=A8=E8=AF=AD=E8=A8=80=E5=A5=91?= =?UTF-8?q?=E7=BA=A6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit # 影响范围 - 入口:Java/Python embedding setup 与 connection 的 embed/embedWithUsage 双 API。 - 调用链:setup -> connection -> provider,以及 Java<->Python resource wrapper 和 resource-cross-language E2E agent。 - 字段/协议:EmbeddingResult.embeddings、token usage 的 prompt_tokens/total_tokens;Tongyi 支持原生 usage.total_tokens 响应。 # 改动影响面 - 普通 embed(...) 恢复直调同名 connection API,保留既有 provider/subclass override 分派;usage API 保持独立路径。 - Java 调 Python 将 dataclass 显式平铺为 Pemja 安全的 list/map/primitive;Python 调 Java 显式调用 embedWithUsage(...) 并还原结果。 - 不受影响范围:现有 embed(...) 返回类型、无 usage provider 的 no-op 语义、向量存储的 legacy embed 调用均不变。 # 功能改进/开发/新增 - 新增双向 wrapper、单/批 usage 结果转换与 Java/Python 回归测试。 - 现有跨语言 embedding agent 改为覆盖 usage API;Tongyi total_tokens 作为 prompt_tokens 回退。 # 验证 - JDK17 Maven API 定向测试:29 tests PASS。 - Python 3.12 定向 pytest:13 passed, 2 skipped。 - Ruff format/check、Spotless check、resource-cross-language test-compile、git diff --check PASS。 - Ollama 不可用,外部服务 E2E 未在本机执行。 --- .../model/BaseEmbeddingModelSetup.java | 8 ++- .../embedding/model/EmbeddingModelUtils.java | 64 +++++++++++++++++ .../PythonEmbeddingModelConnection.java | 29 ++++++++ .../python/PythonEmbeddingModelSetup.java | 29 ++++++++ ...mbeddingModelSetupEmbeddingResultTest.java | 38 ++++++++++ .../PythonEmbeddingModelConnectionTest.java | 37 ++++++++++ .../python/PythonEmbeddingModelSetupTest.java | 36 ++++++++++ .../test/EmbeddingCrossLanguageAgent.java | 8 ++- .../api/embedding_models/embedding_model.py | 4 +- .../tests/test_embedding_result.py | 25 +++++++ .../embedding_model_cross_language_agent.py | 6 +- .../tests/test_tongyi_embedding_model.py | 2 +- .../tongyi_embedding_model.py | 2 + .../runtime/java/java_embedding_model.py | 46 ++++++++++++ .../flink_agents/runtime/python_java_utils.py | 34 +++++++-- .../tests/test_java_embedding_model.py | 70 +++++++++++++++++++ .../runtime/tests/test_python_java_utils.py | 34 ++++++++- 17 files changed, 457 insertions(+), 15 deletions(-) create mode 100644 python/flink_agents/runtime/tests/test_java_embedding_model.py 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 5d829dff2..605c189d5 100644 --- a/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetup.java +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetup.java @@ -104,7 +104,9 @@ public float[] embed(String text) { } public float[] embed(String text, Map parameters) { - return embedWithUsage(text, parameters).getEmbeddings(); + Map params = this.getParameters(); + params.putAll(parameters); + return getConnection().embed(text, params); } public EmbeddingResult embedWithUsage(String text) { @@ -130,7 +132,9 @@ public List embed(List texts) { } public List embed(List texts, Map parameters) { - return embedWithUsage(texts, parameters).getEmbeddings(); + Map params = this.getParameters(); + params.putAll(parameters); + return getConnection().embed(texts, params); } public EmbeddingResult> embedWithUsage(List texts) { diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingModelUtils.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingModelUtils.java index c1c71e589..08128eb45 100644 --- a/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingModelUtils.java +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/EmbeddingModelUtils.java @@ -17,7 +17,9 @@ */ package org.apache.flink.agents.api.embedding.model; +import java.util.ArrayList; import java.util.List; +import java.util.Map; public class EmbeddingModelUtils { public static float[] toFloatArray(List list) { @@ -34,4 +36,66 @@ public static float[] toFloatArray(List list) { } return array; } + + public static EmbeddingResult toSingleEmbeddingResult(Object result) { + Map values = toResultMap(result); + return new EmbeddingResult<>( + toFloatArray(toEmbeddingList(values.get("embeddings"))), + toTokenUsage(values.get("token_usage"))); + } + + public static EmbeddingResult> toBatchEmbeddingResult(Object result) { + Map values = toResultMap(result); + List rawEmbeddings = toEmbeddingList(values.get("embeddings")); + List embeddings = new ArrayList<>(); + for (Object embedding : rawEmbeddings) { + embeddings.add(toFloatArray(toEmbeddingList(embedding))); + } + return new EmbeddingResult<>(embeddings, toTokenUsage(values.get("token_usage"))); + } + + private static Map toResultMap(Object result) { + if (result instanceof Map) { + return (Map) result; + } + throw new IllegalArgumentException( + "Expected Map from Python embed_with_usage method, but got: " + + (result == null ? "null" : result.getClass().getName())); + } + + private static List toEmbeddingList(Object embeddings) { + if (embeddings instanceof List) { + return (List) embeddings; + } + throw new IllegalArgumentException( + "Expected List value in Python embedding result, but got: " + + (embeddings == null ? "null" : embeddings.getClass().getName())); + } + + private static EmbeddingTokenUsage toTokenUsage(Object tokenUsage) { + if (tokenUsage == null) { + return null; + } + if (!(tokenUsage instanceof Map)) { + throw new IllegalArgumentException( + "Expected Map token_usage in Python embedding result, but got: " + + tokenUsage.getClass().getName()); + } + + Map usage = (Map) tokenUsage; + return new EmbeddingTokenUsage( + toLong(usage.get("prompt_tokens"), "prompt_tokens"), + toLong(usage.get("total_tokens"), "total_tokens")); + } + + private static long toLong(Object value, String fieldName) { + if (value instanceof Number) { + return ((Number) value).longValue(); + } + throw new IllegalArgumentException( + "Expected numeric " + + fieldName + + " in Python embedding token usage, but got: " + + (value == null ? "null" : value.getClass().getName())); + } } diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnection.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnection.java index 974e362a1..c9b769628 100644 --- a/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnection.java +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnection.java @@ -19,6 +19,7 @@ import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelConnection; import org.apache.flink.agents.api.embedding.model.EmbeddingModelUtils; +import org.apache.flink.agents.api.embedding.model.EmbeddingResult; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.apache.flink.agents.api.resource.python.PythonResourceAdapter; @@ -41,6 +42,9 @@ public class PythonEmbeddingModelConnection extends BaseEmbeddingModelConnection implements PythonResourceWrapper { + private static final String CALL_EMBED_WITH_USAGE = + "python_java_utils.call_embedding_with_usage"; + private final PyObject embeddingModel; private final PythonResourceAdapter adapter; @@ -118,6 +122,31 @@ public List embed(List texts, Map parameters) { + (results == null ? "null" : results.getClass().getName())); } + @Override + public EmbeddingResult embedWithUsage(String text, Map parameters) { + checkState( + embeddingModel != null, + "EmbeddingModelSetup is not initialized. Cannot perform embed operation."); + + Map kwargs = new HashMap<>(parameters); + kwargs.put("text", text); + Object result = adapter.invoke(CALL_EMBED_WITH_USAGE, embeddingModel, kwargs); + return EmbeddingModelUtils.toSingleEmbeddingResult(result); + } + + @Override + public EmbeddingResult> embedWithUsage( + List texts, Map parameters) { + checkState( + embeddingModel != null, + "EmbeddingModelSetup is not initialized. Cannot perform embed operation."); + + Map kwargs = new HashMap<>(parameters); + kwargs.put("text", texts); + Object result = adapter.invoke(CALL_EMBED_WITH_USAGE, embeddingModel, kwargs); + return EmbeddingModelUtils.toBatchEmbeddingResult(result); + } + @Override public Object getPythonResource() { return embeddingModel; diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetup.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetup.java index f0b9eca4b..fdbd05a83 100644 --- a/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetup.java +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetup.java @@ -19,6 +19,7 @@ import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelSetup; import org.apache.flink.agents.api.embedding.model.EmbeddingModelUtils; +import org.apache.flink.agents.api.embedding.model.EmbeddingResult; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.apache.flink.agents.api.resource.python.PythonResourceAdapter; @@ -41,6 +42,9 @@ */ public class PythonEmbeddingModelSetup extends BaseEmbeddingModelSetup implements PythonResourceWrapper { + private static final String CALL_EMBED_WITH_USAGE = + "python_java_utils.call_embedding_with_usage"; + private final PyObject embeddingModelSetup; private final PythonResourceAdapter adapter; @@ -123,6 +127,31 @@ public List embed(List texts, Map parameters) { + (results == null ? "null" : results.getClass().getName())); } + @Override + public EmbeddingResult embedWithUsage(String text, Map parameters) { + checkState( + embeddingModelSetup != null, + "EmbeddingModelSetup is not initialized. Cannot perform embed operation."); + + Map kwargs = new HashMap<>(parameters); + kwargs.put("text", text); + Object result = adapter.invoke(CALL_EMBED_WITH_USAGE, embeddingModelSetup, kwargs); + return EmbeddingModelUtils.toSingleEmbeddingResult(result); + } + + @Override + public EmbeddingResult> embedWithUsage( + List texts, Map parameters) { + checkState( + embeddingModelSetup != null, + "EmbeddingModelSetup is not initialized. Cannot perform embed operation."); + + Map kwargs = new HashMap<>(parameters); + kwargs.put("text", texts); + Object result = adapter.invoke(CALL_EMBED_WITH_USAGE, embeddingModelSetup, kwargs); + return EmbeddingModelUtils.toBatchEmbeddingResult(result); + } + @Override public Map getParameters() { return Map.of(); diff --git a/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupEmbeddingResultTest.java b/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupEmbeddingResultTest.java index a0d93cc65..fbbf738da 100644 --- 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 @@ -73,4 +73,42 @@ public EmbeddingResult embedWithUsage( assertEquals(7L, result.getTokenUsage().getPromptTokens()); assertEquals(9L, result.getTokenUsage().getTotalTokens()); } + + @Test + void testEmbedDelegatesToExistingConnectionMethod() { + BaseEmbeddingModelSetup setup = + new BaseEmbeddingModelSetup( + new ResourceDescriptor( + "test", Map.of("connection", "connection", "model", "model")), + mock(ResourceContext.class)) { + @Override + public Map getParameters() { + return Collections.emptyMap(); + } + }; + setup.connection = + new BaseEmbeddingModelConnection( + new ResourceDescriptor("connection", Collections.emptyMap()), + mock(ResourceContext.class)) { + @Override + public float[] embed(String text, Map parameters) { + return new float[] {0.1f, 0.2f}; + } + + @Override + public java.util.List embed( + java.util.List texts, Map parameters) { + throw new UnsupportedOperationException(); + } + + @Override + public EmbeddingResult embedWithUsage( + String text, Map parameters) { + return new EmbeddingResult<>( + new float[] {0.3f, 0.4f}, new EmbeddingTokenUsage(7L, 9L)); + } + }; + + assertArrayEquals(new float[] {0.1f, 0.2f}, setup.embed("hello")); + } } diff --git a/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnectionTest.java b/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnectionTest.java index 470d51270..a89bf3313 100644 --- a/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnectionTest.java +++ b/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnectionTest.java @@ -17,6 +17,7 @@ */ package org.apache.flink.agents.api.embedding.model.python; +import org.apache.flink.agents.api.embedding.model.EmbeddingResult; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.apache.flink.agents.api.resource.python.PythonResourceAdapter; @@ -136,6 +137,42 @@ void testEmbedSingleTextWithEmptyParameters() { assertThat(result).hasSize(2); } + @Test + void testEmbedWithUsageMultipleTexts() { + List texts = List.of("first", "second"); + when(mockAdapter.invoke( + eq("python_java_utils.call_embedding_with_usage"), + eq(mockEmbeddingModel), + any(Map.class))) + .thenReturn( + Map.of( + "embeddings", + List.of(List.of(0.1, 0.2), List.of(0.3, 0.4)), + "token_usage", + Map.of("prompt_tokens", 7, "total_tokens", 9))); + + EmbeddingResult> result = + pythonEmbeddingModelConnection.embedWithUsage(texts, Map.of("batch_size", 2)); + + assertThat(result.getEmbeddings()).hasSize(2); + assertThat(result.getEmbeddings().get(0)).containsExactly(0.1f, 0.2f); + assertThat(result.getEmbeddings().get(1)).containsExactly(0.3f, 0.4f); + assertThat(result.getTokenUsage()).isNotNull(); + assertThat(result.getTokenUsage().getPromptTokens()).isEqualTo(7L); + assertThat(result.getTokenUsage().getTotalTokens()).isEqualTo(9L); + verify(mockAdapter) + .invoke( + eq("python_java_utils.call_embedding_with_usage"), + eq(mockEmbeddingModel), + argThat( + kwargs -> { + Map values = (Map) kwargs; + assertThat(values).containsEntry("text", texts); + assertThat(values).containsEntry("batch_size", 2); + return true; + })); + } + @Test void testEmbedSingleTextWithNullEmbeddingModelThrowsException() { PythonEmbeddingModelConnection connectionWithNullModel = diff --git a/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetupTest.java b/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetupTest.java index d8071d7cd..430ddc4c3 100644 --- a/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetupTest.java +++ b/api/src/test/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetupTest.java @@ -17,6 +17,7 @@ */ package org.apache.flink.agents.api.embedding.model.python; +import org.apache.flink.agents.api.embedding.model.EmbeddingResult; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.apache.flink.agents.api.resource.python.PythonResourceAdapter; @@ -143,6 +144,41 @@ void testEmbedSingleTextWithEmptyParameters() { assertThat(result).hasSize(2); } + @Test + void testEmbedWithUsageSingleText() { + String text = "test text"; + Map parameters = Map.of("model", "test-model"); + when(mockAdapter.invoke( + eq("python_java_utils.call_embedding_with_usage"), + eq(mockEmbeddingModelSetup), + any(Map.class))) + .thenReturn( + Map.of( + "embeddings", + List.of(0.1, 0.2), + "token_usage", + Map.of("prompt_tokens", 7, "total_tokens", 9))); + + EmbeddingResult result = + pythonEmbeddingModelSetup.embedWithUsage(text, parameters); + + assertThat(result.getEmbeddings()).containsExactly(0.1f, 0.2f); + assertThat(result.getTokenUsage()).isNotNull(); + assertThat(result.getTokenUsage().getPromptTokens()).isEqualTo(7L); + assertThat(result.getTokenUsage().getTotalTokens()).isEqualTo(9L); + verify(mockAdapter) + .invoke( + eq("python_java_utils.call_embedding_with_usage"), + eq(mockEmbeddingModelSetup), + argThat( + kwargs -> { + Map values = (Map) kwargs; + assertThat(values).containsEntry("text", text); + assertThat(values).containsEntry("model", "test-model"); + return true; + })); + } + @Test void testEmbedSingleTextWithNullEmbeddingModelSetupThrowsException() { PythonEmbeddingModelSetup setupWithNullModel = diff --git a/e2e-test/flink-agents-end-to-end-tests-resource-cross-language/src/test/java/org/apache/flink/agents/resource/test/EmbeddingCrossLanguageAgent.java b/e2e-test/flink-agents-end-to-end-tests-resource-cross-language/src/test/java/org/apache/flink/agents/resource/test/EmbeddingCrossLanguageAgent.java index be5f8d494..14bfbbb4d 100644 --- a/e2e-test/flink-agents-end-to-end-tests-resource-cross-language/src/test/java/org/apache/flink/agents/resource/test/EmbeddingCrossLanguageAgent.java +++ b/e2e-test/flink-agents-end-to-end-tests-resource-cross-language/src/test/java/org/apache/flink/agents/resource/test/EmbeddingCrossLanguageAgent.java @@ -29,6 +29,7 @@ import org.apache.flink.agents.api.annotation.EmbeddingModelSetup; import org.apache.flink.agents.api.context.RunnerContext; import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelSetup; +import org.apache.flink.agents.api.embedding.model.EmbeddingResult; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.apache.flink.agents.api.resource.ResourceName; @@ -105,11 +106,14 @@ public static void testEmbeddingGeneration(Event event, RunnerContext ctx) throw org.apache.flink.agents.api.resource.ResourceType .EMBEDDING_MODEL); - float[] embedding = embeddingModel.embed(text); + EmbeddingResult embeddingResult = embeddingModel.embedWithUsage(text); + float[] embedding = embeddingResult.getEmbeddings(); System.out.printf("[TEST] Generated embedding with dimension: %d%n", embedding.length); validateEmbeddingResult(id, text, embedding); - List embeddings = embeddingModel.embed(List.of(text)); + EmbeddingResult> embeddingsResult = + embeddingModel.embedWithUsage(List.of(text)); + List embeddings = embeddingsResult.getEmbeddings(); validateEmbeddingResults(id, List.of(text), embeddings); // Create a minimal test result to avoid serialization issues diff --git a/python/flink_agents/api/embedding_models/embedding_model.py b/python/flink_agents/api/embedding_models/embedding_model.py index ce96c0095..c31ab5324 100644 --- a/python/flink_agents/api/embedding_models/embedding_model.py +++ b/python/flink_agents/api/embedding_models/embedding_model.py @@ -145,7 +145,9 @@ def embed( A list of floating-point numbers representing the embedding vector. The dimension of the vector depends on the specific embedding model used. """ - return self.embed_with_usage(text, **kwargs).embeddings + 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 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 index 88534dfd3..e42193520 100644 --- a/python/flink_agents/api/embedding_models/tests/test_embedding_result.py +++ b/python/flink_agents/api/embedding_models/tests/test_embedding_result.py @@ -54,6 +54,25 @@ def embed( 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]: @@ -100,3 +119,9 @@ def test_embed_preserves_existing_embedding_only_api() -> None: setup = _make_setup(FakeEmbeddingModelConnection(name="connection")) assert setup.embed("hello") == [0.1, 0.2] + + +def test_embed_preserves_connection_embed_dispatch() -> None: + setup = _make_setup(DispatchAwareEmbeddingModelConnection(name="connection")) + + assert setup.embed("hello") == [0.1, 0.2] diff --git a/python/flink_agents/e2e_tests/e2e_tests_resource_cross_language/embedding_model_cross_language_agent.py b/python/flink_agents/e2e_tests/e2e_tests_resource_cross_language/embedding_model_cross_language_agent.py index dbce4891b..e57fbc07f 100644 --- a/python/flink_agents/e2e_tests/e2e_tests_resource_cross_language/embedding_model_cross_language_agent.py +++ b/python/flink_agents/e2e_tests/e2e_tests_resource_cross_language/embedding_model_cross_language_agent.py @@ -77,7 +77,8 @@ def process_input(event: Event, ctx: RunnerContext) -> None: ) # Test single text embedding - embedding = embeddingModel.embed(input_text) + embedding_result = embeddingModel.embed_with_usage(input_text) + embedding = embedding_result.embeddings print(f"[TEST] Generated embedding with dimension: {len(embedding)}") # Validate single embedding result @@ -98,7 +99,8 @@ def process_input(event: Event, ctx: RunnerContext) -> None: ) # Test batch embedding - embeddings = embeddingModel.embed([input_text]) + embeddings_result = embeddingModel.embed_with_usage([input_text]) + embeddings = embeddings_result.embeddings print(f"[TEST] Generated batch embeddings: count={len(embeddings)}") # Validate batch embedding results diff --git a/python/flink_agents/integrations/embedding_models/tests/test_tongyi_embedding_model.py b/python/flink_agents/integrations/embedding_models/tests/test_tongyi_embedding_model.py index 98508a9b2..6262d4bd6 100644 --- a/python/flink_agents/integrations/embedding_models/tests/test_tongyi_embedding_model.py +++ b/python/flink_agents/integrations/embedding_models/tests/test_tongyi_embedding_model.py @@ -164,8 +164,8 @@ def test_tongyi_embedding_returns_token_usage( status_code=HTTPStatus.OK, output={ "embeddings": [{"embedding": mock_embedding}], - "usage": {"input_tokens": 6, "total_tokens": 6}, }, + usage={"total_tokens": 6}, message="Success", ) mock_call = MagicMock(return_value=mocked_response) 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 bbcc981e4..bea3075bd 100644 --- a/python/flink_agents/integrations/embedding_models/tongyi_embedding_model.py +++ b/python/flink_agents/integrations/embedding_models/tongyi_embedding_model.py @@ -129,6 +129,8 @@ def embed_with_usage( 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( diff --git a/python/flink_agents/runtime/java/java_embedding_model.py b/python/flink_agents/runtime/java/java_embedding_model.py index 2cb15b819..7773515f3 100644 --- a/python/flink_agents/runtime/java/java_embedding_model.py +++ b/python/flink_agents/runtime/java/java_embedding_model.py @@ -19,12 +19,38 @@ 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, ) +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. @@ -64,6 +90,16 @@ def embed( ) return list(result) if isinstance(text, str) else [list(emb) for emb in result] + @override + def embed_with_usage( + self, text: str | Sequence[str], **kwargs: Any + ) -> EmbeddingResult[list[float] | list[list[float]]]: + """Generate embeddings through Java and preserve provider token usage.""" + result = self._j_resource.embedWithUsage( + text if isinstance(text, str) else list(text), kwargs + ) + return _from_java_embedding_result(result, text) + class JavaEmbeddingModelSetupImpl(JavaEmbeddingModelSetup): """Java-based implementation of EmbeddingModelSetup that wraps a Java embedding @@ -121,3 +157,13 @@ def embed( text if isinstance(text, str) else list(text), kwargs ) return list(result) if isinstance(text, str) else [list(emb) for emb in result] + + @override + def embed_with_usage( + self, text: str | Sequence[str], **kwargs: Any + ) -> EmbeddingResult[list[float] | list[list[float]]]: + """Generate embeddings through Java and preserve provider token usage.""" + result = self._j_resource.embedWithUsage( + text if isinstance(text, str) else list(text), kwargs + ) + return _from_java_embedding_result(result, text) diff --git a/python/flink_agents/runtime/python_java_utils.py b/python/flink_agents/runtime/python_java_utils.py index 51f08f810..2dca8ad9f 100644 --- a/python/flink_agents/runtime/python_java_utils.py +++ b/python/flink_agents/runtime/python_java_utils.py @@ -154,7 +154,9 @@ def get_python_tool_metadata( descriptor = PythonFunction(module=module, qualname=qual_name) callable_ = descriptor.as_callable() name = callable_.__name__ - description = (parse(callable_.__doc__).description or "") if callable_.__doc__ else "" + description = ( + (parse(callable_.__doc__).description or "") if callable_.__doc__ else "" + ) callable_injected_args = normalize_injected_args( getattr(callable_, "_injected_args", None) ) @@ -177,9 +179,7 @@ def _dump_injected_args(injected_args: Dict[str, Any]) -> str: ) -def invoke_python_tool( - module: str, qual_name: str, kwargs: Dict[str, Any] -) -> Any: +def invoke_python_tool(module: str, qual_name: str, kwargs: Dict[str, Any]) -> Any: """Invoke a Python callable as a tool, passing the provided keyword arguments. Used by the Java-side ``PythonResourceAdapter.invokePythonTool`` so a Java host can @@ -226,6 +226,32 @@ def from_java_resource(type_name: str, kwargs: Dict[str, Any]) -> Resource: return cls(**kwargs) +def embedding_result_to_java(result: Any) -> Dict[str, Any]: + """Convert an embedding result into Pemja-safe Python primitives.""" + usage = result.token_usage + return { + "embeddings": result.embeddings, + "token_usage": ( + None + if usage is None + else { + "prompt_tokens": usage.prompt_tokens, + "total_tokens": usage.total_tokens, + } + ), + } + + +def call_embedding_with_usage( + embedding_model: Any, kwargs: Dict[str, Any] +) -> Dict[str, Any]: + """Call ``embed_with_usage`` and return Java-safe primitives. + + This avoids PyObject attribute access in the Java caller. + """ + return embedding_result_to_java(embedding_model.embed_with_usage(**kwargs)) + + def normalize_tool_call_id(tool_call: Dict[str, Any]) -> Dict[str, Any]: """Normalize tool call by converting the ID field to string format while preserving all other fields. diff --git a/python/flink_agents/runtime/tests/test_java_embedding_model.py b/python/flink_agents/runtime/tests/test_java_embedding_model.py new file mode 100644 index 000000000..2a3866687 --- /dev/null +++ b/python/flink_agents/runtime/tests/test_java_embedding_model.py @@ -0,0 +1,70 @@ +################################################################################ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +################################################################################# +from unittest.mock import MagicMock + +import pytest + +from flink_agents.runtime.java.java_embedding_model import ( + JavaEmbeddingModelConnectionImpl, + JavaEmbeddingModelSetupImpl, +) + + +class _JavaTokenUsage: + def getPromptTokens(self) -> int: + return 7 + + def getTotalTokens(self) -> int: + return 9 + + +class _JavaEmbeddingResult: + def getEmbeddings(self) -> list[list[float]]: + return [[0.1, 0.2], [0.3, 0.4]] + + def getTokenUsage(self) -> _JavaTokenUsage: + return _JavaTokenUsage() + + +@pytest.mark.parametrize( + ("wrapper_class", "kwargs"), + [ + (JavaEmbeddingModelConnectionImpl, {}), + ( + JavaEmbeddingModelSetupImpl, + {"connection": "connection", "model": "test-model"}, + ), + ], +) +def test_java_embedding_wrappers_preserve_usage( + wrapper_class: type[JavaEmbeddingModelConnectionImpl | JavaEmbeddingModelSetupImpl], + kwargs: dict[str, str], +) -> None: + j_resource = MagicMock() + j_resource.embedWithUsage.return_value = _JavaEmbeddingResult() + wrapper = wrapper_class(j_resource, MagicMock(), **kwargs) + + result = wrapper.embed_with_usage(["first", "second"], batch_size=2) + + assert result.embeddings == [[0.1, 0.2], [0.3, 0.4]] + assert result.token_usage is not None + assert result.token_usage.prompt_tokens == 7 + assert result.token_usage.total_tokens == 9 + j_resource.embedWithUsage.assert_called_once_with( + ["first", "second"], {"batch_size": 2} + ) diff --git a/python/flink_agents/runtime/tests/test_python_java_utils.py b/python/flink_agents/runtime/tests/test_python_java_utils.py index 7aeeafe9c..fa708353b 100644 --- a/python/flink_agents/runtime/tests/test_python_java_utils.py +++ b/python/flink_agents/runtime/tests/test_python_java_utils.py @@ -18,8 +18,15 @@ import json from flink_agents.api.decorators import tool +from flink_agents.api.embedding_models.embedding_model import ( + EmbeddingResult, + EmbeddingTokenUsage, +) from flink_agents.api.tools import InjectedArg -from flink_agents.runtime.python_java_utils import get_python_tool_metadata +from flink_agents.runtime.python_java_utils import ( + call_embedding_with_usage, + get_python_tool_metadata, +) @tool(injected_args={"tenant_id": InjectedArg.from_config("tenant.id")}) @@ -36,6 +43,27 @@ def test_get_python_tool_metadata_merges_callable_injected_args() -> None: schema = json.loads(flat["inputSchema"]) assert set(schema["properties"]) == {"order_id"} injected_args = json.loads(flat["injectedArgs"]) - assert injected_args == { - "tenant_id": {"source": "config", "key": "tenant.id"} + assert injected_args == {"tenant_id": {"source": "config", "key": "tenant.id"}} + + +class _UsageAwareEmbeddingModel: + def embed_with_usage( + self, text: str, **kwargs: object + ) -> EmbeddingResult[list[float]]: + assert text == "hello" + assert kwargs == {"model": "test-model"} + return EmbeddingResult( + embeddings=[0.1, 0.2], + token_usage=EmbeddingTokenUsage(prompt_tokens=7, total_tokens=9), + ) + + +def test_call_embedding_with_usage_returns_pemja_safe_primitives() -> None: + result = call_embedding_with_usage( + _UsageAwareEmbeddingModel(), {"text": "hello", "model": "test-model"} + ) + + assert result == { + "embeddings": [0.1, 0.2], + "token_usage": {"prompt_tokens": 7, "total_tokens": 9}, }