From 9bde2748ceb17117521081f91c21709a9ddb9c90 Mon Sep 17 00:00:00 2001 From: Jinkun Liu Date: Sun, 2 Aug 2026 19:18:17 +0800 Subject: [PATCH 1/5] [api][runtime] Add event attachments with automatic MemoryRef wrapping and resolution --- .../org/apache/flink/agents/api/Event.java | 53 ++++++-- .../apache/flink/agents/api/InputEvent.java | 1 + .../apache/flink/agents/api/OutputEvent.java | 1 + .../flink/agents/api/context/MemoryRef.java | 67 ++++++++++ .../agents/api/event/ChatRequestEvent.java | 1 + .../agents/api/event/ChatResponseEvent.java | 1 + .../event/ContextRetrievalRequestEvent.java | 1 + .../event/ContextRetrievalResponseEvent.java | 1 + .../agents/api/event/ToolRequestEvent.java | 1 + .../agents/api/event/ToolResponseEvent.java | 1 + python/flink_agents/api/events/chat_event.py | 2 + .../api/events/context_retrieval_event.py | 2 + python/flink_agents/api/events/event.py | 19 +++ python/flink_agents/api/events/tool_event.py | 2 + python/flink_agents/api/memory_object.py | 2 +- .../runtime/flink_runner_context.py | 2 + .../runtime/memory/event_attachment_utils.py | 90 +++++++++++++ .../flink_agents/runtime/python_java_utils.py | 7 ++ .../runtime/context/RunnerContextImpl.java | 6 + .../runtime/memory/EventAttachmentUtils.java | 118 ++++++++++++++++++ .../runtime/operator/JavaActionTask.java | 2 + .../python/utils/PythonActionExecutor.java | 5 + 22 files changed, 376 insertions(+), 9 deletions(-) create mode 100644 python/flink_agents/runtime/memory/event_attachment_utils.py create mode 100644 runtime/src/main/java/org/apache/flink/agents/runtime/memory/EventAttachmentUtils.java diff --git a/api/src/main/java/org/apache/flink/agents/api/Event.java b/api/src/main/java/org/apache/flink/agents/api/Event.java index e7fbde464..8dc381156 100644 --- a/api/src/main/java/org/apache/flink/agents/api/Event.java +++ b/api/src/main/java/org/apache/flink/agents/api/Event.java @@ -22,6 +22,7 @@ import com.fasterxml.jackson.annotation.JsonIgnore; import com.fasterxml.jackson.annotation.JsonProperty; import com.fasterxml.jackson.databind.ObjectMapper; +import org.apache.flink.agents.api.context.MemoryRef; import java.io.IOException; import java.util.HashMap; @@ -37,6 +38,7 @@ public class Event { private final UUID id; private final String type; private final Map attributes; + private final Map attachments; /** * Runtime-internal timestamp from the source record. Not part of the cross-language event @@ -46,7 +48,7 @@ public class Event { /** Unified event with user-defined type and attributes. */ public Event(String type, Map attributes) { - this(UUID.randomUUID(), type, attributes); + this(UUID.randomUUID(), type, attributes, new HashMap<>()); } /** Unified event with user-defined type and empty attributes. */ @@ -54,17 +56,26 @@ public Event(String type) { this(type, new HashMap<>()); } - @JsonCreator public Event( @JsonProperty("id") UUID id, @JsonProperty("type") String type, @JsonProperty("attributes") Map attributes) { + this(id, type, attributes, new HashMap<>()); + } + + @JsonCreator + public Event( + @JsonProperty("id") UUID id, + @JsonProperty("type") String type, + @JsonProperty("attributes") Map attributes, + @JsonProperty("attachments") Map attachments) { if (type == null || type.isEmpty()) { throw new IllegalArgumentException("Event 'type' must not be null or empty."); } this.id = id; this.type = type; this.attributes = attributes != null ? attributes : new HashMap<>(); + this.attachments = attachments != null ? attachments : new HashMap<>(); } public UUID getId() { @@ -81,10 +92,18 @@ public Map getAttributes() { return attributes; } + public Map getAttachments() { + return attachments; + } + public Object getAttr(String name) { return attributes.get(name); } + public Object getAttachment(String name) { + return attachments.get(name); + } + public void setAttr(String name, Object value) { attributes.put(name, value); } @@ -105,12 +124,17 @@ public void setSourceTimestamp(long timestamp) { } /** - * Creates a base Event from another Event, copying id, type, and attributes. Subclasses - * override this to reconstruct typed event objects with proper field deserialization. + * Creates a base Event from another Event, copying id, type, attributes, and attachments. + * Subclasses override this to reconstruct typed event objects with proper field + * deserialization. */ public static Event fromEvent(Event event) { Event copy = - new Event(event.getId(), event.getType(), new HashMap<>(event.getAttributes())); + new Event( + event.getId(), + event.getType(), + new HashMap<>(event.getAttributes()), + new HashMap<>(event.attachments)); if (event.hasSourceTimestamp()) { copy.setSourceTimestamp(event.getSourceTimestamp()); } @@ -125,7 +149,19 @@ public static Event fromEvent(Event event) { * @throws IOException if JSON parsing fails or the 'type' field is missing or empty */ public static Event fromJson(String json) throws IOException { - return MAPPER.readValue(json, Event.class); + Event event = MAPPER.readValue(json, Event.class); + for (Map.Entry entry : event.getAttachments().entrySet()) { + Object attachment = entry.getValue(); + if (attachment instanceof Map) { + Map map = (Map) attachment; + if (map.size() == 2 + && map.containsKey(MemoryRef.MEMORY_TYPE_FIELD) + && map.containsKey(MemoryRef.PATH_FIELD)) { + entry.setValue(MAPPER.convertValue(attachment, MemoryRef.class)); + } + } + } + return event; } @Override @@ -135,11 +171,12 @@ public boolean equals(Object o) { Event other = (Event) o; return Objects.equals(this.id, other.id) && Objects.equals(this.getType(), other.getType()) - && Objects.equals(this.attributes, other.attributes); + && Objects.equals(this.attributes, other.attributes) + && Objects.equals(this.attachments, other.attachments); } @Override public int hashCode() { - return Objects.hash(id, getType(), attributes); + return Objects.hash(id, getType(), attributes, attachments); } } diff --git a/api/src/main/java/org/apache/flink/agents/api/InputEvent.java b/api/src/main/java/org/apache/flink/agents/api/InputEvent.java index d07370215..508b7c562 100644 --- a/api/src/main/java/org/apache/flink/agents/api/InputEvent.java +++ b/api/src/main/java/org/apache/flink/agents/api/InputEvent.java @@ -51,6 +51,7 @@ public InputEvent( */ public static InputEvent fromEvent(Event event) { InputEvent result = new InputEvent(event.getId(), new HashMap<>(event.getAttributes())); + result.getAttachments().putAll(event.getAttachments()); if (event.hasSourceTimestamp()) { result.setSourceTimestamp(event.getSourceTimestamp()); } diff --git a/api/src/main/java/org/apache/flink/agents/api/OutputEvent.java b/api/src/main/java/org/apache/flink/agents/api/OutputEvent.java index d7a7e0f4a..0eabefcdc 100644 --- a/api/src/main/java/org/apache/flink/agents/api/OutputEvent.java +++ b/api/src/main/java/org/apache/flink/agents/api/OutputEvent.java @@ -54,6 +54,7 @@ public OutputEvent( */ public static OutputEvent fromEvent(Event event) { OutputEvent result = new OutputEvent(event.getId(), new HashMap<>(event.getAttributes())); + result.getAttachments().putAll(event.getAttachments()); if (event.hasSourceTimestamp()) { result.setSourceTimestamp(event.getSourceTimestamp()); } diff --git a/api/src/main/java/org/apache/flink/agents/api/context/MemoryRef.java b/api/src/main/java/org/apache/flink/agents/api/context/MemoryRef.java index df6f6c188..c21f96935 100644 --- a/api/src/main/java/org/apache/flink/agents/api/context/MemoryRef.java +++ b/api/src/main/java/org/apache/flink/agents/api/context/MemoryRef.java @@ -17,7 +17,21 @@ */ package org.apache.flink.agents.api.context; +import com.fasterxml.jackson.core.JsonGenerator; +import com.fasterxml.jackson.core.JsonParser; +import com.fasterxml.jackson.databind.DeserializationContext; +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.SerializerProvider; +import com.fasterxml.jackson.databind.annotation.JsonDeserialize; +import com.fasterxml.jackson.databind.annotation.JsonSerialize; +import com.fasterxml.jackson.databind.deser.std.StdDeserializer; +import com.fasterxml.jackson.databind.ser.std.StdSerializer; + +import java.io.IOException; import java.io.Serializable; +import java.util.LinkedHashMap; +import java.util.Locale; +import java.util.Map; import java.util.Objects; /** @@ -25,9 +39,14 @@ * lightweight pointer, containing the path of the data, allowing for efficient passing of large * objects between Actions. */ +@JsonSerialize(using = MemoryRef.Serializer.class) +@JsonDeserialize(using = MemoryRef.Deserializer.class) public final class MemoryRef implements Serializable { private static final long serialVersionUID = 1L; + public static final String MEMORY_TYPE_FIELD = "memory_type"; + public static final String PATH_FIELD = "path"; + private final MemoryObject.MemoryType type; private final String path; @@ -67,6 +86,54 @@ public String getPath() { return path; } + public MemoryObject.MemoryType getType() { + return type; + } + + /** Serializes a {@link MemoryRef} to JSON. */ + public static final class Serializer extends StdSerializer { + + public Serializer() { + super(MemoryRef.class); + } + + @Override + public void serialize(MemoryRef value, JsonGenerator generator, SerializerProvider provider) + throws IOException { + Map serialized = new LinkedHashMap<>(); + serialized.put(MEMORY_TYPE_FIELD, value.getType().name().toLowerCase(Locale.ROOT)); + serialized.put(PATH_FIELD, value.getPath()); + generator.writeObject(serialized); + } + } + + /** Deserializes a {@link MemoryRef} from JSON. */ + public static final class Deserializer extends StdDeserializer { + + public Deserializer() { + super(MemoryRef.class); + } + + @Override + public MemoryRef deserialize(JsonParser parser, DeserializationContext context) + throws IOException { + JsonNode node = parser.getCodec().readTree(parser); + JsonNode typeNode = node.get(MEMORY_TYPE_FIELD); + JsonNode pathNode = node.get(PATH_FIELD); + if (typeNode == null || typeNode.isNull() || pathNode == null || pathNode.isNull()) { + throw new IllegalArgumentException( + "MemoryRef JSON must contain non-null '" + + MEMORY_TYPE_FIELD + + "' and '" + + PATH_FIELD + + "' fields."); + } + MemoryObject.MemoryType memoryType = + MemoryObject.MemoryType.valueOf(typeNode.asText().toUpperCase(Locale.ROOT)); + return create(memoryType, pathNode.asText()); + } + } + @Override public boolean equals(Object o) { if (this == o) return true; diff --git a/api/src/main/java/org/apache/flink/agents/api/event/ChatRequestEvent.java b/api/src/main/java/org/apache/flink/agents/api/event/ChatRequestEvent.java index 4e05d1c88..4483dc4a8 100644 --- a/api/src/main/java/org/apache/flink/agents/api/event/ChatRequestEvent.java +++ b/api/src/main/java/org/apache/flink/agents/api/event/ChatRequestEvent.java @@ -98,6 +98,7 @@ private static Map normalizeAttributes(Map attri public static ChatRequestEvent fromEvent(Event event) { ChatRequestEvent result = new ChatRequestEvent(event.getId(), new HashMap<>(event.getAttributes())); + result.getAttachments().putAll(event.getAttachments()); if (event.hasSourceTimestamp()) { result.setSourceTimestamp(event.getSourceTimestamp()); } diff --git a/api/src/main/java/org/apache/flink/agents/api/event/ChatResponseEvent.java b/api/src/main/java/org/apache/flink/agents/api/event/ChatResponseEvent.java index 8dab3b1d8..4c55cc8c9 100644 --- a/api/src/main/java/org/apache/flink/agents/api/event/ChatResponseEvent.java +++ b/api/src/main/java/org/apache/flink/agents/api/event/ChatResponseEvent.java @@ -77,6 +77,7 @@ private static Map normalizeAttributes(Map attri public static ChatResponseEvent fromEvent(Event event) { ChatResponseEvent result = new ChatResponseEvent(event.getId(), new HashMap<>(event.getAttributes())); + result.getAttachments().putAll(event.getAttachments()); if (event.hasSourceTimestamp()) { result.setSourceTimestamp(event.getSourceTimestamp()); } diff --git a/api/src/main/java/org/apache/flink/agents/api/event/ContextRetrievalRequestEvent.java b/api/src/main/java/org/apache/flink/agents/api/event/ContextRetrievalRequestEvent.java index 387319c95..fac6b6bf8 100644 --- a/api/src/main/java/org/apache/flink/agents/api/event/ContextRetrievalRequestEvent.java +++ b/api/src/main/java/org/apache/flink/agents/api/event/ContextRetrievalRequestEvent.java @@ -62,6 +62,7 @@ public static ContextRetrievalRequestEvent fromEvent(Event event) { ContextRetrievalRequestEvent result = new ContextRetrievalRequestEvent( event.getId(), new HashMap<>(event.getAttributes())); + result.getAttachments().putAll(event.getAttachments()); if (event.hasSourceTimestamp()) { result.setSourceTimestamp(event.getSourceTimestamp()); } diff --git a/api/src/main/java/org/apache/flink/agents/api/event/ContextRetrievalResponseEvent.java b/api/src/main/java/org/apache/flink/agents/api/event/ContextRetrievalResponseEvent.java index a71fd7031..f7f6ca008 100644 --- a/api/src/main/java/org/apache/flink/agents/api/event/ContextRetrievalResponseEvent.java +++ b/api/src/main/java/org/apache/flink/agents/api/event/ContextRetrievalResponseEvent.java @@ -85,6 +85,7 @@ public static ContextRetrievalResponseEvent fromEvent(Event event) { ContextRetrievalResponseEvent result = new ContextRetrievalResponseEvent( event.getId(), new HashMap<>(event.getAttributes())); + result.getAttachments().putAll(event.getAttachments()); if (event.hasSourceTimestamp()) { result.setSourceTimestamp(event.getSourceTimestamp()); } diff --git a/api/src/main/java/org/apache/flink/agents/api/event/ToolRequestEvent.java b/api/src/main/java/org/apache/flink/agents/api/event/ToolRequestEvent.java index 9a24ef7ea..8476f63fa 100644 --- a/api/src/main/java/org/apache/flink/agents/api/event/ToolRequestEvent.java +++ b/api/src/main/java/org/apache/flink/agents/api/event/ToolRequestEvent.java @@ -56,6 +56,7 @@ public ToolRequestEvent( public static ToolRequestEvent fromEvent(Event event) { ToolRequestEvent result = new ToolRequestEvent(event.getId(), new HashMap<>(event.getAttributes())); + result.getAttachments().putAll(event.getAttachments()); if (event.hasSourceTimestamp()) { result.setSourceTimestamp(event.getSourceTimestamp()); } diff --git a/api/src/main/java/org/apache/flink/agents/api/event/ToolResponseEvent.java b/api/src/main/java/org/apache/flink/agents/api/event/ToolResponseEvent.java index 896f3bcb7..35a0a6753 100644 --- a/api/src/main/java/org/apache/flink/agents/api/event/ToolResponseEvent.java +++ b/api/src/main/java/org/apache/flink/agents/api/event/ToolResponseEvent.java @@ -100,6 +100,7 @@ private static Map normalizeAttributes(Map attri public static ToolResponseEvent fromEvent(Event event) { ToolResponseEvent result = new ToolResponseEvent(event.getId(), new HashMap<>(event.getAttributes())); + result.getAttachments().putAll(event.getAttachments()); if (event.hasSourceTimestamp()) { result.setSourceTimestamp(event.getSourceTimestamp()); } diff --git a/python/flink_agents/api/events/chat_event.py b/python/flink_agents/api/events/chat_event.py index a3a3deec1..0b74442ab 100644 --- a/python/flink_agents/api/events/chat_event.py +++ b/python/flink_agents/api/events/chat_event.py @@ -83,6 +83,7 @@ def from_event(cls, event: Event) -> "ChatRequestEvent": prompt_args=event.attributes.get("prompt_args"), output_schema=output_schema_raw, ) + result.attachments = dict(event.attachments) result.id = event.id return result @@ -160,6 +161,7 @@ def from_event(cls, event: Event) -> "ChatResponseEvent": retry_count=event.attributes.get("retry_count", 0), total_retry_wait_sec=event.attributes.get("total_retry_wait_sec", 0), ) + result.attachments = dict(event.attachments) result.id = event.id return result diff --git a/python/flink_agents/api/events/context_retrieval_event.py b/python/flink_agents/api/events/context_retrieval_event.py index a8245a4ec..71d7ee1b4 100644 --- a/python/flink_agents/api/events/context_retrieval_event.py +++ b/python/flink_agents/api/events/context_retrieval_event.py @@ -63,6 +63,7 @@ def from_event(cls, event: Event) -> "ContextRetrievalRequestEvent": vector_store=event.attributes["vector_store"], max_results=event.attributes.get("max_results", 3), ) + result.attachments = dict(event.attachments) result.id = event.id return result @@ -124,6 +125,7 @@ def from_event(cls, event: Event) -> "ContextRetrievalResponseEvent": query=event.attributes["query"], documents=documents, ) + result.attachments = dict(event.attachments) result.id = event.id return result diff --git a/python/flink_agents/api/events/event.py b/python/flink_agents/api/events/event.py index 77e7e5faf..d042fc47d 100644 --- a/python/flink_agents/api/events/event.py +++ b/python/flink_agents/api/events/event.py @@ -29,6 +29,8 @@ from pydantic_core import PydanticSerializationError from pyflink.common import Row +from flink_agents.api.memory_reference import MemoryRef + def _reconstruct_row_if_needed(data: Any) -> Any: """Recursively reconstruct pyflink Row objects from their JSON-serialized dicts. @@ -72,11 +74,14 @@ class Event(BaseModel, extra="allow"): Event type string used for routing. Required for all events. attributes : Dict[str, Any] Key-value properties for the event data. + attachments : Dict[str, Any] + Key-value data passed between actions through sensory memory. """ id: UUID = Field(default=None) type: str attributes: Dict[str, Any] = Field(default_factory=dict) + attachments: Dict[str, Any] = Field(default_factory=dict) @staticmethod def __serialize_unknown(field: Any) -> Dict[str, Any]: @@ -138,6 +143,14 @@ def set_attr(self, name: str, value: Any) -> None: """Set an attribute value in the attributes map.""" self.attributes[name] = value + def get_attachment(self, name: str) -> Any: + """Get an attachment value from the attachments map.""" + return self.attachments.get(name) + + def set_attachment(self, name: str, value: Any) -> None: + """Set an attachment value in the attachments map.""" + self.attachments = {**self.attachments, name: value} + @classmethod def from_event(cls, event: "Event") -> "Event": """Reconstruct a typed event from a base Event. @@ -173,6 +186,10 @@ def from_json(cls, json_str: str) -> "Event": event = cls.model_validate(data) for key in list(event.attributes): event.attributes[key] = _reconstruct_row_if_needed(event.attributes[key]) + for key in list(event.attachments): + value = event.attachments[key] + if isinstance(value, dict) and set(value) == {"memory_type", "path"}: + event.attachments[key] = MemoryRef.model_validate(value) return event @@ -200,6 +217,7 @@ def __init__(self, input: Any) -> None: def from_event(cls, event: Event) -> "InputEvent": assert "input" in event.attributes result = InputEvent(input=event.attributes["input"]) + result.attachments = dict(event.attachments) result.id = event.id return result @@ -233,6 +251,7 @@ def __init__(self, output: Any) -> None: def from_event(cls, event: Event) -> "OutputEvent": assert "output" in event.attributes result = OutputEvent(output=event.attributes["output"]) + result.attachments = dict(event.attachments) result.id = event.id return result diff --git a/python/flink_agents/api/events/tool_event.py b/python/flink_agents/api/events/tool_event.py index b7396b6f7..df938612e 100644 --- a/python/flink_agents/api/events/tool_event.py +++ b/python/flink_agents/api/events/tool_event.py @@ -58,6 +58,7 @@ def from_event(cls, event: Event) -> "ToolRequestEvent": model=event.attributes["model"], tool_calls=event.attributes["tool_calls"], ) + result.attachments = dict(event.attachments) result.id = event.id return result @@ -125,6 +126,7 @@ def from_event(cls, event: Event) -> "ToolResponseEvent": ), error=event.attributes.get("error", {}), ) + result.attachments = dict(event.attachments) result.id = event.id return result diff --git a/python/flink_agents/api/memory_object.py b/python/flink_agents/api/memory_object.py index 169d2e169..96452c7b5 100644 --- a/python/flink_agents/api/memory_object.py +++ b/python/flink_agents/api/memory_object.py @@ -91,7 +91,7 @@ def _validate(value: Any, where: str) -> None: class MemoryType(Enum): """Memory types based on MemoryObject.""" - SENSORY = ("sensory",) + SENSORY = "sensory" SHORT_TERM = "short_term" diff --git a/python/flink_agents/runtime/flink_runner_context.py b/python/flink_agents/runtime/flink_runner_context.py index ada5b8e30..341105e23 100644 --- a/python/flink_agents/runtime/flink_runner_context.py +++ b/python/flink_agents/runtime/flink_runner_context.py @@ -45,6 +45,7 @@ ) from flink_agents.runtime.flink_memory_object import FlinkMemoryObject from flink_agents.runtime.flink_metric_group import FlinkMetricGroup +from flink_agents.runtime.memory.event_attachment_utils import store_event_attachments from flink_agents.runtime.memory.internal_base_long_term_memory import ( InternalBaseLongTermMemory, ) @@ -297,6 +298,7 @@ def send_event(self, event: Event) -> None: event : Event The event to be processed by the agent system. """ + store_event_attachments(event, self) event_json = event.model_dump_json() try: self._j_runner_context.sendEventJson(event_json) diff --git a/python/flink_agents/runtime/memory/event_attachment_utils.py b/python/flink_agents/runtime/memory/event_attachment_utils.py new file mode 100644 index 000000000..3532b009a --- /dev/null +++ b/python/flink_agents/runtime/memory/event_attachment_utils.py @@ -0,0 +1,90 @@ +################################################################################ +# 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 __future__ import annotations + +import hashlib +from typing import TYPE_CHECKING, Any + +from flink_agents.api.events.event import OutputEvent +from flink_agents.api.memory_object import validate_memory_value +from flink_agents.api.memory_reference import MemoryRef + +if TYPE_CHECKING: + from uuid import UUID + + from flink_agents.api.events.event import Event + from flink_agents.api.runner_context import RunnerContext + +_ATTACHMENT_ROOT = "__event_attachments__" + + +class EventAttachmentError(RuntimeError): + """Raised when an Event attachment cannot be stored or loaded.""" + + +def _hash_attachment_key(key: str) -> str: + return hashlib.sha256(key.encode("UTF-8")).hexdigest() + + +def build_attachment_path(event_id: UUID, key: str) -> str: + """Build the canonical SensoryMemory path for one attachment.""" + return f"{_ATTACHMENT_ROOT}.{event_id}.{_hash_attachment_key(key)}" + + +def _attachment_context(event: Event, key: str, path: str | None = None) -> str: + suffix = f", path={path}" if path is not None else "" + return f"event_id={event.id}, event_type={event.type}, key={key}{suffix}" + + +def store_event_attachments(event: Event, ctx: RunnerContext) -> None: + """Store concrete attachment values in SensoryMemory and replace them with refs.""" + if not event.attachments: + return + + if event.type == OutputEvent.EVENT_TYPE: + keys = ", ".join(sorted(event.attachments)) + msg = f"Output events cannot carry attachments: {_attachment_context(event, keys)}" + raise EventAttachmentError(msg) + + pending: list[tuple[str, str, Any]] = [] + for key, value in event.attachments.items(): + if isinstance(value, MemoryRef): + continue + path = build_attachment_path(event.id, key) + try: + validate_memory_value(path, value) + except Exception as exc: + msg = f"Invalid event attachment value: {_attachment_context(event, key, path)}" + raise EventAttachmentError(msg) from exc + pending.append((key, path, value)) + + for key, path, value in pending: + event.attachments[key] = ctx.sensory_memory.set(path, value) + + +def load_event_attachments(event: Event, ctx: RunnerContext) -> None: + """Load sensory refs in place immediately before a Python Action runs.""" + for key, value in list(event.attachments.items()): + if not isinstance(value, MemoryRef): + continue + + try: + event.attachments[key] = ctx.sensory_memory.get(value) + except Exception as exc: + msg = f"Failed to load event attachment: {_attachment_context(event, key, value.path)}" + raise EventAttachmentError(msg) from exc diff --git a/python/flink_agents/runtime/python_java_utils.py b/python/flink_agents/runtime/python_java_utils.py index 3c4bb60dc..a0f149061 100644 --- a/python/flink_agents/runtime/python_java_utils.py +++ b/python/flink_agents/runtime/python_java_utils.py @@ -42,6 +42,7 @@ JavaResourceContextWrapper, JavaTool, ) +from flink_agents.runtime.memory.event_attachment_utils import load_event_attachments def convert_to_python_object(bytesObject: bytes) -> Any: @@ -60,6 +61,12 @@ def convert_json_to_python_event(event_json: str) -> Event: return Event.from_json(event_json) +def load_event_attachments_for_action(event: Event, ctx: Any) -> Event: + """Load attachment values before invoking a Python action.""" + load_event_attachments(event, ctx) + return event + + def wrap_to_input_event(bytesObject: bytes) -> str: """Wrap data to python input event and serialize as JSON. diff --git a/runtime/src/main/java/org/apache/flink/agents/runtime/context/RunnerContextImpl.java b/runtime/src/main/java/org/apache/flink/agents/runtime/context/RunnerContextImpl.java index 1395bdc56..1c4f393a3 100644 --- a/runtime/src/main/java/org/apache/flink/agents/runtime/context/RunnerContextImpl.java +++ b/runtime/src/main/java/org/apache/flink/agents/runtime/context/RunnerContextImpl.java @@ -37,6 +37,7 @@ import org.apache.flink.agents.runtime.actionstate.ActionState; import org.apache.flink.agents.runtime.actionstate.CallResult; import org.apache.flink.agents.runtime.memory.CachedMemoryStore; +import org.apache.flink.agents.runtime.memory.EventAttachmentUtils; import org.apache.flink.agents.runtime.memory.InteranlBaseLongTermMemory; import org.apache.flink.agents.runtime.memory.MemoryObjectImpl; import org.apache.flink.agents.runtime.metrics.FlinkAgentsMetricGroupImpl; @@ -149,6 +150,11 @@ public FlinkAgentsMetricGroupImpl getActionMetricGroup() { @Override public void sendEvent(Event event) { mailboxThreadChecker.run(); + try { + EventAttachmentUtils.storeEventAttachments(event, this); + } catch (Exception e) { + throw new IllegalArgumentException("Failed to store event attachments.", e); + } try { JsonUtils.checkSerializable(event); } catch (JsonProcessingException e) { diff --git a/runtime/src/main/java/org/apache/flink/agents/runtime/memory/EventAttachmentUtils.java b/runtime/src/main/java/org/apache/flink/agents/runtime/memory/EventAttachmentUtils.java new file mode 100644 index 000000000..3b10af1bb --- /dev/null +++ b/runtime/src/main/java/org/apache/flink/agents/runtime/memory/EventAttachmentUtils.java @@ -0,0 +1,118 @@ +/* + * 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. + */ +package org.apache.flink.agents.runtime.memory; + +import org.apache.flink.agents.api.Event; +import org.apache.flink.agents.api.OutputEvent; +import org.apache.flink.agents.api.context.MemoryObject; +import org.apache.flink.agents.api.context.MemoryRef; +import org.apache.flink.agents.api.context.RunnerContext; + +import java.nio.charset.StandardCharsets; +import java.security.MessageDigest; +import java.security.NoSuchAlgorithmException; +import java.util.Map; +import java.util.UUID; +import java.util.stream.Collectors; + +/** Stores event attachments in sensory memory while events cross action boundaries. */ +public final class EventAttachmentUtils { + + private static final String ATTACHMENT_ROOT = "__event_attachments__"; + + private EventAttachmentUtils() {} + + /** Stores concrete attachment values and replaces them with sensory-memory references. */ + public static void storeEventAttachments(Event event, RunnerContext context) throws Exception { + if (event.getAttachments().isEmpty()) { + return; + } + + if (OutputEvent.EVENT_TYPE.equals(event.getType())) { + String keys = + event.getAttachments().keySet().stream() + .sorted() + .collect(Collectors.joining(", ")); + throw new IllegalArgumentException( + "Output events cannot carry attachments: event_id=" + + event.getId() + + ", event_type=" + + event.getType() + + ", key=" + + keys); + } + + for (Map.Entry entry : event.getAttachments().entrySet()) { + String key = entry.getKey(); + Object value = entry.getValue(); + if (value instanceof MemoryRef) { + continue; + } + + MemoryRef reference = + context.getSensoryMemory().set(buildAttachmentPath(event.getId(), key), value); + + event.getAttachments().put(key, reference); + } + } + + /** Loads sensory-memory references in place before a Java action is invoked. */ + public static void loadEventAttachments(Event event, RunnerContext context) throws Exception { + for (Map.Entry entry : event.getAttachments().entrySet()) { + Object value = entry.getValue(); + if (!(value instanceof MemoryRef)) { + continue; + } + MemoryRef reference = (MemoryRef) value; + + MemoryObject attachment = context.getSensoryMemory().get(reference); + if (attachment == null) { + throw new IllegalStateException( + "Event attachment does not exist in sensory memory: " + + reference.getPath()); + } + event.getAttachments().put(entry.getKey(), attachment.getValue()); + } + } + + /** Builds the sensory-memory path for one event attachment. */ + public static String buildAttachmentPath(UUID eventId, String key) { + if (eventId == null) { + throw new IllegalArgumentException("Event attachment requires a non-null event id."); + } + if (key == null) { + throw new IllegalArgumentException("Event attachment key must not be null."); + } + return ATTACHMENT_ROOT + "." + eventId + "." + hashAttachmentKey(key); + } + + private static String hashAttachmentKey(String key) { + try { + byte[] digest = + MessageDigest.getInstance("SHA-256") + .digest(key.getBytes(StandardCharsets.UTF_8)); + StringBuilder sb = new StringBuilder(digest.length * 2); + for (byte value : digest) { + sb.append(String.format("%02x", value)); + } + return sb.toString(); + } catch (NoSuchAlgorithmException e) { + throw new IllegalStateException("SHA-256 is not available.", e); + } + } +} diff --git a/runtime/src/main/java/org/apache/flink/agents/runtime/operator/JavaActionTask.java b/runtime/src/main/java/org/apache/flink/agents/runtime/operator/JavaActionTask.java index 11724ce68..867a0df44 100644 --- a/runtime/src/main/java/org/apache/flink/agents/runtime/operator/JavaActionTask.java +++ b/runtime/src/main/java/org/apache/flink/agents/runtime/operator/JavaActionTask.java @@ -21,6 +21,7 @@ import org.apache.flink.agents.plan.JavaFunction; import org.apache.flink.agents.plan.actions.Action; import org.apache.flink.agents.runtime.context.JavaRunnerContextImpl; +import org.apache.flink.agents.runtime.memory.EventAttachmentUtils; import org.apache.flink.agents.runtime.python.utils.PythonActionExecutor; import java.util.Collections; @@ -56,6 +57,7 @@ public ActionTaskResult invoke(ClassLoader userCodeClassLoader, PythonActionExec if (!executionStarted) { runnerContext.checkNoPendingEvents(); + EventAttachmentUtils.loadEventAttachments(event, runnerContext); executionStarted = true; } diff --git a/runtime/src/main/java/org/apache/flink/agents/runtime/python/utils/PythonActionExecutor.java b/runtime/src/main/java/org/apache/flink/agents/runtime/python/utils/PythonActionExecutor.java index 67c80f38d..097bbc39e 100644 --- a/runtime/src/main/java/org/apache/flink/agents/runtime/python/utils/PythonActionExecutor.java +++ b/runtime/src/main/java/org/apache/flink/agents/runtime/python/utils/PythonActionExecutor.java @@ -64,6 +64,8 @@ public class PythonActionExecutor { // =========== PYTHON AND JAVA OBJECT CONVERT =========== private static final String CONVERT_JSON_TO_PYTHON_EVENT = "python_java_utils.convert_json_to_python_event"; + private static final String LOAD_EVENT_ATTACHMENTS_FOR_ACTION = + "python_java_utils.load_event_attachments_for_action"; private static final String WRAP_TO_INPUT_EVENT = "python_java_utils.wrap_to_input_event"; private static final String GET_OUTPUT_FROM_OUTPUT_EVENT = "python_java_utils.get_output_from_output_event"; @@ -134,6 +136,9 @@ public String executePythonFunction(PythonFunction function, Event event, int ha String eventJson = new ObjectMapper().writeValueAsString(event); Object pythonEventObject = interpreter.invoke(CONVERT_JSON_TO_PYTHON_EVENT, eventJson); + pythonEventObject = + interpreter.invoke( + LOAD_EVENT_ATTACHMENTS_FOR_ACTION, pythonEventObject, pythonRunnerContext); try { Object calledResult = function.call(pythonEventObject, pythonRunnerContext); From 1c9da03faba749de5b4ef53bf95c0d53037092c3 Mon Sep 17 00:00:00 2001 From: Jinkun Liu Date: Sun, 2 Aug 2026 19:18:17 +0800 Subject: [PATCH 2/5] [api][runtime] Test event attachments with automatic MemoryRef wrapping and resolution --- .../api/CrossLanguageEventSnapshotTest.java | 56 ++++- .../apache/flink/agents/api/EventTest.java | 33 +++ .../agents/api/context/MemoryRefJsonTest.java | 56 +++++ .../java/chat_request_event.json | 6 + .../java/chat_response_event.json | 6 + .../java/context_retrieval_request_event.json | 6 + .../context_retrieval_response_event.json | 6 + .../java/generic_event_with_attrs.json | 6 + .../java/input_event.json | 6 + .../java/output_event.json | 6 + .../java/tool_request_event.json | 6 + .../java/tool_response_event.json | 6 + .../python/chat_request_event.json | 6 + .../python/chat_response_event.json | 6 + .../context_retrieval_request_event.json | 6 + .../context_retrieval_response_event.json | 6 + .../python/generic_event_with_attrs.json | 6 + .../python/input_event.json | 6 + .../python/output_event.json | 6 + .../python/python_only_subclass_event.json | 6 + .../python/tool_request_event.json | 6 + .../python/tool_response_event.json | 6 + .../test_cross_language_event_snapshots.py | 31 +++ python/flink_agents/api/tests/test_event.py | 21 +- .../event_attachments_test.py | 100 +++++++++ .../tests/test_event_attachment_utils.py | 110 ++++++++++ .../memory/EventAttachmentUtilsTest.java | 201 ++++++++++++++++++ 27 files changed, 711 insertions(+), 11 deletions(-) create mode 100644 api/src/test/java/org/apache/flink/agents/api/context/MemoryRefJsonTest.java create mode 100644 python/flink_agents/e2e_tests/e2e_tests_integration/event_attachments_test.py create mode 100644 python/flink_agents/runtime/tests/test_event_attachment_utils.py create mode 100644 runtime/src/test/java/org/apache/flink/agents/runtime/memory/EventAttachmentUtilsTest.java diff --git a/api/src/test/java/org/apache/flink/agents/api/CrossLanguageEventSnapshotTest.java b/api/src/test/java/org/apache/flink/agents/api/CrossLanguageEventSnapshotTest.java index 5ae8bc7e8..496bfeb3a 100644 --- a/api/src/test/java/org/apache/flink/agents/api/CrossLanguageEventSnapshotTest.java +++ b/api/src/test/java/org/apache/flink/agents/api/CrossLanguageEventSnapshotTest.java @@ -23,6 +23,8 @@ import org.apache.flink.agents.api.agents.OutputSchema; import org.apache.flink.agents.api.chat.messages.ChatMessage; import org.apache.flink.agents.api.chat.messages.MessageRole; +import org.apache.flink.agents.api.context.MemoryObject; +import org.apache.flink.agents.api.context.MemoryRef; import org.apache.flink.agents.api.event.ChatRequestEvent; import org.apache.flink.agents.api.event.ChatResponseEvent; import org.apache.flink.agents.api.event.ContextRetrievalRequestEvent; @@ -65,6 +67,8 @@ class CrossLanguageEventSnapshotTest { private static final String FIXED_TOOL_CALL_ID = "call_aaaa"; private static final String FIXED_TOOL_CALL_ID_NUMERIC = "call_bbbb"; private static final String FIXED_TOOL_CALL_ID_BOOL = "call_cccc"; + private static final String ATTACHMENT_KEY = "payload"; + private static final String ATTACHMENT_PATH = "memory.path"; private static final long FIXED_TIMESTAMP = 1_700_000_000_000L; private static Path snapshotDir; @@ -122,12 +126,28 @@ private static Event readPythonSnapshot(String fileName) throws Exception { return Event.fromJson(Files.readString(pythonSnapshot)); } + private static T withMemoryRefAttachment(T event) { + event.getAttachments() + .put( + ATTACHMENT_KEY, + MemoryRef.create(MemoryObject.MemoryType.SENSORY, ATTACHMENT_PATH)); + return event; + } + + private static void assertMemoryRefAttachment(Event event) { + Object attachment = event.getAttachment(ATTACHMENT_KEY); + assertTrue(attachment instanceof MemoryRef); + MemoryRef reference = (MemoryRef) attachment; + assertEquals(MemoryObject.MemoryType.SENSORY, reference.getType()); + assertEquals(ATTACHMENT_PATH, reference.getPath()); + } + // ── InputEvent ───────────────────────────────────────────────────────── private static InputEvent buildInputEvent() { Map attrs = new HashMap<>(); attrs.put("input", "hello"); - return new InputEvent(FIXED_EVENT_ID, attrs); + return withMemoryRefAttachment(new InputEvent(FIXED_EVENT_ID, attrs)); } @Test @@ -144,12 +164,14 @@ void inputEventJavaSnapshotIsStable() throws Exception { @Test void javaCanDeserializeInputEventFromPythonSnapshot() throws Exception { Event base = readPythonSnapshot("input_event.json"); + assertMemoryRefAttachment(base); InputEvent typed = InputEvent.fromEvent(base); assertEquals( FIXED_EVENT_ID, typed.getId(), "ID lost when deserializing Python InputEvent."); assertEquals(InputEvent.EVENT_TYPE, typed.getType()); assertEquals("hello", typed.getInput(), "InputEvent.input mismatch."); + assertMemoryRefAttachment(typed); } // ── OutputEvent ──────────────────────────────────────────────────────── @@ -157,7 +179,7 @@ void javaCanDeserializeInputEventFromPythonSnapshot() throws Exception { private static OutputEvent buildOutputEvent() { Map attrs = new HashMap<>(); attrs.put("output", "world"); - return new OutputEvent(FIXED_EVENT_ID, attrs); + return withMemoryRefAttachment(new OutputEvent(FIXED_EVENT_ID, attrs)); } @Test @@ -174,12 +196,14 @@ void outputEventJavaSnapshotIsStable() throws Exception { @Test void javaCanDeserializeOutputEventFromPythonSnapshot() throws Exception { Event base = readPythonSnapshot("output_event.json"); + assertMemoryRefAttachment(base); OutputEvent typed = OutputEvent.fromEvent(base); assertEquals( FIXED_EVENT_ID, typed.getId(), "ID lost when deserializing Python OutputEvent."); assertEquals(OutputEvent.EVENT_TYPE, typed.getType()); assertEquals("world", typed.getOutput(), "OutputEvent.output mismatch."); + assertMemoryRefAttachment(typed); } // ── ChatRequestEvent ─────────────────────────────────────────────────── @@ -188,7 +212,7 @@ private static ChatRequestEvent buildChatRequestEvent() { Map attrs = new LinkedHashMap<>(); attrs.put("model", "test-model"); attrs.put("messages", List.of(new ChatMessage(MessageRole.USER, "hello world"))); - return new ChatRequestEvent(FIXED_EVENT_ID, attrs); + return withMemoryRefAttachment(new ChatRequestEvent(FIXED_EVENT_ID, attrs)); } @Test @@ -205,6 +229,7 @@ void chatRequestEventJavaSnapshotIsStable() throws Exception { @Test void javaCanDeserializeChatRequestEventFromPythonSnapshot() throws Exception { Event base = readPythonSnapshot("chat_request_event.json"); + assertMemoryRefAttachment(base); ChatRequestEvent typed = ChatRequestEvent.fromEvent(base); assertEquals(FIXED_EVENT_ID, typed.getId()); @@ -215,6 +240,7 @@ void javaCanDeserializeChatRequestEventFromPythonSnapshot() throws Exception { ChatMessage msg = typed.getMessages().get(0); assertEquals(MessageRole.USER, msg.getRole(), "Role mismatch on Python-produced message."); assertEquals("hello world", msg.getContent()); + assertMemoryRefAttachment(typed); } /** @@ -253,7 +279,7 @@ private static ChatResponseEvent buildChatResponseEvent() { attrs.put("response", new ChatMessage(MessageRole.ASSISTANT, "hi there")); attrs.put("retry_count", 0); attrs.put("total_retry_wait_sec", 0); - return new ChatResponseEvent(FIXED_EVENT_ID, attrs); + return withMemoryRefAttachment(new ChatResponseEvent(FIXED_EVENT_ID, attrs)); } @Test @@ -270,6 +296,7 @@ void chatResponseEventJavaSnapshotIsStable() throws Exception { @Test void javaCanDeserializeChatResponseEventFromPythonSnapshot() throws Exception { Event base = readPythonSnapshot("chat_response_event.json"); + assertMemoryRefAttachment(base); ChatResponseEvent typed = ChatResponseEvent.fromEvent(base); assertEquals(FIXED_EVENT_ID, typed.getId()); @@ -279,6 +306,7 @@ void javaCanDeserializeChatResponseEventFromPythonSnapshot() throws Exception { assertNotNull(response, "response field is null."); assertEquals(MessageRole.ASSISTANT, response.getRole(), "Role mismatch on response."); assertEquals("hi there", response.getContent()); + assertMemoryRefAttachment(typed); } // ── ToolRequestEvent ─────────────────────────────────────────────────── @@ -292,7 +320,7 @@ private static ToolRequestEvent buildToolRequestEvent() { Map attrs = new LinkedHashMap<>(); attrs.put("model", "test-model"); attrs.put("tool_calls", List.of(toolCall)); - return new ToolRequestEvent(FIXED_EVENT_ID, attrs); + return withMemoryRefAttachment(new ToolRequestEvent(FIXED_EVENT_ID, attrs)); } @Test @@ -309,6 +337,7 @@ void toolRequestEventJavaSnapshotIsStable() throws Exception { @Test void javaCanDeserializeToolRequestEventFromPythonSnapshot() throws Exception { Event base = readPythonSnapshot("tool_request_event.json"); + assertMemoryRefAttachment(base); ToolRequestEvent typed = ToolRequestEvent.fromEvent(base); assertEquals(FIXED_EVENT_ID, typed.getId()); @@ -318,6 +347,7 @@ void javaCanDeserializeToolRequestEventFromPythonSnapshot() throws Exception { assertNotNull(toolCalls); assertEquals(1, toolCalls.size()); assertEquals(FIXED_TOOL_CALL_ID, toolCalls.get(0).get("id")); + assertMemoryRefAttachment(typed); } // ── ToolResponseEvent ────────────────────────────────────────────────── @@ -330,7 +360,7 @@ private static ToolResponseEvent buildToolResponseEvent() { attrs.put("error", new HashMap()); attrs.put("external_ids", new HashMap()); attrs.put("timestamp", FIXED_TIMESTAMP); - return new ToolResponseEvent(FIXED_EVENT_ID, attrs); + return withMemoryRefAttachment(new ToolResponseEvent(FIXED_EVENT_ID, attrs)); } @Test @@ -347,6 +377,7 @@ void toolResponseEventJavaSnapshotIsStable() throws Exception { @Test void pythonToolResponseEventRoundTripsScalarResponses() throws Exception { Event base = readPythonSnapshot("tool_response_event.json"); + assertMemoryRefAttachment(base); ToolResponseEvent typed = ToolResponseEvent.fromEvent(base); assertEquals(FIXED_REQUEST_ID, typed.getRequestId()); @@ -370,6 +401,7 @@ void pythonToolResponseEventRoundTripsScalarResponses() throws Exception { assertEquals(Boolean.TRUE, typed.getSuccess().get(FIXED_TOOL_CALL_ID)); assertTrue(typed.getError().isEmpty()); assertFalse(attrs.containsKey("timestamp")); + assertMemoryRefAttachment(typed); } // ── ContextRetrievalRequestEvent ─────────────────────────────────────── @@ -379,7 +411,7 @@ private static ContextRetrievalRequestEvent buildContextRetrievalRequestEvent() attrs.put("query", "what is flink"); attrs.put("vector_store", "test-store"); attrs.put("max_results", 5); - return new ContextRetrievalRequestEvent(FIXED_EVENT_ID, attrs); + return withMemoryRefAttachment(new ContextRetrievalRequestEvent(FIXED_EVENT_ID, attrs)); } @Test @@ -398,6 +430,7 @@ void contextRetrievalRequestEventJavaSnapshotIsStable() throws Exception { @Test void javaCanDeserializeContextRetrievalRequestEventFromPythonSnapshot() throws Exception { Event base = readPythonSnapshot("context_retrieval_request_event.json"); + assertMemoryRefAttachment(base); ContextRetrievalRequestEvent typed = ContextRetrievalRequestEvent.fromEvent(base); assertEquals(FIXED_EVENT_ID, typed.getId()); @@ -405,6 +438,7 @@ void javaCanDeserializeContextRetrievalRequestEventFromPythonSnapshot() throws E assertEquals("what is flink", typed.getQuery()); assertEquals("test-store", typed.getVectorStore()); assertEquals(5, typed.getMaxResults()); + assertMemoryRefAttachment(typed); } // ── ContextRetrievalResponseEvent ────────────────────────────────────── @@ -415,7 +449,7 @@ private static ContextRetrievalResponseEvent buildContextRetrievalResponseEvent( attrs.put("request_id", FIXED_REQUEST_ID); attrs.put("query", "what is flink"); attrs.put("documents", new ArrayList<>(List.of(doc))); - return new ContextRetrievalResponseEvent(FIXED_EVENT_ID, attrs); + return withMemoryRefAttachment(new ContextRetrievalResponseEvent(FIXED_EVENT_ID, attrs)); } @Test @@ -434,6 +468,7 @@ void contextRetrievalResponseEventJavaSnapshotIsStable() throws Exception { @Test void javaCanDeserializeContextRetrievalResponseEventFromPythonSnapshot() throws Exception { Event base = readPythonSnapshot("context_retrieval_response_event.json"); + assertMemoryRefAttachment(base); ContextRetrievalResponseEvent typed = ContextRetrievalResponseEvent.fromEvent(base); assertEquals(FIXED_EVENT_ID, typed.getId()); @@ -445,6 +480,7 @@ void javaCanDeserializeContextRetrievalResponseEventFromPythonSnapshot() throws assertEquals(1, docs.size()); assertEquals("doc content", docs.get(0).getContent()); assertEquals("doc-1", docs.get(0).getId()); + assertMemoryRefAttachment(typed); } // ── Generic Event with primitive attributes (user-authored axis) ─────── @@ -460,7 +496,7 @@ private static Event buildGenericEvent() { attrs.put("k_null", null); attrs.put("k_list", List.of(1, 2, 3)); attrs.put("k_dict", Map.of("nested", "value")); - return new Event(FIXED_EVENT_ID, GENERIC_EVENT_TYPE, attrs); + return withMemoryRefAttachment(new Event(FIXED_EVENT_ID, GENERIC_EVENT_TYPE, attrs)); } @Test @@ -479,6 +515,7 @@ void javaCanDeserializeGenericEventFromPythonSnapshot() throws Exception { Event base = readPythonSnapshot("generic_event_with_attrs.json"); assertEquals(GENERIC_EVENT_TYPE, base.getType()); + assertMemoryRefAttachment(base); Map attrs = base.getAttributes(); assertEquals(42, attrs.get("k_int")); assertTrue(attrs.get("k_int") instanceof Integer); @@ -501,6 +538,7 @@ void javaCanDeserializePythonOnlySubclassEventAsBaseEvent() throws Exception { assertEquals(Event.class, base.getClass()); assertEquals("_my_python_only_event", base.getType()); assertEquals(FIXED_EVENT_ID, base.getId()); + assertMemoryRefAttachment(base); Map attrs = base.getAttributes(); assertEquals("ping", attrs.get("value")); diff --git a/api/src/test/java/org/apache/flink/agents/api/EventTest.java b/api/src/test/java/org/apache/flink/agents/api/EventTest.java index c4a4ca104..ab84003e5 100644 --- a/api/src/test/java/org/apache/flink/agents/api/EventTest.java +++ b/api/src/test/java/org/apache/flink/agents/api/EventTest.java @@ -20,6 +20,8 @@ import com.fasterxml.jackson.databind.JsonNode; import com.fasterxml.jackson.databind.ObjectMapper; +import org.apache.flink.agents.api.context.MemoryObject; +import org.apache.flink.agents.api.context.MemoryRef; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; @@ -131,6 +133,37 @@ void testUnifiedEventJsonDeserialization() throws Exception { assertEquals(1, event.getAttr("x")); } + @Test + void testUnifiedEventJsonRoundTripPreservesAttachments() throws Exception { + Map attachments = new HashMap<>(); + attachments.put("payload", Map.of("path", "memory.path")); + Event original = new Event(UUID.randomUUID(), "MyEvent", new HashMap<>(), attachments); + + Event restored = + objectMapper.readValue(objectMapper.writeValueAsString(original), Event.class); + + assertEquals(attachments, restored.getAttachments()); + } + + @Test + void testUnifiedEventJsonRoundTripPreservesMemoryRefAttachments() throws Exception { + MemoryRef reference = MemoryRef.create(MemoryObject.MemoryType.SENSORY, "memory.path"); + Event original = + new Event( + UUID.randomUUID(), + "MyEvent", + new HashMap<>(), + Map.of("payload", reference)); + + String json = objectMapper.writeValueAsString(original); + JsonNode attachment = objectMapper.readTree(json).get("attachments").get("payload"); + Event restored = Event.fromJson(json); + + assertEquals("sensory", attachment.get(MemoryRef.MEMORY_TYPE_FIELD).asText()); + assertEquals("memory.path", attachment.get(MemoryRef.PATH_FIELD).asText()); + assertEquals(reference, restored.getAttachment("payload")); + } + @Test void testSubclassedEventJsonRoundTrip() throws Exception { InputEvent original = new InputEvent("round trip"); diff --git a/api/src/test/java/org/apache/flink/agents/api/context/MemoryRefJsonTest.java b/api/src/test/java/org/apache/flink/agents/api/context/MemoryRefJsonTest.java new file mode 100644 index 000000000..145d3b01e --- /dev/null +++ b/api/src/test/java/org/apache/flink/agents/api/context/MemoryRefJsonTest.java @@ -0,0 +1,56 @@ +/* + * 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. + */ +package org.apache.flink.agents.api.context; + +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import org.junit.jupiter.api.Test; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; + +class MemoryRefJsonTest { + + private final ObjectMapper objectMapper = new ObjectMapper(); + + @Test + void serializesAndDeserializesMemoryRef() throws Exception { + MemoryRef original = MemoryRef.create(MemoryObject.MemoryType.SENSORY, "memory.path"); + + String json = objectMapper.writeValueAsString(original); + JsonNode node = objectMapper.readTree(json); + MemoryRef restored = objectMapper.readValue(json, MemoryRef.class); + + assertEquals("sensory", node.get("memory_type").asText()); + assertEquals("memory.path", node.get("path").asText()); + assertEquals(original, restored); + assertEquals(MemoryObject.MemoryType.SENSORY, restored.getType()); + } + + @Test + void rejectsInvalidMemoryRef() throws JsonProcessingException { + assertThrows( + IllegalArgumentException.class, + () -> + objectMapper.readValue( + "{\"memory_type\":\"unknown\",\"path\":\"memory.path\"}", + MemoryRef.class), + "No enum constant org.apache.flink.agents.api.context.MemoryObject.MemoryType.UNKNOWN"); + } +} diff --git a/e2e-test/cross-language-event-snapshots/java/chat_request_event.json b/e2e-test/cross-language-event-snapshots/java/chat_request_event.json index 347c47e71..259a7c750 100644 --- a/e2e-test/cross-language-event-snapshots/java/chat_request_event.json +++ b/e2e-test/cross-language-event-snapshots/java/chat_request_event.json @@ -9,5 +9,11 @@ "extra_args" : { } } ] }, + "attachments" : { + "payload" : { + "memory_type" : "sensory", + "path" : "memory.path" + } + }, "type" : "_chat_request_event" } diff --git a/e2e-test/cross-language-event-snapshots/java/chat_response_event.json b/e2e-test/cross-language-event-snapshots/java/chat_response_event.json index 3d5b4793c..eec8c27ba 100644 --- a/e2e-test/cross-language-event-snapshots/java/chat_response_event.json +++ b/e2e-test/cross-language-event-snapshots/java/chat_response_event.json @@ -11,5 +11,11 @@ "retry_count" : 0, "total_retry_wait_sec" : 0 }, + "attachments" : { + "payload" : { + "memory_type" : "sensory", + "path" : "memory.path" + } + }, "type" : "_chat_response_event" } diff --git a/e2e-test/cross-language-event-snapshots/java/context_retrieval_request_event.json b/e2e-test/cross-language-event-snapshots/java/context_retrieval_request_event.json index ead03f8de..94712370e 100644 --- a/e2e-test/cross-language-event-snapshots/java/context_retrieval_request_event.json +++ b/e2e-test/cross-language-event-snapshots/java/context_retrieval_request_event.json @@ -5,5 +5,11 @@ "vector_store" : "test-store", "max_results" : 5 }, + "attachments" : { + "payload" : { + "memory_type" : "sensory", + "path" : "memory.path" + } + }, "type" : "_context_retrieval_request_event" } diff --git a/e2e-test/cross-language-event-snapshots/java/context_retrieval_response_event.json b/e2e-test/cross-language-event-snapshots/java/context_retrieval_response_event.json index 90592d565..0d5594f55 100644 --- a/e2e-test/cross-language-event-snapshots/java/context_retrieval_response_event.json +++ b/e2e-test/cross-language-event-snapshots/java/context_retrieval_response_event.json @@ -13,5 +13,11 @@ "score" : null } ] }, + "attachments" : { + "payload" : { + "memory_type" : "sensory", + "path" : "memory.path" + } + }, "type" : "_context_retrieval_response_event" } diff --git a/e2e-test/cross-language-event-snapshots/java/generic_event_with_attrs.json b/e2e-test/cross-language-event-snapshots/java/generic_event_with_attrs.json index 96b9fc0f5..802e39920 100644 --- a/e2e-test/cross-language-event-snapshots/java/generic_event_with_attrs.json +++ b/e2e-test/cross-language-event-snapshots/java/generic_event_with_attrs.json @@ -11,5 +11,11 @@ "k_dict" : { "nested" : "value" } + }, + "attachments" : { + "payload" : { + "memory_type" : "sensory", + "path" : "memory.path" + } } } diff --git a/e2e-test/cross-language-event-snapshots/java/input_event.json b/e2e-test/cross-language-event-snapshots/java/input_event.json index 8150a1ce6..4c15fde28 100644 --- a/e2e-test/cross-language-event-snapshots/java/input_event.json +++ b/e2e-test/cross-language-event-snapshots/java/input_event.json @@ -3,5 +3,11 @@ "attributes" : { "input" : "hello" }, + "attachments" : { + "payload" : { + "memory_type" : "sensory", + "path" : "memory.path" + } + }, "type" : "_input_event" } diff --git a/e2e-test/cross-language-event-snapshots/java/output_event.json b/e2e-test/cross-language-event-snapshots/java/output_event.json index 3fb4269e8..3643c702d 100644 --- a/e2e-test/cross-language-event-snapshots/java/output_event.json +++ b/e2e-test/cross-language-event-snapshots/java/output_event.json @@ -3,5 +3,11 @@ "attributes" : { "output" : "world" }, + "attachments" : { + "payload" : { + "memory_type" : "sensory", + "path" : "memory.path" + } + }, "type" : "_output_event" } diff --git a/e2e-test/cross-language-event-snapshots/java/tool_request_event.json b/e2e-test/cross-language-event-snapshots/java/tool_request_event.json index 0f8ab2f3a..42b92c97c 100644 --- a/e2e-test/cross-language-event-snapshots/java/tool_request_event.json +++ b/e2e-test/cross-language-event-snapshots/java/tool_request_event.json @@ -10,5 +10,11 @@ } } ] }, + "attachments" : { + "payload" : { + "memory_type" : "sensory", + "path" : "memory.path" + } + }, "type" : "_tool_request_event" } diff --git a/e2e-test/cross-language-event-snapshots/java/tool_response_event.json b/e2e-test/cross-language-event-snapshots/java/tool_response_event.json index 04698abe2..023cbfc8c 100644 --- a/e2e-test/cross-language-event-snapshots/java/tool_response_event.json +++ b/e2e-test/cross-language-event-snapshots/java/tool_response_event.json @@ -18,5 +18,11 @@ "external_ids" : { }, "timestamp" : 1700000000000 }, + "attachments" : { + "payload" : { + "memory_type" : "sensory", + "path" : "memory.path" + } + }, "type" : "_tool_response_event" } diff --git a/e2e-test/cross-language-event-snapshots/python/chat_request_event.json b/e2e-test/cross-language-event-snapshots/python/chat_request_event.json index ac8808231..515a3c41e 100644 --- a/e2e-test/cross-language-event-snapshots/python/chat_request_event.json +++ b/e2e-test/cross-language-event-snapshots/python/chat_request_event.json @@ -13,5 +13,11 @@ ], "prompt_args": {}, "output_schema": null + }, + "attachments": { + "payload": { + "memory_type": "sensory", + "path": "memory.path" + } } } diff --git a/e2e-test/cross-language-event-snapshots/python/chat_response_event.json b/e2e-test/cross-language-event-snapshots/python/chat_response_event.json index bafb28116..11c2b507b 100644 --- a/e2e-test/cross-language-event-snapshots/python/chat_response_event.json +++ b/e2e-test/cross-language-event-snapshots/python/chat_response_event.json @@ -11,5 +11,11 @@ }, "retry_count": 0, "total_retry_wait_sec": 0 + }, + "attachments": { + "payload": { + "memory_type": "sensory", + "path": "memory.path" + } } } diff --git a/e2e-test/cross-language-event-snapshots/python/context_retrieval_request_event.json b/e2e-test/cross-language-event-snapshots/python/context_retrieval_request_event.json index 357ce8bc9..e9ac0ace6 100644 --- a/e2e-test/cross-language-event-snapshots/python/context_retrieval_request_event.json +++ b/e2e-test/cross-language-event-snapshots/python/context_retrieval_request_event.json @@ -5,5 +5,11 @@ "query": "what is flink", "vector_store": "test-store", "max_results": 5 + }, + "attachments": { + "payload": { + "memory_type": "sensory", + "path": "memory.path" + } } } diff --git a/e2e-test/cross-language-event-snapshots/python/context_retrieval_response_event.json b/e2e-test/cross-language-event-snapshots/python/context_retrieval_response_event.json index 95e14f0a6..a2ffe1af0 100644 --- a/e2e-test/cross-language-event-snapshots/python/context_retrieval_response_event.json +++ b/e2e-test/cross-language-event-snapshots/python/context_retrieval_response_event.json @@ -15,5 +15,11 @@ "score": null } ] + }, + "attachments": { + "payload": { + "memory_type": "sensory", + "path": "memory.path" + } } } diff --git a/e2e-test/cross-language-event-snapshots/python/generic_event_with_attrs.json b/e2e-test/cross-language-event-snapshots/python/generic_event_with_attrs.json index cfd461b36..75ef867ac 100644 --- a/e2e-test/cross-language-event-snapshots/python/generic_event_with_attrs.json +++ b/e2e-test/cross-language-event-snapshots/python/generic_event_with_attrs.json @@ -15,5 +15,11 @@ "k_dict": { "nested": "value" } + }, + "attachments": { + "payload": { + "memory_type": "sensory", + "path": "memory.path" + } } } diff --git a/e2e-test/cross-language-event-snapshots/python/input_event.json b/e2e-test/cross-language-event-snapshots/python/input_event.json index db24e4c56..c22e28a6c 100644 --- a/e2e-test/cross-language-event-snapshots/python/input_event.json +++ b/e2e-test/cross-language-event-snapshots/python/input_event.json @@ -3,5 +3,11 @@ "type": "_input_event", "attributes": { "input": "hello" + }, + "attachments": { + "payload": { + "memory_type": "sensory", + "path": "memory.path" + } } } diff --git a/e2e-test/cross-language-event-snapshots/python/output_event.json b/e2e-test/cross-language-event-snapshots/python/output_event.json index f4b48a746..1de411dca 100644 --- a/e2e-test/cross-language-event-snapshots/python/output_event.json +++ b/e2e-test/cross-language-event-snapshots/python/output_event.json @@ -3,5 +3,11 @@ "type": "_output_event", "attributes": { "output": "world" + }, + "attachments": { + "payload": { + "memory_type": "sensory", + "path": "memory.path" + } } } diff --git a/e2e-test/cross-language-event-snapshots/python/python_only_subclass_event.json b/e2e-test/cross-language-event-snapshots/python/python_only_subclass_event.json index a48448c12..7dfd7131c 100644 --- a/e2e-test/cross-language-event-snapshots/python/python_only_subclass_event.json +++ b/e2e-test/cross-language-event-snapshots/python/python_only_subclass_event.json @@ -4,5 +4,11 @@ "attributes": { "value": "ping", "count": 7 + }, + "attachments": { + "payload": { + "memory_type": "sensory", + "path": "memory.path" + } } } diff --git a/e2e-test/cross-language-event-snapshots/python/tool_request_event.json b/e2e-test/cross-language-event-snapshots/python/tool_request_event.json index 2ac1fc511..925d07930 100644 --- a/e2e-test/cross-language-event-snapshots/python/tool_request_event.json +++ b/e2e-test/cross-language-event-snapshots/python/tool_request_event.json @@ -12,5 +12,11 @@ } } ] + }, + "attachments": { + "payload": { + "memory_type": "sensory", + "path": "memory.path" + } } } diff --git a/e2e-test/cross-language-event-snapshots/python/tool_response_event.json b/e2e-test/cross-language-event-snapshots/python/tool_response_event.json index 4ec6d0b0a..da5c73d96 100644 --- a/e2e-test/cross-language-event-snapshots/python/tool_response_event.json +++ b/e2e-test/cross-language-event-snapshots/python/tool_response_event.json @@ -19,5 +19,11 @@ "call_bbbb": null, "call_cccc": null } + }, + "attachments": { + "payload": { + "memory_type": "sensory", + "path": "memory.path" + } } } diff --git a/python/flink_agents/api/tests/test_cross_language_event_snapshots.py b/python/flink_agents/api/tests/test_cross_language_event_snapshots.py index 315301aeb..3c1ec095d 100644 --- a/python/flink_agents/api/tests/test_cross_language_event_snapshots.py +++ b/python/flink_agents/api/tests/test_cross_language_event_snapshots.py @@ -33,6 +33,8 @@ ) from flink_agents.api.events.event import Event, InputEvent, OutputEvent from flink_agents.api.events.tool_event import ToolRequestEvent, ToolResponseEvent +from flink_agents.api.memory_object import MemoryType +from flink_agents.api.memory_reference import MemoryRef from flink_agents.api.vector_stores.vector_store import Document _REPO_ROOT = Path(__file__).resolve().parents[4] @@ -43,6 +45,8 @@ _FIXED_TOOL_CALL_ID = "call_aaaa" _FIXED_TOOL_CALL_ID_NUMERIC = "call_bbbb" _FIXED_TOOL_CALL_ID_BOOL = "call_cccc" +_ATTACHMENT_KEY = "payload" +_ATTACHMENT_PATH = "memory.path" def _regenerate_enabled() -> bool: @@ -50,10 +54,20 @@ def _regenerate_enabled() -> bool: def _force_id(event: Event, fixed_id: UUID) -> Event: + event.set_attachment( + _ATTACHMENT_KEY, MemoryRef.create(MemoryType.SENSORY, _ATTACHMENT_PATH) + ) object.__setattr__(event, "id", fixed_id) return event +def _assert_memory_ref_attachment(event: Event) -> None: + attachment = event.get_attachment(_ATTACHMENT_KEY) + assert isinstance(attachment, MemoryRef) + assert attachment.memory_type == MemoryType.SENSORY + assert attachment.path == _ATTACHMENT_PATH + + def _write_python_snapshot(name: str, event: Event) -> None: target = _SNAPSHOT_DIR / "python" / name target.parent.mkdir(parents=True, exist_ok=True) @@ -103,9 +117,11 @@ def test_input_event_python_snapshot_is_stable() -> None: def test_python_can_deserialize_input_event_from_java_snapshot() -> None: base = _read_java_snapshot("input_event.json") + _assert_memory_ref_attachment(base) typed = InputEvent.from_event(base) assert typed.input == "hello", "InputEvent.input mismatch." assert typed.type == InputEvent.EVENT_TYPE + _assert_memory_ref_attachment(typed) # ── OutputEvent ───────────────────────────────────────────────────────── @@ -127,9 +143,11 @@ def test_output_event_python_snapshot_is_stable() -> None: def test_python_can_deserialize_output_event_from_java_snapshot() -> None: base = _read_java_snapshot("output_event.json") + _assert_memory_ref_attachment(base) typed = OutputEvent.from_event(base) assert typed.output == "world", "OutputEvent.output mismatch." assert typed.type == OutputEvent.EVENT_TYPE + _assert_memory_ref_attachment(typed) # ── ChatRequestEvent ──────────────────────────────────────────────────── @@ -157,12 +175,14 @@ def test_chat_request_event_python_snapshot_is_stable() -> None: def test_python_can_deserialize_chat_request_event_from_java_snapshot() -> None: base = _read_java_snapshot("chat_request_event.json") + _assert_memory_ref_attachment(base) typed = ChatRequestEvent.from_event(base) assert typed.model == "test-model" assert len(typed.messages) == 1 msg = typed.messages[0] assert msg.role == MessageRole.USER, f"Role mismatch: got {msg.role!r}" assert msg.content == "hello world" + _assert_memory_ref_attachment(typed) def test_chat_request_row_type_info_output_schema_is_not_portable_across_languages_known_gap() -> None: @@ -220,6 +240,7 @@ def test_chat_response_event_python_snapshot_is_stable() -> None: def test_python_can_deserialize_chat_response_event_from_java_snapshot() -> None: base = _read_java_snapshot("chat_response_event.json") + _assert_memory_ref_attachment(base) typed = ChatResponseEvent.from_event(base) expected_request_id = str(_FIXED_REQUEST_ID) actual_request_id = ( @@ -231,6 +252,7 @@ def test_python_can_deserialize_chat_response_event_from_java_snapshot() -> None f"Response role mismatch: got {typed.response.role!r}" ) assert typed.response.content == "hi there" + _assert_memory_ref_attachment(typed) # ── ToolRequestEvent ──────────────────────────────────────────────────── @@ -256,10 +278,12 @@ def test_tool_request_event_python_snapshot_is_stable() -> None: def test_python_can_deserialize_tool_request_event_from_java_snapshot() -> None: base = _read_java_snapshot("tool_request_event.json") + _assert_memory_ref_attachment(base) typed = ToolRequestEvent.from_event(base) assert typed.model == "test-model" assert len(typed.tool_calls) == 1 assert typed.tool_calls[0]["id"] == _FIXED_TOOL_CALL_ID + _assert_memory_ref_attachment(typed) # ── ToolResponseEvent ─────────────────────────────────────────────────── @@ -299,6 +323,7 @@ def test_tool_response_event_python_snapshot_is_stable() -> None: def test_python_can_deserialize_java_tool_response_event_status_fields() -> None: base = _read_java_snapshot("tool_response_event.json") + _assert_memory_ref_attachment(base) typed = ToolResponseEvent.from_event(base) assert typed.request_id == _FIXED_REQUEST_ID @@ -310,6 +335,7 @@ def test_python_can_deserialize_java_tool_response_event_status_fields() -> None assert "result" in response_value assert "timestamp" not in typed.attributes + _assert_memory_ref_attachment(typed) # ── ContextRetrievalRequestEvent ──────────────────────────────────────── @@ -342,10 +368,12 @@ def test_context_retrieval_request_event_python_snapshot_is_stable() -> None: def test_python_can_deserialize_context_retrieval_request_event_from_java_snapshot() -> None: base = _read_java_snapshot("context_retrieval_request_event.json") + _assert_memory_ref_attachment(base) typed = ContextRetrievalRequestEvent.from_event(base) assert typed.query == "what is flink" assert typed.vector_store == "test-store" assert typed.max_results == 5 + _assert_memory_ref_attachment(typed) # ── ContextRetrievalResponseEvent ─────────────────────────────────────── @@ -379,6 +407,7 @@ def test_context_retrieval_response_event_python_snapshot_is_stable() -> None: def test_python_can_deserialize_context_retrieval_response_event_from_java_snapshot() -> None: base = _read_java_snapshot("context_retrieval_response_event.json") + _assert_memory_ref_attachment(base) typed = ContextRetrievalResponseEvent.from_event(base) expected_request_id = str(_FIXED_REQUEST_ID) actual_request_id = ( @@ -389,6 +418,7 @@ def test_python_can_deserialize_context_retrieval_response_event_from_java_snaps assert len(typed.documents) == 1 assert typed.documents[0].content == "doc content" assert typed.documents[0].id == "doc-1" + _assert_memory_ref_attachment(typed) # ── Generic Event with primitive attributes (user-authored axis) ─────── @@ -429,6 +459,7 @@ def test_python_can_deserialize_generic_event_from_java_snapshot() -> None: base = _read_java_snapshot("generic_event_with_attrs.json") assert base.type == _GENERIC_EVENT_TYPE + _assert_memory_ref_attachment(base) assert base.attributes["k_int"] == 42 assert isinstance(base.attributes["k_int"], int) assert base.attributes["k_float"] == 1.5 diff --git a/python/flink_agents/api/tests/test_event.py b/python/flink_agents/api/tests/test_event.py index 110106ea4..043fea2d7 100644 --- a/python/flink_agents/api/tests/test_event.py +++ b/python/flink_agents/api/tests/test_event.py @@ -52,7 +52,11 @@ def test_input_event_ignore_row_unserializable() -> None: def test_event_row_with_non_serializable_fails() -> None: with pytest.raises(ValidationError): - Event(type="test", row_field=Row({"a": 1}), non_serializable_field=Type[InputEvent]) + Event( + type="test", + row_field=Row({"a": 1}), + non_serializable_field=Type[InputEvent], + ) def test_event_multiple_rows_serializable() -> None: @@ -157,9 +161,14 @@ def test_output_event_from_event() -> None: def test_unified_event_creation() -> None: """Test creating a unified event with type and attributes.""" - event = Event(type="MyEvent", attributes={"field1": "test", "field2": 42}) + event = Event( + type="MyEvent", + attributes={"field1": "test", "field2": 42}, + attachments={"field3": "0105", "field4": 2004}, + ) assert event.type == "MyEvent" assert event.attributes == {"field1": "test", "field2": 42} + assert event.attachments == {"field3": "0105", "field4": 2004} assert event.get_type() == "MyEvent" @@ -177,6 +186,14 @@ def test_unified_event_get_attr_set_attr() -> None: assert event.get_attr("missing") is None +def test_unified_event_get_attachment_set_attachment() -> None: + """Test get_attachment and set_attachment convenience methods.""" + event = Event(type="TestEvent") + event.set_attachment("key", "value") + assert event.get_attachment("key") == "value" + assert event.get_attachment("missing") is None + + def test_unified_event_from_json() -> None: """Test deserializing a unified event from JSON.""" data = {"type": "MyEvent", "attributes": {"x": 1}} diff --git a/python/flink_agents/e2e_tests/e2e_tests_integration/event_attachments_test.py b/python/flink_agents/e2e_tests/e2e_tests_integration/event_attachments_test.py new file mode 100644 index 000000000..bef9a2f0a --- /dev/null +++ b/python/flink_agents/e2e_tests/e2e_tests_integration/event_attachments_test.py @@ -0,0 +1,100 @@ +################################################################################ +# 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. +################################################################################# +import os +import sys +import sysconfig +from pathlib import Path +from typing import Any + +from pyflink.common import Configuration +from pyflink.datastream import KeySelector, StreamExecutionEnvironment + +from flink_agents.api.agents.agent import Agent +from flink_agents.api.decorators import action +from flink_agents.api.events.event import Event, InputEvent, OutputEvent +from flink_agents.api.execution_environment import AgentsExecutionEnvironment +from flink_agents.api.runner_context import RunnerContext + +current_dir = Path(__file__).parent +os.environ["PYTHONPATH"] = ( + f"{current_dir.parent.parent.parent}:{sysconfig.get_paths()['purelib']}" +) + + +class _KeySelector(KeySelector): + def get_key(self, value: dict[str, Any]) -> str: + return str(value["key"]) + + +class EventAttachmentsAgent(Agent): + @action(InputEvent.EVENT_TYPE) + @staticmethod + def send_attachment(event: Event, ctx: RunnerContext) -> None: + value = InputEvent.from_event(event).input + ctx.send_event( + Event( + type="AttachmentStep", + attributes={"kind": "inline"}, + attachments={ + "payload": { + "value": value, + "items": [1, 2, 3], + } + }, + ) + ) + + @action("AttachmentStep") + @staticmethod + def receive_attachment(event: Event, ctx: RunnerContext) -> None: + print(f"received attachments: {event.attachments}") + ctx.send_event( + OutputEvent( + output={ + "kind": event.get_attr("kind"), + "payload": event.get_attachment("payload"), + } + ) + ) + + +def test_python_event_attachments_roundtrip_on_flink() -> None: + config = Configuration() + config.set_string("python.pythonpath", os.environ["PYTHONPATH"]) + env = StreamExecutionEnvironment.get_execution_environment(config) + env.set_python_executable(sys.executable) + env.set_parallelism(1) + input_stream = env.from_collection( + [{"key": "k1", "value": {"message": "hello"}}] + ) + agents_env = AgentsExecutionEnvironment.get_execution_environment(env=env) + output = ( + agents_env.from_datastream(input_stream, _KeySelector()) + .apply(EventAttachmentsAgent()) + .to_datastream() + ) + + assert list(output.execute_and_collect()) == [ + { + "kind": "inline", + "payload": { + "value": {"key": "k1", "value": {"message": "hello"}}, + "items": [1, 2, 3], + }, + } + ] diff --git a/python/flink_agents/runtime/tests/test_event_attachment_utils.py b/python/flink_agents/runtime/tests/test_event_attachment_utils.py new file mode 100644 index 000000000..46ba3238e --- /dev/null +++ b/python/flink_agents/runtime/tests/test_event_attachment_utils.py @@ -0,0 +1,110 @@ +################################################################################ +# 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. +################################################################################# +import uuid + +import pytest + +from flink_agents.api.events.event import Event, OutputEvent +from flink_agents.api.memory_object import MemoryType +from flink_agents.api.memory_reference import MemoryRef +from flink_agents.runtime.memory.event_attachment_utils import ( + EventAttachmentError, + build_attachment_path, + load_event_attachments, + store_event_attachments, +) +from flink_agents.runtime.tests.local_memory_object import LocalMemoryObject + + +class MockRunnerContext: + def __init__(self, memory: LocalMemoryObject) -> None: + self._memory = memory + + @property + def sensory_memory(self) -> LocalMemoryObject: + return self._memory + + +def test_store_event_attachments() -> None: + sensory_memory = LocalMemoryObject(MemoryType.SENSORY, {}) + ctx = MockRunnerContext(sensory_memory) + event_id = uuid.uuid4() + payload = {"value": "original"} + event = Event.model_construct( + id=event_id, + type="AttachmentStep", + attributes={}, + attachments={"payload": payload}, + ) + + store_event_attachments(event, ctx) + + attachment = event.get_attachment("payload") + assert isinstance(attachment, MemoryRef) + assert attachment.path == build_attachment_path(event_id, "payload") + assert sensory_memory.get(attachment) == payload + + +def test_store_rejects_output_event_attachments_before_storing_them() -> None: + sensory_memory = LocalMemoryObject(MemoryType.SENSORY, {}) + ctx = MockRunnerContext(sensory_memory) + event_id = uuid.uuid4() + attachments = {"zeta": {"value": 2}, "alpha": {"value": 1}} + event = Event.model_construct( + id=event_id, + type=OutputEvent.EVENT_TYPE, + attributes={"output": "result"}, + attachments=dict(attachments), + ) + + with pytest.raises(EventAttachmentError) as exc_info: + store_event_attachments(event, ctx) + + message = str(exc_info.value) + assert message.startswith("Output events cannot carry attachments:") + + +def test_load_event_attachments() -> None: + sensory_memory = LocalMemoryObject(MemoryType.SENSORY, {}) + ctx = MockRunnerContext(sensory_memory) + event_id = uuid.uuid4() + payload = {"value": "original"} + reference = sensory_memory.set(build_attachment_path(event_id, "payload"), payload) + event = Event.model_construct( + id=event_id, + type="AttachmentStep", + attributes={}, + attachments={"payload": reference}, + ) + + load_event_attachments(event, ctx) + + assert event.get_attachment("payload") == payload + + +def test_build_attachment_path() -> None: + event_id = uuid.UUID("00000000-0000-0000-0000-000000000001") + + path = build_attachment_path(event_id, "payload") + + assert ( + path + == "__event_attachments__." + + str(event_id) + + ".239f59ed55e737c77147cf55ad0c1b030b6d7ee748a7426952f9b852d5a935e5" + ) diff --git a/runtime/src/test/java/org/apache/flink/agents/runtime/memory/EventAttachmentUtilsTest.java b/runtime/src/test/java/org/apache/flink/agents/runtime/memory/EventAttachmentUtilsTest.java new file mode 100644 index 000000000..ecfba0b74 --- /dev/null +++ b/runtime/src/test/java/org/apache/flink/agents/runtime/memory/EventAttachmentUtilsTest.java @@ -0,0 +1,201 @@ +/* + * 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. + */ +package org.apache.flink.agents.runtime.memory; + +import org.apache.flink.agents.api.Event; +import org.apache.flink.agents.api.OutputEvent; +import org.apache.flink.agents.api.configuration.ReadableConfiguration; +import org.apache.flink.agents.api.context.DurableCallable; +import org.apache.flink.agents.api.context.MemoryObject; +import org.apache.flink.agents.api.context.MemoryRef; +import org.apache.flink.agents.api.context.RunnerContext; +import org.apache.flink.agents.api.memory.BaseLongTermMemory; +import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup; +import org.apache.flink.agents.api.resource.Resource; +import org.apache.flink.agents.api.resource.ResourceType; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +import java.util.HashMap; +import java.util.LinkedList; +import java.util.Map; +import java.util.UUID; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +class EventAttachmentUtilsTest { + + private MemoryObject sensoryMemory; + private RunnerContext context; + + @BeforeEach + void setUp() throws Exception { + sensoryMemory = + new MemoryObjectImpl( + MemoryObject.MemoryType.SENSORY, + new CachedMemoryStore(new ForTestMemoryMapState<>()), + MemoryObjectImpl.ROOT_KEY, + new LinkedList<>()); + context = new MockRunnerContext(sensoryMemory); + } + + @Test + void storesEventAttachments() throws Exception { + UUID eventId = UUID.randomUUID(); + Map payload = Map.of("value", "original"); + Event event = + new Event( + eventId, + "AttachmentStep", + Map.of(), + new HashMap<>(Map.of("payload", payload))); + + EventAttachmentUtils.storeEventAttachments(event, context); + + Object attachment = event.getAttachment("payload"); + assertTrue(attachment instanceof MemoryRef); + MemoryRef reference = (MemoryRef) attachment; + assertEquals( + EventAttachmentUtils.buildAttachmentPath(eventId, "payload"), reference.getPath()); + assertEquals(payload, sensoryMemory.get(reference).getValue()); + } + + @Test + void rejectsOutputEventAttachmentsBeforeStoringThem() throws Exception { + UUID eventId = UUID.randomUUID(); + Map attachments = + Map.of("zeta", Map.of("value", 2), "alpha", Map.of("value", 1)); + Event event = + new Event( + eventId, + OutputEvent.EVENT_TYPE, + Map.of("output", "result"), + new HashMap<>(attachments)); + + IllegalArgumentException error = + assertThrows( + IllegalArgumentException.class, + () -> EventAttachmentUtils.storeEventAttachments(event, context)); + + assertTrue(error.getMessage().startsWith("Output events cannot carry attachments:")); + } + + @Test + void loadsEventAttachments() throws Exception { + UUID eventId = UUID.randomUUID(); + Map payload = Map.of("value", "original"); + MemoryRef reference = + sensoryMemory.set( + EventAttachmentUtils.buildAttachmentPath(eventId, "payload"), payload); + Event event = + new Event( + eventId, + "AttachmentStep", + Map.of(), + new HashMap<>(Map.of("payload", reference))); + + EventAttachmentUtils.loadEventAttachments(event, context); + + assertEquals(payload, event.getAttachment("payload")); + } + + @Test + void buildsAttachmentPath() { + UUID eventId = UUID.fromString("00000000-0000-0000-0000-000000000001"); + + String path = EventAttachmentUtils.buildAttachmentPath(eventId, "payload"); + + assertEquals( + "__event_attachments__." + + eventId + + ".239f59ed55e737c77147cf55ad0c1b030b6d7ee748a7426952f9b852d5a935e5", + path); + } + + /** Mock RunnerContext for testing resolve(). */ + static class MockRunnerContext implements RunnerContext { + private final MemoryObject memoryObject; + + MockRunnerContext(MemoryObject memoryObject) { + this.memoryObject = memoryObject; + } + + @Override + public MemoryObject getShortTermMemory() { + return null; + } + + @Override + public BaseLongTermMemory getLongTermMemory() throws Exception { + return null; + } + + @Override + public MemoryObject getSensoryMemory() { + return memoryObject; + } + + @Override + public void sendEvent(org.apache.flink.agents.api.Event event) {} + + @Override + public FlinkAgentsMetricGroup getAgentMetricGroup() { + return null; + } + + @Override + public FlinkAgentsMetricGroup getActionMetricGroup() { + return null; + } + + @Override + public Resource getResource(String name, ResourceType type) throws Exception { + return null; + } + + @Override + public ReadableConfiguration getConfig() { + return null; + } + + @Override + public Map getActionConfig() { + return Map.of(); + } + + @Override + public Object getActionConfigValue(String key) { + return null; + } + + @Override + public T durableExecute(DurableCallable callable) throws Exception { + return callable.call(); + } + + @Override + public T durableExecuteAsync(DurableCallable callable) throws Exception { + return callable.call(); + } + + @Override + public void close() throws Exception {} + } +} From 600b90158d6645e04c0ec989cd5cf1d7f9bdeb2c Mon Sep 17 00:00:00 2001 From: Jinkun Liu Date: Sun, 9 Aug 2026 14:42:14 +0800 Subject: [PATCH 3/5] [Fix] Preserve MemoryRef when restoring ActionState --- .../org/apache/flink/agents/api/Event.java | 46 +++++++++++++------ .../flink/agents/api/context/MemoryRef.java | 3 ++ .../apache/flink/agents/api/EventTest.java | 9 ++-- .../agents/api/context/MemoryRefJsonTest.java | 1 + .../java/chat_request_event.json | 1 + .../java/chat_response_event.json | 1 + .../java/context_retrieval_request_event.json | 1 + .../context_retrieval_response_event.json | 1 + .../java/generic_event_with_attrs.json | 1 + .../java/input_event.json | 1 + .../java/output_event.json | 1 + .../java/tool_request_event.json | 1 + .../java/tool_response_event.json | 1 + .../python/chat_request_event.json | 1 + .../python/chat_response_event.json | 1 + .../context_retrieval_request_event.json | 1 + .../context_retrieval_response_event.json | 1 + .../python/generic_event_with_attrs.json | 1 + .../python/input_event.json | 1 + .../python/output_event.json | 1 + .../python/python_only_subclass_event.json | 1 + .../python/tool_request_event.json | 1 + .../python/tool_response_event.json | 1 + python/flink_agents/api/events/event.py | 20 ++++++-- python/flink_agents/api/memory_reference.py | 12 ++++- python/flink_agents/api/tests/test_event.py | 15 ++++++ .../actionstate/ActionStateSerdeTest.java | 8 ++++ 27 files changed, 108 insertions(+), 25 deletions(-) diff --git a/api/src/main/java/org/apache/flink/agents/api/Event.java b/api/src/main/java/org/apache/flink/agents/api/Event.java index 8dc381156..674d44fde 100644 --- a/api/src/main/java/org/apache/flink/agents/api/Event.java +++ b/api/src/main/java/org/apache/flink/agents/api/Event.java @@ -21,7 +21,12 @@ import com.fasterxml.jackson.annotation.JsonCreator; import com.fasterxml.jackson.annotation.JsonIgnore; import com.fasterxml.jackson.annotation.JsonProperty; +import com.fasterxml.jackson.core.JsonParser; +import com.fasterxml.jackson.databind.DeserializationContext; +import com.fasterxml.jackson.databind.JsonDeserializer; +import com.fasterxml.jackson.databind.JsonNode; import com.fasterxml.jackson.databind.ObjectMapper; +import com.fasterxml.jackson.databind.annotation.JsonDeserialize; import org.apache.flink.agents.api.context.MemoryRef; import java.io.IOException; @@ -38,6 +43,10 @@ public class Event { private final UUID id; private final String type; private final Map attributes; + + // Keep the annotation on the field as well as the creator parameter so it also applies when + // Jackson constructs Event subclasses whose creators do not declare attachments. + @JsonDeserialize(contentUsing = AttachmentValueDeserializer.class) private final Map attachments; /** @@ -68,7 +77,9 @@ public Event( @JsonProperty("id") UUID id, @JsonProperty("type") String type, @JsonProperty("attributes") Map attributes, - @JsonProperty("attachments") Map attachments) { + @JsonProperty("attachments") + @JsonDeserialize(contentUsing = AttachmentValueDeserializer.class) + Map attachments) { if (type == null || type.isEmpty()) { throw new IllegalArgumentException("Event 'type' must not be null or empty."); } @@ -100,12 +111,16 @@ public Object getAttr(String name) { return attributes.get(name); } + public void setAttr(String name, Object value) { + attributes.put(name, value); + } + public Object getAttachment(String name) { return attachments.get(name); } - public void setAttr(String name, Object value) { - attributes.put(name, value); + public void setAttachment(String name, Object value) { + attachments.put(name, value); } @JsonIgnore @@ -149,19 +164,22 @@ public static Event fromEvent(Event event) { * @throws IOException if JSON parsing fails or the 'type' field is missing or empty */ public static Event fromJson(String json) throws IOException { - Event event = MAPPER.readValue(json, Event.class); - for (Map.Entry entry : event.getAttachments().entrySet()) { - Object attachment = entry.getValue(); - if (attachment instanceof Map) { - Map map = (Map) attachment; - if (map.size() == 2 - && map.containsKey(MemoryRef.MEMORY_TYPE_FIELD) - && map.containsKey(MemoryRef.PATH_FIELD)) { - entry.setValue(MAPPER.convertValue(attachment, MemoryRef.class)); - } + return MAPPER.readValue(json, Event.class); + } + + /** Deserializes one attachment value, preserving explicitly tagged memory references. */ + static final class AttachmentValueDeserializer extends JsonDeserializer { + + @Override + public Object deserialize(JsonParser parser, DeserializationContext context) + throws IOException { + JsonNode node = parser.getCodec().readTree(parser); + if (node.isObject() + && MemoryRef.TYPE_VALUE.equals(node.path(MemoryRef.TYPE_FIELD).asText())) { + return parser.getCodec().treeToValue(node, MemoryRef.class); } + return parser.getCodec().treeToValue(node, Object.class); } - return event; } @Override diff --git a/api/src/main/java/org/apache/flink/agents/api/context/MemoryRef.java b/api/src/main/java/org/apache/flink/agents/api/context/MemoryRef.java index c21f96935..b01a76c05 100644 --- a/api/src/main/java/org/apache/flink/agents/api/context/MemoryRef.java +++ b/api/src/main/java/org/apache/flink/agents/api/context/MemoryRef.java @@ -44,6 +44,8 @@ public final class MemoryRef implements Serializable { private static final long serialVersionUID = 1L; + public static final String TYPE_FIELD = "@type"; + public static final String TYPE_VALUE = "memory_ref"; public static final String MEMORY_TYPE_FIELD = "memory_type"; public static final String PATH_FIELD = "path"; @@ -101,6 +103,7 @@ public Serializer() { public void serialize(MemoryRef value, JsonGenerator generator, SerializerProvider provider) throws IOException { Map serialized = new LinkedHashMap<>(); + serialized.put(TYPE_FIELD, TYPE_VALUE); serialized.put(MEMORY_TYPE_FIELD, value.getType().name().toLowerCase(Locale.ROOT)); serialized.put(PATH_FIELD, value.getPath()); generator.writeObject(serialized); diff --git a/api/src/test/java/org/apache/flink/agents/api/EventTest.java b/api/src/test/java/org/apache/flink/agents/api/EventTest.java index ab84003e5..f7c3d80b3 100644 --- a/api/src/test/java/org/apache/flink/agents/api/EventTest.java +++ b/api/src/test/java/org/apache/flink/agents/api/EventTest.java @@ -136,7 +136,7 @@ void testUnifiedEventJsonDeserialization() throws Exception { @Test void testUnifiedEventJsonRoundTripPreservesAttachments() throws Exception { Map attachments = new HashMap<>(); - attachments.put("payload", Map.of("path", "memory.path")); + attachments.put("payload", Map.of("memory_type", "sensory", "path", "memory.path")); Event original = new Event(UUID.randomUUID(), "MyEvent", new HashMap<>(), attachments); Event restored = @@ -157,10 +157,11 @@ void testUnifiedEventJsonRoundTripPreservesMemoryRefAttachments() throws Excepti String json = objectMapper.writeValueAsString(original); JsonNode attachment = objectMapper.readTree(json).get("attachments").get("payload"); - Event restored = Event.fromJson(json); + Event restored = objectMapper.readValue(json, Event.class); - assertEquals("sensory", attachment.get(MemoryRef.MEMORY_TYPE_FIELD).asText()); - assertEquals("memory.path", attachment.get(MemoryRef.PATH_FIELD).asText()); + assertEquals("memory_ref", attachment.get("@type").asText()); + assertEquals("sensory", attachment.get("memory_type").asText()); + assertEquals("memory.path", attachment.get("path").asText()); assertEquals(reference, restored.getAttachment("payload")); } diff --git a/api/src/test/java/org/apache/flink/agents/api/context/MemoryRefJsonTest.java b/api/src/test/java/org/apache/flink/agents/api/context/MemoryRefJsonTest.java index 145d3b01e..ee8359b19 100644 --- a/api/src/test/java/org/apache/flink/agents/api/context/MemoryRefJsonTest.java +++ b/api/src/test/java/org/apache/flink/agents/api/context/MemoryRefJsonTest.java @@ -37,6 +37,7 @@ void serializesAndDeserializesMemoryRef() throws Exception { JsonNode node = objectMapper.readTree(json); MemoryRef restored = objectMapper.readValue(json, MemoryRef.class); + assertEquals("memory_ref", node.get("@type").asText()); assertEquals("sensory", node.get("memory_type").asText()); assertEquals("memory.path", node.get("path").asText()); assertEquals(original, restored); diff --git a/e2e-test/cross-language-event-snapshots/java/chat_request_event.json b/e2e-test/cross-language-event-snapshots/java/chat_request_event.json index 259a7c750..ee9980969 100644 --- a/e2e-test/cross-language-event-snapshots/java/chat_request_event.json +++ b/e2e-test/cross-language-event-snapshots/java/chat_request_event.json @@ -11,6 +11,7 @@ }, "attachments" : { "payload" : { + "@type" : "memory_ref", "memory_type" : "sensory", "path" : "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/java/chat_response_event.json b/e2e-test/cross-language-event-snapshots/java/chat_response_event.json index eec8c27ba..a7c308118 100644 --- a/e2e-test/cross-language-event-snapshots/java/chat_response_event.json +++ b/e2e-test/cross-language-event-snapshots/java/chat_response_event.json @@ -13,6 +13,7 @@ }, "attachments" : { "payload" : { + "@type" : "memory_ref", "memory_type" : "sensory", "path" : "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/java/context_retrieval_request_event.json b/e2e-test/cross-language-event-snapshots/java/context_retrieval_request_event.json index 94712370e..8664bb11c 100644 --- a/e2e-test/cross-language-event-snapshots/java/context_retrieval_request_event.json +++ b/e2e-test/cross-language-event-snapshots/java/context_retrieval_request_event.json @@ -7,6 +7,7 @@ }, "attachments" : { "payload" : { + "@type" : "memory_ref", "memory_type" : "sensory", "path" : "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/java/context_retrieval_response_event.json b/e2e-test/cross-language-event-snapshots/java/context_retrieval_response_event.json index 0d5594f55..5f36c60d2 100644 --- a/e2e-test/cross-language-event-snapshots/java/context_retrieval_response_event.json +++ b/e2e-test/cross-language-event-snapshots/java/context_retrieval_response_event.json @@ -15,6 +15,7 @@ }, "attachments" : { "payload" : { + "@type" : "memory_ref", "memory_type" : "sensory", "path" : "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/java/generic_event_with_attrs.json b/e2e-test/cross-language-event-snapshots/java/generic_event_with_attrs.json index 802e39920..6e1d1d049 100644 --- a/e2e-test/cross-language-event-snapshots/java/generic_event_with_attrs.json +++ b/e2e-test/cross-language-event-snapshots/java/generic_event_with_attrs.json @@ -14,6 +14,7 @@ }, "attachments" : { "payload" : { + "@type" : "memory_ref", "memory_type" : "sensory", "path" : "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/java/input_event.json b/e2e-test/cross-language-event-snapshots/java/input_event.json index 4c15fde28..e27af580c 100644 --- a/e2e-test/cross-language-event-snapshots/java/input_event.json +++ b/e2e-test/cross-language-event-snapshots/java/input_event.json @@ -5,6 +5,7 @@ }, "attachments" : { "payload" : { + "@type" : "memory_ref", "memory_type" : "sensory", "path" : "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/java/output_event.json b/e2e-test/cross-language-event-snapshots/java/output_event.json index 3643c702d..39f04d3a7 100644 --- a/e2e-test/cross-language-event-snapshots/java/output_event.json +++ b/e2e-test/cross-language-event-snapshots/java/output_event.json @@ -5,6 +5,7 @@ }, "attachments" : { "payload" : { + "@type" : "memory_ref", "memory_type" : "sensory", "path" : "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/java/tool_request_event.json b/e2e-test/cross-language-event-snapshots/java/tool_request_event.json index 42b92c97c..19ec5055e 100644 --- a/e2e-test/cross-language-event-snapshots/java/tool_request_event.json +++ b/e2e-test/cross-language-event-snapshots/java/tool_request_event.json @@ -12,6 +12,7 @@ }, "attachments" : { "payload" : { + "@type" : "memory_ref", "memory_type" : "sensory", "path" : "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/java/tool_response_event.json b/e2e-test/cross-language-event-snapshots/java/tool_response_event.json index 023cbfc8c..4dc85549f 100644 --- a/e2e-test/cross-language-event-snapshots/java/tool_response_event.json +++ b/e2e-test/cross-language-event-snapshots/java/tool_response_event.json @@ -20,6 +20,7 @@ }, "attachments" : { "payload" : { + "@type" : "memory_ref", "memory_type" : "sensory", "path" : "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/python/chat_request_event.json b/e2e-test/cross-language-event-snapshots/python/chat_request_event.json index 515a3c41e..2ca0dcc0b 100644 --- a/e2e-test/cross-language-event-snapshots/python/chat_request_event.json +++ b/e2e-test/cross-language-event-snapshots/python/chat_request_event.json @@ -16,6 +16,7 @@ }, "attachments": { "payload": { + "@type": "memory_ref", "memory_type": "sensory", "path": "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/python/chat_response_event.json b/e2e-test/cross-language-event-snapshots/python/chat_response_event.json index 11c2b507b..190b4eba4 100644 --- a/e2e-test/cross-language-event-snapshots/python/chat_response_event.json +++ b/e2e-test/cross-language-event-snapshots/python/chat_response_event.json @@ -14,6 +14,7 @@ }, "attachments": { "payload": { + "@type": "memory_ref", "memory_type": "sensory", "path": "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/python/context_retrieval_request_event.json b/e2e-test/cross-language-event-snapshots/python/context_retrieval_request_event.json index e9ac0ace6..cc5e5e8d1 100644 --- a/e2e-test/cross-language-event-snapshots/python/context_retrieval_request_event.json +++ b/e2e-test/cross-language-event-snapshots/python/context_retrieval_request_event.json @@ -8,6 +8,7 @@ }, "attachments": { "payload": { + "@type": "memory_ref", "memory_type": "sensory", "path": "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/python/context_retrieval_response_event.json b/e2e-test/cross-language-event-snapshots/python/context_retrieval_response_event.json index a2ffe1af0..81e4c0b9b 100644 --- a/e2e-test/cross-language-event-snapshots/python/context_retrieval_response_event.json +++ b/e2e-test/cross-language-event-snapshots/python/context_retrieval_response_event.json @@ -18,6 +18,7 @@ }, "attachments": { "payload": { + "@type": "memory_ref", "memory_type": "sensory", "path": "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/python/generic_event_with_attrs.json b/e2e-test/cross-language-event-snapshots/python/generic_event_with_attrs.json index 75ef867ac..48e7ae0da 100644 --- a/e2e-test/cross-language-event-snapshots/python/generic_event_with_attrs.json +++ b/e2e-test/cross-language-event-snapshots/python/generic_event_with_attrs.json @@ -18,6 +18,7 @@ }, "attachments": { "payload": { + "@type": "memory_ref", "memory_type": "sensory", "path": "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/python/input_event.json b/e2e-test/cross-language-event-snapshots/python/input_event.json index c22e28a6c..431476483 100644 --- a/e2e-test/cross-language-event-snapshots/python/input_event.json +++ b/e2e-test/cross-language-event-snapshots/python/input_event.json @@ -6,6 +6,7 @@ }, "attachments": { "payload": { + "@type": "memory_ref", "memory_type": "sensory", "path": "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/python/output_event.json b/e2e-test/cross-language-event-snapshots/python/output_event.json index 1de411dca..ab1d30728 100644 --- a/e2e-test/cross-language-event-snapshots/python/output_event.json +++ b/e2e-test/cross-language-event-snapshots/python/output_event.json @@ -6,6 +6,7 @@ }, "attachments": { "payload": { + "@type": "memory_ref", "memory_type": "sensory", "path": "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/python/python_only_subclass_event.json b/e2e-test/cross-language-event-snapshots/python/python_only_subclass_event.json index 7dfd7131c..971c60133 100644 --- a/e2e-test/cross-language-event-snapshots/python/python_only_subclass_event.json +++ b/e2e-test/cross-language-event-snapshots/python/python_only_subclass_event.json @@ -7,6 +7,7 @@ }, "attachments": { "payload": { + "@type": "memory_ref", "memory_type": "sensory", "path": "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/python/tool_request_event.json b/e2e-test/cross-language-event-snapshots/python/tool_request_event.json index 925d07930..6782b920e 100644 --- a/e2e-test/cross-language-event-snapshots/python/tool_request_event.json +++ b/e2e-test/cross-language-event-snapshots/python/tool_request_event.json @@ -15,6 +15,7 @@ }, "attachments": { "payload": { + "@type": "memory_ref", "memory_type": "sensory", "path": "memory.path" } diff --git a/e2e-test/cross-language-event-snapshots/python/tool_response_event.json b/e2e-test/cross-language-event-snapshots/python/tool_response_event.json index da5c73d96..e2c307fcf 100644 --- a/e2e-test/cross-language-event-snapshots/python/tool_response_event.json +++ b/e2e-test/cross-language-event-snapshots/python/tool_response_event.json @@ -22,6 +22,7 @@ }, "attachments": { "payload": { + "@type": "memory_ref", "memory_type": "sensory", "path": "memory.path" } diff --git a/python/flink_agents/api/events/event.py b/python/flink_agents/api/events/event.py index d042fc47d..81cc454eb 100644 --- a/python/flink_agents/api/events/event.py +++ b/python/flink_agents/api/events/event.py @@ -25,7 +25,7 @@ from typing_extensions import override from uuid import UUID -from pydantic import BaseModel, Field, model_validator +from pydantic import BaseModel, Field, field_validator, model_validator from pydantic_core import PydanticSerializationError from pyflink.common import Row @@ -83,6 +83,20 @@ class Event(BaseModel, extra="allow"): attributes: Dict[str, Any] = Field(default_factory=dict) attachments: Dict[str, Any] = Field(default_factory=dict) + @field_validator("attachments", mode="before") + @classmethod + def _deserialize_memory_ref_attachments(cls, attachments: Any) -> Any: + """Restore explicitly tagged memory-reference attachment values.""" + if not isinstance(attachments, dict): + return attachments + return { + key: MemoryRef.model_validate(value) + if isinstance(value, dict) + and value.get(MemoryRef.TYPE_FIELD) == MemoryRef.TYPE_VALUE + else value + for key, value in attachments.items() + } + @staticmethod def __serialize_unknown(field: Any) -> Dict[str, Any]: """Handle serialization of unknown types, specifically Row objects.""" @@ -186,10 +200,6 @@ def from_json(cls, json_str: str) -> "Event": event = cls.model_validate(data) for key in list(event.attributes): event.attributes[key] = _reconstruct_row_if_needed(event.attributes[key]) - for key in list(event.attachments): - value = event.attachments[key] - if isinstance(value, dict) and set(value) == {"memory_type", "path"}: - event.attachments[key] = MemoryRef.model_validate(value) return event diff --git a/python/flink_agents/api/memory_reference.py b/python/flink_agents/api/memory_reference.py index 5c37793dc..d48e614bf 100644 --- a/python/flink_agents/api/memory_reference.py +++ b/python/flink_agents/api/memory_reference.py @@ -17,9 +17,9 @@ ################################################################################# from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, ClassVar -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, model_serializer from flink_agents.api.memory_object import MemoryType @@ -30,6 +30,9 @@ class MemoryRef(BaseModel): """Reference to a specific data item in the Short-Term Memory.""" + TYPE_FIELD: ClassVar[str] = "@type" + TYPE_VALUE: ClassVar[str] = "memory_ref" + memory_type: MemoryType = MemoryType.SHORT_TERM path: str @@ -53,6 +56,11 @@ def create(memory_type: MemoryType, path: str) -> MemoryRef: """ return MemoryRef(memory_type=memory_type, path=path) + @model_serializer(mode="wrap") + def _serialize_with_discriminator(self, handler: Any) -> dict[str, Any]: + """Serialize this reference with its language-neutral type discriminator.""" + return {self.TYPE_FIELD: self.TYPE_VALUE, **handler(self)} + def resolve(self, ctx: RunnerContext) -> Any: """Resolve the reference to get the actual data. diff --git a/python/flink_agents/api/tests/test_event.py b/python/flink_agents/api/tests/test_event.py index 043fea2d7..106eaa27b 100644 --- a/python/flink_agents/api/tests/test_event.py +++ b/python/flink_agents/api/tests/test_event.py @@ -24,6 +24,8 @@ from pyflink.common import Row from flink_agents.api.events.event import Event, InputEvent, OutputEvent +from flink_agents.api.memory_object import MemoryType +from flink_agents.api.memory_reference import MemoryRef def test_event_init_serializable() -> None: @@ -220,6 +222,19 @@ def test_unified_event_serialization_roundtrip() -> None: assert restored.attributes == {"a": 1, "b": "two"} +def test_unified_event_serialization_roundtrip_with_memory_ref_attachment() -> None: + """Test that tagged memory-reference attachments retain their concrete type.""" + reference = MemoryRef.create(MemoryType.SENSORY, "memory.path") + original = Event(type="RoundTrip", attachments={"payload": reference}) + + parsed = json.loads(original.model_dump_json()) + restored = Event.model_validate(parsed) + + assert parsed["attachments"]["payload"][MemoryRef.TYPE_FIELD] == MemoryRef.TYPE_VALUE + assert restored.get_attachment("payload") == reference + assert isinstance(restored.get_attachment("payload"), MemoryRef) + + def test_unified_event_serialization_roundtrip_with_row() -> None: """Test that unified events with Row fields survive JSON roundtrip.""" original = Event( diff --git a/runtime/src/test/java/org/apache/flink/agents/runtime/actionstate/ActionStateSerdeTest.java b/runtime/src/test/java/org/apache/flink/agents/runtime/actionstate/ActionStateSerdeTest.java index f8dbc6fa8..c06d005b7 100644 --- a/runtime/src/test/java/org/apache/flink/agents/runtime/actionstate/ActionStateSerdeTest.java +++ b/runtime/src/test/java/org/apache/flink/agents/runtime/actionstate/ActionStateSerdeTest.java @@ -22,6 +22,8 @@ import org.apache.flink.agents.api.OutputEvent; import org.apache.flink.agents.api.chat.messages.ChatMessage; import org.apache.flink.agents.api.chat.messages.MessageRole; +import org.apache.flink.agents.api.context.MemoryObject; +import org.apache.flink.agents.api.context.MemoryRef; import org.apache.flink.agents.api.context.MemoryUpdate; import org.apache.flink.agents.api.event.ChatRequestEvent; import org.apache.flink.agents.api.event.ChatResponseEvent; @@ -51,6 +53,8 @@ public void testActionStateSerializationDeserialization() throws Exception { // Create test data InputEvent inputEvent = new InputEvent("test input"); inputEvent.setAttr("testAttr", "testValue"); + MemoryRef reference = MemoryRef.create(MemoryObject.MemoryType.SENSORY, "attachment.path"); + inputEvent.setAttachment("payload", reference); OutputEvent outputEvent = new OutputEvent("test output"); outputEvent.setAttr("outputAttr", 123); @@ -79,6 +83,10 @@ public void testActionStateSerializationDeserialization() throws Exception { InputEvent deserializedInputEvent = (InputEvent) deserializedState.getTaskEvent(); assertEquals("test input", deserializedInputEvent.getInput()); assertEquals("testValue", deserializedInputEvent.getAttr("testAttr")); + Object restoredAttachment = deserializedInputEvent.getAttachment("payload"); + assertInstanceOf(MemoryRef.class, restoredAttachment); + assertEquals(reference, restoredAttachment); + assertEquals(MemoryObject.MemoryType.SENSORY, ((MemoryRef) restoredAttachment).getType()); // Verify memoryUpdates assertEquals(1, deserializedState.getSensoryMemoryUpdates().size()); From e4ffa8eaab35730cff4bdcd15b17ad40fbca9efe Mon Sep 17 00:00:00 2001 From: Jinkun Liu Date: Sun, 9 Aug 2026 17:03:26 +0800 Subject: [PATCH 4/5] [Fix] Avoid JSON-serializing raw attachments before offload --- python/flink_agents/api/events/event.py | 36 +++++-------------- .../tests/test_event_attachment_utils.py | 15 ++++++++ 2 files changed, 24 insertions(+), 27 deletions(-) diff --git a/python/flink_agents/api/events/event.py b/python/flink_agents/api/events/event.py index 81cc454eb..644bada5c 100644 --- a/python/flink_agents/api/events/event.py +++ b/python/flink_agents/api/events/event.py @@ -15,7 +15,6 @@ # See the License for the specific language governing permissions and # limitations under the License. ################################################################################# -import hashlib import json from typing import Any, ClassVar, Dict @@ -23,7 +22,7 @@ from typing import override except ImportError: from typing_extensions import override -from uuid import UUID +from uuid import UUID, uuid4 from pydantic import BaseModel, Field, field_validator, model_validator from pydantic_core import PydanticSerializationError @@ -68,8 +67,7 @@ class Event(BaseModel, extra="allow"): Attributes: ---------- id : UUID - Unique identifier for the event, generated deterministically based on - event content. + Unique identifier for the event, generated randomly when not supplied. type : str Event type string used for routing. Required for all events. attributes : Dict[str, Any] @@ -78,7 +76,7 @@ class Event(BaseModel, extra="allow"): Key-value data passed between actions through sensory memory. """ - id: UUID = Field(default=None) + id: UUID = Field(default_factory=uuid4) type: str attributes: Dict[str, Any] = Field(default_factory=dict) attachments: Dict[str, Any] = Field(default_factory=dict) @@ -117,33 +115,17 @@ def model_dump_json(self, **kwargs: Any) -> str: kwargs["fallback"] = self.__serialize_unknown return super().model_dump_json(**kwargs) - def _generate_content_based_id(self) -> UUID: - """Generate a deterministic UUID based on event content using MD5 hash. - - Similar to Java's UUID.nameUUIDFromBytes(), uses MD5 for version 3 UUID. - """ - # Serialize content excluding 'id' to avoid circular dependency - content_json = super().model_dump_json( - exclude={"id"}, fallback=self.__serialize_unknown - ) - md5_hash = hashlib.md5(content_json.encode()).digest() - return UUID(bytes=md5_hash, version=3) - @model_validator(mode="after") - def validate_and_set_id(self) -> "Event": - """Validate that fields are serializable and generate content-based ID.""" - if self.id is None: - object.__setattr__(self, "id", self._generate_content_based_id()) - self.model_dump_json() + def validate_serializable_fields(self) -> "Event": + """Validate JSON event fields without serializing raw attachments.""" + self.model_dump_json(exclude={"attachments"}) return self def __setattr__(self, name: str, value: Any) -> None: super().__setattr__(name, value) - # Ensure added property can be serialized. - self.model_dump_json() - # Regenerate ID if content changed (but not if setting 'id' itself) - if name != "id": - object.__setattr__(self, "id", self._generate_content_based_id()) + # Raw attachments are offloaded to sensory memory before sending. Validate every + # other field here without serializing those payloads. + self.model_dump_json(exclude={"attachments"}) def get_type(self) -> str: """Return the event type string used for routing.""" diff --git a/python/flink_agents/runtime/tests/test_event_attachment_utils.py b/python/flink_agents/runtime/tests/test_event_attachment_utils.py index 46ba3238e..f0716279d 100644 --- a/python/flink_agents/runtime/tests/test_event_attachment_utils.py +++ b/python/flink_agents/runtime/tests/test_event_attachment_utils.py @@ -60,6 +60,21 @@ def test_store_event_attachments() -> None: assert sensory_memory.get(attachment) == payload +def test_store_offloads_non_utf8_bytes_from_regular_event() -> None: + """Test that raw bytes are offloaded before JSON serialization is required.""" + sensory_memory = LocalMemoryObject(MemoryType.SENSORY, {}) + ctx = MockRunnerContext(sensory_memory) + payload = b"\xff\x00" + event = Event(type="AttachmentStep", attachments={"payload": payload}) + + store_event_attachments(event, ctx) + + attachment = event.get_attachment("payload") + assert isinstance(attachment, MemoryRef) + assert sensory_memory.get(attachment) == payload + event.model_dump_json() + + def test_store_rejects_output_event_attachments_before_storing_them() -> None: sensory_memory = LocalMemoryObject(MemoryType.SENSORY, {}) ctx = MockRunnerContext(sensory_memory) From e2aa09e04c7b199e326c97b7d844f2b6b8a0d7e3 Mon Sep 17 00:00:00 2001 From: Jinkun Liu Date: Sun, 9 Aug 2026 17:50:35 +0800 Subject: [PATCH 5/5] [Fix] Make the Event own a mutable attachment map & attributes map --- .../java/org/apache/flink/agents/api/Event.java | 4 ++-- .../runtime/memory/EventAttachmentUtilsTest.java | 13 +++++++++++++ 2 files changed, 15 insertions(+), 2 deletions(-) diff --git a/api/src/main/java/org/apache/flink/agents/api/Event.java b/api/src/main/java/org/apache/flink/agents/api/Event.java index 674d44fde..9c31e63b6 100644 --- a/api/src/main/java/org/apache/flink/agents/api/Event.java +++ b/api/src/main/java/org/apache/flink/agents/api/Event.java @@ -85,8 +85,8 @@ public Event( } this.id = id; this.type = type; - this.attributes = attributes != null ? attributes : new HashMap<>(); - this.attachments = attachments != null ? attachments : new HashMap<>(); + this.attributes = attributes != null ? new HashMap<>(attributes) : new HashMap<>(); + this.attachments = attachments != null ? new HashMap<>(attachments) : new HashMap<>(); } public UUID getId() { diff --git a/runtime/src/test/java/org/apache/flink/agents/runtime/memory/EventAttachmentUtilsTest.java b/runtime/src/test/java/org/apache/flink/agents/runtime/memory/EventAttachmentUtilsTest.java index ecfba0b74..323d88a6e 100644 --- a/runtime/src/test/java/org/apache/flink/agents/runtime/memory/EventAttachmentUtilsTest.java +++ b/runtime/src/test/java/org/apache/flink/agents/runtime/memory/EventAttachmentUtilsTest.java @@ -77,6 +77,19 @@ void storesEventAttachments() throws Exception { assertEquals(payload, sensoryMemory.get(reference).getValue()); } + @Test + void storesAttachmentsFromImmutableMap() throws Exception { + UUID eventId = UUID.randomUUID(); + Map payload = Map.of("value", "original"); + Map attachments = Map.of("payload", payload); + Event event = new Event(eventId, "AttachmentStep", Map.of(), attachments); + + EventAttachmentUtils.storeEventAttachments(event, context); + + assertTrue(event.getAttachment("payload") instanceof MemoryRef); + assertEquals(payload, attachments.get("payload")); + } + @Test void rejectsOutputEventAttachmentsBeforeStoringThem() throws Exception { UUID eventId = UUID.randomUUID();