fix: Use AsyncSession in memory (#4665)

This commit is contained in:
Christophe Bornet 2024-12-06 17:25:59 +01:00 • committed by GitHub
commit 79b03ba133
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
26 changed files with 610 additions and 173 deletions

View file

@ -1,4 +1,3 @@
import asyncio
import re
from abc import abstractmethod
from typing import TYPE_CHECKING, cast
@ -168,7 +167,7 @@ class LCAgentComponent(Component):
)
except ExceptionWithMessageError as e:
msg_id = e.agent_message.id
await asyncio.to_thread(delete_message, id_=msg_id)
await delete_message(id_=msg_id)
self._send_message_event(e.agent_message, category="remove_message")
raise
except Exception:

View file

@ -1,5 +1,4 @@
# Add helper functions for each event type
import asyncio
from collections.abc import AsyncIterator
from time import perf_counter
from typing import Any, Protocol
@ -53,7 +52,7 @@ def _calculate_duration(start_time: float) -> int:
return result
def handle_on_chain_start(
async def handle_on_chain_start(
event: dict[str, Any], agent_message: Message, send_message_method: SendMessageFunctionType, start_time: float
) -> tuple[Message, float]:
# Create content blocks if they don't exist
@ -75,7 +74,7 @@ def handle_on_chain_start(
header={"title": "Input", "icon": "MessageSquare"},
)
agent_message.content_blocks[0].contents.append(text_content)
agent_message = send_message_method(message=agent_message)
agent_message = await send_message_method(message=agent_message)
start_time = perf_counter()
return agent_message, start_time
@ -91,7 +90,7 @@ def _extract_output_text(output: str | list) -> str:
return text
def handle_on_chain_end(
async def handle_on_chain_end(
event: dict[str, Any], agent_message: Message, send_message_method: SendMessageFunctionType, start_time: float
) -> tuple[Message, float]:
data_output = event["data"].get("output")
@ -110,12 +109,12 @@ def handle_on_chain_end(
header={"title": "Output", "icon": "MessageSquare"},
)
agent_message.content_blocks[0].contents.append(text_content)
agent_message = send_message_method(message=agent_message)
agent_message = await send_message_method(message=agent_message)
start_time = perf_counter()
return agent_message, start_time
def handle_on_tool_start(
async def handle_on_tool_start(
event: dict[str, Any],
agent_message: Message,
tool_blocks_map: dict[str, ToolContent],
@ -149,12 +148,12 @@ def handle_on_tool_start(
tool_blocks_map[tool_key] = tool_content
agent_message.content_blocks[0].contents.append(tool_content)
agent_message = send_message_method(message=agent_message)
agent_message = await send_message_method(message=agent_message)
tool_blocks_map[tool_key] = agent_message.content_blocks[0].contents[-1]
return agent_message, new_start_time
def handle_on_tool_end(
async def handle_on_tool_end(
event: dict[str, Any],
agent_message: Message,
tool_blocks_map: dict[str, ToolContent],
@ -172,13 +171,13 @@ def handle_on_tool_end(
tool_content.duration = duration
tool_content.header = {"title": f"Executed **{tool_content.name}**", "icon": "Hammer"}
agent_message = send_message_method(message=agent_message)
agent_message = await send_message_method(message=agent_message)
new_start_time = perf_counter() # Get new start time for next operation
return agent_message, new_start_time
return agent_message, start_time
def handle_on_tool_error(
async def handle_on_tool_error(
event: dict[str, Any],
agent_message: Message,
tool_blocks_map: dict[str, ToolContent],
@ -194,12 +193,12 @@ def handle_on_tool_error(
tool_content.error = event["data"].get("error", "Unknown error")
tool_content.duration = _calculate_duration(start_time)
tool_content.header = {"title": f"Error using **{tool_content.name}**", "icon": "Hammer"}
agent_message = send_message_method(message=agent_message)
agent_message = await send_message_method(message=agent_message)
start_time = perf_counter()
return agent_message, start_time
def handle_on_chain_stream(
async def handle_on_chain_stream(
event: dict[str, Any],
agent_message: Message,
send_message_method: SendMessageFunctionType,
@ -211,13 +210,13 @@ def handle_on_chain_stream(
if output and isinstance(output, str | list):
agent_message.text = _extract_output_text(output)
agent_message.properties.state = "complete"
agent_message = send_message_method(message=agent_message)
agent_message = await send_message_method(message=agent_message)
start_time = perf_counter()
return agent_message, start_time
class ToolEventHandler(Protocol):
def __call__(
async def __call__(
self,
event: dict[str, Any],
agent_message: Message,
@ -228,7 +227,7 @@ class ToolEventHandler(Protocol):
class ChainEventHandler(Protocol):
def __call__(
async def __call__(
self,
event: dict[str, Any],
agent_message: Message,
@ -265,7 +264,7 @@ async def process_agent_events(
agent_message.properties.icon = "Bot"
agent_message.properties.state = "partial"
# Store the initial message
agent_message = await asyncio.to_thread(send_message_method, message=agent_message)
agent_message = await send_message_method(message=agent_message)
try:
# Create a mapping of run_ids to tool contents
tool_blocks_map: dict[str, ToolContent] = {}
@ -273,14 +272,14 @@ async def process_agent_events(
async for event in agent_executor:
if event["event"] in TOOL_EVENT_HANDLERS:
tool_handler = TOOL_EVENT_HANDLERS[event["event"]]
agent_message, start_time = tool_handler(
agent_message, start_time = await tool_handler(
event, agent_message, tool_blocks_map, send_message_method, start_time
)
elif event["event"] in CHAIN_EVENT_HANDLERS:
chain_handler = CHAIN_EVENT_HANDLERS[event["event"]]
agent_message, start_time = chain_handler(event, agent_message, send_message_method, start_time)
agent_message, start_time = await chain_handler(event, agent_message, send_message_method, start_time)
agent_message.properties.state = "complete"
except Exception as e:
raise ExceptionWithMessageError(agent_message) from e
return Message(**agent_message.model_dump())
return await Message.create(**agent_message.model_dump())

View file

@ -1,7 +1,8 @@
import asyncio
from typing import cast
from langflow.custom import Component
from langflow.memory import store_message
from langflow.memory import astore_message
from langflow.schema import Data
from langflow.schema.message import Message
@ -10,7 +11,7 @@ class ChatComponent(Component):
display_name = "Chat Component"
description = "Use as base for chat components."
def build_with_data(
async def build_with_data(
self,
*,
sender: str | None = "User",
@ -20,13 +21,13 @@ class ChatComponent(Component):
session_id: str | None = None,
return_message: bool = False,
) -> str | Message:
message = self._create_message(input_value, sender, sender_name, files, session_id)
message = await asyncio.to_thread(self._create_message, input_value, sender, sender_name, files, session_id)
message_text = message.text if not return_message else message
self.status = message_text
if session_id and isinstance(message, Message) and isinstance(message.text, str):
flow_id = self.graph.flow_id if hasattr(self, "graph") else None
messages = store_message(message, flow_id=flow_id)
messages = await astore_message(message, flow_id=flow_id)
self.status = messages
self._send_messages_events(messages)

View file

@ -56,7 +56,7 @@ def build_description(component: Component, output: Output) -> str:
return f"{output.method}({args}) - {component.description}"
def send_message_noop(
async def send_message_noop(
message: Message,
text: str | None = None, # noqa: ARG001
background_color: str | None = None, # noqa: ARG001

View file

@ -66,7 +66,7 @@ class AgentComponent(ToolCallingAgentComponent):
if llm_model is None:
msg = "No language model selected"
raise ValueError(msg)
self.chat_history = self.get_memory_data()
self.chat_history = await self.get_memory_data()
if self.add_current_date_tool:
if not isinstance(self.tools, list): # type: ignore[has-type]
@ -92,12 +92,12 @@ class AgentComponent(ToolCallingAgentComponent):
agent = self.create_agent_runnable()
return await self.run_agent(agent)
def get_memory_data(self):
async def get_memory_data(self):
memory_kwargs = {
component_input.name: getattr(self, f"{component_input.name}") for component_input in self.memory_inputs
}
return MemoryComponent().set(**memory_kwargs).retrieve_messages()
return await MemoryComponent().set(**memory_kwargs).retrieve_messages()
def get_llm(self):
if isinstance(self.agent_llm, str):

View file

@ -1,5 +1,5 @@
from langflow.custom import CustomComponent
from langflow.memory import get_messages, store_message
from langflow.memory import aget_messages, astore_message
from langflow.schema.message import Message
@ -13,12 +13,12 @@ class StoreMessageComponent(CustomComponent):
"message": {"display_name": "Message"},
}
def build(
async def build(
self,
message: Message,
) -> Message:
flow_id = self.graph.flow_id if hasattr(self, "graph") else None
store_message(message, flow_id=flow_id)
self.status = get_messages()
await astore_message(message, flow_id=flow_id)
self.status = await aget_messages()
return message

View file

@ -5,7 +5,7 @@ from langflow.field_typing import BaseChatMemory
from langflow.helpers.data import data_to_text
from langflow.inputs import HandleInput
from langflow.io import DropdownInput, IntInput, MessageTextInput, MultilineInput, Output
from langflow.memory import LCBuiltinChatMemory, get_messages
from langflow.memory import LCBuiltinChatMemory, aget_messages
from langflow.schema import Data
from langflow.schema.message import Message
from langflow.utils.constants import MESSAGE_SENDER_AI, MESSAGE_SENDER_USER
@ -74,7 +74,7 @@ class MemoryComponent(Component):
Output(display_name="Text", name="messages_text", method="retrieve_messages_as_text"),
]
def retrieve_messages(self) -> Data:
async def retrieve_messages(self) -> Data:
sender = self.sender
sender_name = self.sender_name
session_id = self.session_id
@ -88,7 +88,7 @@ class MemoryComponent(Component):
# override session_id
self.memory.session_id = session_id
stored = self.memory.messages
stored = await self.memory.aget_messages()
# langchain memories are supposed to return messages in ascending order
if order == "DESC":
stored = stored[::-1]
@ -99,7 +99,7 @@ class MemoryComponent(Component):
expected_type = MESSAGE_SENDER_AI if sender == MESSAGE_SENDER_AI else MESSAGE_SENDER_USER
stored = [m for m in stored if m.type == expected_type]
else:
stored = get_messages(
stored = await aget_messages(
sender=sender,
sender_name=sender_name,
session_id=session_id,
@ -109,8 +109,8 @@ class MemoryComponent(Component):
self.status = stored
return stored
def retrieve_messages_as_text(self) -> Message:
stored_text = data_to_text(self.template, self.retrieve_messages())
async def retrieve_messages_as_text(self) -> Message:
stored_text = data_to_text(self.template, await self.retrieve_messages())
self.status = stored_text
return Message(text=stored_text)

View file

@ -1,7 +1,7 @@
from langflow.custom import Component
from langflow.inputs import HandleInput, MessageInput
from langflow.inputs.inputs import MessageTextInput
from langflow.memory import get_messages, store_message
from langflow.memory import aget_messages, astore_message
from langflow.schema.message import Message
from langflow.template import Output
from langflow.utils.constants import MESSAGE_SENDER_AI, MESSAGE_SENDER_NAME_AI
@ -47,7 +47,7 @@ class StoreMessageComponent(Component):
Output(display_name="Stored Messages", name="stored_messages", method="store_message"),
]
def store_message(self) -> Message:
async def store_message(self) -> Message:
message = self.message
message.session_id = self.session_id or message.session_id
@ -58,13 +58,15 @@ class StoreMessageComponent(Component):
# override session_id
self.memory.session_id = message.session_id
lc_message = message.to_lc_message()
self.memory.add_messages([lc_message])
stored = self.memory.messages
await self.memory.aadd_messages([lc_message])
stored = await self.memory.aget_messages()
stored = [Message.from_lc_message(m) for m in stored]
if message.sender:
stored = [m for m in stored if m.sender == message.sender]
else:
store_message(message, flow_id=self.graph.flow_id)
stored = get_messages(session_id=message.session_id, sender_name=message.sender_name, sender=message.sender)
await astore_message(message, flow_id=self.graph.flow_id)
stored = await aget_messages(
session_id=message.session_id, sender_name=message.sender_name, sender=message.sender
)
self.status = stored
return stored

View file

@ -78,11 +78,12 @@ class ChatInput(ChatComponent):
Output(display_name="Message", name="message", method="message_response"),
]
def message_response(self) -> Message:
async def message_response(self) -> Message:
_background_color = self.background_color
_text_color = self.text_color
_icon = self.chat_icon
message = Message(
message = await Message.create(
text=self.input_value,
sender=self.sender,
sender_name=self.sender_name,
@ -91,7 +92,7 @@ class ChatInput(ChatComponent):
properties={"background_color": _background_color, "text_color": _text_color, "icon": _icon},
)
if self.session_id and isinstance(message, Message) and self.should_store_message:
stored_message = self.send_message(
stored_message = await self.send_message(
message,
)
self.message.value = stored_message

View file

@ -90,7 +90,7 @@ class ChatOutput(ChatComponent):
source_dict["source"] = source
return Source(**source_dict)
def message_response(self) -> Message:
async def message_response(self) -> Message:
_source, _icon, _display_name, _source_id = self.get_properties_from_source_component()
_background_color = self.background_color
_text_color = self.text_color
@ -106,7 +106,7 @@ class ChatOutput(ChatComponent):
message.properties.background_color = _background_color
message.properties.text_color = _text_color
if self.session_id and isinstance(message, Message) and self.should_store_message:
stored_message = self.send_message(
stored_message = await self.send_message(
message,
)
self.message.value = stored_message

View file

@ -24,7 +24,7 @@ from langflow.exceptions.component import StreamingError
from langflow.field_typing import Tool # noqa: TCH001 Needed by _add_toolkit_output
from langflow.graph.state.model import create_state_model
from langflow.helpers.custom import format_type
from langflow.memory import delete_message, store_message, update_messages
from langflow.memory import astore_message, aupdate_messages, delete_message
from langflow.schema.artifact import get_artifact_type, post_process_raw
from langflow.schema.data import Data
from langflow.schema.message import ErrorMessage, Message
@ -847,7 +847,7 @@ class Component(CustomComponent):
return await self._build_with_tracing()
return await self._build_without_tracing()
except StreamingError as e:
self.send_error(
await self.send_error(
exception=e.cause,
session_id=session_id,
trace_name=getattr(self, "trace_name", None),
@ -855,7 +855,7 @@ class Component(CustomComponent):
)
raise e.cause # noqa: B904
except Exception as e:
self.send_error(
await self.send_error(
exception=e,
session_id=session_id,
source=Source(id=self._id, display_name=self.display_name, source=self.display_name),
@ -1016,10 +1016,10 @@ class Component(CustomComponent):
)
)
def send_message(self, message: Message, id_: str | None = None):
async def send_message(self, message: Message, id_: str | None = None):
if (hasattr(self, "graph") and self.graph.session_id) and (message is not None and not message.session_id):
message.session_id = self.graph.session_id
stored_message = self._store_message(message)
stored_message = await self._store_message(message)
self._stored_message_id = stored_message.id
try:
@ -1029,22 +1029,22 @@ class Component(CustomComponent):
and message is not None
and isinstance(message.text, AsyncIterator | Iterator)
):
complete_message = self._stream_message(message.text, stored_message)
complete_message = await self._stream_message(message.text, stored_message)
stored_message.text = complete_message
stored_message = self._update_stored_message(stored_message)
stored_message = await self._update_stored_message(stored_message)
else:
# Only send message event for non-streaming messages
self._send_message_event(stored_message, id_=id_)
except Exception:
# remove the message from the database
delete_message(stored_message.id)
await delete_message(stored_message.id)
raise
self.status = stored_message
return stored_message
def _store_message(self, message: Message) -> Message:
async def _store_message(self, message: Message) -> Message:
flow_id = self.graph.flow_id if hasattr(self, "graph") else None
messages = store_message(message, flow_id=flow_id)
messages = await astore_message(message, flow_id=flow_id)
if len(messages) != 1:
msg = "Only one message can be stored at a time."
raise ValueError(msg)
@ -1073,21 +1073,21 @@ class Component(CustomComponent):
and not isinstance(original_message.text, str)
)
def _update_stored_message(self, stored_message: Message) -> Message:
message_tables = update_messages(stored_message)
async def _update_stored_message(self, stored_message: Message) -> Message:
message_tables = await aupdate_messages(stored_message)
if len(message_tables) != 1:
msg = "Only one message can be updated at a time."
raise ValueError(msg)
message_table = message_tables[0]
return Message(**message_table.model_dump())
return await Message.create(**message_table.model_dump())
def _stream_message(self, iterator: AsyncIterator | Iterator, message: Message) -> str:
async def _stream_message(self, iterator: AsyncIterator | Iterator, message: Message) -> str:
if not isinstance(iterator, AsyncIterator | Iterator):
msg = "The message must be an iterator or an async iterator."
raise TypeError(msg)
if isinstance(iterator, AsyncIterator):
return run_until_complete(self._handle_async_iterator(iterator, message.id, message))
return await self._handle_async_iterator(iterator, message.id, message)
try:
complete_message = ""
first_chunk = True
@ -1129,7 +1129,7 @@ class Component(CustomComponent):
)
return complete_message
def send_error(
async def send_error(
self,
exception: Exception,
session_id: str,
@ -1145,7 +1145,7 @@ class Component(CustomComponent):
trace_name=trace_name,
source=source,
)
self.send_message(error_message)
await self.send_message(error_message)
return error_message
def _append_tool_to_outputs_map(self):

View file

@ -401,7 +401,7 @@ class InterfaceVertex(ComponentVertex):
type=ArtifactType.OBJECT.value,
).model_dump()
message = Message(
message = await Message.create(
text=complete_message,
sender=self.params.get("sender", ""),
sender_name=self.params.get("sender_name", ""),
@ -434,7 +434,7 @@ class InterfaceVertex(ComponentVertex):
and hasattr(self.custom_component, "should_store_message")
and hasattr(self.custom_component, "store_message")
):
self.custom_component.store_message(message)
await self.custom_component.store_message(message)
await log_vertex_build(
flow_id=self.graph.flow_id,
vertex_id=self.id,

View file

@ -7,13 +7,40 @@ from langchain_core.messages import BaseMessage
from loguru import logger
from sqlalchemy import delete
from sqlmodel import Session, col, select
from sqlmodel.ext.asyncio.session import AsyncSession
from langflow.schema.message import Message
from langflow.services.database.models.message.model import MessageRead, MessageTable
from langflow.services.deps import session_scope
from langflow.services.deps import async_session_scope, session_scope
from langflow.utils.constants import MESSAGE_SENDER_AI, MESSAGE_SENDER_USER
def _get_variable_query(
sender: str | None = None,
sender_name: str | None = None,
session_id: str | None = None,
order_by: str | None = "timestamp",
order: str | None = "DESC",
flow_id: UUID | None = None,
limit: int | None = None,
):
stmt = select(MessageTable).where(MessageTable.error == False) # noqa: E712
if sender:
stmt = stmt.where(MessageTable.sender == sender)
if sender_name:
stmt = stmt.where(MessageTable.sender_name == sender_name)
if session_id:
stmt = stmt.where(MessageTable.session_id == session_id)
if flow_id:
stmt = stmt.where(MessageTable.flow_id == flow_id)
if order_by:
col = getattr(MessageTable, order_by).desc() if order == "DESC" else getattr(MessageTable, order_by).asc()
stmt = stmt.order_by(col)
if limit:
stmt = stmt.limit(limit)
return stmt
def get_messages(
sender: str | None = None,
sender_name: str | None = None,
@ -38,24 +65,40 @@ def get_messages(
List[Data]: A list of Data objects representing the retrieved messages.
"""
with session_scope() as session:
stmt = select(MessageTable).where(MessageTable.error == False) # noqa: E712
if sender:
stmt = stmt.where(MessageTable.sender == sender)
if sender_name:
stmt = stmt.where(MessageTable.sender_name == sender_name)
if session_id:
stmt = stmt.where(MessageTable.session_id == session_id)
if flow_id:
stmt = stmt.where(MessageTable.flow_id == flow_id)
if order_by:
col = getattr(MessageTable, order_by).desc() if order == "DESC" else getattr(MessageTable, order_by).asc()
stmt = stmt.order_by(col)
if limit:
stmt = stmt.limit(limit)
stmt = _get_variable_query(sender, sender_name, session_id, order_by, order, flow_id, limit)
messages = session.exec(stmt)
return [Message(**d.model_dump()) for d in messages]
async def aget_messages(
sender: str | None = None,
sender_name: str | None = None,
session_id: str | None = None,
order_by: str | None = "timestamp",
order: str | None = "DESC",
flow_id: UUID | None = None,
limit: int | None = None,
) -> list[Message]:
"""Retrieves messages from the monitor service based on the provided filters.
Args:
sender (Optional[str]): The sender of the messages (e.g., "Machine" or "User")
sender_name (Optional[str]): The name of the sender.
session_id (Optional[str]): The session ID associated with the messages.
order_by (Optional[str]): The field to order the messages by. Defaults to "timestamp".
order (Optional[str]): The order in which to retrieve the messages. Defaults to "DESC".
flow_id (Optional[UUID]): The flow ID associated with the messages.
limit (Optional[int]): The maximum number of messages to retrieve.
Returns:
List[Data]: A list of Data objects representing the retrieved messages.
"""
async with async_session_scope() as session:
stmt = _get_variable_query(sender, sender_name, session_id, order_by, order, flow_id, limit)
messages = await session.exec(stmt)
return [await Message.create(**d.model_dump()) for d in messages]
def add_messages(messages: Message | list[Message], flow_id: str | None = None):
"""Add a message to the monitor service."""
if not isinstance(messages, list):
@ -76,6 +119,26 @@ def add_messages(messages: Message | list[Message], flow_id: str | None = None):
raise
async def aadd_messages(messages: Message | list[Message], flow_id: str | None = None):
"""Add a message to the monitor service."""
if not isinstance(messages, list):
messages = [messages]
if not all(isinstance(message, Message) for message in messages):
types = ", ".join([str(type(message)) for message in messages])
msg = f"The messages must be instances of Message. Found: {types}"
raise ValueError(msg)
try:
messages_models = [MessageTable.from_message(msg, flow_id=flow_id) for msg in messages]
async with async_session_scope() as session:
messages_models = await aadd_messagetables(messages_models, session)
return [await Message.create(**message.model_dump()) for message in messages_models]
except Exception as e:
logger.exception(e)
raise
def update_messages(messages: Message | list[Message]) -> list[Message]:
if not isinstance(messages, list):
messages = [messages]
@ -95,6 +158,25 @@ def update_messages(messages: Message | list[Message]) -> list[Message]:
return [MessageRead.model_validate(message, from_attributes=True) for message in updated_messages]
async def aupdate_messages(messages: Message | list[Message]) -> list[Message]:
if not isinstance(messages, list):
messages = [messages]
async with async_session_scope() as session:
updated_messages: list[MessageTable] = []
for message in messages:
msg = await session.get(MessageTable, message.id)
if msg:
msg.sqlmodel_update(message.model_dump(exclude_unset=True, exclude_none=True))
session.add(msg)
await session.commit()
await session.refresh(msg)
updated_messages.append(msg)
else:
logger.warning(f"Message with id {message.id} not found")
return [MessageRead.model_validate(message, from_attributes=True) for message in updated_messages]
def add_messagetables(messages: list[MessageTable], session: Session):
for message in messages:
try:
@ -115,6 +197,27 @@ def add_messagetables(messages: list[MessageTable], session: Session):
return [MessageRead.model_validate(message, from_attributes=True) for message in new_messages]
async def aadd_messagetables(messages: list[MessageTable], session: AsyncSession):
try:
for message in messages:
session.add(message)
await session.commit()
for message in messages:
await session.refresh(message)
except Exception as e:
logger.exception(e)
raise
new_messages = []
for msg in messages:
msg.properties = json.loads(msg.properties) if isinstance(msg.properties, str) else msg.properties # type: ignore[arg-type]
msg.content_blocks = [json.loads(j) if isinstance(j, str) else j for j in msg.content_blocks] # type: ignore[arg-type]
msg.category = msg.category or ""
new_messages.append(msg)
return [MessageRead.model_validate(message, from_attributes=True) for message in new_messages]
def delete_messages(session_id: str) -> None:
"""Delete messages from the monitor service based on the provided session ID.
@ -129,17 +232,32 @@ def delete_messages(session_id: str) -> None:
)
def delete_message(id_: str) -> None:
async def adelete_messages(session_id: str) -> None:
"""Delete messages from the monitor service based on the provided session ID.
Args:
session_id (str): The session ID associated with the messages to delete.
"""
async with async_session_scope() as session:
stmt = (
delete(MessageTable)
.where(col(MessageTable.session_id) == session_id)
.execution_options(synchronize_session="fetch")
)
await session.exec(stmt)
async def delete_message(id_: str) -> None:
"""Delete a message from the monitor service based on the provided ID.
Args:
id_ (str): The ID of the message to delete.
"""
with session_scope() as session:
message = session.get(MessageTable, id_)
async with async_session_scope() as session:
message = await session.get(MessageTable, id_)
if message:
session.delete(message)
session.commit()
await session.delete(message)
await session.commit()
def store_message(
@ -182,6 +300,35 @@ def store_message(
return add_messages([message], flow_id=flow_id)
async def astore_message(
message: Message,
flow_id: str | None = None,
) -> list[Message]:
"""Stores a message in the memory.
Args:
message (Message): The message to store.
flow_id (Optional[str]): The flow ID associated with the message.
When running from the CustomComponent you can access this using `self.graph.flow_id`.
Returns:
List[Message]: A list of data containing the stored message.
Raises:
ValueError: If any of the required parameters (session_id, sender, sender_name) is not provided.
"""
if not message:
logger.warning("No message provided.")
return []
if not message.session_id or not message.sender or not message.sender_name:
msg = "All of session_id, sender, and sender_name must be provided."
raise ValueError(msg)
if hasattr(message, "id") and message.id:
return await aupdate_messages([message])
return await aadd_messages([message], flow_id=flow_id)
class LCBuiltinChatMemory(BaseChatMessageHistory):
def __init__(
self,
@ -198,11 +345,26 @@ class LCBuiltinChatMemory(BaseChatMessageHistory):
)
return [m.to_lc_message() for m in messages if not m.error] # Exclude error messages
async def aget_messages(self) -> list[BaseMessage]:
messages = await aget_messages(
session_id=self.session_id,
)
return [m.to_lc_message() for m in messages if not m.error] # Exclude error messages
def add_messages(self, messages: Sequence[BaseMessage]) -> None:
for lc_message in messages:
message = Message.from_lc_message(lc_message)
message.session_id = self.session_id
store_message(message, flow_id=self.flow_id)
async def aadd_messages(self, messages: Sequence[BaseMessage]) -> None:
for lc_message in messages:
message = Message.from_lc_message(lc_message)
message.session_id = self.session_id
await astore_message(message, flow_id=self.flow_id)
def clear(self) -> None:
delete_messages(self.session_id)
async def aclear(self) -> None:
await adelete_messages(self.session_id)

View file

@ -14,7 +14,7 @@ class LogFunctionType(Protocol):
class SendMessageFunctionType(Protocol):
def __call__(
async def __call__(
self,
message: Message | None = None,
text: str | None = None,

View file

@ -1,5 +1,6 @@
from __future__ import annotations
import asyncio
import json
import re
import traceback
@ -267,6 +268,13 @@ class Message(Data):
instance.messages = instance.prompt.get("kwargs", {}).get("messages", [])
return instance
@classmethod
async def create(cls, **kwargs):
"""If files are present, create the message in a separate thread as is_image_file is blocking."""
if "files" in kwargs:
return await asyncio.to_thread(cls, **kwargs)
return cls(**kwargs)
class DefaultModel(BaseModel):
class Config: