feat: code validation, endpoint and tests

This commit is contained in:
Gabriel Almeida 2023-03-28 19:17:26 -03:00
commit b8a41037ec
11 changed files with 238 additions and 30 deletions

View file

@ -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": []}

View file

@ -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))

View file

@ -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):

View file

@ -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:

View file

@ -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
"""

View file

@ -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.

View file

@ -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=[]), "<string>", "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

View file

@ -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")

View file

@ -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)

View file

@ -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 ':' (<unknown>, 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 ':' (<unknown>, line 4)"]},
}

View file

@ -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 ':' (<unknown>, line 4)"]},
}