This is an automated email from the ASF dual-hosted git repository.
wenjin272 pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/flink-agents.git
The following commit(s) were added to refs/heads/main by this push:
new cddd1e6c [api][plan][python] Use request-scoped metric groups for chat
token metrics (#861)
cddd1e6c is described below
commit cddd1e6ca4deee2a9340cc3ed752289767c13b74
Author: John <[email protected]>
AuthorDate: Wed Aug 12 19:13:23 2026 +0800
[api][plan][python] Use request-scoped metric groups for chat token metrics
(#861)
* [api][plan][python] Use request-scoped metric groups for chat token
metrics
* [api][plan][python] Address request-scoped metrics review
---------
Co-authored-by: 80347547 <[email protected]>
---
.../agents/api/chat/model/BaseChatModelSetup.java | 17 +++---
.../model/BaseChatModelSetupTokenMetricsTest.java | 43 +++++++-------
.../flink/agents/plan/actions/ChatModelAction.java | 14 ++++-
.../plan/actions/ChatModelActionRetryTest.java | 29 +++++++++
.../agents/plan/actions/ChatModelActionTest.java | 49 ++++++++++++----
python/flink_agents/api/chat_models/chat_model.py | 11 +++-
.../api/chat_models/tests/test_token_metrics.py | 68 +++++++++++++++-------
.../flink_agents/plan/actions/chat_model_action.py | 5 +-
8 files changed, 169 insertions(+), 67 deletions(-)
diff --git
a/api/src/main/java/org/apache/flink/agents/api/chat/model/BaseChatModelSetup.java
b/api/src/main/java/org/apache/flink/agents/api/chat/model/BaseChatModelSetup.java
index 6c59f6ef..0c61bf6d 100644
---
a/api/src/main/java/org/apache/flink/agents/api/chat/model/BaseChatModelSetup.java
+++
b/api/src/main/java/org/apache/flink/agents/api/chat/model/BaseChatModelSetup.java
@@ -114,18 +114,21 @@ public abstract class BaseChatModelSetup extends Resource
{
public abstract Map<String, Object> getParameters();
/**
- * Record token usage metrics for the given model on this setup's bound
metric group.
+ * Record token usage metrics for the given model on the provided metric
group.
*
+ * @param metricGroup the non-null metric group captured when the request
was initiated
* @param modelName the name of the model used
* @param promptTokens the number of prompt tokens
* @param completionTokens the number of completion tokens
*/
- public void recordTokenMetrics(String modelName, long promptTokens, long
completionTokens) {
- FlinkAgentsMetricGroup metricGroup = getMetricGroup();
- if (metricGroup == null) {
- return;
- }
- FlinkAgentsMetricGroup modelGroup = metricGroup.getSubGroup("model",
modelName);
+ public void recordTokenMetrics(
+ FlinkAgentsMetricGroup metricGroup,
+ String modelName,
+ long promptTokens,
+ long completionTokens) {
+ FlinkAgentsMetricGroup modelGroup =
+ Preconditions.checkNotNull(metricGroup, "Metric group must not
be null.")
+ .getSubGroup("model", modelName);
modelGroup.getCounter("promptTokens").inc(promptTokens);
modelGroup.getCounter("completionTokens").inc(completionTokens);
}
diff --git
a/api/src/test/java/org/apache/flink/agents/api/chat/model/BaseChatModelSetupTokenMetricsTest.java
b/api/src/test/java/org/apache/flink/agents/api/chat/model/BaseChatModelSetupTokenMetricsTest.java
index 8e47105f..abc7f2c0 100644
---
a/api/src/test/java/org/apache/flink/agents/api/chat/model/BaseChatModelSetupTokenMetricsTest.java
+++
b/api/src/test/java/org/apache/flink/agents/api/chat/model/BaseChatModelSetupTokenMetricsTest.java
@@ -35,11 +35,9 @@ import java.util.Collections;
import java.util.HashMap;
import java.util.Map;
-import static org.junit.jupiter.api.Assertions.assertDoesNotThrow;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verify;
-import static org.mockito.Mockito.verifyNoInteractions;
import static org.mockito.Mockito.when;
/** Test cases for BaseChatModelSetup token metrics functionality. */
@@ -85,9 +83,7 @@ class BaseChatModelSetupTokenMetricsTest {
@Test
@DisplayName("Test token metrics are recorded when metric group is set")
void testRecordTokenMetricsWithMetricGroup() {
- setup.setMetricGroup(mockMetricGroup);
-
- setup.recordTokenMetrics("gpt-4", 100, 50);
+ setup.recordTokenMetrics(mockMetricGroup, "gpt-4", 100, 50);
verify(mockMetricGroup).getSubGroup("model", "gpt-4");
verify(mockModelGroup).getCounter("promptTokens");
@@ -97,18 +93,26 @@ class BaseChatModelSetupTokenMetricsTest {
}
@Test
- @DisplayName("Test token metrics are not recorded when metric group is
null")
- void testRecordTokenMetricsWithoutMetricGroup() {
- assertDoesNotThrow(() -> setup.recordTokenMetrics("gpt-4", 100, 50));
+ @DisplayName("Test token metrics use the request-scoped metric group")
+ void testRecordTokenMetricsWithRequestScopedMetricGroup() {
+ TestMetricGroup actionA = new TestMetricGroup();
+ TestMetricGroup actionB = new TestMetricGroup();
+
+ setup.setMetricGroup(actionB);
+ setup.recordTokenMetrics(actionA, "gpt-4", 100, 50);
- verifyNoInteractions(mockMetricGroup);
+ TestMetricGroup actionAModelGroup = (TestMetricGroup)
actionA.getSubGroup("model", "gpt-4");
+ assertEquals(100,
actionAModelGroup.counters.get("promptTokens").getCount());
+ assertEquals(50,
actionAModelGroup.counters.get("completionTokens").getCount());
+
+ TestMetricGroup actionBModelGroup = (TestMetricGroup)
actionB.getSubGroup("model", "gpt-4");
+ assertEquals(0,
actionBModelGroup.getCounter("promptTokens").getCount());
+ assertEquals(0,
actionBModelGroup.getCounter("completionTokens").getCount());
}
@Test
@DisplayName("Test token metrics hierarchy: metricGroup -> modelName ->
counters")
void testTokenMetricsHierarchy() {
- setup.setMetricGroup(mockMetricGroup);
-
FlinkAgentsMetricGroup mockGpt35Group =
mock(FlinkAgentsMetricGroup.class);
Counter mockGpt35PromptCounter = mock(Counter.class);
Counter mockGpt35CompletionCounter = mock(Counter.class);
@@ -117,8 +121,8 @@ class BaseChatModelSetupTokenMetricsTest {
when(mockGpt35Group.getCounter("promptTokens")).thenReturn(mockGpt35PromptCounter);
when(mockGpt35Group.getCounter("completionTokens")).thenReturn(mockGpt35CompletionCounter);
- setup.recordTokenMetrics("gpt-4", 100, 50);
- setup.recordTokenMetrics("gpt-3.5-turbo", 200, 100);
+ setup.recordTokenMetrics(mockMetricGroup, "gpt-4", 100, 50);
+ setup.recordTokenMetrics(mockMetricGroup, "gpt-3.5-turbo", 200, 100);
verify(mockMetricGroup).getSubGroup("model", "gpt-4");
verify(mockMetricGroup).getSubGroup("model", "gpt-3.5-turbo");
@@ -187,9 +191,8 @@ class BaseChatModelSetupTokenMetricsTest {
@DisplayName("Value-based: token counters are accessible under model
key-value group")
void testTokenMetricsUnderModelKeyValueGroup() {
TestMetricGroup root = new TestMetricGroup();
- setup.setMetricGroup(root);
- setup.recordTokenMetrics("gpt-4", 100, 50);
+ setup.recordTokenMetrics(root, "gpt-4", 100, 50);
TestMetricGroup modelGroup = (TestMetricGroup)
root.getSubGroup("model", "gpt-4");
assertEquals(100, modelGroup.counters.get("promptTokens").getCount());
@@ -200,10 +203,9 @@ class BaseChatModelSetupTokenMetricsTest {
@DisplayName("Value-based: different models have independent counters")
void testDifferentModelsHaveIndependentCounters() {
TestMetricGroup root = new TestMetricGroup();
- setup.setMetricGroup(root);
- setup.recordTokenMetrics("gpt-4", 100, 50);
- setup.recordTokenMetrics("gpt-3.5-turbo", 200, 80);
+ setup.recordTokenMetrics(root, "gpt-4", 100, 50);
+ setup.recordTokenMetrics(root, "gpt-3.5-turbo", 200, 80);
TestMetricGroup gpt4 = (TestMetricGroup) root.getSubGroup("model",
"gpt-4");
TestMetricGroup gpt35 = (TestMetricGroup) root.getSubGroup("model",
"gpt-3.5-turbo");
@@ -218,10 +220,9 @@ class BaseChatModelSetupTokenMetricsTest {
@DisplayName("Value-based: counters accumulate across multiple calls")
void testCountersAccumulate() {
TestMetricGroup root = new TestMetricGroup();
- setup.setMetricGroup(root);
- setup.recordTokenMetrics("gpt-4", 100, 50);
- setup.recordTokenMetrics("gpt-4", 150, 75);
+ setup.recordTokenMetrics(root, "gpt-4", 100, 50);
+ setup.recordTokenMetrics(root, "gpt-4", 150, 75);
TestMetricGroup modelGroup = (TestMetricGroup)
root.getSubGroup("model", "gpt-4");
assertEquals(250, modelGroup.counters.get("promptTokens").getCount());
diff --git
a/plan/src/main/java/org/apache/flink/agents/plan/actions/ChatModelAction.java
b/plan/src/main/java/org/apache/flink/agents/plan/actions/ChatModelAction.java
index df28c1d4..e8c663f3 100644
---
a/plan/src/main/java/org/apache/flink/agents/plan/actions/ChatModelAction.java
+++
b/plan/src/main/java/org/apache/flink/agents/plan/actions/ChatModelAction.java
@@ -182,7 +182,13 @@ public class ChatModelAction {
}
}
- static void recordChatTokenMetrics(BaseChatModelSetup chatModel,
ChatMessage response) {
+ static void recordChatTokenMetrics(
+ BaseChatModelSetup chatModel,
+ ChatMessage response,
+ @Nullable FlinkAgentsMetricGroup requestMetricGroup) {
+ if (requestMetricGroup == null) {
+ return;
+ }
Map<String, Object> extraArgs = response.getExtraArgs();
Object modelName = extraArgs.get("model_name");
Object promptTokens = extraArgs.get("promptTokens");
@@ -194,7 +200,8 @@ public class ChatModelAction {
long prompt = ((Number) promptTokens).longValue();
long completion = ((Number) completionTokens).longValue();
if (prompt > 0 && completion > 0) {
- chatModel.recordTokenMetrics(modelName.toString(), prompt,
completion);
+ chatModel.recordTokenMetrics(
+ requestMetricGroup, modelName.toString(), prompt,
completion);
}
}
}
@@ -322,6 +329,7 @@ public class ChatModelAction {
throws Exception {
BaseChatModelSetup chatModel =
(BaseChatModelSetup) ctx.getResource(model,
ResourceType.CHAT_MODEL);
+ FlinkAgentsMetricGroup requestMetricGroup = ctx.getActionMetricGroup();
boolean chatAsync =
ctx.getConfig().get(AgentExecutionOptions.CHAT_ASYNC);
@@ -372,7 +380,7 @@ public class ChatModelAction {
chatAsync
? ctx.durableExecuteAsync(callable)
: ctx.durableExecute(callable);
- recordChatTokenMetrics(chatModel, response);
+ recordChatTokenMetrics(chatModel, response,
requestMetricGroup);
// only generate structured output for final response.
if (outputSchema != null && response.getToolCalls().isEmpty())
{
response = generateStructuredOutput(response,
outputSchema);
diff --git
a/plan/src/test/java/org/apache/flink/agents/plan/actions/ChatModelActionRetryTest.java
b/plan/src/test/java/org/apache/flink/agents/plan/actions/ChatModelActionRetryTest.java
index 8c139505..f93a0836 100644
---
a/plan/src/test/java/org/apache/flink/agents/plan/actions/ChatModelActionRetryTest.java
+++
b/plan/src/test/java/org/apache/flink/agents/plan/actions/ChatModelActionRetryTest.java
@@ -127,6 +127,35 @@ class ChatModelActionRetryTest {
verify(mockActionMetricGroup, never()).getSubGroup(anyString(),
anyString());
}
+ @Test
+ void chatRecordsTokenMetricsWithRequestScopedMetricGroup() throws
Exception {
+ configureRetryStrategy(0, 0);
+ FlinkAgentsMetricGroup actionA = mock(FlinkAgentsMetricGroup.class);
+ FlinkAgentsMetricGroup actionB = mock(FlinkAgentsMetricGroup.class);
+ when(mockCtx.getActionMetricGroup()).thenReturn(actionA, actionB);
+
+ ChatMessage response =
+ new ChatMessage(
+ MessageRole.ASSISTANT,
+ "hello",
+ Map.of(
+ "model_name", "provider-model",
+ "promptTokens", 100L,
+ "completionTokens", 50L));
+ when(mockChatModel.chat(any(), any(), any())).thenReturn(response);
+
+ ChatModelAction.chat(
+ UUID.randomUUID(),
+ "test-model",
+ List.of(new ChatMessage(MessageRole.USER, "hi")),
+ Map.of(),
+ null,
+ mockCtx);
+
+ verify(mockChatModel).recordTokenMetrics(actionA, "provider-model",
100L, 50L);
+ verify(mockChatModel, never()).recordTokenMetrics(actionB,
"provider-model", 100L, 50L);
+ }
+
@Test
void chatRetriesWithExponentialBackoff() throws Exception {
// 1 second base interval; fail once then succeed -> wait 1s (1 * 2^0)
diff --git
a/plan/src/test/java/org/apache/flink/agents/plan/actions/ChatModelActionTest.java
b/plan/src/test/java/org/apache/flink/agents/plan/actions/ChatModelActionTest.java
index 85c263a6..46485b5a 100644
---
a/plan/src/test/java/org/apache/flink/agents/plan/actions/ChatModelActionTest.java
+++
b/plan/src/test/java/org/apache/flink/agents/plan/actions/ChatModelActionTest.java
@@ -20,17 +20,20 @@ package org.apache.flink.agents.plan.actions;
import org.apache.flink.agents.api.chat.messages.ChatMessage;
import org.apache.flink.agents.api.chat.messages.MessageRole;
import org.apache.flink.agents.api.chat.model.BaseChatModelSetup;
+import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup;
import org.junit.jupiter.api.Test;
import java.util.HashMap;
import java.util.Map;
import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.anyLong;
import static org.mockito.ArgumentMatchers.anyString;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.never;
import static org.mockito.Mockito.verify;
+import static org.mockito.Mockito.verifyNoInteractions;
/** Tests for {@link ChatModelAction}. */
class ChatModelActionTest {
@@ -42,71 +45,95 @@ class ChatModelActionTest {
@Test
void testRecordChatTokenMetricsRecordsWhenAllKeysPresent() {
BaseChatModelSetup setup = mock(BaseChatModelSetup.class);
+ FlinkAgentsMetricGroup requestMetricGroup =
mock(FlinkAgentsMetricGroup.class);
Map<String, Object> extraArgs = new HashMap<>();
extraArgs.put("model_name", "m");
extraArgs.put("promptTokens", 100L);
extraArgs.put("completionTokens", 50L);
- ChatModelAction.recordChatTokenMetrics(setup, responseWith(extraArgs));
+ ChatModelAction.recordChatTokenMetrics(setup, responseWith(extraArgs),
requestMetricGroup);
- verify(setup).recordTokenMetrics("m", 100L, 50L);
+ verify(setup).recordTokenMetrics(requestMetricGroup, "m", 100L, 50L);
}
@Test
void testRecordChatTokenMetricsHandlesIntegerTokenValues() {
BaseChatModelSetup setup = mock(BaseChatModelSetup.class);
+ FlinkAgentsMetricGroup requestMetricGroup =
mock(FlinkAgentsMetricGroup.class);
Map<String, Object> extraArgs = new HashMap<>();
extraArgs.put("model_name", "m");
extraArgs.put("promptTokens", 100);
extraArgs.put("completionTokens", 50);
- ChatModelAction.recordChatTokenMetrics(setup, responseWith(extraArgs));
+ ChatModelAction.recordChatTokenMetrics(setup, responseWith(extraArgs),
requestMetricGroup);
- verify(setup).recordTokenMetrics("m", 100L, 50L);
+ verify(setup).recordTokenMetrics(requestMetricGroup, "m", 100L, 50L);
+ }
+
+ @Test
+ void testRecordChatTokenMetricsSkipsWhenMetricGroupMissing() {
+ BaseChatModelSetup setup = mock(BaseChatModelSetup.class);
+ Map<String, Object> extraArgs = new HashMap<>();
+ extraArgs.put("model_name", "m");
+ extraArgs.put("promptTokens", 100L);
+ extraArgs.put("completionTokens", 50L);
+
+ ChatModelAction.recordChatTokenMetrics(setup, responseWith(extraArgs),
null);
+
+ verifyNoInteractions(setup);
}
@Test
void testRecordChatTokenMetricsSkipsWhenTokenValueNonNumeric() {
BaseChatModelSetup setup = mock(BaseChatModelSetup.class);
+ FlinkAgentsMetricGroup requestMetricGroup =
mock(FlinkAgentsMetricGroup.class);
Map<String, Object> extraArgs = new HashMap<>();
extraArgs.put("model_name", "m");
extraArgs.put("promptTokens", "100");
extraArgs.put("completionTokens", 50L);
- ChatModelAction.recordChatTokenMetrics(setup, responseWith(extraArgs));
+ ChatModelAction.recordChatTokenMetrics(setup, responseWith(extraArgs),
requestMetricGroup);
- verify(setup, never()).recordTokenMetrics(anyString(), anyLong(),
anyLong());
+ verify(setup, never())
+ .recordTokenMetrics(
+ any(FlinkAgentsMetricGroup.class), anyString(),
anyLong(), anyLong());
}
@Test
void testRecordChatTokenMetricsSkipsWhenKeyMissing() {
BaseChatModelSetup setup = mock(BaseChatModelSetup.class);
+ FlinkAgentsMetricGroup requestMetricGroup =
mock(FlinkAgentsMetricGroup.class);
Map<String, Object> extraArgs = new HashMap<>();
extraArgs.put("model_name", "m");
extraArgs.put("completionTokens", 50L);
- ChatModelAction.recordChatTokenMetrics(setup, responseWith(extraArgs));
+ ChatModelAction.recordChatTokenMetrics(setup, responseWith(extraArgs),
requestMetricGroup);
- verify(setup, never()).recordTokenMetrics(anyString(), anyLong(),
anyLong());
+ verify(setup, never())
+ .recordTokenMetrics(
+ any(FlinkAgentsMetricGroup.class), anyString(),
anyLong(), anyLong());
}
@Test
void testRecordChatTokenMetricsSkipsZeroTokensOrEmptyModel() {
BaseChatModelSetup setup = mock(BaseChatModelSetup.class);
+ FlinkAgentsMetricGroup requestMetricGroup =
mock(FlinkAgentsMetricGroup.class);
Map<String, Object> zeroPrompt = new HashMap<>();
zeroPrompt.put("model_name", "m");
zeroPrompt.put("promptTokens", 0L);
zeroPrompt.put("completionTokens", 50L);
- ChatModelAction.recordChatTokenMetrics(setup,
responseWith(zeroPrompt));
+ ChatModelAction.recordChatTokenMetrics(setup,
responseWith(zeroPrompt), requestMetricGroup);
Map<String, Object> emptyModel = new HashMap<>();
emptyModel.put("model_name", "");
emptyModel.put("promptTokens", 100L);
emptyModel.put("completionTokens", 50L);
- ChatModelAction.recordChatTokenMetrics(setup,
responseWith(emptyModel));
+ ChatModelAction.recordChatTokenMetrics(setup,
responseWith(emptyModel), requestMetricGroup);
- verify(setup, never()).recordTokenMetrics(anyString(), anyLong(),
anyLong());
+ verify(setup, never())
+ .recordTokenMetrics(
+ any(FlinkAgentsMetricGroup.class), anyString(),
anyLong(), anyLong());
}
@Test
diff --git a/python/flink_agents/api/chat_models/chat_model.py
b/python/flink_agents/api/chat_models/chat_model.py
index d40f238e..c3a9d786 100644
--- a/python/flink_agents/api/chat_models/chat_model.py
+++ b/python/flink_agents/api/chat_models/chat_model.py
@@ -29,6 +29,7 @@ from flink_agents.api.chat_message import (
MessageRole,
find_first_system_message,
)
+from flink_agents.api.metric_group import MetricGroup
from flink_agents.api.prompts.prompt import Prompt
from flink_agents.api.resource import Resource, ResourceType
from flink_agents.api.skills import BASH_TOOL, LOAD_SKILL_TOOL
@@ -415,7 +416,11 @@ class BaseChatModelSetup(Resource):
return connection.chat(messages, tools=self._get_tools(),
**merged_kwargs)
def _record_token_metrics(
- self, model_name: str, prompt_tokens: int, completion_tokens: int
+ self,
+ model_name: str,
+ prompt_tokens: int,
+ completion_tokens: int,
+ metric_group: MetricGroup | None,
) -> None:
"""Record token usage metrics for the given model.
@@ -427,8 +432,10 @@ class BaseChatModelSetup(Resource):
The number of prompt tokens
completion_tokens : int
The number of completion tokens
+ metric_group : MetricGroup | None
+ The metric group captured when the request was initiated. If None,
token
+ metrics are skipped.
"""
- metric_group = self.metric_group
if metric_group is None:
return
diff --git a/python/flink_agents/api/chat_models/tests/test_token_metrics.py
b/python/flink_agents/api/chat_models/tests/test_token_metrics.py
index e8195151..15455836 100644
--- a/python/flink_agents/api/chat_models/tests/test_token_metrics.py
+++ b/python/flink_agents/api/chat_models/tests/test_token_metrics.py
@@ -46,10 +46,16 @@ class TestChatModelSetup(BaseChatModelSetup):
return ChatMessage(role=MessageRole.ASSISTANT, content="Test response")
def test_record_token_metrics(
- self, model_name: str, prompt_tokens: int, completion_tokens: int
+ self,
+ model_name: str,
+ prompt_tokens: int,
+ completion_tokens: int,
+ metric_group: MetricGroup | None,
) -> None:
"""Expose protected method for testing."""
- self._record_token_metrics(model_name, prompt_tokens,
completion_tokens)
+ self._record_token_metrics(
+ model_name, prompt_tokens, completion_tokens, metric_group
+ )
class _MockCounter(Counter):
@@ -104,39 +110,60 @@ class TestBaseChatModelTokenMetrics:
chat_model = TestChatModelSetup(connection="mock", model="mock-model")
mock_metric_group = _MockMetricGroup()
- # Set the metric group
- chat_model.set_metric_group(mock_metric_group)
-
# Record token metrics
- chat_model.test_record_token_metrics("gpt-4", 100, 50)
+ chat_model.test_record_token_metrics("gpt-4", 100, 50,
mock_metric_group)
# Verify the metrics were recorded
model_group = mock_metric_group.get_sub_group("model", "gpt-4")
assert model_group.get_counter("promptTokens").get_count() == 100
assert model_group.get_counter("completionTokens").get_count() == 50
- def test_record_token_metrics_without_metric_group(self) -> None:
- """Test token metrics are not recorded when metric group is null."""
+ def test_record_token_metrics_skips_when_metric_group_missing(self) ->
None:
+ """Test token metrics are not recorded when metric group is missing."""
chat_model = TestChatModelSetup(connection="mock", model="mock-model")
+ bound_metric_group = _MockMetricGroup()
+ chat_model.set_metric_group(bound_metric_group)
+
+ chat_model.test_record_token_metrics("gpt-4", 100, 50, None)
+
+ model_group = bound_metric_group.get_sub_group("model", "gpt-4")
+ assert model_group.get_counter("promptTokens").get_count() == 0
+ assert model_group.get_counter("completionTokens").get_count() == 0
+
+ def test_record_token_metrics_with_request_scoped_metric_group(self) ->
None:
+ """Token metrics use the metric group captured when the request
started."""
+ chat_model = TestChatModelSetup(connection="mock", model="mock-model")
+ action_a_metric_group = _MockMetricGroup()
+ action_b_metric_group = _MockMetricGroup()
+
+ chat_model.set_metric_group(action_a_metric_group)
+ request_metric_group = chat_model.metric_group
- # Do not set metric group (should be None by default)
- # Record token metrics - should not throw
- chat_model.test_record_token_metrics("gpt-4", 100, 50)
- # No exception should be raised
+ chat_model.set_metric_group(action_b_metric_group)
+ chat_model.test_record_token_metrics(
+ "gpt-4", 100, 50, metric_group=request_metric_group
+ )
+
+ action_a_model_group = action_a_metric_group.get_sub_group("model",
"gpt-4")
+ assert action_a_model_group.get_counter("promptTokens").get_count() ==
100
+ assert
action_a_model_group.get_counter("completionTokens").get_count() == 50
+
+ action_b_model_group = action_b_metric_group.get_sub_group("model",
"gpt-4")
+ assert action_b_model_group.get_counter("promptTokens").get_count() == 0
+ assert
action_b_model_group.get_counter("completionTokens").get_count() == 0
def test_token_metrics_hierarchy(self) -> None:
"""Test token metrics hierarchy: actionMetricGroup -> modelName ->
counters."""
chat_model = TestChatModelSetup(connection="mock", model="mock-model")
mock_metric_group = _MockMetricGroup()
- # Set the metric group
- chat_model.set_metric_group(mock_metric_group)
-
# Record for gpt-4
- chat_model.test_record_token_metrics("gpt-4", 100, 50)
+ chat_model.test_record_token_metrics("gpt-4", 100, 50,
mock_metric_group)
# Record for gpt-3.5-turbo
- chat_model.test_record_token_metrics("gpt-3.5-turbo", 200, 100)
+ chat_model.test_record_token_metrics(
+ "gpt-3.5-turbo", 200, 100, mock_metric_group
+ )
# Verify each model has its own counters
gpt4_group = mock_metric_group.get_sub_group("model", "gpt-4")
@@ -152,12 +179,9 @@ class TestBaseChatModelTokenMetrics:
chat_model = TestChatModelSetup(connection="mock", model="mock-model")
mock_metric_group = _MockMetricGroup()
- # Set the metric group
- chat_model.set_metric_group(mock_metric_group)
-
# Record multiple times for the same model
- chat_model.test_record_token_metrics("gpt-4", 100, 50)
- chat_model.test_record_token_metrics("gpt-4", 150, 75)
+ chat_model.test_record_token_metrics("gpt-4", 100, 50,
mock_metric_group)
+ chat_model.test_record_token_metrics("gpt-4", 150, 75,
mock_metric_group)
# Verify the metrics accumulated
model_group = mock_metric_group.get_sub_group("model", "gpt-4")
diff --git a/python/flink_agents/plan/actions/chat_model_action.py
b/python/flink_agents/plan/actions/chat_model_action.py
index de67881b..1065c785 100644
--- a/python/flink_agents/plan/actions/chat_model_action.py
+++ b/python/flink_agents/plan/actions/chat_model_action.py
@@ -290,6 +290,7 @@ async def chat(
chat_model = cast(
"BaseChatModelSetup", ctx.get_resource(model, ResourceType.CHAT_MODEL)
)
+ request_metric_group = ctx.action_metric_group
chat_async = ctx.config.get(AgentExecutionOptions.CHAT_ASYNC)
@@ -326,7 +327,8 @@ async def chat(
)
if (
- response.extra_args.get("model_name")
+ request_metric_group is not None
+ and response.extra_args.get("model_name")
and response.extra_args.get("promptTokens")
and response.extra_args.get("completionTokens")
):
@@ -334,6 +336,7 @@ async def chat(
response.extra_args["model_name"],
response.extra_args["promptTokens"],
response.extra_args["completionTokens"],
+ request_metric_group,
)
if output_schema is not None and len(response.tool_calls) == 0:
response = _generate_structured_output(response, output_schema)