Refactor tests and add new test file for data components

This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-03-08 13:42:28 -03:00
commit 1363f387e9
10 changed files with 340 additions and 192 deletions

View file

@ -2,6 +2,7 @@ from pathlib import Path
from tempfile import tempdir from tempfile import tempdir
import pytest import pytest
from langflow.__main__ import app from langflow.__main__ import app
from langflow.services import deps from langflow.services import deps

View file

@ -3,19 +3,21 @@ import types
from uuid import uuid4 from uuid import uuid4
import pytest import pytest
from fastapi import HTTPException
from langchain_core.documents import Document from langchain_core.documents import Document
from langflow.interface.custom.base import CustomComponent from langflow.interface.custom.base import CustomComponent
from langflow.interface.custom.code_parser.code_parser import CodeParser, CodeSyntaxError from langflow.interface.custom.code_parser.code_parser import (
CodeParser,
CodeSyntaxError,
)
from langflow.interface.custom.custom_component.component import ( from langflow.interface.custom.custom_component.component import (
Component, Component,
ComponentCodeNullError, ComponentCodeNullError,
ComponentFunctionEntrypointNameNullError,
) )
from langflow.services.database.models.flow import Flow, FlowCreate from langflow.services.database.models.flow import Flow, FlowCreate
code_default = """ code_default = """
from langflow import Prompt from langflow.field_typing import Prompt
from langflow.interface.custom.custom_component import CustomComponent from langflow.interface.custom.custom_component import CustomComponent
from langchain.llms.base import BaseLLM from langchain.llms.base import BaseLLM
@ -96,23 +98,15 @@ def test_component_code_null_error():
component.get_function() component.get_function()
def test_component_function_entrypoint_name_null_error():
"""
Test the get_function method raises the ComponentFunctionEntrypointNameNullError
when the function_entrypoint_name is empty.
"""
component = Component(code=code_default, function_entrypoint_name="")
with pytest.raises(ComponentFunctionEntrypointNameNullError):
component.get_function()
def test_custom_component_init(): def test_custom_component_init():
""" """
Test the initialization of the CustomComponent class. Test the initialization of the CustomComponent class.
""" """
function_entrypoint_name = "build" function_entrypoint_name = "build"
custom_component = CustomComponent(code=code_default, function_entrypoint_name=function_entrypoint_name) custom_component = CustomComponent(
code=code_default, function_entrypoint_name=function_entrypoint_name
)
assert custom_component.code == code_default assert custom_component.code == code_default
assert custom_component.function_entrypoint_name == function_entrypoint_name assert custom_component.function_entrypoint_name == function_entrypoint_name
@ -121,7 +115,9 @@ def test_custom_component_build_template_config():
""" """
Test the build_template_config property of the CustomComponent class. Test the build_template_config property of the CustomComponent class.
""" """
custom_component = CustomComponent(code=code_default, function_entrypoint_name="build") custom_component = CustomComponent(
code=code_default, function_entrypoint_name="build"
)
config = custom_component.build_template_config() config = custom_component.build_template_config()
assert isinstance(config, dict) assert isinstance(config, dict)
@ -130,8 +126,10 @@ def test_custom_component_get_function():
""" """
Test the get_function property of the CustomComponent class. Test the get_function property of the CustomComponent class.
""" """
custom_component = CustomComponent(code="def build(): pass", function_entrypoint_name="build") custom_component = CustomComponent(
my_function = custom_component.get_function code="def build(): pass", function_entrypoint_name="build"
)
my_function = custom_component.get_function()
assert isinstance(my_function, types.FunctionType) assert isinstance(my_function, types.FunctionType)
@ -215,7 +213,9 @@ def test_custom_component_get_function_entrypoint_args():
Test the get_function_entrypoint_args Test the get_function_entrypoint_args
property of the CustomComponent class. property of the CustomComponent class.
""" """
custom_component = CustomComponent(code=code_default, function_entrypoint_name="build") custom_component = CustomComponent(
code=code_default, function_entrypoint_name="build"
)
args = custom_component.get_function_entrypoint_args args = custom_component.get_function_entrypoint_args
assert len(args) == 4 assert len(args) == 4
assert args[0]["name"] == "self" assert args[0]["name"] == "self"
@ -229,7 +229,9 @@ def test_custom_component_get_function_entrypoint_return_type():
property of the CustomComponent class. property of the CustomComponent class.
""" """
custom_component = CustomComponent(code=code_default, function_entrypoint_name="build") custom_component = CustomComponent(
code=code_default, function_entrypoint_name="build"
)
return_type = custom_component.get_function_entrypoint_return_type return_type = custom_component.get_function_entrypoint_return_type
assert return_type == [Document] assert return_type == [Document]
@ -238,7 +240,9 @@ def test_custom_component_get_main_class_name():
""" """
Test the get_main_class_name property of the CustomComponent class. Test the get_main_class_name property of the CustomComponent class.
""" """
custom_component = CustomComponent(code=code_default, function_entrypoint_name="build") custom_component = CustomComponent(
code=code_default, function_entrypoint_name="build"
)
class_name = custom_component.get_main_class_name class_name = custom_component.get_main_class_name
assert class_name == "YourComponent" assert class_name == "YourComponent"
@ -248,7 +252,9 @@ def test_custom_component_get_function_valid():
Test the get_function property of the CustomComponent Test the get_function property of the CustomComponent
class with valid code and function_entrypoint_name. class with valid code and function_entrypoint_name.
""" """
custom_component = CustomComponent(code="def build(): pass", function_entrypoint_name="build") custom_component = CustomComponent(
code="def build(): pass", function_entrypoint_name="build"
)
my_function = custom_component.get_function my_function = custom_component.get_function
assert callable(my_function) assert callable(my_function)
@ -283,7 +289,9 @@ def test_code_parser_parse_callable_details_no_args():
parser = CodeParser("") parser = CodeParser("")
node = ast.FunctionDef( node = ast.FunctionDef(
name="test", name="test",
args=ast.arguments(args=[], vararg=None, kwonlyargs=[], kw_defaults=[], kwarg=None, defaults=[]), args=ast.arguments(
args=[], vararg=None, kwonlyargs=[], kw_defaults=[], kwarg=None, defaults=[]
),
body=[], body=[],
decorator_list=[], decorator_list=[],
returns=None, returns=None,
@ -329,7 +337,9 @@ def test_code_parser_parse_function_def_not_init():
parser = CodeParser("") parser = CodeParser("")
stmt = ast.FunctionDef( stmt = ast.FunctionDef(
name="test", name="test",
args=ast.arguments(args=[], vararg=None, kwonlyargs=[], kw_defaults=[], kwarg=None, defaults=[]), args=ast.arguments(
args=[], vararg=None, kwonlyargs=[], kw_defaults=[], kwarg=None, defaults=[]
),
body=[], body=[],
decorator_list=[], decorator_list=[],
returns=None, returns=None,
@ -347,7 +357,9 @@ def test_code_parser_parse_function_def_init():
parser = CodeParser("") parser = CodeParser("")
stmt = ast.FunctionDef( stmt = ast.FunctionDef(
name="__init__", name="__init__",
args=ast.arguments(args=[], vararg=None, kwonlyargs=[], kw_defaults=[], kwarg=None, defaults=[]), args=ast.arguments(
args=[], vararg=None, kwonlyargs=[], kw_defaults=[], kwarg=None, defaults=[]
),
body=[], body=[],
decorator_list=[], decorator_list=[],
returns=None, returns=None,
@ -373,7 +385,7 @@ def test_custom_component_class_template_validation_no_code():
raises the HTTPException when the code is None. raises the HTTPException when the code is None.
""" """
custom_component = CustomComponent(code=None, function_entrypoint_name="build") custom_component = CustomComponent(code=None, function_entrypoint_name="build")
with pytest.raises(HTTPException): with pytest.raises(TypeError):
custom_component.get_function() custom_component.get_function()
@ -382,7 +394,9 @@ def test_custom_component_get_code_tree_syntax_error():
Test the get_code_tree method of the CustomComponent class Test the get_code_tree method of the CustomComponent class
raises the CodeSyntaxError when given incorrect syntax. raises the CodeSyntaxError when given incorrect syntax.
""" """
custom_component = CustomComponent(code="import os as", function_entrypoint_name="build") custom_component = CustomComponent(
code="import os as", function_entrypoint_name="build"
)
with pytest.raises(CodeSyntaxError): with pytest.raises(CodeSyntaxError):
custom_component.get_code_tree(custom_component.code) custom_component.get_code_tree(custom_component.code)
@ -436,7 +450,9 @@ def test_custom_component_build_not_implemented():
Test the build method of the CustomComponent Test the build method of the CustomComponent
class raises the NotImplementedError. class raises the NotImplementedError.
""" """
custom_component = CustomComponent(code="def build(): pass", function_entrypoint_name="build") custom_component = CustomComponent(
code="def build(): pass", function_entrypoint_name="build"
)
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
custom_component.build() custom_component.build()
@ -444,7 +460,7 @@ def test_custom_component_build_not_implemented():
def test_build_config_no_code(): def test_build_config_no_code():
component = CustomComponent(code=None) component = CustomComponent(code=None)
assert component.get_function_entrypoint_args == "" assert component.get_function_entrypoint_args == []
assert component.get_function_entrypoint_return_type == [] assert component.get_function_entrypoint_return_type == []
@ -470,7 +486,9 @@ def test_flow(db):
} }
# Create flow # Create flow
flow = FlowCreate(id=uuid4(), name="Test Flow", description="Fixture flow", data=flow_data) flow = FlowCreate(
id=uuid4(), name="Test Flow", description="Fixture flow", data=flow_data
)
# Add to database # Add to database
db.add(flow) db.add(flow)

View file

@ -0,0 +1,93 @@
import httpx
import pytest
import respx
from httpx import Response
from langflow.components import (
data,
) # Adjust the import according to your project structure
@pytest.fixture
def api_request():
# This fixture provides an instance of APIRequest for each test case
return data.APIRequest()
@pytest.mark.asyncio
@respx.mock
async def test_successful_get_request(api_request):
# Mocking a successful GET request
url = "https://example.com/api/test"
method = "GET"
mock_response = {"success": True}
respx.get(url).mock(return_value=Response(200, json=mock_response))
# Making the request
result = await api_request.make_request(
client=httpx.AsyncClient(), method=method, url=url
)
# Assertions
assert result.data["status_code"] == 200
assert result.data["result"] == mock_response
@pytest.mark.asyncio
@respx.mock
async def test_failed_request(api_request):
# Mocking a failed GET request
url = "https://example.com/api/test"
method = "GET"
respx.get(url).mock(return_value=Response(404))
# Making the request
result = await api_request.make_request(
client=httpx.AsyncClient(), method=method, url=url
)
# Assertions
assert result.data["status_code"] == 404
@pytest.mark.asyncio
@respx.mock
async def test_timeout(api_request):
# Mocking a timeout
url = "https://example.com/api/timeout"
method = "GET"
respx.get(url).mock(
side_effect=httpx.TimeoutException(message="Timeout", request=None)
)
# Making the request
result = await api_request.make_request(
client=httpx.AsyncClient(), method=method, url=url, timeout=1
)
# Assertions
assert result.data["status_code"] == 408
assert result.data["error"] == "Request timed out"
@pytest.mark.asyncio
@respx.mock
async def test_build_with_multiple_urls(api_request):
# This test depends on having a working internet connection and accessible URLs
# It's better to mock these requests using respx or a similar library
# Setup for multiple URLs
method = "GET"
urls = ["https://example.com/api/one", "https://example.com/api/two"]
# You would mock these requests similarly to the single request tests
for url in urls:
respx.get(url).mock(return_value=Response(200, json={"success": True}))
# Do I have to mock the async client?
#
# Execute the build method
results = await api_request.build(method=method, urls=urls)
# Assertions
assert len(results) == len(urls)

View file

@ -3,12 +3,14 @@ from uuid import UUID, uuid4
import orjson import orjson
import pytest import pytest
from fastapi.testclient import TestClient from fastapi.testclient import TestClient
from sqlmodel import Session
from langflow.api.v1.schemas import FlowListCreate from langflow.api.v1.schemas import FlowListCreate
from langflow.initial_setup.setup import load_starter_projects
from langflow.services.database.models.base import orjson_dumps from langflow.services.database.models.base import orjson_dumps
from langflow.services.database.models.flow import Flow, FlowCreate, FlowUpdate from langflow.services.database.models.flow import Flow, FlowCreate, FlowUpdate
from langflow.services.database.utils import session_getter from langflow.services.database.utils import session_getter
from langflow.services.deps import get_db_service from langflow.services.deps import get_db_service
from sqlmodel import Session
@pytest.fixture(scope="module") @pytest.fixture(scope="module")
@ -25,7 +27,9 @@ def json_style():
) )
def test_create_flow(client: TestClient, json_flow: str, active_user, logged_in_headers): def test_create_flow(
client: TestClient, json_flow: str, active_user, logged_in_headers
):
flow = orjson.loads(json_flow) flow = orjson.loads(json_flow)
data = flow["data"] data = flow["data"]
flow = FlowCreate(name="Test Flow", description="description", data=data) flow = FlowCreate(name="Test Flow", description="description", data=data)
@ -35,7 +39,9 @@ def test_create_flow(client: TestClient, json_flow: str, active_user, logged_in_
assert response.json()["data"] == flow.data assert response.json()["data"] == flow.data
# flow is optional so we can create a flow without a flow # flow is optional so we can create a flow without a flow
flow = FlowCreate(name="Test Flow") flow = FlowCreate(name="Test Flow")
response = client.post("api/v1/flows/", json=flow.dict(exclude_unset=True), headers=logged_in_headers) response = client.post(
"api/v1/flows/", json=flow.dict(exclude_unset=True), headers=logged_in_headers
)
assert response.status_code == 201 assert response.status_code == 201
assert response.json()["name"] == flow.name assert response.json()["name"] == flow.name
assert response.json()["data"] == flow.data assert response.json()["data"] == flow.data
@ -76,7 +82,9 @@ def test_read_flow(client: TestClient, json_flow: str, active_user, logged_in_he
assert response.json()["data"] == flow.data assert response.json()["data"] == flow.data
def test_update_flow(client: TestClient, json_flow: str, active_user, logged_in_headers): def test_update_flow(
client: TestClient, json_flow: str, active_user, logged_in_headers
):
flow = orjson.loads(json_flow) flow = orjson.loads(json_flow)
data = flow["data"] data = flow["data"]
@ -89,7 +97,9 @@ def test_update_flow(client: TestClient, json_flow: str, active_user, logged_in_
description="updated description", description="updated description",
data=data, data=data,
) )
response = client.patch(f"api/v1/flows/{flow_id}", json=updated_flow.dict(), headers=logged_in_headers) response = client.patch(
f"api/v1/flows/{flow_id}", json=updated_flow.dict(), headers=logged_in_headers
)
assert response.status_code == 200 assert response.status_code == 200
assert response.json()["name"] == updated_flow.name assert response.json()["name"] == updated_flow.name
@ -97,7 +107,9 @@ def test_update_flow(client: TestClient, json_flow: str, active_user, logged_in_
# assert response.json()["data"] == updated_flow.data # assert response.json()["data"] == updated_flow.data
def test_delete_flow(client: TestClient, json_flow: str, active_user, logged_in_headers): def test_delete_flow(
client: TestClient, json_flow: str, active_user, logged_in_headers
):
flow = orjson.loads(json_flow) flow = orjson.loads(json_flow)
data = flow["data"] data = flow["data"]
flow = FlowCreate(name="Test Flow", description="description", data=data) flow = FlowCreate(name="Test Flow", description="description", data=data)
@ -108,7 +120,9 @@ def test_delete_flow(client: TestClient, json_flow: str, active_user, logged_in_
assert response.json()["message"] == "Flow deleted successfully" assert response.json()["message"] == "Flow deleted successfully"
def test_create_flows(client: TestClient, session: Session, json_flow: str, logged_in_headers): def test_create_flows(
client: TestClient, session: Session, json_flow: str, logged_in_headers
):
flow = orjson.loads(json_flow) flow = orjson.loads(json_flow)
data = flow["data"] data = flow["data"]
# Create test data # Create test data
@ -119,7 +133,9 @@ def test_create_flows(client: TestClient, session: Session, json_flow: str, logg
] ]
) )
# Make request to endpoint # Make request to endpoint
response = client.post("api/v1/flows/batch/", json=flow_list.dict(), headers=logged_in_headers) response = client.post(
"api/v1/flows/batch/", json=flow_list.dict(), headers=logged_in_headers
)
# Check response status code # Check response status code
assert response.status_code == 201 assert response.status_code == 201
# Check response data # Check response data
@ -133,7 +149,9 @@ def test_create_flows(client: TestClient, session: Session, json_flow: str, logg
assert response_data[1]["data"] == data assert response_data[1]["data"] == data
def test_upload_file(client: TestClient, session: Session, json_flow: str, logged_in_headers): def test_upload_file(
client: TestClient, session: Session, json_flow: str, logged_in_headers
):
flow = orjson.loads(json_flow) flow = orjson.loads(json_flow)
data = flow["data"] data = flow["data"]
# Create test data # Create test data
@ -182,7 +200,7 @@ def test_download_file(
with session_getter(db_manager) as session: with session_getter(db_manager) as session:
for flow in flow_list.flows: for flow in flow_list.flows:
flow.user_id = active_user.id flow.user_id = active_user.id
db_flow = Flow.from_orm(flow) db_flow = Flow.model_validate(flow, from_attributes=True)
session.add(db_flow) session.add(db_flow)
session.commit() session.commit()
# Make request to endpoint # Make request to endpoint
@ -191,7 +209,9 @@ def test_download_file(
assert response.status_code == 200, response.json() assert response.status_code == 200, response.json()
# Check response data # Check response data
response_data = response.json()["flows"] response_data = response.json()["flows"]
assert len(response_data) == 2, response_data starter_projects = load_starter_projects()
number_of_projects = len(starter_projects) + len(flow_list.flows)
assert len(response_data) == number_of_projects, response_data
assert response_data[0]["name"] == "Flow 1" assert response_data[0]["name"] == "Flow 1"
assert response_data[0]["description"] == "description" assert response_data[0]["description"] == "description"
assert response_data[0]["data"] == data assert response_data[0]["data"] == data
@ -200,7 +220,9 @@ def test_download_file(
assert response_data[1]["data"] == data assert response_data[1]["data"] == data
def test_create_flow_with_invalid_data(client: TestClient, active_user, logged_in_headers): def test_create_flow_with_invalid_data(
client: TestClient, active_user, logged_in_headers
):
flow = {"name": "a" * 256, "data": "Invalid flow data"} flow = {"name": "a" * 256, "data": "Invalid flow data"}
response = client.post("api/v1/flows/", json=flow, headers=logged_in_headers) response = client.post("api/v1/flows/", json=flow, headers=logged_in_headers)
assert response.status_code == 422 assert response.status_code == 422
@ -212,19 +234,29 @@ def test_get_nonexistent_flow(client: TestClient, active_user, logged_in_headers
assert response.status_code == 404 assert response.status_code == 404
def test_update_flow_idempotency(client: TestClient, json_flow: str, active_user, logged_in_headers): def test_update_flow_idempotency(
client: TestClient, json_flow: str, active_user, logged_in_headers
):
flow_data = orjson.loads(json_flow) flow_data = orjson.loads(json_flow)
data = flow_data["data"] data = flow_data["data"]
flow_data = FlowCreate(name="Test Flow", description="description", data=data) flow_data = FlowCreate(name="Test Flow", description="description", data=data)
response = client.post("api/v1/flows/", json=flow_data.dict(), headers=logged_in_headers) response = client.post(
"api/v1/flows/", json=flow_data.dict(), headers=logged_in_headers
)
flow_id = response.json()["id"] flow_id = response.json()["id"]
updated_flow = FlowCreate(name="Updated Flow", description="description", data=data) updated_flow = FlowCreate(name="Updated Flow", description="description", data=data)
response1 = client.put(f"api/v1/flows/{flow_id}", json=updated_flow.dict(), headers=logged_in_headers) response1 = client.put(
response2 = client.put(f"api/v1/flows/{flow_id}", json=updated_flow.dict(), headers=logged_in_headers) f"api/v1/flows/{flow_id}", json=updated_flow.dict(), headers=logged_in_headers
)
response2 = client.put(
f"api/v1/flows/{flow_id}", json=updated_flow.dict(), headers=logged_in_headers
)
assert response1.json() == response2.json() assert response1.json() == response2.json()
def test_update_nonexistent_flow(client: TestClient, json_flow: str, active_user, logged_in_headers): def test_update_nonexistent_flow(
client: TestClient, json_flow: str, active_user, logged_in_headers
):
flow_data = orjson.loads(json_flow) flow_data = orjson.loads(json_flow)
data = flow_data["data"] data = flow_data["data"]
uuid = uuid4() uuid = uuid4()
@ -233,7 +265,9 @@ def test_update_nonexistent_flow(client: TestClient, json_flow: str, active_user
description="description", description="description",
data=data, data=data,
) )
response = client.patch(f"api/v1/flows/{uuid}", json=updated_flow.dict(), headers=logged_in_headers) response = client.patch(
f"api/v1/flows/{uuid}", json=updated_flow.dict(), headers=logged_in_headers
)
assert response.status_code == 404 assert response.status_code == 404
@ -243,7 +277,8 @@ def test_delete_nonexistent_flow(client: TestClient, active_user, logged_in_head
assert response.status_code == 404 assert response.status_code == 404
def test_read_empty_flows(client: TestClient, active_user, logged_in_headers): def test_read_only_starter_projects(client: TestClient, active_user, logged_in_headers):
response = client.get("api/v1/flows/", headers=logged_in_headers) response = client.get("api/v1/flows/", headers=logged_in_headers)
starter_projects = load_starter_projects()
assert response.status_code == 200 assert response.status_code == 200
assert len(response.json()) == 0 assert len(response.json()) == len(starter_projects)

View file

@ -4,7 +4,6 @@ import pytest
from fastapi.testclient import TestClient from fastapi.testclient import TestClient
from langflow.interface.custom.directory_reader.directory_reader import DirectoryReader from langflow.interface.custom.directory_reader.directory_reader import DirectoryReader
from langflow.interface.tools.constants import CUSTOM_TOOLS
from langflow.services.auth.utils import get_password_hash from langflow.services.auth.utils import get_password_hash
from langflow.services.database.models.api_key.model import ApiKey from langflow.services.database.models.api_key.model import ApiKey
from langflow.services.database.utils import session_getter from langflow.services.database.utils import session_getter
@ -29,7 +28,10 @@ def poll_task_status(client, headers, href, max_attempts=20, sleep_time=1):
href, href,
headers=headers, headers=headers,
) )
if task_status_response.status_code == 200 and task_status_response.json()["status"] == "SUCCESS": if (
task_status_response.status_code == 200
and task_status_response.json()["status"] == "SUCCESS"
):
return task_status_response.json() return task_status_response.json()
time.sleep(sleep_time) time.sleep(sleep_time)
return None # Return None if task did not complete in time return None # Return None if task did not complete in time
@ -123,7 +125,11 @@ def created_api_key(active_user):
) )
db_manager = get_db_service() db_manager = get_db_service()
with session_getter(db_manager) as session: with session_getter(db_manager) as session:
if existing_api_key := session.query(ApiKey).filter(ApiKey.api_key == api_key.api_key).first(): if (
existing_api_key := session.query(ApiKey)
.filter(ApiKey.api_key == api_key.api_key)
.first()
):
return existing_api_key return existing_api_key
session.add(api_key) session.add(api_key)
session.commit() session.commit()
@ -289,13 +295,17 @@ def test_get_all(client: TestClient, logged_in_headers):
dir_reader = DirectoryReader(settings.COMPONENTS_PATH[0]) dir_reader = DirectoryReader(settings.COMPONENTS_PATH[0])
files = dir_reader.get_files() files = dir_reader.get_files()
# json_response is a dict of dicts # json_response is a dict of dicts
all_names = [component_name for _, components in response.json().items() for component_name in components] all_names = [
component_name
for _, components in response.json().items()
for component_name in components
]
json_response = response.json() json_response = response.json()
# We need to test the custom nodes # We need to test the custom nodes
assert len(all_names) > len(files) assert len(all_names) > len(files)
assert "Prompt" in json_response["prompts"] assert "ChatInput" in json_response["inputs"]
# All CUSTOM_TOOLS(dict) should be in the response assert "Prompt" in json_response["inputs"]
assert all(tool in json_response["tools"] for tool in CUSTOM_TOOLS.keys()) assert "ChatOutput" in json_response["outputs"]
def test_post_validate_code(client: TestClient): def test_post_validate_code(client: TestClient):
@ -414,35 +424,46 @@ def test_various_prompts(client, prompt, expected_input_variables):
def test_get_vertices_flow_not_found(client, logged_in_headers): def test_get_vertices_flow_not_found(client, logged_in_headers):
response = client.get("/api/v1/build/nonexistent_id/vertices", headers=logged_in_headers) response = client.get(
assert response.status_code == 500 # Or whatever status code you've set for invalid ID "/api/v1/build/nonexistent_id/vertices", headers=logged_in_headers
)
assert (
response.status_code == 500
) # Or whatever status code you've set for invalid ID
def test_get_vertices(client, added_flow_with_prompt_and_history, logged_in_headers): def test_get_vertices(client, added_flow_with_prompt_and_history, logged_in_headers):
flow_id = added_flow_with_prompt_and_history["id"] flow_id = added_flow_with_prompt_and_history["id"]
response = client.get(f"/api/v1/build/{flow_id}/vertices", headers=logged_in_headers) response = client.get(
f"/api/v1/build/{flow_id}/vertices", headers=logged_in_headers
)
assert response.status_code == 200 assert response.status_code == 200
assert "ids" in response.json() assert "ids" in response.json()
# The response should contain the list in this order # The response should contain the list in this order
# ['ConversationBufferMemory-Lu2Nb', 'PromptTemplate-5Q0W8', 'ChatOpenAI-vy7fV', 'LLMChain-UjBh1'] # ['ConversationBufferMemory-Lu2Nb', 'PromptTemplate-5Q0W8', 'ChatOpenAI-vy7fV', 'LLMChain-UjBh1']
# The important part is before the - (ConversationBufferMemory, PromptTemplate, ChatOpenAI, LLMChain) # The important part is before the - (ConversationBufferMemory, PromptTemplate, ChatOpenAI, LLMChain)
ids = [inner_id.split("-")[0] for _id in response.json()["ids"] for inner_id in _id] ids = [_id.split("-")[0] for _id in response.json()["ids"]]
assert ids == [ assert ids == [
"ChatOpenAI", "ChatOpenAI",
"PromptTemplate", "PromptTemplate",
"ConversationBufferMemory", "ConversationBufferMemory",
"LLMChain",
] ]
def test_build_vertex_invalid_flow_id(client, logged_in_headers): def test_build_vertex_invalid_flow_id(client, logged_in_headers):
response = client.post("/api/v1/build/nonexistent_id/vertices/vertex_id", headers=logged_in_headers) response = client.post(
"/api/v1/build/nonexistent_id/vertices/vertex_id", headers=logged_in_headers
)
assert response.status_code == 500 assert response.status_code == 500
def test_build_vertex_invalid_vertex_id(client, added_flow_with_prompt_and_history, logged_in_headers): def test_build_vertex_invalid_vertex_id(
client, added_flow_with_prompt_and_history, logged_in_headers
):
flow_id = added_flow_with_prompt_and_history["id"] flow_id = added_flow_with_prompt_and_history["id"]
response = client.post(f"/api/v1/build/{flow_id}/vertices/invalid_vertex_id", headers=logged_in_headers) response = client.post(
f"/api/v1/build/{flow_id}/vertices/invalid_vertex_id", headers=logged_in_headers
)
assert response.status_code == 500 assert response.status_code == 500

View file

@ -1,6 +1,7 @@
from unittest.mock import MagicMock from unittest.mock import MagicMock
import pytest import pytest
from langflow.services.deps import get_storage_service from langflow.services.deps import get_storage_service
from langflow.services.storage.service import StorageService from langflow.services.storage.service import StorageService
from langflow.services.storage.utils import build_content_type_from_extension from langflow.services.storage.utils import build_content_type_from_extension
@ -24,10 +25,15 @@ def test_upload_file(client, mock_storage_service, created_api_key, flow):
client.app.dependency_overrides[get_storage_service] = lambda: mock_storage_service client.app.dependency_overrides[get_storage_service] = lambda: mock_storage_service
response = client.post( response = client.post(
f"api/v1/files/upload/{flow.id}", files={"file": ("test.txt", b"test content")}, headers=headers f"api/v1/files/upload/{flow.id}",
files={"file": ("test.txt", b"test content")},
headers=headers,
) )
assert response.status_code == 201 assert response.status_code == 201
assert response.json() == {"flowId": str(flow.id), "file_path": f"{flow.id}/test.txt"} assert response.json() == {
"flowId": str(flow.id),
"file_path": f"{flow.id}/test.txt",
}
def test_download_file(client, mock_storage_service, created_api_key, flow): def test_download_file(client, mock_storage_service, created_api_key, flow):
@ -64,9 +70,16 @@ def test_file_operations(client, created_api_key, flow):
file_content = b"Hello, world!" file_content = b"Hello, world!"
# Step 1: Upload the file # Step 1: Upload the file
response = client.post(f"api/v1/files/upload/{flow_id}", files={"file": (file_name, file_content)}, headers=headers) response = client.post(
f"api/v1/files/upload/{flow_id}",
files={"file": (file_name, file_content)},
headers=headers,
)
assert response.status_code == 201 assert response.status_code == 201
assert response.json() == {"flowId": str(flow_id), "file_path": f"{flow_id}/{file_name}"} assert response.json() == {
"flowId": str(flow_id),
"file_path": f"{flow_id}/{file_name}",
}
# Step 2: List files in the folder # Step 2: List files in the folder
response = client.get(f"api/v1/files/list/{flow_id}", headers=headers) response = client.get(f"api/v1/files/list/{flow_id}", headers=headers)
@ -75,13 +88,19 @@ def test_file_operations(client, created_api_key, flow):
# Step 3: Download the file and verify its content # Step 3: Download the file and verify its content
mime_type = build_content_type_from_extension(file_name.split(".")[-1]) mime_type = build_content_type_from_extension(file_name.split(".")[-1])
response = client.get(f"api/v1/files/download/{flow_id}/{file_name}", headers=headers) response = client.get(
f"api/v1/files/download/{flow_id}/{file_name}", headers=headers
)
assert response.status_code == 200 assert response.status_code == 200
assert response.content == file_content assert response.content == file_content
assert mime_type in response.headers["content-type"] # the headers are application/octet-stream
assert response.headers["content-type"] == "application/octet-stream"
# mime_type is inside media_type
# Step 4: Delete the file # Step 4: Delete the file
response = client.delete(f"api/v1/files/delete/{flow_id}/{file_name}", headers=headers) response = client.delete(
f"api/v1/files/delete/{flow_id}/{file_name}", headers=headers
)
assert response.status_code == 200 assert response.status_code == 200
assert response.json() == {"message": f"File {file_name} deleted successfully"} assert response.json() == {"message": f"File {file_name} deleted successfully"}

View file

@ -1,14 +1,10 @@
import copy import copy
import json import json
import os
import pickle import pickle
from pathlib import Path
from typing import Type, Union from typing import Type, Union
import pytest import pytest
from langchain.agents import AgentExecutor
from langchain.chains.base import Chain
from langchain.llms.fake import FakeListLLM
from langflow.graph import Graph from langflow.graph import Graph
from langflow.graph.edge.base import Edge from langflow.graph.edge.base import Edge
from langflow.graph.graph.utils import ( from langflow.graph.graph.utils import (
@ -21,8 +17,7 @@ from langflow.graph.graph.utils import (
update_template, update_template,
) )
from langflow.graph.vertex.base import Vertex from langflow.graph.vertex.base import Vertex
from langflow.graph.vertex.types import FileToolVertex, LLMVertex, ToolkitVertex from langflow.initial_setup.setup import load_starter_projects
from langflow.processing.process import get_result_and_thought
from langflow.utils.payload import get_root_vertex from langflow.utils.payload import get_root_vertex
# Test cases for the graph module # Test cases for the graph module
@ -44,7 +39,13 @@ def sample_nodes():
return [ return [
{ {
"id": "node1", "id": "node1",
"data": {"node": {"template": {"some_field": {"show": True, "advanced": False, "name": "Name1"}}}}, "data": {
"node": {
"template": {
"some_field": {"show": True, "advanced": False, "name": "Name1"}
}
}
},
}, },
{ {
"id": "node2", "id": "node2",
@ -62,7 +63,11 @@ def sample_nodes():
}, },
{ {
"id": "node3", "id": "node3",
"data": {"node": {"template": {"unrelated_field": {"show": True, "advanced": True}}}}, "data": {
"node": {
"template": {"unrelated_field": {"show": True, "advanced": True}}
}
},
}, },
] ]
@ -147,9 +152,15 @@ def test_get_node_neighbors_basic(basic_graph):
# Root Node is an Agent, it requires an LLMChain and tools # Root Node is an Agent, it requires an LLMChain and tools
# We need to check if there is a Chain in the one of the neighbors' # We need to check if there is a Chain in the one of the neighbors'
# data attribute in the type key # data attribute in the type key
assert any("ConversationBufferMemory" in neighbor.data["type"] for neighbor, val in neighbors.items() if val) assert any(
"ConversationBufferMemory" in neighbor.data["type"]
for neighbor, val in neighbors.items()
if val
)
assert any("OpenAI" in neighbor.data["type"] for neighbor, val in neighbors.items() if val) assert any(
"OpenAI" in neighbor.data["type"] for neighbor, val in neighbors.items() if val
)
def test_get_node(basic_graph): def test_get_node(basic_graph):
@ -231,59 +242,6 @@ def test_build_params(basic_graph):
assert "memory" in root.params assert "memory" in root.params
@pytest.mark.asyncio
async def test_build(basic_graph):
"""Test Node's build method"""
await assert_agent_was_built(basic_graph)
async def assert_agent_was_built(graph):
"""Assert that the agent was built"""
assert isinstance(graph, Graph)
# Now we test the build method
# Build the Agent
result = await graph.build()
# The agent should be a AgentExecutor
assert isinstance(result, Chain)
def test_llm_node_build(basic_graph):
llm_node = get_node_by_type(basic_graph, LLMVertex)
assert llm_node is not None
built_object = llm_node.build()
assert built_object is not None
def test_toolkit_node_build(client, openapi_graph):
# Write a file to the disk
file_path = "api-with-examples.yaml"
with open(file_path, "w") as f:
f.write("openapi: 3.0.0")
toolkit_node = get_node_by_type(openapi_graph, ToolkitVertex)
assert toolkit_node is not None
built_object = toolkit_node.build()
assert built_object is not None
# Remove the file
os.remove(file_path)
assert not Path(file_path).exists()
def test_file_tool_node_build(client, openapi_graph):
file_path = "api-with-examples.yaml"
with open(file_path, "w") as f:
f.write("openapi: 3.0.0")
assert Path(file_path).exists()
file_tool_node = get_node_by_type(openapi_graph, FileToolVertex)
assert file_tool_node is not None
built_object = file_tool_node.build()
assert built_object is not None
# Remove the file
os.remove(file_path)
assert not Path(file_path).exists()
# def test_wrapper_node_build(openapi_graph): # def test_wrapper_node_build(openapi_graph):
# wrapper_node = get_node_by_type(openapi_graph, WrapperVertex) # wrapper_node = get_node_by_type(openapi_graph, WrapperVertex)
# assert wrapper_node is not None # assert wrapper_node is not None
@ -291,29 +249,6 @@ def test_file_tool_node_build(client, openapi_graph):
# assert built_object is not None # assert built_object is not None
@pytest.mark.asyncio
async def test_get_result_and_thought(basic_graph):
"""Test the get_result_and_thought method"""
responses = [
"Final Answer: I am a response",
]
message = {"input": "Hello"}
# Find the node that is an LLMNode and change the
# _built_object to a FakeListLLM
llm_node = get_node_by_type(basic_graph, LLMVertex)
assert llm_node is not None
llm_node._built_object = FakeListLLM(responses=responses)
llm_node._built = True
langchain_object = await basic_graph.build()
# assert all nodes are built
assert all(node._built for node in basic_graph.vertices)
# now build again and check if FakeListLLM was used
# Get the result and thought
result = get_result_and_thought(langchain_object, message)
assert isinstance(result, dict)
def test_find_last_node(grouped_chat_json_flow): def test_find_last_node(grouped_chat_json_flow):
grouped_chat_data = json.loads(grouped_chat_json_flow).get("data") grouped_chat_data = json.loads(grouped_chat_json_flow).get("data")
nodes, edges = grouped_chat_data["nodes"], grouped_chat_data["edges"] nodes, edges = grouped_chat_data["nodes"], grouped_chat_data["edges"]
@ -324,7 +259,9 @@ def test_find_last_node(grouped_chat_json_flow):
def test_ungroup_node(grouped_chat_json_flow): def test_ungroup_node(grouped_chat_json_flow):
grouped_chat_data = json.loads(grouped_chat_json_flow).get("data") grouped_chat_data = json.loads(grouped_chat_json_flow).get("data")
group_node = grouped_chat_data["nodes"][2] # Assuming the first node is a group node group_node = grouped_chat_data["nodes"][
2
] # Assuming the first node is a group node
base_flow = copy.deepcopy(grouped_chat_data) base_flow = copy.deepcopy(grouped_chat_data)
ungroup_node(group_node["data"], base_flow) ungroup_node(group_node["data"], base_flow)
# after ungroup_node is called, the base_flow and grouped_chat_data should be different # after ungroup_node is called, the base_flow and grouped_chat_data should be different
@ -376,9 +313,14 @@ def test_process_flow_one_group(one_grouped_chat_json_flow):
assert "edges" in processed_flow assert "edges" in processed_flow
# Now get the node that has ChatOpenAI in its id # Now get the node that has ChatOpenAI in its id
chat_openai_node = next((node for node in processed_flow["nodes"] if "ChatOpenAI" in node["id"]), None) chat_openai_node = next(
(node for node in processed_flow["nodes"] if "ChatOpenAI" in node["id"]), None
)
assert chat_openai_node is not None assert chat_openai_node is not None
assert chat_openai_node["data"]["node"]["template"]["openai_api_key"]["value"] == "test" assert (
chat_openai_node["data"]["node"]["template"]["openai_api_key"]["value"]
== "test"
)
def test_process_flow_vector_store_grouped(vector_store_grouped_json_flow): def test_process_flow_vector_store_grouped(vector_store_grouped_json_flow):
@ -427,11 +369,17 @@ def test_update_template(sample_template, sample_nodes):
assert node1_updated["data"]["node"]["template"]["some_field"]["show"] is True assert node1_updated["data"]["node"]["template"]["some_field"]["show"] is True
assert node1_updated["data"]["node"]["template"]["some_field"]["advanced"] is False assert node1_updated["data"]["node"]["template"]["some_field"]["advanced"] is False
assert node1_updated["data"]["node"]["template"]["some_field"]["display_name"] == "Name1" assert (
node1_updated["data"]["node"]["template"]["some_field"]["display_name"]
== "Name1"
)
assert node2_updated["data"]["node"]["template"]["other_field"]["show"] is False assert node2_updated["data"]["node"]["template"]["other_field"]["show"] is False
assert node2_updated["data"]["node"]["template"]["other_field"]["advanced"] is True assert node2_updated["data"]["node"]["template"]["other_field"]["advanced"] is True
assert node2_updated["data"]["node"]["template"]["other_field"]["display_name"] == "DisplayName2" assert (
node2_updated["data"]["node"]["template"]["other_field"]["display_name"]
== "DisplayName2"
)
# Ensure node3 remains unchanged # Ensure node3 remains unchanged
assert node3_updated == sample_nodes[2] assert node3_updated == sample_nodes[2]
@ -462,7 +410,9 @@ def test_set_new_target_handle():
"data": { "data": {
"node": { "node": {
"flow": True, "flow": True,
"template": {"field_1": {"proxy": {"field": "new_field", "id": "new_id"}}}, "template": {
"field_1": {"proxy": {"field": "new_field", "id": "new_id"}}
},
} }
} }
} }
@ -482,30 +432,30 @@ def test_update_source_handle():
"nodes": [{"id": "some_node"}, {"id": "last_node"}], "nodes": [{"id": "some_node"}, {"id": "last_node"}],
"edges": [{"source": "some_node"}], "edges": [{"source": "some_node"}],
} }
updated_edge = update_source_handle(new_edge, flow_data["nodes"], flow_data["edges"]) updated_edge = update_source_handle(
new_edge, flow_data["nodes"], flow_data["edges"]
)
assert updated_edge["source"] == "last_node" assert updated_edge["source"] == "last_node"
assert updated_edge["data"]["sourceHandle"]["id"] == "last_node" assert updated_edge["data"]["sourceHandle"]["id"] == "last_node"
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_pickle_graph(json_vector_store): async def test_pickle_graph(json_vector_store):
loaded_json = json.loads(json_vector_store) starter_projects = load_starter_projects()
graph = Graph.from_payload(loaded_json) data = starter_projects[0]["data"]
graph = Graph.from_payload(data)
assert isinstance(graph, Graph) assert isinstance(graph, Graph)
first_result = await graph.build()
assert isinstance(first_result, AgentExecutor)
pickled = pickle.dumps(graph) pickled = pickle.dumps(graph)
assert pickled is not None assert pickled is not None
unpickled = pickle.loads(pickled) unpickled = pickle.loads(pickled)
assert unpickled is not None assert unpickled is not None
result = await unpickled.build()
assert isinstance(result, AgentExecutor)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_pickle_each_vertex(json_vector_store): async def test_pickle_each_vertex(json_vector_store):
loaded_json = json.loads(json_vector_store) starter_projects = load_starter_projects()
graph = Graph.from_payload(loaded_json) data = starter_projects[0]["data"]
graph = Graph.from_payload(data)
assert isinstance(graph, Graph) assert isinstance(graph, Graph)
for vertex in graph.vertices: for vertex in graph.vertices:
await vertex.build() await vertex.build()

