diff --git a/tests/test_cli.py b/tests/test_cli.py index efde059f6..60595f48a 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -2,6 +2,7 @@ from pathlib import Path from tempfile import tempdir import pytest + from langflow.__main__ import app from langflow.services import deps diff --git a/tests/test_custom_component.py b/tests/test_custom_component.py index a47bb8d68..1a06348eb 100644 --- a/tests/test_custom_component.py +++ b/tests/test_custom_component.py @@ -3,19 +3,21 @@ import types from uuid import uuid4 import pytest -from fastapi import HTTPException from langchain_core.documents import Document + 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 ( Component, ComponentCodeNullError, - ComponentFunctionEntrypointNameNullError, ) from langflow.services.database.models.flow import Flow, FlowCreate code_default = """ -from langflow import Prompt +from langflow.field_typing import Prompt from langflow.interface.custom.custom_component import CustomComponent from langchain.llms.base import BaseLLM @@ -96,23 +98,15 @@ def test_component_code_null_error(): 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(): """ Test the initialization of the CustomComponent class. """ 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.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. """ - 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() assert isinstance(config, dict) @@ -130,8 +126,10 @@ def test_custom_component_get_function(): """ Test the get_function property of the CustomComponent class. """ - custom_component = CustomComponent(code="def build(): pass", function_entrypoint_name="build") - my_function = custom_component.get_function + custom_component = CustomComponent( + code="def build(): pass", function_entrypoint_name="build" + ) + my_function = custom_component.get_function() assert isinstance(my_function, types.FunctionType) @@ -215,7 +213,9 @@ def test_custom_component_get_function_entrypoint_args(): Test the get_function_entrypoint_args 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 assert len(args) == 4 assert args[0]["name"] == "self" @@ -229,7 +229,9 @@ def test_custom_component_get_function_entrypoint_return_type(): 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 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. """ - 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 assert class_name == "YourComponent" @@ -248,7 +252,9 @@ def test_custom_component_get_function_valid(): Test the get_function property of the CustomComponent 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 assert callable(my_function) @@ -283,7 +289,9 @@ def test_code_parser_parse_callable_details_no_args(): parser = CodeParser("") node = ast.FunctionDef( 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=[], decorator_list=[], returns=None, @@ -329,7 +337,9 @@ def test_code_parser_parse_function_def_not_init(): parser = CodeParser("") stmt = ast.FunctionDef( 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=[], decorator_list=[], returns=None, @@ -347,7 +357,9 @@ def test_code_parser_parse_function_def_init(): parser = CodeParser("") stmt = ast.FunctionDef( 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=[], decorator_list=[], returns=None, @@ -373,7 +385,7 @@ def test_custom_component_class_template_validation_no_code(): raises the HTTPException when the code is None. """ custom_component = CustomComponent(code=None, function_entrypoint_name="build") - with pytest.raises(HTTPException): + with pytest.raises(TypeError): 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 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): 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 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): custom_component.build() @@ -444,7 +460,7 @@ def test_custom_component_build_not_implemented(): def test_build_config_no_code(): component = CustomComponent(code=None) - assert component.get_function_entrypoint_args == "" + assert component.get_function_entrypoint_args == [] assert component.get_function_entrypoint_return_type == [] @@ -470,7 +486,9 @@ def test_flow(db): } # 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 db.add(flow) diff --git a/tests/test_data_components.py b/tests/test_data_components.py new file mode 100644 index 000000000..6aa9a4768 --- /dev/null +++ b/tests/test_data_components.py @@ -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) diff --git a/tests/test_database.py b/tests/test_database.py index a5ed533d4..f04aaf9ff 100644 --- a/tests/test_database.py +++ b/tests/test_database.py @@ -3,12 +3,14 @@ from uuid import UUID, uuid4 import orjson import pytest from fastapi.testclient import TestClient +from sqlmodel import Session + 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.flow import Flow, FlowCreate, FlowUpdate from langflow.services.database.utils import session_getter from langflow.services.deps import get_db_service -from sqlmodel import Session @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) data = flow["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 # flow is optional so we can create a flow without a 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.json()["name"] == flow.name 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 -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) data = flow["data"] @@ -89,7 +97,9 @@ def test_update_flow(client: TestClient, json_flow: str, active_user, logged_in_ description="updated description", 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.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 -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) data = flow["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" -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) data = flow["data"] # Create test data @@ -119,7 +133,9 @@ def test_create_flows(client: TestClient, session: Session, json_flow: str, logg ] ) # 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 assert response.status_code == 201 # 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 -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) data = flow["data"] # Create test data @@ -182,7 +200,7 @@ def test_download_file( with session_getter(db_manager) as session: for flow in flow_list.flows: 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.commit() # Make request to endpoint @@ -191,7 +209,9 @@ def test_download_file( assert response.status_code == 200, response.json() # Check response data 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]["description"] == "description" assert response_data[0]["data"] == data @@ -200,7 +220,9 @@ def test_download_file( 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"} response = client.post("api/v1/flows/", json=flow, headers=logged_in_headers) 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 -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) data = flow_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"] 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) - response2 = client.put(f"api/v1/flows/{flow_id}", json=updated_flow.dict(), headers=logged_in_headers) + response1 = client.put( + 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() -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) data = flow_data["data"] uuid = uuid4() @@ -233,7 +265,9 @@ def test_update_nonexistent_flow(client: TestClient, json_flow: str, active_user description="description", 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 @@ -243,7 +277,8 @@ def test_delete_nonexistent_flow(client: TestClient, active_user, logged_in_head 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) + starter_projects = load_starter_projects() assert response.status_code == 200 - assert len(response.json()) == 0 + assert len(response.json()) == len(starter_projects) diff --git a/tests/test_endpoints.py b/tests/test_endpoints.py index 6a2f9cff4..455a3c0f6 100644 --- a/tests/test_endpoints.py +++ b/tests/test_endpoints.py @@ -4,7 +4,6 @@ import pytest from fastapi.testclient import TestClient 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.database.models.api_key.model import ApiKey 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, 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() time.sleep(sleep_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() 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 session.add(api_key) session.commit() @@ -289,13 +295,17 @@ def test_get_all(client: TestClient, logged_in_headers): dir_reader = DirectoryReader(settings.COMPONENTS_PATH[0]) files = dir_reader.get_files() # 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() # We need to test the custom nodes assert len(all_names) > len(files) - assert "Prompt" in json_response["prompts"] - # All CUSTOM_TOOLS(dict) should be in the response - assert all(tool in json_response["tools"] for tool in CUSTOM_TOOLS.keys()) + assert "ChatInput" in json_response["inputs"] + assert "Prompt" in json_response["inputs"] + assert "ChatOutput" in json_response["outputs"] 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): - response = client.get("/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 + response = client.get( + "/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): 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 "ids" in response.json() # The response should contain the list in this order # ['ConversationBufferMemory-Lu2Nb', 'PromptTemplate-5Q0W8', 'ChatOpenAI-vy7fV', 'LLMChain-UjBh1'] # 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 == [ "ChatOpenAI", "PromptTemplate", "ConversationBufferMemory", - "LLMChain", ] 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 -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"] - 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 diff --git a/tests/test_files.py b/tests/test_files.py index ad66f1497..0b3086b6f 100644 --- a/tests/test_files.py +++ b/tests/test_files.py @@ -1,6 +1,7 @@ from unittest.mock import MagicMock import pytest + from langflow.services.deps import get_storage_service from langflow.services.storage.service import StorageService 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 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.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): @@ -64,9 +70,16 @@ def test_file_operations(client, created_api_key, flow): file_content = b"Hello, world!" # 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.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 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 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.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 - 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.json() == {"message": f"File {file_name} deleted successfully"} diff --git a/tests/test_graph.py b/tests/test_graph.py index 6cc19a101..3605d4c2a 100644 --- a/tests/test_graph.py +++ b/tests/test_graph.py @@ -1,14 +1,10 @@ import copy import json -import os import pickle -from pathlib import Path from typing import Type, Union 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.edge.base import Edge from langflow.graph.graph.utils import ( @@ -21,8 +17,7 @@ from langflow.graph.graph.utils import ( update_template, ) from langflow.graph.vertex.base import Vertex -from langflow.graph.vertex.types import FileToolVertex, LLMVertex, ToolkitVertex -from langflow.processing.process import get_result_and_thought +from langflow.initial_setup.setup import load_starter_projects from langflow.utils.payload import get_root_vertex # Test cases for the graph module @@ -44,7 +39,13 @@ def sample_nodes(): return [ { "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", @@ -62,7 +63,11 @@ def sample_nodes(): }, { "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 # We need to check if there is a Chain in the one of the neighbors' # 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): @@ -231,59 +242,6 @@ def test_build_params(basic_graph): 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): # wrapper_node = get_node_by_type(openapi_graph, WrapperVertex) # 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 -@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): grouped_chat_data = json.loads(grouped_chat_json_flow).get("data") 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): 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) ungroup_node(group_node["data"], base_flow) # 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 # 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["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): @@ -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"]["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"]["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 assert node3_updated == sample_nodes[2] @@ -462,7 +410,9 @@ def test_set_new_target_handle(): "data": { "node": { "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"}], "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["data"]["sourceHandle"]["id"] == "last_node" @pytest.mark.asyncio async def test_pickle_graph(json_vector_store): - loaded_json = json.loads(json_vector_store) - graph = Graph.from_payload(loaded_json) + starter_projects = load_starter_projects() + data = starter_projects[0]["data"] + graph = Graph.from_payload(data) assert isinstance(graph, Graph) - first_result = await graph.build() - assert isinstance(first_result, AgentExecutor) pickled = pickle.dumps(graph) assert pickled is not None unpickled = pickle.loads(pickled) assert unpickled is not None - result = await unpickled.build() - assert isinstance(result, AgentExecutor) @pytest.mark.asyncio async def test_pickle_each_vertex(json_vector_store): - loaded_json = json.loads(json_vector_store) - graph = Graph.from_payload(loaded_json) + starter_projects = load_starter_projects() + data = starter_projects[0]["data"] + graph = Graph.from_payload(data) assert isinstance(graph, Graph) for vertex in graph.vertices: await vertex.build() diff --git a/tests/test_loading.py b/tests/test_loading.py index 74aba58db..b95791139 100644 --- a/tests/test_loading.py +++ b/tests/test_loading.py @@ -17,4 +17,3 @@ def test_load_flow_from_json_with_tweaks(): loaded = load_flow_from_json(pytest.BASIC_EXAMPLE_PATH, tweaks=tweaks) assert loaded is not None assert isinstance(loaded, Graph) - assert loaded.llm.model_name == "gpt-3.5-turbo-16k-0613" diff --git a/tests/test_process.py b/tests/test_process.py index 1b75f9406..f5bae0569 100644 --- a/tests/test_process.py +++ b/tests/test_process.py @@ -268,9 +268,13 @@ async def test_load_langchain_object_with_cached_session(client, basic_graph_dat # Provide a non-existent session_id session_service = get_session_service() 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 - 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 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_id1 = "non-existent-session-id" 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 session_service.clear_session(session_id) # 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) # 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 session_service = get_session_service() 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 - 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 diff --git a/tests/test_setup_superuser.py b/tests/test_setup_superuser.py index d3033b728..c7d343818 100644 --- a/tests/test_setup_superuser.py +++ b/tests/test_setup_superuser.py @@ -1,7 +1,9 @@ from unittest.mock import MagicMock, patch -from langflow.services.database.models.user.model import User -from langflow.services.settings.constants import DEFAULT_SUPERUSER, DEFAULT_SUPERUSER_PASSWORD +from langflow.services.settings.constants import ( + DEFAULT_SUPERUSER, + DEFAULT_SUPERUSER_PASSWORD, +) from langflow.services.utils import teardown_superuser # @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_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.auth_settings.AUTO_LOGIN = True 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) - mock_session.query.assert_called_once_with(User) - 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() + mock_session.query.assert_not_called() @patch("langflow.services.deps.get_settings_service") @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" mock_settings_service = MagicMock() mock_settings_service.auth_settings.AUTO_LOGIN = False