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 605c189d5..2fd7e5c6f 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; @@ -93,9 +94,52 @@ public String getModel() { return model; } + /** + * Record embedding token usage metrics for the given model on this setup's bound metric group. + * + *

Mirrors {@code BaseChatModelSetup#recordTokenMetrics} but records input-side tokens only, + * since embeddings have no completion tokens. Counters are placed under the same {@code model} + * key-value group used by chat metrics, so embedding and chat usage for a model share one + * dimension. + * + *

Unlike the chat path, embedding calls do not run inside a plan action that hands in a + * request-scoped metric group (vector-store, RAG, and direct calls reach this setup directly), + * so the resource-bound metric group injected via {@link #setMetricGroup} is used instead. + * + * @param modelName the name of the model used + * @param promptTokens the number of prompt tokens + * @param totalTokens the total number of tokens reported by the provider + */ + public void recordTokenMetrics(String modelName, long promptTokens, long totalTokens) { + Preconditions.checkArgument( + modelName != null && !modelName.isBlank(), "Model name must not be null or blank."); + FlinkAgentsMetricGroup metricGroup = getMetricGroup(); + if (metricGroup == null) { + return; + } + FlinkAgentsMetricGroup modelGroup = metricGroup.getSubGroup("model", modelName); + modelGroup.getCounter("promptTokens").inc(promptTokens); + modelGroup.getCounter("totalTokens").inc(totalTokens); + } + + /** + * Record the provider-reported embedding token usage, if any, onto this setup's bound metric + * group. Called from {@link #embedWithUsage} so direct calls and vector-store/RAG paths are + * both covered without each provider repeating the recording. + */ + protected void recordTokenUsage(@Nullable EmbeddingTokenUsage tokenUsage) { + if (tokenUsage == null || model == null || model.isBlank()) { + return; + } + recordTokenMetrics(model, tokenUsage.getPromptTokens(), tokenUsage.getTotalTokens()); + } + /** * Generate embeddings for the given text. * + *

Token usage metrics are only recorded by {@link #embedWithUsage}; this method discards + * provider usage because it is not returned. + * * @param text The input text to generate embeddings for * @return An array of floating-point values representing the text embeddings */ @@ -117,7 +161,9 @@ public EmbeddingResult embedWithUsage(String text, Map Map params = this.getParameters(); params.putAll(parameters); BaseEmbeddingModelConnection currentConnection = getConnection(); - return currentConnection.embedWithUsage(text, params); + EmbeddingResult result = currentConnection.embedWithUsage(text, params); + recordTokenUsage(result.getTokenUsage()); + return result; } /** @@ -146,6 +192,8 @@ public EmbeddingResult> embedWithUsage( Map params = this.getParameters(); params.putAll(parameters); BaseEmbeddingModelConnection currentConnection = getConnection(); - return currentConnection.embedWithUsage(texts, params); + EmbeddingResult> result = currentConnection.embedWithUsage(texts, params); + recordTokenUsage(result.getTokenUsage()); + return result; } } 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 b6febbc58..2460b3b24 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 @@ -137,7 +137,10 @@ public EmbeddingResult embedWithUsage(String text, Map Map kwargs = new HashMap<>(parameters); kwargs.put("text", text); Object result = adapter.invoke(CALL_EMBED_WITH_USAGE, embeddingModelSetup, kwargs); - return EmbeddingModelUtils.toSingleEmbeddingResult(result); + EmbeddingResult embeddingResult = + EmbeddingModelUtils.toSingleEmbeddingResult(result); + recordTokenUsage(embeddingResult.getTokenUsage()); + return embeddingResult; } @Override @@ -150,7 +153,10 @@ public EmbeddingResult> embedWithUsage( Map kwargs = new HashMap<>(parameters); kwargs.put("text", texts); Object result = adapter.invoke(CALL_EMBED_WITH_USAGE, embeddingModelSetup, kwargs); - return EmbeddingModelUtils.toBatchEmbeddingResult(result); + EmbeddingResult> embeddingResult = + EmbeddingModelUtils.toBatchEmbeddingResult(result); + recordTokenUsage(embeddingResult.getTokenUsage()); + return embeddingResult; } @Override 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..fe989b4c4 --- /dev/null +++ b/api/src/test/java/org/apache/flink/agents/api/embedding/model/BaseEmbeddingModelSetupTokenMetricsTest.java @@ -0,0 +1,318 @@ +/* + * 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 of + * 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.DisplayName; +import org.junit.jupiter.api.Test; + +import java.util.Collections; +import java.util.HashMap; +import java.util.List; +import java.util.Map; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.mockito.Mockito.mock; + +/** + * Test cases for embedding token usage metrics recorded by {@link BaseEmbeddingModelSetup}. Mirrors + * {@code BaseChatModelSetupTokenMetricsTest}: embedding providers already populate {@link + * EmbeddingTokenUsage} on the returned {@link EmbeddingResult}, but nothing records it until this + * setup reads it back at the {@code embedWithUsage} chokepoint. + */ +class BaseEmbeddingModelSetupTokenMetricsTest { + + /** Value-based metric group that mirrors the one in the chat token-metrics test. */ + 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, k -> new TestMetricGroup()); + } + + @Override + public FlinkAgentsMetricGroup getSubGroup(String key, String value) { + return subGroups.computeIfAbsent(key + "=" + value, k -> new TestMetricGroup()); + } + + @Override + public Counter getCounter(String name) { + return counters.computeIfAbsent(name, k -> new SimpleCounter()); + } + + @Override + public UpdatableGauge getGauge(String name) { + return null; + } + + @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; + } + } + + private static final float[] VEC = new float[] {0.1f, 0.2f}; + + /** + * Builds a setup bound to a connection that reports the given usage on single-text embed, with + * the given model name in its descriptor (may be {@code null} to exercise the guard). + */ + private static BaseEmbeddingModelSetup setupWithSingleUsageAndModel( + EmbeddingTokenUsage usage, String model) { + BaseEmbeddingModelSetup setup = + new BaseEmbeddingModelSetup( + new ResourceDescriptor("test", descriptorArgs(model)), + mock(ResourceContext.class)) { + @Override + public Map getParameters() { + return new HashMap<>(); + } + }; + setup.connection = + new BaseEmbeddingModelConnection( + new ResourceDescriptor("conn", Collections.emptyMap()), + mock(ResourceContext.class)) { + @Override + public float[] embed(String text, Map parameters) { + return VEC; + } + + @Override + public List embed(List texts, Map parameters) { + throw new UnsupportedOperationException(); + } + + @Override + public EmbeddingResult embedWithUsage( + String text, Map parameters) { + return new EmbeddingResult<>(VEC, usage); + } + }; + return setup; + } + + /** Builds a setup bound to a connection that reports the given usage on single-text embed. */ + private static BaseEmbeddingModelSetup setupWithSingleUsage(EmbeddingTokenUsage usage) { + return setupWithSingleUsageAndModel(usage, "bedrock-text"); + } + + /** Descriptor args with an optional model (omitted when null/blank so it stays unset). */ + private static Map descriptorArgs(String model) { + Map args = new HashMap<>(); + args.put("connection", "conn"); + if (model != null && !model.isBlank()) { + args.put("model", model); + } + return args; + } + + /** Builds a setup whose connection reports the given usage on batch embed. */ + private static BaseEmbeddingModelSetup setupWithBatchUsage(EmbeddingTokenUsage usage) { + BaseEmbeddingModelSetup setup = + new BaseEmbeddingModelSetup( + new ResourceDescriptor( + "test", Map.of("connection", "conn", "model", "bedrock-text")), + mock(ResourceContext.class)) { + @Override + public Map getParameters() { + return new HashMap<>(); + } + }; + setup.connection = + new BaseEmbeddingModelConnection( + new ResourceDescriptor("conn", Collections.emptyMap()), + mock(ResourceContext.class)) { + @Override + public float[] embed(String text, Map parameters) { + throw new UnsupportedOperationException(); + } + + @Override + public List embed(List texts, Map parameters) { + return Collections.singletonList(VEC); + } + + @Override + public EmbeddingResult> embedWithUsage( + List texts, Map parameters) { + return new EmbeddingResult<>(Collections.singletonList(VEC), usage); + } + }; + return setup; + } + + private static TestMetricGroup modelGroup(TestMetricGroup root, String model) { + return (TestMetricGroup) root.getSubGroup("model", model); + } + + @Test + @DisplayName("recordTokenMetrics records prompt and total tokens under the model group") + void testRecordTokenMetricsUnderModelGroup() { + BaseEmbeddingModelSetup setup = setupWithSingleUsage(null); + TestMetricGroup root = new TestMetricGroup(); + setup.setMetricGroup(root); + + setup.recordTokenMetrics("bedrock-text", 100, 210); + + TestMetricGroup model = modelGroup(root, "bedrock-text"); + assertEquals(100, model.counters.get("promptTokens").getCount()); + assertEquals(210, model.counters.get("totalTokens").getCount()); + } + + @Test + @DisplayName("embedWithUsage single records provider usage onto the model group") + void testEmbedWithUsageSingleRecordsUsage() { + BaseEmbeddingModelSetup setup = setupWithSingleUsage(new EmbeddingTokenUsage(100, 210)); + TestMetricGroup root = new TestMetricGroup(); + setup.setMetricGroup(root); + + EmbeddingResult result = setup.embedWithUsage("hello"); + + // result is still returned with its usage intact + assertEquals(100, result.getTokenUsage().getPromptTokens()); + assertEquals(210, result.getTokenUsage().getTotalTokens()); + // ...and the same usage was recorded as metrics + TestMetricGroup model = modelGroup(root, "bedrock-text"); + assertEquals(100, model.counters.get("promptTokens").getCount()); + assertEquals(210, model.counters.get("totalTokens").getCount()); + } + + @Test + @DisplayName("embedWithUsage batch records provider usage onto the model group") + void testEmbedWithUsageBatchRecordsUsage() { + BaseEmbeddingModelSetup setup = setupWithBatchUsage(new EmbeddingTokenUsage(100, 210)); + TestMetricGroup root = new TestMetricGroup(); + setup.setMetricGroup(root); + + EmbeddingResult> result = + setup.embedWithUsage(Collections.singletonList("hello")); + + assertEquals(100, result.getTokenUsage().getPromptTokens()); + assertEquals(210, result.getTokenUsage().getTotalTokens()); + TestMetricGroup model = modelGroup(root, "bedrock-text"); + assertEquals(100, model.counters.get("promptTokens").getCount()); + assertEquals(210, model.counters.get("totalTokens").getCount()); + } + + @Test + @DisplayName("embedWithUsage records nothing when the provider reports no usage") + void testEmbedWithUsageNullUsageRecordsNothing() { + BaseEmbeddingModelSetup setup = setupWithSingleUsage(null); + TestMetricGroup root = new TestMetricGroup(); + setup.setMetricGroup(root); + + setup.embedWithUsage("hello"); + + // model group exists only if a counter was requested; absent means nothing was recorded + assertFalse(root.subGroups.containsKey("model=bedrock-text")); + } + + @Test + @DisplayName("recordTokenMetrics is a no-op when no metric group is bound") + void testRecordTokenMetricsWithoutMetricGroup() { + BaseEmbeddingModelSetup setup = setupWithSingleUsage(null); + // no setMetricGroup call -> getMetricGroup() returns null + + // must not throw + setup.recordTokenMetrics("bedrock-text", 100, 210); + } + + @Test + @DisplayName("embedWithUsage records nothing when no metric group is bound") + void testEmbedWithUsageWithoutMetricGroupRecordsNothingButReturnsUsage() { + BaseEmbeddingModelSetup setup = setupWithSingleUsage(new EmbeddingTokenUsage(100, 210)); + // no setMetricGroup call + + EmbeddingResult result = setup.embedWithUsage("hello"); + + // usage still flows back to the caller; only metrics are skipped + assertEquals(100, result.getTokenUsage().getPromptTokens()); + } + + @Test + @DisplayName("counters accumulate across multiple embedding calls") + void testCountersAccumulate() { + BaseEmbeddingModelSetup setup = setupWithSingleUsage(new EmbeddingTokenUsage(100, 210)); + TestMetricGroup root = new TestMetricGroup(); + setup.setMetricGroup(root); + + setup.embedWithUsage("a"); + setup.embedWithUsage("b"); + + TestMetricGroup model = modelGroup(root, "bedrock-text"); + assertEquals(200, model.counters.get("promptTokens").getCount()); + assertEquals(420, model.counters.get("totalTokens").getCount()); + } + + @Test + @DisplayName("embedWithUsage records nothing when the setup has no model name") + void testEmbedWithUsageNullModelRecordsNothing() { + // descriptor omits "model" -> getArgument("model") returns null + BaseEmbeddingModelSetup setup = + setupWithSingleUsageAndModel(new EmbeddingTokenUsage(100, 210), null); + TestMetricGroup root = new TestMetricGroup(); + setup.setMetricGroup(root); + + EmbeddingResult result = setup.embedWithUsage("hello"); + + // usage still flows back ... + assertEquals(100, result.getTokenUsage().getPromptTokens()); + // ... but no model group is created without a model name to key on + assertFalse(root.subGroups.containsKey("model=null")); + } + + @Test + @DisplayName("recordTokenMetrics rejects a null or blank model name") + void testRecordTokenMetricsRejectsBlankModelName() { + BaseEmbeddingModelSetup setup = setupWithSingleUsage(null); + setup.setMetricGroup(new TestMetricGroup()); + + assertThrows( + IllegalArgumentException.class, () -> setup.recordTokenMetrics(null, 100, 210)); + assertThrows(IllegalArgumentException.class, () -> setup.recordTokenMetrics("", 100, 210)); + assertThrows( + IllegalArgumentException.class, () -> setup.recordTokenMetrics(" ", 100, 210)); + } +} diff --git a/python/flink_agents/api/embedding_models/embedding_model.py b/python/flink_agents/api/embedding_models/embedding_model.py index c31ab5324..51218697b 100644 --- a/python/flink_agents/api/embedding_models/embedding_model.py +++ b/python/flink_agents/api/embedding_models/embedding_model.py @@ -155,4 +155,50 @@ def embed_with_usage( """Generate embeddings and return provider token usage when available.""" merged_kwargs = self.model_kwargs.copy() merged_kwargs.update(kwargs) - return self._get_connection().embed_with_usage(text, **merged_kwargs) + result = self._get_connection().embed_with_usage(text, **merged_kwargs) + self._record_token_usage(result.token_usage) + return result + + def _record_token_metrics( + self, model_name: str, prompt_tokens: int, total_tokens: int + ) -> None: + """Record embedding token usage metrics for the given model. + + Mirrors ``BaseChatModelSetup._record_token_metrics`` but records input-side + tokens only, since embeddings have no completion tokens. Counters are placed + under the same ``model`` key-value group used by chat metrics, so embedding + and chat usage for a model share one dimension. + + Unlike the chat path, embedding calls do not run inside a plan action that + hands in a request-scoped metric group (vector-store, RAG, and direct calls + reach this setup directly), so the resource-bound metric group injected via + ``set_metric_group`` is used instead. + + Parameters + ---------- + model_name : str + The name of the model used + prompt_tokens : int + The number of prompt tokens + total_tokens : int + The total number of tokens reported by the provider + """ + metric_group = self.metric_group + if metric_group is None: + return + + model_group = metric_group.get_sub_group("model", model_name) + model_group.get_counter("promptTokens").inc(prompt_tokens) + model_group.get_counter("totalTokens").inc(total_tokens) + + def _record_token_usage(self, token_usage: EmbeddingTokenUsage | None) -> None: + """Record the provider-reported embedding token usage, if any. + + Called from ``embed_with_usage`` so direct calls and vector-store/RAG paths + are both covered without each provider repeating the recording. + """ + if token_usage is None or not self.model: + return + self._record_token_metrics( + self.model, token_usage.prompt_tokens, token_usage.total_tokens + ) 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..7c12d584a --- /dev/null +++ b/python/flink_agents/api/embedding_models/tests/test_token_metrics.py @@ -0,0 +1,200 @@ +################################################################################ +# 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. +################################################################################# +"""Test cases for BaseEmbeddingModelSetup token usage metrics. + +Mirrors ``test_token_metrics.py`` for the chat side: embedding providers already +populate ``EmbeddingTokenUsage`` on the returned ``EmbeddingResult``, but nothing +records it until this setup reads it back inside ``embed_with_usage``. +""" + +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.metric_group import Counter, MetricGroup +from flink_agents.api.resource import Resource, ResourceType +from flink_agents.api.resource_context import ResourceContext + + +class _EmbeddingConnectionWithUsage(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 _EmbeddingConnectionWithoutUsage(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 _TestEmbeddingModelSetup(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) -> _TestEmbeddingModelSetup: + 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 = _TestEmbeddingModelSetup( + name="embedding", + connection="mock-connection", + model="mock-model", + resource_context=ctx, + ) + setup.open() + return setup + + +def test_embed_with_usage_records_token_metrics() -> None: + """embed_with_usage records provider usage onto the model metric group.""" + setup = _make_setup(_EmbeddingConnectionWithUsage(name="connection")) + mock_metric_group = _MockMetricGroup() + setup.set_metric_group(mock_metric_group) + + result = setup.embed_with_usage("hello") + + # usage still flows back to the caller ... + assert result.token_usage == EmbeddingTokenUsage(prompt_tokens=7, total_tokens=9) + # ... and the same usage was recorded as metrics + model_group = mock_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_embed_with_usage_records_token_metrics_batch() -> None: + """embed_with_usage records provider usage for batch inputs too.""" + setup = _make_setup(_EmbeddingConnectionWithUsage(name="connection")) + mock_metric_group = _MockMetricGroup() + setup.set_metric_group(mock_metric_group) + + setup.embed_with_usage(["hello", "world"]) + + model_group = mock_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_embed_with_usage_without_usage_records_nothing() -> None: + """When the provider reports no usage, no metrics are recorded.""" + setup = _make_setup(_EmbeddingConnectionWithoutUsage(name="connection")) + mock_metric_group = _MockMetricGroup() + setup.set_metric_group(mock_metric_group) + + setup.embed_with_usage("hello") + + # model group is only created when a counter is requested; absent means nothing recorded + assert "model=mock-model" not in mock_metric_group._sub_groups + + +def test_embed_with_usage_without_metric_group_returns_usage() -> None: + """Without a bound metric group, usage still flows back; only metrics are skipped.""" + setup = _make_setup(_EmbeddingConnectionWithUsage(name="connection")) + # no set_metric_group call + + result = setup.embed_with_usage("hello") + + assert result.token_usage == EmbeddingTokenUsage(prompt_tokens=7, total_tokens=9) + + +def test_token_metrics_accumulate() -> None: + """Counters accumulate across multiple embedding calls.""" + setup = _make_setup(_EmbeddingConnectionWithUsage(name="connection")) + mock_metric_group = _MockMetricGroup() + setup.set_metric_group(mock_metric_group) + + setup.embed_with_usage("a") + setup.embed_with_usage("b") + + model_group = mock_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_token_metrics_without_metric_group_is_noop() -> None: + """record_token_metrics must not throw when no metric group is bound.""" + setup = _make_setup(_EmbeddingConnectionWithUsage(name="connection")) + + # no set_metric_group call -> metric_group is None + setup._record_token_metrics("mock-model", 7, 9) + # no exception raised