mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-22 09:35:07 +02:00
alternative design
This commit is contained in:
+26
-18
@@ -1,3 +1,4 @@
|
||||
from itertools import filterfalse
|
||||
import uuid
|
||||
from typing import Annotated, Literal, TypedDict, Union
|
||||
|
||||
@@ -7,19 +8,17 @@ from langchain_core.messages import (
|
||||
convert_to_messages,
|
||||
message_chunk_to_message,
|
||||
)
|
||||
from langchain_core.messages.base import BaseMessage
|
||||
|
||||
from langgraph.graph.state import StateGraph
|
||||
|
||||
Messages = Union[list[MessageLikeRepresentation], MessageLikeRepresentation]
|
||||
|
||||
|
||||
class DeleteMessage(BaseMessage):
|
||||
class MessageModifier(TypedDict):
|
||||
id: str
|
||||
type: Literal["delete"] = "delete"
|
||||
action: Literal["delete"] = "delete"
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
return super().__init__("delete-message", **kwargs)
|
||||
|
||||
Message = Union[MessageLikeRepresentation, MessageModifier]
|
||||
Messages = Union[list[MessageLikeRepresentation], MessageLikeRepresentation]
|
||||
|
||||
|
||||
def add_messages(left: Messages, right: Messages) -> Messages:
|
||||
@@ -73,6 +72,10 @@ def add_messages(left: Messages, right: Messages) -> Messages:
|
||||
left = [left]
|
||||
if not isinstance(right, list):
|
||||
right = [right]
|
||||
|
||||
is_modifier = lambda m: isinstance(m, dict) and m.get("action")
|
||||
message_modifiers = list(filter(is_modifier, right))
|
||||
right = list(filterfalse(is_modifier, right))
|
||||
# coerce to message
|
||||
left = [message_chunk_to_message(m) for m in convert_to_messages(left)]
|
||||
right = [message_chunk_to_message(m) for m in convert_to_messages(right)]
|
||||
@@ -84,21 +87,26 @@ def add_messages(left: Messages, right: Messages) -> Messages:
|
||||
if m.id is None:
|
||||
m.id = str(uuid.uuid4())
|
||||
# merge
|
||||
left_idx_by_id = {m.id: i for i, m in enumerate(left)}
|
||||
existing_idx_by_id = {m.id: i for i, m in enumerate(left)}
|
||||
merged = left.copy()
|
||||
for m in right:
|
||||
if (existing_idx := left_idx_by_id.get(m.id)) is not None:
|
||||
if isinstance(m, DeleteMessage):
|
||||
del merged[existing_idx]
|
||||
else:
|
||||
merged[existing_idx] = m
|
||||
if (existing_idx := existing_idx_by_id.get(m.id)) is not None:
|
||||
merged[existing_idx] = m
|
||||
else:
|
||||
if isinstance(m, DeleteMessage):
|
||||
raise ValueError(
|
||||
f"Attempting to delete a message with an ID that doesn't exist ('{m.id}')"
|
||||
)
|
||||
|
||||
existing_idx_by_id[m.id] = len(merged)
|
||||
merged.append(m)
|
||||
|
||||
for modifier in message_modifiers:
|
||||
if (existing_idx := existing_idx_by_id.get(modifier["id"])) is None:
|
||||
raise ValueError(
|
||||
f"Attempting to modify a message with an ID that doesn't exist ('{modifier['id']}')"
|
||||
)
|
||||
|
||||
if modifier["action"] == "delete":
|
||||
del merged[existing_idx]
|
||||
else:
|
||||
raise ValueError(f"Unsupported modifier action '{modifier['action']}'")
|
||||
|
||||
return merged
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user