From d2875cc5760598c8d3e62afdb28c6a17c9403f11 Mon Sep 17 00:00:00 2001 From: Tat Dat Duong Date: Fri, 16 May 2025 06:00:43 -0700 Subject: [PATCH 1/3] feat(graph): add push_message method to push manually to `messages` / `message-tuple` stream --- libs/langgraph/langgraph/graph/message.py | 41 +++++++++++++++++++++++ 1 file changed, 41 insertions(+) diff --git a/libs/langgraph/langgraph/graph/message.py b/libs/langgraph/langgraph/graph/message.py index ee601d6c5..0282f1fb6 100644 --- a/libs/langgraph/langgraph/graph/message.py +++ b/libs/langgraph/langgraph/graph/message.py @@ -294,3 +294,44 @@ def _format_messages(messages: Sequence[BaseMessage]) -> list[BaseMessage]: return list(messages) else: return convert_to_messages(convert_to_openai_messages(messages)) + + +def push_message( + message: Union[MessageLikeRepresentation, BaseMessageChunk], +) -> AnyMessage: + """Write a message manually to the `messages` / `messages-tuple` stream mode.""" + + from langchain_core.callbacks.base import ( + BaseCallbackHandler, + BaseCallbackManager, + ) + + from langgraph.config import get_config + from langgraph.constants import NS_SEP + from langgraph.pregel.messages import StreamMessagesHandler + + config = get_config() + message = next(x for x in convert_to_messages([message])) + + if message.id is None: + raise ValueError("Message ID is required") + + if isinstance(config["callbacks"], BaseCallbackManager): + manager = cast(BaseCallbackManager, config["callbacks"]) + handlers = manager.handlers + elif isinstance(config["callbacks"], list) and all( + isinstance(x, BaseCallbackHandler) for x in config["callbacks"] + ): + handlers = config["callbacks"] + + if stream_handler := next( + (x for x in handlers if isinstance(x, StreamMessagesHandler)), None + ): + metadata = config["metadata"] + message_meta = ( + tuple(cast(str, metadata["langgraph_checkpoint_ns"]).split(NS_SEP)), + metadata, + ) + stream_handler._emit(message_meta, message, dedupe=False) + + return message From 7a6bdb34410356a007abd7cfe71ee239c19bc53b Mon Sep 17 00:00:00 2001 From: Tat Dat Duong Date: Thu, 22 May 2025 00:23:26 +0200 Subject: [PATCH 2/3] Remove redundant cast --- libs/langgraph/langgraph/graph/message.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/libs/langgraph/langgraph/graph/message.py b/libs/langgraph/langgraph/graph/message.py index 0282f1fb6..16077eeb7 100644 --- a/libs/langgraph/langgraph/graph/message.py +++ b/libs/langgraph/langgraph/graph/message.py @@ -317,7 +317,7 @@ def push_message( raise ValueError("Message ID is required") if isinstance(config["callbacks"], BaseCallbackManager): - manager = cast(BaseCallbackManager, config["callbacks"]) + manager = config["callbacks"] handlers = manager.handlers elif isinstance(config["callbacks"], list) and all( isinstance(x, BaseCallbackHandler) for x in config["callbacks"] From ab924c72fa90a058dc503dbf31439f4f32563cfe Mon Sep 17 00:00:00 2001 From: Tat Dat Duong Date: Thu, 22 May 2025 00:49:23 +0200 Subject: [PATCH 3/3] Add message state test --- libs/langgraph/tests/test_messages_state.py | 37 ++++++++++++++++++++- 1 file changed, 36 insertions(+), 1 deletion(-) diff --git a/libs/langgraph/tests/test_messages_state.py b/libs/langgraph/tests/test_messages_state.py index 882b0cfed..4315cd25e 100644 --- a/libs/langgraph/tests/test_messages_state.py +++ b/libs/langgraph/tests/test_messages_state.py @@ -15,7 +15,7 @@ from pydantic import BaseModel from typing_extensions import TypedDict from langgraph.graph import add_messages -from langgraph.graph.message import REMOVE_ALL_MESSAGES, MessagesState +from langgraph.graph.message import REMOVE_ALL_MESSAGES, MessagesState, push_message from langgraph.graph.state import END, START, StateGraph from tests.messages import _AnyIdHumanMessage @@ -332,3 +332,38 @@ def test_remove_all_messages(): assert result == [ _AnyIdHumanMessage(content="Updated hi there"), ] + + +def test_push_messages_in_graph(): + class MessagesState(TypedDict): + messages: Annotated[list[AnyMessage], add_messages] + + def chat(_: MessagesState) -> MessagesState: + with pytest.raises(ValueError, match="Message ID is required"): + push_message(AIMessage(content="No ID")) + + return { + "messages": [ + push_message(AIMessage(content="First", id="1")), + push_message(HumanMessage(content="Second", id="2")), + push_message(AIMessage(content="Third", id="3")), + ] + } + + builder = StateGraph(MessagesState) + builder.add_node(chat) + builder.add_edge(START, "chat") + + graph = builder.compile() + + messages, values = [], None + for event, chunk in graph.stream( + {"messages": []}, stream_mode=["messages", "values"] + ): + if event == "values": + values = chunk + elif event == "messages": + message, _ = chunk + messages.append(message) + + assert values["messages"] == messages