alternative design

This commit is contained in:
vbarda
2024-06-12 15:43:10 -04:00
parent 4239e6fdff
commit 2ad6009160
+26 -18
View File
@@ -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