Skip to content
Merged
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 @@ -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;
Expand Down Expand Up @@ -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<ChatMessage> messages, List<Tool> tools, Map<String, Object> modelParams) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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<String, Object> getParameters() {
return Map.of();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -132,4 +133,15 @@ public Map<String, Object> getParameters() {
public Object getPythonResource() {
return embeddingModelSetup;
}

@Override
public PythonResourceAdapter getPythonResourceAdapter() {
return adapter;
}

@Override
public void setMetricGroup(FlinkAgentsMetricGroup metricGroup) {
super.setMetricGroup(metricGroup);
setPythonResourceMetricGroup(metricGroup);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -120,6 +121,14 @@ public interface PythonResourceAdapter {
*/
Object callMethod(Object obj, String methodName, Map<String, Object> 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.
*
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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);
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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;
Expand Down
13 changes: 13 additions & 0 deletions python/flink_agents/runtime/java/java_chat_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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,
Expand Down Expand Up @@ -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]:
Expand Down
13 changes: 13 additions & 0 deletions python/flink_agents/runtime/java/java_embedding_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,9 @@
JavaEmbeddingModelConnection,
JavaEmbeddingModelSetup,
)
from flink_agents.runtime.java.java_resource_wrapper import (
set_java_resource_metric_group,
)


class JavaEmbeddingModelConnectionImpl(JavaEmbeddingModelConnection):
Expand All @@ -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]]:
Expand Down Expand Up @@ -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.
Expand Down
16 changes: 16 additions & 0 deletions python/flink_agents/runtime/java/java_resource_wrapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down
8 changes: 8 additions & 0 deletions python/flink_agents/runtime/java/java_vector_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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]:
Expand Down
Loading
Loading