Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -109,18 +109,21 @@ public void open() throws Exception {
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);
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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. */
Expand Down Expand Up @@ -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");
Expand All @@ -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);
Expand All @@ -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");
Expand Down Expand Up @@ -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());
Expand All @@ -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");
Expand All @@ -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());
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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<String, Object> extraArgs = response.getExtraArgs();
Object modelName = extraArgs.get("model_name");
Object promptTokens = extraArgs.get("promptTokens");
Expand All @@ -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);
}
}
}
Expand Down Expand Up @@ -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);

Expand Down Expand Up @@ -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);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand All @@ -42,71 +45,95 @@ private static ChatMessage responseWith(Map<String, Object> extraArgs) {
@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
Expand Down
11 changes: 9 additions & 2 deletions python/flink_agents/api/chat_models/chat_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.

Expand All @@ -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:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

One parity question: Python _record_token_metrics(..., metric_group=None) is a no-op, while Java recordTokenMetrics(null, ...) fails via checkNotNull. Should we align the lower-level helper contract?

return

Expand Down
Loading
Loading