From b8a41037ec909ab0e96a64c9b9ba85f99e0c5e8e Mon Sep 17 00:00:00 2001 From: Gabriel Almeida Date: Tue, 28 Mar 2023 19:17:26 -0300 Subject: [PATCH] feat: code validation, endpoint and tests --- src/backend/langflow/api/base.py | 24 +++++++ src/backend/langflow/api/endpoints.py | 15 +++-- .../langflow/interface/custom_types.py | 17 ++--- src/backend/langflow/interface/loading.py | 4 +- src/backend/langflow/utils/constants.py | 1 + src/backend/langflow/utils/util.py | 15 ----- src/backend/langflow/utils/validate.py | 63 ++++++++++++++++++ tests/conftest.py | 5 ++ tests/test_custom_types.py | 19 ++++++ tests/test_endpoints.py | 66 +++++++++++++++++++ tests/test_validate_code.py | 39 +++++++++++ 11 files changed, 238 insertions(+), 30 deletions(-) create mode 100644 src/backend/langflow/api/base.py create mode 100644 src/backend/langflow/utils/validate.py create mode 100644 tests/test_custom_types.py create mode 100644 tests/test_validate_code.py diff --git a/src/backend/langflow/api/base.py b/src/backend/langflow/api/base.py new file mode 100644 index 000000000..931bfd2ab --- /dev/null +++ b/src/backend/langflow/api/base.py @@ -0,0 +1,24 @@ +from pydantic import BaseModel, validator +import json + + +class Code(BaseModel): + code: str + + @validator("code") + def validate_code(cls, v): + return v + + +# Build ValidationResponse class for {"imports": {"errors": []}, "function": {"errors": []}} +class ValidationResponse(BaseModel): + imports: dict + function: dict + + @validator("imports") + def validate_imports(cls, v): + return v or {"errors": []} + + @validator("function") + def validate_function(cls, v): + return v or {"errors": []} diff --git a/src/backend/langflow/api/endpoints.py b/src/backend/langflow/api/endpoints.py index a7c26c20e..13b6bee7e 100644 --- a/src/backend/langflow/api/endpoints.py +++ b/src/backend/langflow/api/endpoints.py @@ -1,10 +1,12 @@ from typing import Any, Dict -from fastapi import APIRouter, HTTPException +from fastapi import APIRouter, HTTPException, Response +from langflow.api.base import Code, ValidationResponse from langflow.interface.custom_types import PythonFunction from langflow.interface.run import process_graph from langflow.interface.types import build_langchain_types_dict +from langflow.utils.validate import validate_code # build router @@ -24,10 +26,13 @@ def get_load(data: Dict[str, Any]): return HTTPException(status_code=500, detail=str(e)) -@router.post("/validate", status_code=200) -def validate_code(data: PythonFunction): +@router.post("/validate", status_code=200, response_model=ValidationResponse) +def post_validate_code(code: Code): try: - # if the data var gets here then it is valid python code - return {"valid": True} + errors = validate_code(code.code) + return ValidationResponse( + imports=errors.get("imports", {}), + function=errors.get("function", {}), + ) except Exception as e: return HTTPException(status_code=500, detail=str(e)) diff --git a/src/backend/langflow/interface/custom_types.py b/src/backend/langflow/interface/custom_types.py index 41e745212..a8d99c521 100644 --- a/src/backend/langflow/interface/custom_types.py +++ b/src/backend/langflow/interface/custom_types.py @@ -1,11 +1,12 @@ -from typing import Callable -from langflow.utils import util +from typing import Callable, Optional +from langflow.utils import validate from pydantic import BaseModel, validator - class Function(BaseModel): code: str + function: Optional[Callable] = None + imports: Optional[str] = None # Eval code and store the function def __init__(self, **data): @@ -14,16 +15,16 @@ class Function(BaseModel): # Validate the function @validator("code") def validate_func(cls, v): - # Validate with LangChain's tool decorator - func = util.eval_function(v) - if not isinstance(func, Callable): - raise ValueError("Function must be a callable") + try: + validate.eval_function(v) + except Exception as e: + raise e return v def get_function(self): """Get the function""" - return util.eval_function(self.code) + return validate.eval_function(self.code) class PythonFunction(Function): diff --git a/src/backend/langflow/interface/loading.py b/src/backend/langflow/interface/loading.py index 38c84144a..bfb919f5d 100644 --- a/src/backend/langflow/interface/loading.py +++ b/src/backend/langflow/interface/loading.py @@ -22,7 +22,7 @@ from langchain.llms.base import BaseLLM from langchain.llms.loading import load_llm_from_config from langflow.interface.types import get_type_list -from langflow.utils import payload, util +from langflow.utils import payload, util, validate def instantiate_class(node_type: str, base_type: str, params: Dict) -> Any: @@ -43,7 +43,7 @@ def instantiate_class(node_type: str, base_type: str, params: Dict) -> Any: # as the instance function_string = params["code"] if isinstance(function_string, str): - return util.eval_function(function_string) + return validate.eval_function(function_string) raise ValueError("Function should be a string") else: if "tools" not in params: diff --git a/src/backend/langflow/utils/constants.py b/src/backend/langflow/utils/constants.py index c4a4eae18..2d101ab98 100644 --- a/src/backend/langflow/utils/constants.py +++ b/src/backend/langflow/utils/constants.py @@ -10,5 +10,6 @@ CHAT_OPENAI_MODELS = ["gpt-3.5-turbo", "gpt-4", "gpt-4-32k"] DEFAULT_PYTHON_FUNCTION = """ def python_function(text: str) -> str: + \"\"\"This is a default python function that returns the input text\"\"\" return text """ diff --git a/src/backend/langflow/utils/util.py b/src/backend/langflow/utils/util.py index fef223762..8d3b00e2a 100644 --- a/src/backend/langflow/utils/util.py +++ b/src/backend/langflow/utils/util.py @@ -123,21 +123,6 @@ def build_template_from_class( } -def eval_function(function_string: str): - # Create an empty dictionary to serve as a separate namespace - namespace: Dict = {} - - # Execute the code string in the new namespace - exec(function_string, namespace) - function_object = next( - (obj for name, obj in namespace.items() if isinstance(obj, types.FunctionType)), - None, - ) - if function_object is None: - raise ValueError("Function string does not contain a function") - return function_object - - def get_base_classes(cls): """Get the base classes of a class. These are used to determine the output of the nodes. diff --git a/src/backend/langflow/utils/validate.py b/src/backend/langflow/utils/validate.py new file mode 100644 index 000000000..3dfde9a3b --- /dev/null +++ b/src/backend/langflow/utils/validate.py @@ -0,0 +1,63 @@ +import ast +import importlib +import types +from typing import Dict + + +def validate_code(code): + # Initialize the errors dictionary + errors = {"imports": {"errors": []}, "function": {"errors": []}} + + # Parse the code string into an abstract syntax tree (AST) + try: + tree = ast.parse(code) + except Exception as e: + errors["function"]["errors"].append(str(e)) + return errors + + # Add a dummy type_ignores field to the AST + if not hasattr(ast, "TypeIgnore"): + + class TypeIgnore(ast.AST): + _fields = () + + ast.TypeIgnore = TypeIgnore + tree.type_ignores = [] + + # Evaluate the import statements + for node in tree.body: + if isinstance(node, ast.Import): + for alias in node.names: + try: + importlib.import_module(alias.name) + except ModuleNotFoundError as e: + errors["imports"]["errors"].append(str(e)) + + # Evaluate the function definition + for node in tree.body: + if isinstance(node, ast.FunctionDef): + code_obj = compile( + ast.Module(body=[node], type_ignores=[]), "", "exec" + ) + try: + exec(code_obj) + except Exception as e: + errors["function"]["errors"].append(str(e)) + + # Return the errors dictionary + return errors + + +def eval_function(function_string: str): + # Create an empty dictionary to serve as a separate namespace + namespace: Dict = {} + + # Execute the code string in the new namespace + exec(function_string, namespace) + function_object = next( + (obj for name, obj in namespace.items() if isinstance(obj, types.FunctionType)), + None, + ) + if function_object is None: + raise ValueError("Function string does not contain a function") + return function_object diff --git a/tests/conftest.py b/tests/conftest.py index 37657d8c7..5fe1de280 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -11,6 +11,11 @@ def pytest_configure(): Path(__file__).parent.absolute() / "data" / "complex_example.json" ) + pytest.CODE_WITH_SYNTAX_ERROR = """ +def get_text(): + retun "Hello World" + """ + # Create client fixture for FastAPI @pytest.fixture(scope="module") diff --git a/tests/test_custom_types.py b/tests/test_custom_types.py new file mode 100644 index 000000000..31978ea2d --- /dev/null +++ b/tests/test_custom_types.py @@ -0,0 +1,19 @@ +# Test this: +from typing import Optional, Callable +from langflow.interface.custom_types import PythonFunction +from langflow.utils import util, constants +from pydantic import BaseModel, ValidationError, validator +from langchain.agents import tool +import pytest + + +def test_python_function(): + """Test Python function""" + func = PythonFunction(code=constants.DEFAULT_PYTHON_FUNCTION) + assert func.get_function()("text") == "text" + # the tool decorator should raise an error if + # the function is not str -> str + + # This raises ValidationError + with pytest.raises(SyntaxError): + func = PythonFunction(code=pytest.CODE_WITH_SYNTAX_ERROR) diff --git a/tests/test_endpoints.py b/tests/test_endpoints.py index 9bd8cb785..7276d1727 100644 --- a/tests/test_endpoints.py +++ b/tests/test_endpoints.py @@ -1,5 +1,7 @@ from langflow.interface.listing import CUSTOM_TOOLS +from langflow.utils.constants import DEFAULT_PYTHON_FUNCTION from fastapi.testclient import TestClient +import pytest def test_get_all(client: TestClient): @@ -10,3 +12,67 @@ def test_get_all(client: TestClient): assert "ZeroShotPrompt" 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()) + + +def test_post_validate_code(client: TestClient): + # Test case with a valid import and function + code1 = """ +import math + +def square(x): + return x ** 2 +""" + response1 = client.post("/validate", json={"code": code1}) + assert response1.status_code == 200 + assert response1.json() == {"imports": {"errors": []}, "function": {"errors": []}} + + # Test case with an invalid import and valid function + code2 = """ +import non_existent_module + +def square(x): + return x ** 2 +""" + response2 = client.post("/validate", json={"code": code2}) + assert response2.status_code == 200 + assert response2.json() == { + "imports": {"errors": ["No module named 'non_existent_module'"]}, + "function": {"errors": []}, + } + + # Test case with a valid import and invalid function syntax + code3 = """ +import math + +def square(x) + return x ** 2 +""" + response3 = client.post("/validate", json={"code": code3}) + assert response3.status_code == 200 + assert response3.json() == { + "imports": {"errors": []}, + "function": {"errors": ["expected ':' (, line 4)"]}, + } + + # Test case with invalid JSON payload + response4 = client.post("/validate", json={"invalid_key": code1}) + assert response4.status_code == 422 + + # Test case with an empty code string + response5 = client.post("/validate", json={"code": ""}) + assert response5.status_code == 200 + assert response5.json() == {"imports": {"errors": []}, "function": {"errors": []}} + + # Test case with a syntax error in the code + code6 = """ +import math + +def square(x) + return x ** 2 +""" + response6 = client.post("/validate", json={"code": code6}) + assert response6.status_code == 200 + assert response6.json() == { + "imports": {"errors": []}, + "function": {"errors": ["expected ':' (, line 4)"]}, + } diff --git a/tests/test_validate_code.py b/tests/test_validate_code.py new file mode 100644 index 000000000..5fdfa1199 --- /dev/null +++ b/tests/test_validate_code.py @@ -0,0 +1,39 @@ +from langflow.utils.validate import validate_code + + +def test_validate_code(): + # Test case with a valid import and function + code1 = """ +import math + +def square(x): + return x ** 2 +""" + errors1 = validate_code(code1) + assert errors1 == {"imports": {"errors": []}, "function": {"errors": []}} + + # Test case with an invalid import and valid function + code2 = """ +import non_existent_module + +def square(x): + return x ** 2 +""" + errors2 = validate_code(code2) + assert errors2 == { + "imports": {"errors": ["No module named 'non_existent_module'"]}, + "function": {"errors": []}, + } + + # Test case with a valid import and invalid function syntax + code3 = """ +import math + +def square(x) + return x ** 2 +""" + errors3 = validate_code(code3) + assert errors3 == { + "imports": {"errors": []}, + "function": {"errors": ["expected ':' (, line 4)"]}, + }