feat: code validation, endpoint and tests
This commit is contained in:
parent
6e27468daf
commit
b8a41037ec
11 changed files with 238 additions and 30 deletions
24
src/backend/langflow/api/base.py
Normal file
24
src/backend/langflow/api/base.py
Normal 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": []}
|
||||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
63
src/backend/langflow/utils/validate.py
Normal file
63
src/backend/langflow/utils/validate.py
Normal 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
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
19
tests/test_custom_types.py
Normal file
19
tests/test_custom_types.py
Normal 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)
|
||||
|
|
@ -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)"]},
|
||||
}
|
||||
|
|
|
|||
39
tests/test_validate_code.py
Normal file
39
tests/test_validate_code.py
Normal 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)"]},
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue