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

@ -29,6 +29,7 @@ from langflow.services.database.models.vertex_builds.crud import delete_vertex_b
from langflow.services.database.utils import session_getter
from langflow.services.deps import get_db_service
from loguru import logger
from sqlalchemy.ext.asyncio import create_async_engine
from sqlalchemy.orm import selectinload
from sqlmodel import Session, SQLModel, create_engine, select
from sqlmodel.ext.asyncio.session import AsyncSession
@ -151,6 +152,17 @@ def session_fixture():
SQLModel.metadata.drop_all(engine) # Add this line to clean up tables
@pytest.fixture
async def async_session():
engine = create_async_engine("sqlite+aiosqlite://", connect_args={"check_same_thread": False}, poolclass=StaticPool)
async with engine.begin() as conn:
await conn.run_sync(SQLModel.metadata.create_all)
async with AsyncSession(engine) as session:
yield session
async with engine.begin() as conn:
await conn.run_sync(SQLModel.metadata.drop_all)
class Config:
broker_url = "redis://localhost:6379/0"
result_backend = "redis://localhost:6379/0"

View file

@ -1,5 +1,5 @@
from langflow.components.inputs import ChatInput
from langflow.memory import get_messages
from langflow.memory import aget_messages
from langflow.schema.message import Message
from tests.integration.utils import run_single_component
@ -38,7 +38,7 @@ async def test_do_not_store_messages():
assert outputs["message"].text == "hello"
assert outputs["message"].session_id == session_id
assert len(get_messages(session_id=session_id)) == 1
assert len(await aget_messages(session_id=session_id)) == 1
session_id = "test-session-id-another"
outputs = await run_single_component(
@ -48,4 +48,4 @@ async def test_do_not_store_messages():
assert outputs["message"].text == "hello"
assert outputs["message"].session_id == session_id
assert len(get_messages(session_id=session_id)) == 0
assert len(await aget_messages(session_id=session_id)) == 0

View file

@ -1,5 +1,5 @@
from langflow.components.outputs import ChatOutput
from langflow.memory import get_messages
from langflow.memory import aget_messages
from langflow.schema.message import Message
from tests.integration.utils import run_single_component
@ -29,7 +29,7 @@ async def test_do_not_store_message():
assert isinstance(outputs["message"], Message)
assert outputs["message"].text == "hello"
assert len(get_messages(session_id=session_id)) == 1
assert len(await aget_messages(session_id=session_id)) == 1
session_id = "test-session-id-another"
outputs = await run_single_component(
@ -38,4 +38,4 @@ async def test_do_not_store_message():
assert isinstance(outputs["message"], Message)
assert outputs["message"].text == "hello"
assert len(get_messages(session_id=session_id)) == 0
assert len(await aget_messages(session_id=session_id)) == 0

View file

@ -1,6 +1,6 @@
from collections.abc import AsyncIterator
from typing import Any
from unittest.mock import MagicMock
from unittest.mock import AsyncMock
from langchain_core.agents import AgentFinish
from langflow.base.agents.agent import process_agent_events
@ -26,7 +26,7 @@ async def create_event_iterator(events: list[dict[str, Any]]) -> AsyncIterator[d
async def test_chain_start_event():
"""Test handling of on_chain_start event."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
events = [
{"event": "on_chain_start", "data": {"input": {"input": "test input", "chat_history": []}}, "start_time": 0}
@ -51,7 +51,7 @@ async def test_chain_start_event():
async def test_chain_end_event():
"""Test handling of on_chain_end event."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
# Create a mock AgentFinish output
output = AgentFinish(return_values={"output": "final output"}, log="test log")
@ -77,7 +77,7 @@ async def test_chain_end_event():
async def test_tool_start_event():
"""Test handling of on_tool_start event."""
send_message = MagicMock()
send_message = AsyncMock()
# Set up the send_message mock to return the modified message
def update_message(message):
@ -116,7 +116,7 @@ async def test_tool_start_event():
async def test_tool_end_event():
"""Test handling of on_tool_end event."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
events = [
{
@ -151,7 +151,7 @@ async def test_tool_end_event():
async def test_tool_error_event():
"""Test handling of on_tool_error event."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
events = [
{
@ -187,7 +187,7 @@ async def test_tool_error_event():
async def test_chain_stream_event():
"""Test handling of on_chain_stream event."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
events = [{"event": "on_chain_stream", "data": {"chunk": {"output": "streamed output"}}, "start_time": 0}]
agent_message = Message(
@ -205,7 +205,7 @@ async def test_chain_stream_event():
async def test_multiple_events():
"""Test handling of multiple events in sequence."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
# Create a mock AgentFinish output instead of MockOutput
output = AgentFinish(return_values={"output": "final output"}, log="test log")
@ -248,7 +248,7 @@ async def test_multiple_events():
async def test_unknown_event():
"""Test handling of unknown event type."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
agent_message = Message(
sender=MESSAGE_SENDER_AI,
sender_name="Agent",
@ -273,7 +273,7 @@ async def test_unknown_event():
async def test_handle_on_chain_start_with_input():
"""Test handle_on_chain_start with input."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
agent_message = Message(
sender=MESSAGE_SENDER_AI,
sender_name="Agent",
@ -282,7 +282,7 @@ async def test_handle_on_chain_start_with_input():
)
event = {"event": "on_chain_start", "data": {"input": {"input": "test input", "chat_history": []}}, "start_time": 0}
updated_message, start_time = handle_on_chain_start(event, agent_message, send_message, 0.0)
updated_message, start_time = await handle_on_chain_start(event, agent_message, send_message, 0.0)
assert updated_message.properties.icon == "Bot"
assert len(updated_message.content_blocks) == 1
@ -292,7 +292,7 @@ async def test_handle_on_chain_start_with_input():
async def test_handle_on_chain_start_no_input():
"""Test handle_on_chain_start without input."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
agent_message = Message(
sender=MESSAGE_SENDER_AI,
sender_name="Agent",
@ -301,7 +301,7 @@ async def test_handle_on_chain_start_no_input():
)
event = {"event": "on_chain_start", "data": {}, "start_time": 0}
updated_message, start_time = handle_on_chain_start(event, agent_message, send_message, 0.0)
updated_message, start_time = await handle_on_chain_start(event, agent_message, send_message, 0.0)
assert updated_message.properties.icon == "Bot"
assert len(updated_message.content_blocks) == 1
@ -311,7 +311,7 @@ async def test_handle_on_chain_start_no_input():
async def test_handle_on_chain_end_with_output():
"""Test handle_on_chain_end with output."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
agent_message = Message(
sender=MESSAGE_SENDER_AI,
sender_name="Agent",
@ -322,7 +322,7 @@ async def test_handle_on_chain_end_with_output():
output = AgentFinish(return_values={"output": "final output"}, log="test log")
event = {"event": "on_chain_end", "data": {"output": output}, "start_time": 0}
updated_message, start_time = handle_on_chain_end(event, agent_message, send_message, 0.0)
updated_message, start_time = await handle_on_chain_end(event, agent_message, send_message, 0.0)
assert updated_message.properties.icon == "Bot"
assert updated_message.properties.state == "complete"
@ -332,7 +332,7 @@ async def test_handle_on_chain_end_with_output():
async def test_handle_on_chain_end_no_output():
"""Test handle_on_chain_end without output key in data."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
agent_message = Message(
sender=MESSAGE_SENDER_AI,
sender_name="Agent",
@ -341,7 +341,7 @@ async def test_handle_on_chain_end_no_output():
)
event = {"event": "on_chain_end", "data": {}, "start_time": 0}
updated_message, start_time = handle_on_chain_end(event, agent_message, send_message, 0.0)
updated_message, start_time = await handle_on_chain_end(event, agent_message, send_message, 0.0)
assert updated_message.properties.icon == "Bot"
assert updated_message.properties.state == "partial"
@ -351,7 +351,7 @@ async def test_handle_on_chain_end_no_output():
async def test_handle_on_chain_end_empty_data():
"""Test handle_on_chain_end with empty data."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
agent_message = Message(
sender=MESSAGE_SENDER_AI,
sender_name="Agent",
@ -360,7 +360,7 @@ async def test_handle_on_chain_end_empty_data():
)
event = {"event": "on_chain_end", "data": {"output": None}, "start_time": 0}
updated_message, start_time = handle_on_chain_end(event, agent_message, send_message, 0.0)
updated_message, start_time = await handle_on_chain_end(event, agent_message, send_message, 0.0)
assert updated_message.properties.icon == "Bot"
assert updated_message.properties.state == "partial"
@ -370,7 +370,7 @@ async def test_handle_on_chain_end_empty_data():
async def test_handle_on_chain_end_with_empty_return_values():
"""Test handle_on_chain_end with empty return_values."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
agent_message = Message(
sender=MESSAGE_SENDER_AI,
sender_name="Agent",
@ -384,7 +384,7 @@ async def test_handle_on_chain_end_with_empty_return_values():
event = {"event": "on_chain_end", "data": {"output": MockOutputEmptyReturnValues()}, "start_time": 0}
updated_message, start_time = handle_on_chain_end(event, agent_message, send_message, 0.0)
updated_message, start_time = await handle_on_chain_end(event, agent_message, send_message, 0.0)
assert updated_message.properties.icon == "Bot"
assert updated_message.properties.state == "partial"
@ -394,7 +394,7 @@ async def test_handle_on_chain_end_with_empty_return_values():
async def test_handle_on_tool_start():
"""Test handle_on_tool_start event."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
tool_blocks_map = {}
agent_message = Message(
sender=MESSAGE_SENDER_AI,
@ -410,7 +410,7 @@ async def test_handle_on_tool_start():
"start_time": 0,
}
updated_message, start_time = handle_on_tool_start(event, agent_message, tool_blocks_map, send_message, 0.0)
updated_message, start_time = await handle_on_tool_start(event, agent_message, tool_blocks_map, send_message, 0.0)
assert len(updated_message.content_blocks) == 1
assert len(updated_message.content_blocks[0].contents) > 0
@ -426,7 +426,7 @@ async def test_handle_on_tool_start():
async def test_handle_on_tool_end():
"""Test handle_on_tool_end event."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
tool_blocks_map = {}
agent_message = Message(
sender=MESSAGE_SENDER_AI,
@ -441,7 +441,7 @@ async def test_handle_on_tool_end():
"run_id": "test_run",
"data": {"input": {"query": "tool input"}},
}
agent_message, _ = handle_on_tool_start(start_event, agent_message, tool_blocks_map, send_message, 0.0)
agent_message, _ = await handle_on_tool_start(start_event, agent_message, tool_blocks_map, send_message, 0.0)
end_event = {
"event": "on_tool_end",
@ -451,7 +451,7 @@ async def test_handle_on_tool_end():
"start_time": 0,
}
updated_message, start_time = handle_on_tool_end(end_event, agent_message, tool_blocks_map, send_message, 0.0)
updated_message, start_time = await handle_on_tool_end(end_event, agent_message, tool_blocks_map, send_message, 0.0)
f"{end_event['name']}_{end_event['run_id']}"
tool_content = updated_message.content_blocks[0].contents[-1]
@ -463,7 +463,7 @@ async def test_handle_on_tool_end():
async def test_handle_on_tool_error():
"""Test handle_on_tool_error event."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
tool_blocks_map = {}
agent_message = Message(
sender=MESSAGE_SENDER_AI,
@ -478,7 +478,7 @@ async def test_handle_on_tool_error():
"run_id": "test_run",
"data": {"input": {"query": "tool input"}},
}
agent_message, _ = handle_on_tool_start(start_event, agent_message, tool_blocks_map, send_message, 0.0)
agent_message, _ = await handle_on_tool_start(start_event, agent_message, tool_blocks_map, send_message, 0.0)
error_event = {
"event": "on_tool_error",
@ -488,7 +488,9 @@ async def test_handle_on_tool_error():
"start_time": 0,
}
updated_message, start_time = handle_on_tool_error(error_event, agent_message, tool_blocks_map, send_message, 0.0)
updated_message, start_time = await handle_on_tool_error(
error_event, agent_message, tool_blocks_map, send_message, 0.0
)
tool_content = updated_message.content_blocks[0].contents[-1]
assert tool_content.name == "test_tool"
@ -500,7 +502,7 @@ async def test_handle_on_tool_error():
async def test_handle_on_chain_stream_with_output():
"""Test handle_on_chain_stream with output."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
agent_message = Message(
sender=MESSAGE_SENDER_AI,
sender_name="Agent",
@ -512,7 +514,7 @@ async def test_handle_on_chain_stream_with_output():
"data": {"chunk": {"output": "streamed output"}},
}
updated_message, start_time = handle_on_chain_stream(event, agent_message, send_message, 0.0)
updated_message, start_time = await handle_on_chain_stream(event, agent_message, send_message, 0.0)
assert updated_message.text == "streamed output"
assert updated_message.properties.state == "complete"
@ -521,7 +523,7 @@ async def test_handle_on_chain_stream_with_output():
async def test_handle_on_chain_stream_no_output():
"""Test handle_on_chain_stream without output."""
send_message = MagicMock(side_effect=lambda message: message)
send_message = AsyncMock(side_effect=lambda message: message)
agent_message = Message(
sender=MESSAGE_SENDER_AI,
sender_name="Agent",
@ -534,7 +536,7 @@ async def test_handle_on_chain_stream_no_output():
"data": {"chunk": {}},
}
updated_message, start_time = handle_on_chain_stream(event, agent_message, send_message, 0.0)
updated_message, start_time = await handle_on_chain_stream(event, agent_message, send_message, 0.0)
assert updated_message.text == ""
assert updated_message.properties.state == "partial"

View file

@ -1,3 +1,5 @@
import asyncio
import pytest
from langflow.components.inputs import ChatInput, TextInputComponent
from langflow.schema.message import Message
@ -36,10 +38,10 @@ class TestChatInput(ComponentTestBaseWithClient):
{"version": "1.0.19", "module": "inputs", "file_name": "ChatInput"},
]
def test_message_response(self, component_class, default_kwargs):
async def test_message_response(self, component_class, default_kwargs):
"""Test that the message_response method returns a valid Message object."""
component = component_class(**default_kwargs)
message = component.message_response()
message = await component.message_response()
assert isinstance(message, Message)
assert message.text == default_kwargs["input_value"]
@ -58,7 +60,7 @@ class TestChatInput(ComponentTestBaseWithClient):
"targets": [],
}
def test_message_response_ai_sender(self, component_class):
async def test_message_response_ai_sender(self, component_class):
"""Test message response with AI sender type."""
kwargs = {
"input_value": "I am an AI assistant",
@ -67,13 +69,13 @@ class TestChatInput(ComponentTestBaseWithClient):
"session_id": "test_session_123",
}
component = component_class(**kwargs)
message = component.message_response()
message = await component.message_response()
assert isinstance(message, Message)
assert message.sender == MESSAGE_SENDER_AI
assert message.sender_name == "AI Assistant"
def test_message_response_without_session(self, component_class):
async def test_message_response_without_session(self, component_class):
"""Test message response without session ID."""
kwargs = {
"input_value": "Test message",
@ -82,16 +84,16 @@ class TestChatInput(ComponentTestBaseWithClient):
"session_id": "", # Empty session ID
}
component = component_class(**kwargs)
message = component.message_response()
message = await component.message_response()
assert isinstance(message, Message)
assert message.session_id == ""
def test_message_response_with_files(self, component_class, tmp_path):
async def test_message_response_with_files(self, component_class, tmp_path):
"""Test message response with file attachments."""
# Create a temporary test file
test_file = tmp_path / "test.txt"
test_file.write_text("Test content")
await asyncio.to_thread(test_file.write_text, "Test content")
kwargs = {
"input_value": "Message with file",
@ -101,13 +103,13 @@ class TestChatInput(ComponentTestBaseWithClient):
"files": [str(test_file)],
}
component = component_class(**kwargs)
message = component.message_response()
message = await component.message_response()
assert isinstance(message, Message)
assert len(message.files) == 1
assert message.files[0] == str(test_file)
def test_message_storage_disabled(self, component_class):
async def test_message_storage_disabled(self, component_class):
"""Test message response when storage is disabled."""
kwargs = {
"input_value": "Test message",
@ -117,7 +119,7 @@ class TestChatInput(ComponentTestBaseWithClient):
"session_id": "test_session_123",
}
component = component_class(**kwargs)
message = component.message_response()
message = await component.message_response()
assert isinstance(message, Message)
# The message should still be created but not stored

View file

@ -52,7 +52,7 @@ async def test_component_message_sending():
)
# Send the message
sent_message = await asyncio.to_thread(component.send_message, message)
sent_message = await component.send_message(message)
# Verify the message was sent
assert sent_message.id is not None
@ -85,7 +85,7 @@ async def test_component_tool_output():
)
# Send the message
sent_message = await asyncio.to_thread(component.send_message, message)
sent_message = await component.send_message(message)
# Verify the message was stored and processed
assert sent_message.id is not None
@ -112,8 +112,7 @@ async def test_component_error_handling():
msg = "Test error"
raise CustomError(msg)
except CustomError as e:
sent_message = await asyncio.to_thread(
component.send_error,
sent_message = await component.send_error(
exception=e,
session_id="test_session",
trace_name="test_trace",
@ -227,7 +226,7 @@ async def test_component_streaming_message():
)
# Send the streaming message
sent_message = await asyncio.to_thread(component.send_message, message)
sent_message = await component.send_message(message)
# Verify the message
assert sent_message.id is not None

View file

@ -55,6 +55,7 @@ async def test_graph_with_edge():
async def test_graph_functional():
chat_input = ChatInput(_id="chat_input")
chat_input.set(should_store_message=False)
chat_output = ChatOutput(input_value="test", _id="chat_output")
chat_output.set(sender_name=chat_input.message_response)
graph = await asyncio.to_thread(Graph, chat_input, chat_output)

View file

@ -75,9 +75,9 @@ def test_graph_functional_start_graph_state_update():
def test_graph_state_model_serialization():
chat_input = ChatInput(_id="chat_input")
chat_input.set(input_value="Test Sender Name")
chat_input.set(input_value="Test Sender Name", should_store_message=False)
chat_output = ChatOutput(input_value="test", _id="chat_output")
chat_output.set(sender_name=chat_input.message_response)
chat_output.set(sender_name=chat_input.message_response, should_store_message=False)
graph = Graph(chat_input, chat_output)
graph.prepare()

View file

@ -2,7 +2,7 @@ import json
from uuid import UUID
import pytest
from langflow.memory import get_messages
from langflow.memory import aget_messages
from langflow.services.database.models.flow import FlowCreate, FlowUpdate
from orjson import orjson
@ -14,7 +14,7 @@ async def test_build_flow(client, json_memory_chatbot_no_llm, logged_in_headers)
async with client.stream("POST", f"api/v1/build/{flow_id}/flow", json={}, headers=logged_in_headers) as r:
await consume_and_assert_stream(r)
check_messages(flow_id)
await check_messages(flow_id)
@pytest.mark.benchmark
@ -28,7 +28,7 @@ async def test_build_flow_from_request_data(client, json_memory_chatbot_no_llm,
) as r:
await consume_and_assert_stream(r)
check_messages(flow_id)
await check_messages(flow_id)
async def test_build_flow_with_frozen_path(client, json_memory_chatbot_no_llm, logged_in_headers):
@ -47,11 +47,11 @@ async def test_build_flow_with_frozen_path(client, json_memory_chatbot_no_llm, l
async with client.stream("POST", f"api/v1/build/{flow_id}/flow", json={}, headers=logged_in_headers) as r:
await consume_and_assert_stream(r)
check_messages(flow_id)
await check_messages(flow_id)
def check_messages(flow_id):
messages = get_messages(flow_id=UUID(flow_id), order="ASC")
async def check_messages(flow_id):
messages = await aget_messages(flow_id=UUID(flow_id), order="ASC")
assert len(messages) == 2
assert messages[0].session_id == flow_id
assert messages[0].sender == "User"

View file

@ -3,8 +3,14 @@ from uuid import UUID, uuid4
import pytest
from langflow.memory import (
aadd_messages,
aadd_messagetables,
add_messages,
add_messagetables,
adelete_messages,
aget_messages,
astore_message,
aupdate_messages,
delete_messages,
get_messages,
store_message,
@ -18,29 +24,29 @@ from langflow.schema.properties import Properties, Source
# Assuming you have these imports available
from langflow.services.database.models.message import MessageCreate, MessageRead
from langflow.services.database.models.message.model import MessageTable
from langflow.services.deps import session_scope
from langflow.services.deps import async_session_scope
from langflow.services.tracing.utils import convert_to_langchain_type
@pytest.fixture
def created_message():
with session_scope() as session:
async def created_message():
async with async_session_scope() as session:
message = MessageCreate(text="Test message", sender="User", sender_name="User", session_id="session_id")
messagetable = MessageTable.model_validate(message, from_attributes=True)
messagetables = add_messagetables([messagetable], session)
messagetables = await aadd_messagetables([messagetable], session)
return MessageRead.model_validate(messagetables[0], from_attributes=True)
@pytest.fixture
def created_messages(session): # noqa: ARG001
with session_scope() as _session:
async def created_messages(session): # noqa: ARG001
async with async_session_scope() as _session:
messages = [
MessageCreate(text="Test message 1", sender="User", sender_name="User", session_id="session_id2"),
MessageCreate(text="Test message 2", sender="User", sender_name="User", session_id="session_id2"),
MessageCreate(text="Test message 3", sender="User", sender_name="User", session_id="session_id2"),
]
messagetables = [MessageTable.model_validate(message, from_attributes=True) for message in messages]
messagetables = add_messagetables(messagetables, _session)
messagetables = await aadd_messagetables(messagetables, _session)
return [MessageRead.model_validate(messagetable, from_attributes=True) for messagetable in messagetables]
@ -58,6 +64,20 @@ def test_get_messages():
assert messages[1].text == "Test message 2"
@pytest.mark.usefixtures("client")
async def test_aget_messages():
await aadd_messages(
[
Message(text="Test message 1", sender="User", sender_name="User", session_id="session_id2"),
Message(text="Test message 2", sender="User", sender_name="User", session_id="session_id2"),
]
)
messages = await aget_messages(sender="User", session_id="session_id2", limit=2)
assert len(messages) == 2
assert messages[0].text == "Test message 1"
assert messages[1].text == "Test message 2"
@pytest.mark.usefixtures("client")
def test_add_messages():
message = Message(text="New Test message", sender="User", sender_name="User", session_id="new_session_id")
@ -66,6 +86,14 @@ def test_add_messages():
assert messages[0].text == "New Test message"
@pytest.mark.usefixtures("client")
async def test_aadd_messages():
message = Message(text="New Test message", sender="User", sender_name="User", session_id="new_session_id")
messages = await aadd_messages(message)
assert len(messages) == 1
assert messages[0].text == "New Test message"
@pytest.mark.usefixtures("client")
def test_add_messagetables(session):
messages = [MessageTable(text="New Test message", sender="User", sender_name="User", session_id="new_session_id")]
@ -75,17 +103,53 @@ def test_add_messagetables(session):
@pytest.mark.usefixtures("client")
def test_delete_messages(session):
session_id = "session_id2"
async def test_aadd_messagetables(async_session):
messages = [MessageTable(text="New Test message", sender="User", sender_name="User", session_id="new_session_id")]
added_messages = await aadd_messagetables(messages, async_session)
assert len(added_messages) == 1
assert added_messages[0].text == "New Test message"
@pytest.mark.usefixtures("client")
def test_delete_messages():
session_id = "new_session_id"
message = Message(text="New Test message", sender="User", sender_name="User", session_id=session_id)
add_messages([message])
messages = get_messages(sender="User", session_id=session_id)
assert len(messages) == 1
delete_messages(session_id)
messages = session.query(MessageTable).filter(MessageTable.session_id == session_id).all()
messages = get_messages(sender="User", session_id=session_id)
assert len(messages) == 0
@pytest.mark.usefixtures("client")
async def test_adelete_messages():
session_id = "new_session_id"
message = Message(text="New Test message", sender="User", sender_name="User", session_id=session_id)
await aadd_messages([message])
messages = await aget_messages(sender="User", session_id=session_id)
assert len(messages) == 1
await adelete_messages(session_id)
messages = await aget_messages(sender="User", session_id=session_id)
assert len(messages) == 0
@pytest.mark.usefixtures("client")
def test_store_message():
message = Message(text="Stored message", sender="User", sender_name="User", session_id="stored_session_id")
stored_messages = store_message(message)
session_id = "stored_session_id"
message = Message(text="Stored message", sender="User", sender_name="User", session_id=session_id)
store_message(message)
stored_messages = get_messages(sender="User", session_id=session_id)
assert len(stored_messages) == 1
assert stored_messages[0].text == "Stored message"
@pytest.mark.usefixtures("client")
async def test_astore_message():
session_id = "stored_session_id"
message = Message(text="Stored message", sender="User", sender_name="User", session_id=session_id)
await astore_message(message)
stored_messages = await aget_messages(sender="User", session_id=session_id)
assert len(stored_messages) == 1
assert stored_messages[0].text == "Stored message"
@ -298,3 +362,188 @@ def test_update_message_with_nested_properties(created_message):
assert updated[0].properties.allow_markdown is True
assert updated[0].properties.state == "complete"
assert updated[0].properties.targets == []
@pytest.mark.usefixtures("client")
async def test_aupdate_single_message(created_message):
# Modify the message
created_message.text = "Updated message"
updated = await aupdate_messages(created_message)
assert len(updated) == 1
assert updated[0].text == "Updated message"
assert updated[0].id == created_message.id
@pytest.mark.usefixtures("client")
async def test_aupdate_multiple_messages(created_messages):
# Modify the messages
for i, message in enumerate(created_messages):
message.text = f"Updated message {i}"
updated = await aupdate_messages(created_messages)
assert len(updated) == len(created_messages)
for i, message in enumerate(updated):
assert message.text == f"Updated message {i}"
assert message.id == created_messages[i].id
@pytest.mark.usefixtures("client")
async def test_aupdate_nonexistent_message():
# Create a message with a non-existent UUID
message = MessageRead(
id=uuid4(), # Generate a random UUID that won't exist in the database
text="Test message",
sender="User",
sender_name="User",
session_id="session_id",
flow_id=uuid4(),
)
updated = await aupdate_messages(message)
assert len(updated) == 0
@pytest.mark.usefixtures("client")
async def test_aupdate_mixed_messages(created_messages):
# Create a mix of existing and non-existing messages
nonexistent_message = MessageRead(
id=uuid4(), # Generate a random UUID that won't exist in the database
text="Test message",
sender="User",
sender_name="User",
session_id="session_id",
flow_id=uuid4(),
)
messages_to_update = created_messages[:1] + [nonexistent_message]
created_messages[0].text = "Updated existing message"
updated = await aupdate_messages(messages_to_update)
assert len(updated) == 1
assert updated[0].text == "Updated existing message"
assert updated[0].id == created_messages[0].id
assert isinstance(updated[0].id, UUID) # Verify ID is UUID type
@pytest.mark.usefixtures("client")
async def test_aupdate_message_with_timestamp(created_message):
# Set a specific timestamp
new_timestamp = datetime(2024, 1, 1, 12, 0, 0, tzinfo=timezone.utc)
created_message.timestamp = new_timestamp
created_message.text = "Updated message with timestamp"
updated = await aupdate_messages(created_message)
assert len(updated) == 1
assert updated[0].text == "Updated message with timestamp"
# Compare timestamps without timezone info since DB doesn't preserve it
assert updated[0].timestamp.replace(tzinfo=None) == new_timestamp.replace(tzinfo=None)
assert updated[0].id == created_message.id
@pytest.mark.usefixtures("client")
async def test_aupdate_multiple_messages_with_timestamps(created_messages):
# Modify messages with different timestamps
for i, message in enumerate(created_messages):
message.text = f"Updated message {i}"
message.timestamp = datetime(2024, 1, 1, i, 0, 0, tzinfo=timezone.utc)
updated = await aupdate_messages(created_messages)
assert len(updated) == len(created_messages)
for i, message in enumerate(updated):
assert message.text == f"Updated message {i}"
# Compare timestamps without timezone info
expected_timestamp = datetime(2024, 1, 1, i, 0, 0, tzinfo=timezone.utc)
assert message.timestamp.replace(tzinfo=None) == expected_timestamp.replace(tzinfo=None)
assert message.id == created_messages[i].id
@pytest.mark.usefixtures("client")
async def test_aupdate_message_with_content_blocks(created_message):
# Create a content block using proper models
text_content = TextContent(
type="text", text="Test content", duration=5, header={"title": "Test Header", "icon": "TestIcon"}
)
tool_content = ToolContent(type="tool_use", name="test_tool", tool_input={"param": "value"}, duration=10)
content_block = ContentBlock(title="Test Block", contents=[text_content, tool_content], allow_markdown=True)
created_message.content_blocks = [content_block]
created_message.text = "Message with content blocks"
updated = await aupdate_messages(created_message)
assert len(updated) == 1
assert updated[0].text == "Message with content blocks"
assert len(updated[0].content_blocks) == 1
# Verify the content block structure
updated_block = updated[0].content_blocks[0]
assert updated_block.title == "Test Block"
assert len(updated_block.contents) == 2
# Verify text content
text_content = updated_block.contents[0]
assert text_content.type == "text"
assert text_content.text == "Test content"
assert text_content.duration == 5
assert text_content.header["title"] == "Test Header"
# Verify tool content
tool_content = updated_block.contents[1]
assert tool_content.type == "tool_use"
assert tool_content.name == "test_tool"
assert tool_content.tool_input == {"param": "value"}
assert tool_content.duration == 10
@pytest.mark.usefixtures("client")
async def test_aupdate_message_with_nested_properties(created_message):
# Create a text content with nested properties
text_content = TextContent(
type="text", text="Test content", header={"title": "Test Header", "icon": "TestIcon"}, duration=15
)
content_block = ContentBlock(
title="Test Properties",
contents=[text_content],
allow_markdown=True,
media_url=["http://example.com/image.jpg"],
)
# Set properties according to the Properties model structure
created_message.properties = Properties(
text_color="blue",
background_color="white",
edited=False,
source=Source(id="test_id", display_name="Test Source", source="test"),
icon="TestIcon",
allow_markdown=True,
state="complete",
targets=[],
)
created_message.text = "Message with nested properties"
created_message.content_blocks = [content_block]
updated = await aupdate_messages(created_message)
assert len(updated) == 1
assert updated[0].text == "Message with nested properties"
# Verify the properties were properly serialized and stored
assert updated[0].properties.text_color == "blue"
assert updated[0].properties.background_color == "white"
assert updated[0].properties.edited is False
assert updated[0].properties.source.id == "test_id"
assert updated[0].properties.source.display_name == "Test Source"
assert updated[0].properties.source.source == "test"
assert updated[0].properties.icon == "TestIcon"
assert updated[0].properties.allow_markdown is True
assert updated[0].properties.state == "complete"
assert updated[0].properties.targets == []

View file

@ -2,33 +2,33 @@ from uuid import UUID
import pytest
from httpx import AsyncClient
from langflow.memory import add_messagetables
from langflow.memory import aadd_messagetables
# Assuming you have these imports available
from langflow.services.database.models.message import MessageCreate, MessageRead, MessageUpdate
from langflow.services.database.models.message.model import MessageTable
from langflow.services.deps import session_scope
from langflow.services.deps import async_session_scope
@pytest.fixture
def created_message():
with session_scope() as session:
async def created_message():
async with async_session_scope() as session:
message = MessageCreate(text="Test message", sender="User", sender_name="User", session_id="session_id")
messagetable = MessageTable.model_validate(message, from_attributes=True)
messagetables = add_messagetables([messagetable], session)
messagetables = await aadd_messagetables([messagetable], session)
return MessageRead.model_validate(messagetables[0], from_attributes=True)
@pytest.fixture
def created_messages(session): # noqa: ARG001
with session_scope() as _session:
async def created_messages(session): # noqa: ARG001
async with async_session_scope() as _session:
messages = [
MessageCreate(text="Test message 1", sender="User", sender_name="User", session_id="session_id2"),
MessageCreate(text="Test message 2", sender="User", sender_name="User", session_id="session_id2"),
MessageCreate(text="Test message 3", sender="User", sender_name="User", session_id="session_id2"),
]
messagetables = [MessageTable.model_validate(message, from_attributes=True) for message in messages]
return add_messagetables(messagetables, _session)
return await aadd_messagetables(messagetables, _session)
@pytest.mark.api_key_required