View file

@ -17,4 +17,3 @@ def test_load_flow_from_json_with_tweaks():
loaded = load_flow_from_json(pytest.BASIC_EXAMPLE_PATH, tweaks=tweaks) loaded = load_flow_from_json(pytest.BASIC_EXAMPLE_PATH, tweaks=tweaks)
assert loaded is not None assert loaded is not None
assert isinstance(loaded, Graph) assert isinstance(loaded, Graph)
assert loaded.llm.model_name == "gpt-3.5-turbo-16k-0613"

View file

@ -268,9 +268,13 @@ async def test_load_langchain_object_with_cached_session(client, basic_graph_dat
# Provide a non-existent session_id # Provide a non-existent session_id
session_service = get_session_service() session_service = get_session_service()
session_id1 = "non-existent-session-id" session_id1 = "non-existent-session-id"
graph1, artifacts1 = await session_service.load_session(session_id1, basic_graph_data) graph1, artifacts1 = await session_service.load_session(
session_id1, basic_graph_data
)
# Use the new session_id to get the langchain_object again # Use the new session_id to get the langchain_object again
graph2, artifacts2 = await session_service.load_session(session_id1, basic_graph_data) graph2, artifacts2 = await session_service.load_session(
session_id1, basic_graph_data
)
assert graph1 == graph2 assert graph1 == graph2
assert artifacts1 == artifacts2 assert artifacts1 == artifacts2
@ -282,11 +286,15 @@ async def test_load_langchain_object_with_no_cached_session(client, basic_graph_
session_service = get_session_service() session_service = get_session_service()
session_id1 = "non-existent-session-id" session_id1 = "non-existent-session-id"
session_id = session_service.build_key(session_id1, basic_graph_data) session_id = session_service.build_key(session_id1, basic_graph_data)
graph1, artifacts1 = await session_service.load_session(session_id, basic_graph_data) graph1, artifacts1 = await session_service.load_session(
session_id, data_graph=basic_graph_data, flow_id="flow_id"
)
# Clear the cache # Clear the cache
session_service.clear_session(session_id) session_service.clear_session(session_id)
# Use the new session_id to get the langchain_object again # Use the new session_id to get the langchain_object again
graph2, artifacts2 = await session_service.load_session(session_id, basic_graph_data) graph2, artifacts2 = await session_service.load_session(
session_id, data_graph=basic_graph_data, flow_id="flow_id"
)
assert id(graph1) != id(graph2) assert id(graph1) != id(graph2)
# Since the cache was cleared, objects should be different # Since the cache was cleared, objects should be different
@ -297,8 +305,12 @@ async def test_load_langchain_object_without_session_id(client, basic_graph_data
# Provide a non-existent session_id # Provide a non-existent session_id
session_service = get_session_service() session_service = get_session_service()
session_id1 = None session_id1 = None
graph1, artifacts1 = await session_service.load_session(session_id1, basic_graph_data) graph1, artifacts1 = await session_service.load_session(
session_id1, data_graph=basic_graph_data, flow_id="flow_id"
)
# Use the new session_id to get the langchain_object again # Use the new session_id to get the langchain_object again
graph2, artifacts2 = await session_service.load_session(session_id1, basic_graph_data) graph2, artifacts2 = await session_service.load_session(
session_id1, data_graph=basic_graph_data, flow_id="flow_id"
)
assert graph1 == graph2 assert graph1 == graph2

View file

@ -1,7 +1,9 @@
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
from langflow.services.database.models.user.model import User from langflow.services.settings.constants import (
from langflow.services.settings.constants import DEFAULT_SUPERUSER, DEFAULT_SUPERUSER_PASSWORD DEFAULT_SUPERUSER,
DEFAULT_SUPERUSER_PASSWORD,
)
from langflow.services.utils import teardown_superuser from langflow.services.utils import teardown_superuser
# @patch("langflow.services.deps.get_session") # @patch("langflow.services.deps.get_session")
@ -89,7 +91,9 @@ from langflow.services.utils import teardown_superuser
@patch("langflow.services.deps.get_settings_service") @patch("langflow.services.deps.get_settings_service")
@patch("langflow.services.deps.get_session") @patch("langflow.services.deps.get_session")
def test_teardown_superuser_default_superuser(mock_get_session, mock_get_settings_service): def test_teardown_superuser_default_superuser(
mock_get_session, mock_get_settings_service
):
mock_settings_service = MagicMock() mock_settings_service = MagicMock()
mock_settings_service.auth_settings.AUTO_LOGIN = True mock_settings_service.auth_settings.AUTO_LOGIN = True
mock_settings_service.auth_settings.SUPERUSER = DEFAULT_SUPERUSER mock_settings_service.auth_settings.SUPERUSER = DEFAULT_SUPERUSER
@ -104,18 +108,14 @@ def test_teardown_superuser_default_superuser(mock_get_session, mock_get_setting
teardown_superuser(mock_settings_service, mock_session) teardown_superuser(mock_settings_service, mock_session)
mock_session.query.assert_called_once_with(User) mock_session.query.assert_not_called()
actual_expr = mock_session.query.return_value.filter.call_args[0][0]
expected_expr = User.username == DEFAULT_SUPERUSER
assert str(actual_expr) == str(expected_expr)
mock_session.delete.assert_called_once_with(mock_user)
mock_session.commit.assert_called_once()
@patch("langflow.services.deps.get_settings_service") @patch("langflow.services.deps.get_settings_service")
@patch("langflow.services.deps.get_session") @patch("langflow.services.deps.get_session")
def test_teardown_superuser_no_default_superuser(mock_get_session, mock_get_settings_service): def test_teardown_superuser_no_default_superuser(
mock_get_session, mock_get_settings_service
):
ADMIN_USER_NAME = "admin_user" ADMIN_USER_NAME = "admin_user"
mock_settings_service = MagicMock() mock_settings_service = MagicMock()
mock_settings_service.auth_settings.AUTO_LOGIN = False mock_settings_service.auth_settings.AUTO_LOGIN = False