From 7bc3a8e908864b64359f6447e1aa1674fd866ca5 Mon Sep 17 00:00:00 2001 From: Goutam Adwant <8672451+goutamadwant@users.noreply.github.com> Date: Wed, 22 Jul 2026 13:14:02 +0200 Subject: [PATCH] [api][runtime][python] Propagate metric groups to cross-language resources (#860) --- .../python/PythonChatModelConnection.java | 12 ++ .../model/python/PythonChatModelSetup.java | 12 ++ .../PythonEmbeddingModelConnection.java | 12 ++ .../python/PythonEmbeddingModelSetup.java | 12 ++ .../python/PythonResourceAdapter.java | 9 ++ .../python/PythonResourceWrapper.java | 24 +++ .../python/PythonVectorStore.java | 12 ++ .../python/PythonChatModelSetupTest.java | 10 ++ .../plan/resource/python/PythonMCPPrompt.java | 12 ++ .../plan/resource/python/PythonMCPServer.java | 12 ++ .../plan/resource/python/PythonMCPTool.java | 12 ++ .../runtime/java/java_chat_model.py | 13 ++ .../runtime/java/java_embedding_model.py | 13 ++ .../runtime/java/java_resource_wrapper.py | 16 ++ .../runtime/java/java_vector_store.py | 8 + .../flink_agents/runtime/python_java_utils.py | 10 ++ .../tests/test_cross_language_metric_group.py | 149 ++++++++++++++++++ .../utils/PythonResourceAdapterImpl.java | 8 + .../utils/PythonResourceAdapterImplTest.java | 12 ++ 19 files changed, 368 insertions(+) create mode 100644 python/flink_agents/runtime/tests/test_cross_language_metric_group.py diff --git a/api/src/main/java/org/apache/flink/agents/api/chat/model/python/PythonChatModelConnection.java b/api/src/main/java/org/apache/flink/agents/api/chat/model/python/PythonChatModelConnection.java index 2a362f7a7..8ce3662e3 100644 --- a/api/src/main/java/org/apache/flink/agents/api/chat/model/python/PythonChatModelConnection.java +++ b/api/src/main/java/org/apache/flink/agents/api/chat/model/python/PythonChatModelConnection.java @@ -19,6 +19,7 @@ import org.apache.flink.agents.api.chat.messages.ChatMessage; import org.apache.flink.agents.api.chat.model.BaseChatModelConnection; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.apache.flink.agents.api.resource.python.PythonResourceAdapter; @@ -65,6 +66,17 @@ public Object getPythonResource() { return chatModel; } + @Override + public PythonResourceAdapter getPythonResourceAdapter() { + return adapter; + } + + @Override + public void setMetricGroup(FlinkAgentsMetricGroup metricGroup) { + super.setMetricGroup(metricGroup); + setPythonResourceMetricGroup(metricGroup); + } + @Override public ChatMessage chat( List messages, List tools, Map modelParams) { diff --git a/api/src/main/java/org/apache/flink/agents/api/chat/model/python/PythonChatModelSetup.java b/api/src/main/java/org/apache/flink/agents/api/chat/model/python/PythonChatModelSetup.java index ad2117d36..846105afb 100644 --- a/api/src/main/java/org/apache/flink/agents/api/chat/model/python/PythonChatModelSetup.java +++ b/api/src/main/java/org/apache/flink/agents/api/chat/model/python/PythonChatModelSetup.java @@ -19,6 +19,7 @@ import org.apache.flink.agents.api.chat.messages.ChatMessage; import org.apache.flink.agents.api.chat.model.BaseChatModelSetup; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.apache.flink.agents.api.resource.python.PythonResourceAdapter; @@ -89,6 +90,17 @@ public Object getPythonResource() { return chatModelSetup; } + @Override + public PythonResourceAdapter getPythonResourceAdapter() { + return adapter; + } + + @Override + public void setMetricGroup(FlinkAgentsMetricGroup metricGroup) { + super.setMetricGroup(metricGroup); + setPythonResourceMetricGroup(metricGroup); + } + @Override public Map getParameters() { return Map.of(); diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnection.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnection.java index 974e362a1..785ed03a8 100644 --- a/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnection.java +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelConnection.java @@ -19,6 +19,7 @@ import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelConnection; import org.apache.flink.agents.api.embedding.model.EmbeddingModelUtils; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.apache.flink.agents.api.resource.python.PythonResourceAdapter; @@ -123,6 +124,17 @@ public Object getPythonResource() { return embeddingModel; } + @Override + public PythonResourceAdapter getPythonResourceAdapter() { + return adapter; + } + + @Override + public void setMetricGroup(FlinkAgentsMetricGroup metricGroup) { + super.setMetricGroup(metricGroup); + setPythonResourceMetricGroup(metricGroup); + } + @Override public void close() throws Exception { this.embeddingModel.close(); diff --git a/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetup.java b/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetup.java index f0b9eca4b..c0efd5c35 100644 --- a/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetup.java +++ b/api/src/main/java/org/apache/flink/agents/api/embedding/model/python/PythonEmbeddingModelSetup.java @@ -19,6 +19,7 @@ import org.apache.flink.agents.api.embedding.model.BaseEmbeddingModelSetup; import org.apache.flink.agents.api.embedding.model.EmbeddingModelUtils; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.apache.flink.agents.api.resource.python.PythonResourceAdapter; @@ -132,4 +133,15 @@ public Map getParameters() { public Object getPythonResource() { return embeddingModelSetup; } + + @Override + public PythonResourceAdapter getPythonResourceAdapter() { + return adapter; + } + + @Override + public void setMetricGroup(FlinkAgentsMetricGroup metricGroup) { + super.setMetricGroup(metricGroup); + setPythonResourceMetricGroup(metricGroup); + } } diff --git a/api/src/main/java/org/apache/flink/agents/api/resource/python/PythonResourceAdapter.java b/api/src/main/java/org/apache/flink/agents/api/resource/python/PythonResourceAdapter.java index 51a135490..9c4ad0bb6 100644 --- a/api/src/main/java/org/apache/flink/agents/api/resource/python/PythonResourceAdapter.java +++ b/api/src/main/java/org/apache/flink/agents/api/resource/python/PythonResourceAdapter.java @@ -19,6 +19,7 @@ package org.apache.flink.agents.api.resource.python; import org.apache.flink.agents.api.chat.messages.ChatMessage; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.tools.Tool; import org.apache.flink.agents.api.vectorstores.Document; import org.apache.flink.agents.api.vectorstores.VectorStoreQuery; @@ -120,6 +121,14 @@ public interface PythonResourceAdapter { */ Object callMethod(Object obj, String methodName, Map kwargs); + /** + * Binds a Java metric group to a Python resource. + * + * @param pythonResource the Python resource object + * @param metricGroup the Java metric group to expose through Python's metric group API + */ + default void setMetricGroup(Object pythonResource, FlinkAgentsMetricGroup metricGroup) {} + /** * Invokes a method with the specified name and arguments. * diff --git a/api/src/main/java/org/apache/flink/agents/api/resource/python/PythonResourceWrapper.java b/api/src/main/java/org/apache/flink/agents/api/resource/python/PythonResourceWrapper.java index c69cf59b8..7bd3a343d 100644 --- a/api/src/main/java/org/apache/flink/agents/api/resource/python/PythonResourceWrapper.java +++ b/api/src/main/java/org/apache/flink/agents/api/resource/python/PythonResourceWrapper.java @@ -17,6 +17,8 @@ */ package org.apache.flink.agents.api.resource.python; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; + /** * Wrapper interface for Python resource objects. This interface provides a unified way to access * the underlying Python resource from Java objects that encapsulate Python functionality. @@ -29,4 +31,26 @@ public interface PythonResourceWrapper { * @return the wrapped Python resource object */ Object getPythonResource(); + + /** + * Returns the adapter that owns the wrapped Python resource. + * + * @return the Python resource adapter, or null if metric forwarding is unsupported + */ + default PythonResourceAdapter getPythonResourceAdapter() { + return null; + } + + /** + * Binds the current Java metric group to the wrapped Python resource. + * + * @param metricGroup the metric group to bind + */ + default void setPythonResourceMetricGroup(FlinkAgentsMetricGroup metricGroup) { + PythonResourceAdapter adapter = getPythonResourceAdapter(); + Object pythonResource = getPythonResource(); + if (adapter != null && pythonResource != null) { + adapter.setMetricGroup(pythonResource, metricGroup); + } + } } diff --git a/api/src/main/java/org/apache/flink/agents/api/vectorstores/python/PythonVectorStore.java b/api/src/main/java/org/apache/flink/agents/api/vectorstores/python/PythonVectorStore.java index 69025cc11..21bb18931 100644 --- a/api/src/main/java/org/apache/flink/agents/api/vectorstores/python/PythonVectorStore.java +++ b/api/src/main/java/org/apache/flink/agents/api/vectorstores/python/PythonVectorStore.java @@ -18,6 +18,7 @@ package org.apache.flink.agents.api.vectorstores.python; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.apache.flink.agents.api.resource.python.PythonResourceAdapter; @@ -233,4 +234,15 @@ public void updateEmbedding( public Object getPythonResource() { return vectorStore; } + + @Override + public PythonResourceAdapter getPythonResourceAdapter() { + return adapter; + } + + @Override + public void setMetricGroup(FlinkAgentsMetricGroup metricGroup) { + super.setMetricGroup(metricGroup); + setPythonResourceMetricGroup(metricGroup); + } } diff --git a/api/src/test/java/org/apache/flink/agents/api/chat/model/python/PythonChatModelSetupTest.java b/api/src/test/java/org/apache/flink/agents/api/chat/model/python/PythonChatModelSetupTest.java index 4327fb336..42463b083 100644 --- a/api/src/test/java/org/apache/flink/agents/api/chat/model/python/PythonChatModelSetupTest.java +++ b/api/src/test/java/org/apache/flink/agents/api/chat/model/python/PythonChatModelSetupTest.java @@ -18,6 +18,7 @@ package org.apache.flink.agents.api.chat.model.python; import org.apache.flink.agents.api.chat.messages.ChatMessage; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; import org.apache.flink.agents.api.resource.python.PythonResourceAdapter; @@ -161,4 +162,13 @@ void testImplementsPythonResourceWrapper() { .isInstanceOf( org.apache.flink.agents.api.resource.python.PythonResourceWrapper.class); } + + @Test + void testSetMetricGroupPropagatesToPythonResource() { + FlinkAgentsMetricGroup metricGroup = mock(FlinkAgentsMetricGroup.class); + + pythonChatModelSetup.setMetricGroup(metricGroup); + + verify(mockAdapter).setMetricGroup(mockChatModelSetup, metricGroup); + } } diff --git a/plan/src/main/java/org/apache/flink/agents/plan/resource/python/PythonMCPPrompt.java b/plan/src/main/java/org/apache/flink/agents/plan/resource/python/PythonMCPPrompt.java index 625a89fd7..cb4a39351 100644 --- a/plan/src/main/java/org/apache/flink/agents/plan/resource/python/PythonMCPPrompt.java +++ b/plan/src/main/java/org/apache/flink/agents/plan/resource/python/PythonMCPPrompt.java @@ -19,6 +19,7 @@ import org.apache.flink.agents.api.chat.messages.ChatMessage; import org.apache.flink.agents.api.chat.messages.MessageRole; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.prompt.Prompt; import org.apache.flink.agents.api.resource.python.PythonResourceAdapter; import org.apache.flink.agents.api.resource.python.PythonResourceWrapper; @@ -47,6 +48,17 @@ public Object getPythonResource() { return prompt; } + @Override + public PythonResourceAdapter getPythonResourceAdapter() { + return adapter; + } + + @Override + public void setMetricGroup(FlinkAgentsMetricGroup metricGroup) { + super.setMetricGroup(metricGroup); + setPythonResourceMetricGroup(metricGroup); + } + public String getName() { if (name == null) { name = prompt.getAttr("name").toString(); diff --git a/plan/src/main/java/org/apache/flink/agents/plan/resource/python/PythonMCPServer.java b/plan/src/main/java/org/apache/flink/agents/plan/resource/python/PythonMCPServer.java index a6268347d..47009c05a 100644 --- a/plan/src/main/java/org/apache/flink/agents/plan/resource/python/PythonMCPServer.java +++ b/plan/src/main/java/org/apache/flink/agents/plan/resource/python/PythonMCPServer.java @@ -17,6 +17,7 @@ */ package org.apache.flink.agents.plan.resource.python; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.resource.Resource; import org.apache.flink.agents.api.resource.ResourceContext; import org.apache.flink.agents.api.resource.ResourceDescriptor; @@ -84,6 +85,17 @@ public Object getPythonResource() { return server; } + @Override + public PythonResourceAdapter getPythonResourceAdapter() { + return adapter; + } + + @Override + public void setMetricGroup(FlinkAgentsMetricGroup metricGroup) { + super.setMetricGroup(metricGroup); + setPythonResourceMetricGroup(metricGroup); + } + @Override public ResourceType getResourceType() { return ResourceType.MCP_SERVER; diff --git a/plan/src/main/java/org/apache/flink/agents/plan/resource/python/PythonMCPTool.java b/plan/src/main/java/org/apache/flink/agents/plan/resource/python/PythonMCPTool.java index 89a5435d7..9cbd0c585 100644 --- a/plan/src/main/java/org/apache/flink/agents/plan/resource/python/PythonMCPTool.java +++ b/plan/src/main/java/org/apache/flink/agents/plan/resource/python/PythonMCPTool.java @@ -17,6 +17,7 @@ */ package org.apache.flink.agents.plan.resource.python; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.resource.python.PythonResourceAdapter; import org.apache.flink.agents.api.resource.python.PythonResourceWrapper; import org.apache.flink.agents.api.tools.Tool; @@ -75,6 +76,17 @@ public Object getPythonResource() { return tool; } + @Override + public PythonResourceAdapter getPythonResourceAdapter() { + return adapter; + } + + @Override + public void setMetricGroup(FlinkAgentsMetricGroup metricGroup) { + super.setMetricGroup(metricGroup); + setPythonResourceMetricGroup(metricGroup); + } + @Override public ToolType getToolType() { return ToolType.MCP; diff --git a/python/flink_agents/runtime/java/java_chat_model.py b/python/flink_agents/runtime/java/java_chat_model.py index 10bc169c1..3f385a9ad 100644 --- a/python/flink_agents/runtime/java/java_chat_model.py +++ b/python/flink_agents/runtime/java/java_chat_model.py @@ -26,6 +26,9 @@ ) from flink_agents.api.resource import ResourceType from flink_agents.api.tools.tool import Tool +from flink_agents.runtime.java.java_resource_wrapper import ( + set_java_resource_metric_group, +) class JavaChatModelConnectionImpl(JavaChatModelConnection): @@ -51,6 +54,11 @@ def __init__(self, j_resource: Any, j_resource_adapter: Any, **kwargs: Any) -> N self._j_resource = j_resource self._j_resource_adapter = j_resource_adapter + @override + def set_metric_group(self, metric_group: Any) -> None: + super().set_metric_group(metric_group) + set_java_resource_metric_group(self._j_resource, metric_group) + @override def chat( self, @@ -114,6 +122,11 @@ def __init__(self, j_resource: Any, j_resource_adapter: Any, **kwargs: Any) -> N self._j_resource = j_resource self._j_resource_adapter = j_resource_adapter + @override + def set_metric_group(self, metric_group: Any) -> None: + super().set_metric_group(metric_group) + set_java_resource_metric_group(self._j_resource, metric_group) + @property @override def model_kwargs(self) -> Dict[str, Any]: diff --git a/python/flink_agents/runtime/java/java_embedding_model.py b/python/flink_agents/runtime/java/java_embedding_model.py index 2cb15b819..b2ea48724 100644 --- a/python/flink_agents/runtime/java/java_embedding_model.py +++ b/python/flink_agents/runtime/java/java_embedding_model.py @@ -23,6 +23,9 @@ JavaEmbeddingModelConnection, JavaEmbeddingModelSetup, ) +from flink_agents.runtime.java.java_resource_wrapper import ( + set_java_resource_metric_group, +) class JavaEmbeddingModelConnectionImpl(JavaEmbeddingModelConnection): @@ -48,6 +51,11 @@ def __init__(self, j_resource: Any, j_resource_adapter: Any, **kwargs: Any) -> N self._j_resource = j_resource self._j_resource_adapter = j_resource_adapter + @override + def set_metric_group(self, metric_group: Any) -> None: + super().set_metric_group(metric_group) + set_java_resource_metric_group(self._j_resource, metric_group) + def embed( self, text: str | Sequence[str], **kwargs: Any ) -> list[float] | list[list[float]]: @@ -92,6 +100,11 @@ def __init__(self, j_resource: Any, j_resource_adapter: Any, **kwargs: Any) -> N self._j_resource = j_resource self._j_resource_adapter = j_resource_adapter + @override + def set_metric_group(self, metric_group: Any) -> None: + super().set_metric_group(metric_group) + set_java_resource_metric_group(self._j_resource, metric_group) + @property def model_kwargs(self) -> Dict[str, Any]: """Return embedding model settings. diff --git a/python/flink_agents/runtime/java/java_resource_wrapper.py b/python/flink_agents/runtime/java/java_resource_wrapper.py index 886e4c84c..87094b4a6 100644 --- a/python/flink_agents/runtime/java/java_resource_wrapper.py +++ b/python/flink_agents/runtime/java/java_resource_wrapper.py @@ -27,6 +27,22 @@ from flink_agents.api.tools.tool import Tool, ToolMetadata, ToolType +def set_java_resource_metric_group(j_resource: Any, metric_group: Any) -> None: + """Bind the underlying Java metric group to a wrapped Java resource.""" + if j_resource is None: + return + from flink_agents.runtime.flink_metric_group import FlinkMetricGroup + + if metric_group is None: + j_metric_group = None + elif isinstance(metric_group, FlinkMetricGroup): + j_metric_group = metric_group._j_metric_group + else: + msg = "Java resource metric groups must be FlinkMetricGroup or None." + raise TypeError(msg) + j_resource.setMetricGroup(j_metric_group) + + class JavaTool(Tool): """Java Tool that carries tool metadata and can be recognized by PythonChatModel.""" diff --git a/python/flink_agents/runtime/java/java_vector_store.py b/python/flink_agents/runtime/java/java_vector_store.py index 301d7a539..8c3399cad 100644 --- a/python/flink_agents/runtime/java/java_vector_store.py +++ b/python/flink_agents/runtime/java/java_vector_store.py @@ -27,6 +27,9 @@ Document, _maybe_cast_to_list, ) +from flink_agents.runtime.java.java_resource_wrapper import ( + set_java_resource_metric_group, +) from flink_agents.runtime.python_java_utils import from_java_document @@ -60,6 +63,11 @@ def __init__(self, j_resource: Any, j_resource_adapter: Any, **kwargs: Any) -> N self._j_resource = j_resource self._j_resource_adapter = j_resource_adapter + @override + def set_metric_group(self, metric_group: Any) -> None: + super().set_metric_group(metric_group) + set_java_resource_metric_group(self._j_resource, metric_group) + @property @override def store_kwargs(self) -> Dict[str, Any]: diff --git a/python/flink_agents/runtime/python_java_utils.py b/python/flink_agents/runtime/python_java_utils.py index 51f08f810..66b9092f7 100644 --- a/python/flink_agents/runtime/python_java_utils.py +++ b/python/flink_agents/runtime/python_java_utils.py @@ -400,3 +400,13 @@ def call_method(obj: Any, method_name: str, kwargs: Dict[str, Any]) -> Any: method = getattr(obj, method_name) return method(**kwargs) + + +def set_metric_group(obj: Resource, j_metric_group: Any) -> None: + """Bind a Java metric group to a Python resource.""" + from flink_agents.runtime.flink_metric_group import FlinkMetricGroup + + metric_group = ( + FlinkMetricGroup(j_metric_group) if j_metric_group is not None else None + ) + obj.set_metric_group(metric_group) diff --git a/python/flink_agents/runtime/tests/test_cross_language_metric_group.py b/python/flink_agents/runtime/tests/test_cross_language_metric_group.py new file mode 100644 index 000000000..ffb4faad3 --- /dev/null +++ b/python/flink_agents/runtime/tests/test_cross_language_metric_group.py @@ -0,0 +1,149 @@ +################################################################################ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +################################################################################# +from typing import Any + +import pytest + +from flink_agents.api.metric_group import MetricGroup +from flink_agents.runtime.flink_metric_group import FlinkMetricGroup +from flink_agents.runtime.java.java_chat_model import ( + JavaChatModelConnectionImpl, + JavaChatModelSetupImpl, +) +from flink_agents.runtime.java.java_embedding_model import ( + JavaEmbeddingModelConnectionImpl, + JavaEmbeddingModelSetupImpl, +) +from flink_agents.runtime.java.java_resource_wrapper import ( + set_java_resource_metric_group, +) +from flink_agents.runtime.python_java_utils import set_metric_group + + +class _JavaResource: + def __init__(self) -> None: + self.metric_group: Any = None + + def setMetricGroup(self, metric_group: Any) -> None: + self.metric_group = metric_group + + +class _PythonResource: + def __init__(self) -> None: + self.metric_group: Any = None + + def set_metric_group(self, metric_group: Any) -> None: + self.metric_group = metric_group + + +class _CustomMetricGroup(MetricGroup): + def get_sub_group( + self, name: str, value: str | None = None + ) -> "_CustomMetricGroup": + return self + + def get_counter(self, name: str) -> Any: + raise NotImplementedError + + def get_meter(self, name: str) -> Any: + raise NotImplementedError + + def get_histogram(self, name: str, window_size: int = 100) -> Any: + raise NotImplementedError + + def get_gauge(self, name: str) -> Any: + raise NotImplementedError + + +class _JavaMetricGroup: + def __init__(self) -> None: + self.j_metric_group = object() + + +@pytest.mark.parametrize( + "resource", + [ + JavaChatModelConnectionImpl( + j_resource=_JavaResource(), j_resource_adapter=None + ), + JavaChatModelSetupImpl( + j_resource=_JavaResource(), + j_resource_adapter=None, + connection="connection", + model="model", + ), + JavaEmbeddingModelConnectionImpl( + j_resource=_JavaResource(), j_resource_adapter=None + ), + JavaEmbeddingModelSetupImpl( + j_resource=_JavaResource(), + j_resource_adapter=None, + connection="connection", + model="model", + ), + ], +) +def test_java_resource_wrappers_forward_metric_group(resource): + java_metric_group = _JavaMetricGroup() + metric_group = FlinkMetricGroup(java_metric_group.j_metric_group) + + resource.set_metric_group(metric_group) + + assert resource.metric_group is metric_group + assert resource._j_resource.metric_group is java_metric_group.j_metric_group + + +def test_set_java_resource_metric_group_unwraps_flink_metric_group(): + java_resource = _JavaResource() + java_metric_group = _JavaMetricGroup() + metric_group = FlinkMetricGroup(java_metric_group.j_metric_group) + + set_java_resource_metric_group(java_resource, metric_group) + + assert java_resource.metric_group is java_metric_group.j_metric_group + + +def test_set_java_resource_metric_group_accepts_none(): + java_resource = _JavaResource() + + set_java_resource_metric_group(java_resource, None) + + assert java_resource.metric_group is None + + +def test_set_java_resource_metric_group_rejects_non_flink_metric_group(): + with pytest.raises(TypeError, match="FlinkMetricGroup or None"): + set_java_resource_metric_group(_JavaResource(), _CustomMetricGroup()) + + +def test_set_metric_group_wraps_java_metric_group(): + python_resource = _PythonResource() + java_metric_group = _JavaMetricGroup() + + set_metric_group(python_resource, java_metric_group.j_metric_group) + + assert isinstance(python_resource.metric_group, FlinkMetricGroup) + assert python_resource.metric_group._j_metric_group is java_metric_group.j_metric_group + + +def test_set_metric_group_forwards_none(): + python_resource = _PythonResource() + + set_metric_group(python_resource, None) + + assert python_resource.metric_group is None diff --git a/runtime/src/main/java/org/apache/flink/agents/runtime/python/utils/PythonResourceAdapterImpl.java b/runtime/src/main/java/org/apache/flink/agents/runtime/python/utils/PythonResourceAdapterImpl.java index d685d293f..5cf4aec02 100644 --- a/runtime/src/main/java/org/apache/flink/agents/runtime/python/utils/PythonResourceAdapterImpl.java +++ b/runtime/src/main/java/org/apache/flink/agents/runtime/python/utils/PythonResourceAdapterImpl.java @@ -19,6 +19,7 @@ import org.apache.flink.agents.api.chat.messages.ChatMessage; import org.apache.flink.agents.api.chat.messages.MessageRole; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.prompt.Prompt; import org.apache.flink.agents.api.resource.Resource; import org.apache.flink.agents.api.resource.ResourceContext; @@ -53,6 +54,8 @@ public class PythonResourceAdapterImpl implements PythonResourceAdapter { static final String CALL_METHOD = PYTHON_MODULE_PREFIX + "call_method"; + static final String SET_METRIC_GROUP = PYTHON_MODULE_PREFIX + "set_metric_group"; + static final String CREATE_RESOURCE = PYTHON_MODULE_PREFIX + "create_resource"; static final String FROM_JAVA_RESOURCE = PYTHON_MODULE_PREFIX + "from_java_resource"; @@ -200,6 +203,11 @@ public Object callMethod(Object obj, String methodName, Map kwar return interpreter.invoke(CALL_METHOD, obj, methodName, kwargs); } + @Override + public void setMetricGroup(Object pythonResource, FlinkAgentsMetricGroup metricGroup) { + interpreter.invoke(SET_METRIC_GROUP, pythonResource, metricGroup); + } + @Override public Object invoke(String name, Object... args) { return interpreter.invoke(name, args); diff --git a/runtime/src/test/java/org/apache/flink/agents/runtime/python/utils/PythonResourceAdapterImplTest.java b/runtime/src/test/java/org/apache/flink/agents/runtime/python/utils/PythonResourceAdapterImplTest.java index f8821bfb2..e46e3ea68 100644 --- a/runtime/src/test/java/org/apache/flink/agents/runtime/python/utils/PythonResourceAdapterImplTest.java +++ b/runtime/src/test/java/org/apache/flink/agents/runtime/python/utils/PythonResourceAdapterImplTest.java @@ -18,6 +18,7 @@ package org.apache.flink.agents.runtime.python.utils; import org.apache.flink.agents.api.chat.model.python.PythonChatModelSetup; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; import org.apache.flink.agents.api.prompt.Prompt; import org.apache.flink.agents.api.resource.Resource; import org.apache.flink.agents.api.resource.ResourceContext; @@ -181,6 +182,17 @@ void testCallMethod() { .invoke(PythonResourceAdapterImpl.CALL_METHOD, obj, methodName, kwargs); } + @Test + void testSetMetricGroup() { + Object pythonResource = new Object(); + FlinkAgentsMetricGroup metricGroup = mock(FlinkAgentsMetricGroup.class); + + pythonResourceAdapter.setMetricGroup(pythonResource, metricGroup); + + verify(mockInterpreter) + .invoke(PythonResourceAdapterImpl.SET_METRIC_GROUP, pythonResource, metricGroup); + } + @Test void testInvoke() { String name = "test_function";