fix: Use AsyncSession in memory (#4665)
This commit is contained in:
parent
156597d3d1
commit
79b03ba133
26 changed files with 610 additions and 173 deletions
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue