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 34b2b2994..cdfc63143 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 @@ -109,18 +109,21 @@ public void open() throws Exception { public abstract Map 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 8e47105f0..abc7f2c0e 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.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 @@ void setUp() { @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 @@ void testRecordTokenMetricsWithMetricGroup() { } @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 @@ void testTokenMetricsHierarchy() { 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 @@ public Histogram getHistogram(String name, int windowSize) { @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 @@ void testTokenMetricsUnderModelKeyValueGroup() { @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 @@ void testDifferentModelsHaveIndependentCounters() { @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 df28c1d41..e8c663f34 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 @@ private static void recordRetryMetrics( } } - static void recordChatTokenMetrics(BaseChatModelSetup chatModel, ChatMessage response) { + static void recordChatTokenMetrics( + BaseChatModelSetup chatModel, + ChatMessage response, + @Nullable FlinkAgentsMetricGroup requestMetricGroup) { + if (requestMetricGroup == null) { + return; + } Map extraArgs = response.getExtraArgs(); Object modelName = extraArgs.get("model_name"); Object promptTokens = extraArgs.get("promptTokens"); @@ -194,7 +200,8 @@ static void recordChatTokenMetrics(BaseChatModelSetup chatModel, ChatMessage res 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 static void chat( 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 ChatMessage call() throws Exception { 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 8c1395059..f93a0836f 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 @@ void chatSucceedsWithoutRetry_retryCountIsZero() throws Exception { 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 85c263a66..46485b5a1 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 @@ 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 @@ private static ChatMessage responseWith(Map extraArgs) { @Test void testRecordChatTokenMetricsRecordsWhenAllKeysPresent() { BaseChatModelSetup setup = mock(BaseChatModelSetup.class); + FlinkAgentsMetricGroup requestMetricGroup = mock(FlinkAgentsMetricGroup.class); Map 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 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 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 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 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 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 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 ac4a814aa..84a8d3462 100644 --- a/python/flink_agents/api/chat_models/chat_model.py +++ b/python/flink_agents/api/chat_models/chat_model.py @@ -27,6 +27,7 @@ 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 @@ -261,7 +262,11 @@ def chat( ) 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. @@ -273,8 +278,10 @@ def _record_token_metrics( 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 e81951511..15455836c 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 @@ def chat(self, messages: Sequence[ChatMessage], **kwargs: Any) -> ChatMessage: 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 @@ def test_record_token_metrics_with_metric_group(self) -> None: 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 @@ def test_token_metrics_accumulation(self) -> None: 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 e4572056e..64057f536 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)