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)

Reply via email to