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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
import org.apache.flink.agents.api.OutputEvent;
import org.apache.flink.agents.api.annotation.Action;
import org.apache.flink.agents.api.chat.messages.ChatMessage;
import org.apache.flink.agents.api.chat.messages.ChatResult;
import org.apache.flink.agents.api.chat.messages.MessageRole;
import org.apache.flink.agents.api.context.RunnerContext;
import org.apache.flink.agents.api.event.ChatRequestEvent;
Expand Down Expand Up @@ -199,11 +200,11 @@ public static void startAction(Event event, RunnerContext ctx) {
@Action(EventType.ChatResponseEvent)
public static void stopAction(Event event, RunnerContext ctx) {
ChatResponseEvent chatResponse = ChatResponseEvent.fromEvent(event);
ChatMessage response = chatResponse.getResponse();
ChatResult response = chatResponse.getResponse();

Object output;
if (response.getExtraArgs().containsKey(STRUCTURED_OUTPUT)) {
output = response.getExtraArgs().get(STRUCTURED_OUTPUT);
if (chatResponse.getStructuredOutput() != null) {
output = chatResponse.getStructuredOutput();
} else {
output = response.getText();
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,187 +18,139 @@

package org.apache.flink.agents.api.chat.messages;

import com.fasterxml.jackson.annotation.JsonCreator;
import com.fasterxml.jackson.annotation.JsonIgnore;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.fasterxml.jackson.core.type.TypeReference;
import com.fasterxml.jackson.databind.ObjectMapper;

import java.util.ArrayList;
import java.util.Collections;
import java.util.HashMap;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import java.util.Set;
import java.util.stream.Collectors;

/**
* Chat message class that represents all message types (user, system, assistant, tool) with
* different roles.
* A message in a conversation, consisting of a role, an ordered list of content blocks, and
* optional metadata.
*
* <p>Message content is an ordered list of typed {@link ContentBlock}s ({@link TextBlock} plus the
* media blocks); a text-only message simply carries one {@link TextBlock}. The string convenience
* constructors and factories preserve the text-message experience, and {@link #getText()} is the
* ordered concatenation of the text blocks.
* <p>The role identifies the source or purpose of the message. System messages contain text
* instructions, user messages can combine text and media, and assistant messages can include text,
* reasoning, and tool calls. A tool message contains exactly one {@link ToolResultBlock}, which
* associates the result with its tool call.
*
* <p>Blocks are immutable and the message snapshots every block list it is handed, so {@link
* #getBlocks()} is an unmodifiable view: the content changes only by replacing it through {@link
* #setBlocks(List)} or {@link #setText(String)}.
* <p>Content blocks preserve their order within the message, allowing different kinds of content to
* be represented together. {@link #getText()} concatenates only the top-level text blocks;
* reasoning and content nested inside tool results are excluded. {@link #getToolCalls()} returns
* the tool calls in block order. Tool call IDs must be unique within a message.
*
* <p>Metadata holds additional message-level information, such as provider-specific attributes. It
* is separate from the content blocks and is not included by {@link #getText()}.
*/
public class ChatMessage {
public final class ChatMessage {
private static final String ROLE_FIELD = "role";
private static final String BLOCKS_FIELD = "blocks";
private static final String METADATA_FIELD = "metadata";

private static final ObjectMapper MAPPER = new ObjectMapper();
private final MessageRole role;
private final List<ContentBlock> blocks;
private final Map<String, Object> metadata;

private MessageRole role;
private List<ContentBlock> blocks;

@JsonProperty("tool_calls")
private List<Map<String, Object>> toolCalls;

@JsonProperty("extra_args")
private Map<String, Object> extraArgs;

/** Default constructor with SYSTEM role */
public ChatMessage() {
this(MessageRole.SYSTEM, Collections.emptyList(), null, null);
}

/** Constructor with role and text content */
public ChatMessage(MessageRole role, String text) {
this(role, blocksOf(text), null, null);
}

/** Constructor with role and content blocks */
public ChatMessage(MessageRole role, List<ContentBlock> blocks) {
this(role, blocks, null, null);
}

public ChatMessage(MessageRole role, String text, Map<String, Object> extraArgs) {
this(role, blocksOf(text), null, extraArgs);
}

public ChatMessage(MessageRole role, String text, List<Map<String, Object>> toolCalls) {
this(role, blocksOf(text), toolCalls, null);
}

@JsonCreator
public ChatMessage(
MessageRole role,
String text,
List<Map<String, Object>> toolCalls,
Map<String, Object> extraArgs) {
this(role, blocksOf(text), toolCalls, extraArgs);
@JsonProperty(ROLE_FIELD) MessageRole role,
@JsonProperty(BLOCKS_FIELD) List<ContentBlock> blocks,
@JsonProperty(METADATA_FIELD) Map<String, Object> metadata) {
this.role = Objects.requireNonNull(role, ROLE_FIELD);
this.blocks = List.copyOf(Objects.requireNonNull(blocks, BLOCKS_FIELD));
this.metadata = metadata == null ? new HashMap<>() : metadata;
validate();
}

/** Full constructor */
public ChatMessage(
MessageRole role,
List<ContentBlock> blocks,
List<Map<String, Object>> toolCalls,
Map<String, Object> extraArgs) {
this.role = role != null ? role : MessageRole.SYSTEM;
this.blocks = snapshotOf(blocks);
this.toolCalls = toolCalls != null ? toolCalls : new ArrayList<>();
this.extraArgs = extraArgs != null ? new HashMap<>(extraArgs) : new HashMap<>();
public ChatMessage(MessageRole role, List<ContentBlock> blocks) {
this(role, blocks, null);
}

/** An empty or null text becomes an empty block list rather than an empty text block. */
private static List<ContentBlock> blocksOf(String text) {
return text == null || text.isEmpty()
? Collections.emptyList()
: Collections.singletonList(new TextBlock(text));
public ChatMessage(MessageRole role, String text) {
this(role, text == null || text.isEmpty() ? List.of() : List.of(new TextBlock(text)));
}

/** An unmodifiable copy — since blocks are immutable, this freezes the content. */
private static List<ContentBlock> snapshotOf(List<ContentBlock> blocks) {
return List.copyOf(Objects.requireNonNull(blocks, "blocks must not be null"));
private void validate() {
if (role == MessageRole.TOOL
&& (blocks.size() != 1 || !(blocks.get(0) instanceof ToolResultBlock))) {
throw new IllegalArgumentException(
"A TOOL message requires exactly one ToolResultBlock");
}
Set<String> ids = new HashSet<>();
for (ContentBlock b : blocks) {
if (role != MessageRole.TOOL && b instanceof ToolResultBlock) {
throw new IllegalArgumentException("ToolResultBlock requires TOOL role");
}
if (role == MessageRole.SYSTEM && !(b instanceof TextBlock)) {
throw new IllegalArgumentException("SYSTEM messages accept only text");
}
if (role == MessageRole.USER && !(b instanceof TextBlock || b instanceof MediaBlock)) {
throw new IllegalArgumentException("USER messages accept only text and media");
}
if (b instanceof ToolCallBlock && !ids.add(((ToolCallBlock) b).getCallId())) {
throw new IllegalArgumentException("Duplicate tool call ID in one message");
}
}
}

public MessageRole getRole() {
return role;
}

public void setRole(MessageRole role) {
this.role = role;
}

/** The content as an unmodifiable list of immutable blocks. */
public List<ContentBlock> getBlocks() {
return blocks;
}

public void setBlocks(List<ContentBlock> blocks) {
this.blocks = snapshotOf(blocks);
public Map<String, Object> getMetadata() {
return metadata;
}

/** Replaces the content with a single text block (empty text clears the content). */
@JsonIgnore
public void setText(String text) {
this.blocks = blocksOf(text);
}

@JsonProperty("tool_calls")
public List<Map<String, Object>> getToolCalls() {
return toolCalls;
}

@JsonProperty("tool_calls")
public void setToolCalls(List<Map<String, Object>> toolCalls) {
this.toolCalls = toolCalls;
}

@JsonProperty("extra_args")
public Map<String, Object> getExtraArgs() {
return extraArgs;
}

@JsonProperty("extra_args")
public void setExtraArgs(Map<String, Object> extraArgs) {
this.extraArgs = extraArgs != null ? extraArgs : new HashMap<>();
public String getText() {
return blocks.stream()
.filter(b -> b instanceof TextBlock)
.map(b -> ((TextBlock) b).getText())
.collect(Collectors.joining());
}

/**
* The content blocks as plain maps in the serialized (snake_case, discriminated) shape — the
* same representation {@code tool_calls} uses. This is how blocks cross the Python bridge,
* which exchanges JSON-friendly lists and maps rather than typed Java objects.
*/
@JsonIgnore
public List<Map<String, Object>> getBlocksAsMaps() {
public List<ToolCallBlock> getToolCalls() {
return blocks.stream()
.map(
block ->
MAPPER.<Map<String, Object>>convertValue(
block, new TypeReference<Map<String, Object>>() {}))
.collect(Collectors.toList());
.filter(b -> b instanceof ToolCallBlock)
.map(b -> (ToolCallBlock) b)
.collect(Collectors.toUnmodifiableList());
}

/** Replaces the content with blocks given as plain maps — see {@link #getBlocksAsMaps()}. */
@JsonIgnore
public void setBlocksFromMaps(List<Map<String, Object>> blockMaps) {
setBlocks(
Objects.requireNonNull(blockMaps, "blocks must not be null").stream()
.map(map -> MAPPER.convertValue(map, ContentBlock.class))
.collect(Collectors.toList()));
public ChatMessage withBlocks(List<ContentBlock> replacement) {
return new ChatMessage(role, replacement, metadata);
}

/** The text projection: the ordered concatenation of this message's {@link TextBlock}s. */
@JsonIgnore
public String getText() {
return blocks.stream()
.filter(block -> block instanceof TextBlock)
.map(block -> ((TextBlock) block).getText())
.collect(Collectors.joining());
public ChatMessage withMetadata(Map<String, Object> replacement) {
return new ChatMessage(role, blocks, replacement);
}

@JsonIgnore
public Map<String, Object> getMetadata() {
return this.extraArgs;
public Map<String, Object> toMap() {
return MAPPER.convertValue(this, new TypeReference<Map<String, Object>>() {});
}

@JsonIgnore
public MessageRole getMessageType() {
return this.role;
public static ChatMessage fromMap(Map<String, Object> value) {
Map<String, Object> fields = new HashMap<>(value);
fields.putIfAbsent(METADATA_FIELD, Map.of());
if (!fields.containsKey(BLOCKS_FIELD)) {
fields.put(BLOCKS_FIELD, List.of());
}
return MAPPER.convertValue(fields, ChatMessage.class);
}

// Static factory methods for convenience
public static ChatMessage user(String text) {
return new ChatMessage(MessageRole.USER, text);
}
Expand All @@ -215,46 +167,38 @@ public static ChatMessage assistant(String text) {
return new ChatMessage(MessageRole.ASSISTANT, text);
}

public static ChatMessage assistant(String text, List<Map<String, Object>> toolCalls) {
return new ChatMessage(MessageRole.ASSISTANT, text, toolCalls, new HashMap<>());
public static ChatMessage assistant(List<ContentBlock> blocks) {
return new ChatMessage(MessageRole.ASSISTANT, blocks);
}

public static ChatMessage tool(String text) {
return new ChatMessage(MessageRole.TOOL, text);
public static ChatMessage tool(ToolResultBlock result) {
return new ChatMessage(MessageRole.TOOL, List.of(result));
}

public static ChatMessage tool(List<ContentBlock> blocks) {
return new ChatMessage(MessageRole.TOOL, blocks);
public static int findFirstSystemMessage(List<ChatMessage> messages) {
for (int i = 0; i < messages.size(); i++)
if (messages.get(i).role == MessageRole.SYSTEM) {
return i;
}
return -1;
}

@Override
public boolean equals(Object o) {
if (this == o) return true;
if (!(o instanceof ChatMessage)) return false;
ChatMessage that = (ChatMessage) o;
return Objects.equals(role, that.role)
&& Objects.equals(blocks, that.blocks)
&& Objects.equals(toolCalls, that.toolCalls)
&& Objects.equals(extraArgs, that.extraArgs);
if (!(o instanceof ChatMessage)) {
return false;
}
ChatMessage m = (ChatMessage) o;
return role == m.role && blocks.equals(m.blocks) && metadata.equals(m.metadata);
}

@Override
public int hashCode() {
return Objects.hash(role, blocks, toolCalls, extraArgs);
return Objects.hash(role, blocks, metadata);
}

@Override
public String toString() {
return role.getValue() + ": " + getText();
}

/** Return the index of the first system message in the list, or -1 if none. */
public static int findFirstSystemMessage(List<ChatMessage> messages) {
for (int i = 0; i < messages.size(); i++) {
if (messages.get(i).getRole() == MessageRole.SYSTEM) {
return i;
}
}
return -1;
}
}
Loading
Loading