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 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.custom_types import PythonFunction
|
||||||
|
|
||||||
from langflow.interface.run import process_graph
|
from langflow.interface.run import process_graph
|
||||||
from langflow.interface.types import build_langchain_types_dict
|
from langflow.interface.types import build_langchain_types_dict
|
||||||
|
from langflow.utils.validate import validate_code
|
||||||
|
|
||||||
|
|
||||||
# build router
|
# build router
|
||||||
|
|
@ -24,10 +26,13 @@ def get_load(data: Dict[str, Any]):
|
||||||
return HTTPException(status_code=500, detail=str(e))
|
return HTTPException(status_code=500, detail=str(e))
|
||||||
|
|
||||||
|
|
||||||
@router.post("/validate", status_code=200)
|
@router.post("/validate", status_code=200, response_model=ValidationResponse)
|
||||||
def validate_code(data: PythonFunction):
|
def post_validate_code(code: Code):
|
||||||
try:
|
try:
|
||||||
# if the data var gets here then it is valid python code
|
errors = validate_code(code.code)
|
||||||
return {"valid": True}
|
return ValidationResponse(
|
||||||
|
imports=errors.get("imports", {}),
|
||||||
|
function=errors.get("function", {}),
|
||||||
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
return HTTPException(status_code=500, detail=str(e))
|
return HTTPException(status_code=500, detail=str(e))
|
||||||
|
|
|
||||||
|
|
@ -1,11 +1,12 @@
|
||||||
from typing import Callable
|
from typing import Callable, Optional
|
||||||
from langflow.utils import util
|
from langflow.utils import validate
|
||||||
from pydantic import BaseModel, validator
|
from pydantic import BaseModel, validator
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class Function(BaseModel):
|
class Function(BaseModel):
|
||||||
code: str
|
code: str
|
||||||
|
function: Optional[Callable] = None
|
||||||
|
imports: Optional[str] = None
|
||||||
|
|
||||||
# Eval code and store the function
|
# Eval code and store the function
|
||||||
def __init__(self, **data):
|
def __init__(self, **data):
|
||||||
|
|
@ -14,16 +15,16 @@ class Function(BaseModel):
|
||||||
# Validate the function
|
# Validate the function
|
||||||
@validator("code")
|
@validator("code")
|
||||||
def validate_func(cls, v):
|
def validate_func(cls, v):
|
||||||
# Validate with LangChain's tool decorator
|
try:
|
||||||
func = util.eval_function(v)
|
validate.eval_function(v)
|
||||||
if not isinstance(func, Callable):
|
except Exception as e:
|
||||||
raise ValueError("Function must be a callable")
|
raise e
|
||||||
|
|
||||||
return v
|
return v
|
||||||
|
|
||||||
def get_function(self):
|
def get_function(self):
|
||||||
"""Get the function"""
|
"""Get the function"""
|
||||||
return util.eval_function(self.code)
|
return validate.eval_function(self.code)
|
||||||
|
|
||||||
|
|
||||||
class PythonFunction(Function):
|
class PythonFunction(Function):
|
||||||
|
|
|
||||||
|
|
@ -22,7 +22,7 @@ from langchain.llms.base import BaseLLM
|
||||||
from langchain.llms.loading import load_llm_from_config
|
from langchain.llms.loading import load_llm_from_config
|
||||||
|
|
||||||
from langflow.interface.types import get_type_list
|
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:
|
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
|
# as the instance
|
||||||
function_string = params["code"]
|
function_string = params["code"]
|
||||||
if isinstance(function_string, str):
|
if isinstance(function_string, str):
|
||||||
return util.eval_function(function_string)
|
return validate.eval_function(function_string)
|
||||||
raise ValueError("Function should be a string")
|
raise ValueError("Function should be a string")
|
||||||
else:
|
else:
|
||||||
if "tools" not in params:
|
if "tools" not in params:
|
||||||
|
|
|
||||||
|
|
@ -10,5 +10,6 @@ CHAT_OPENAI_MODELS = ["gpt-3.5-turbo", "gpt-4", "gpt-4-32k"]
|
||||||
|
|
||||||
DEFAULT_PYTHON_FUNCTION = """
|
DEFAULT_PYTHON_FUNCTION = """
|
||||||
def python_function(text: str) -> str:
|
def python_function(text: str) -> str:
|
||||||
|
\"\"\"This is a default python function that returns the input text\"\"\"
|
||||||
return 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):
|
def get_base_classes(cls):
|
||||||
"""Get the base classes of a class.
|
"""Get the base classes of a class.
|
||||||
These are used to determine the output of the nodes.
|
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"
|
Path(__file__).parent.absolute() / "data" / "complex_example.json"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
pytest.CODE_WITH_SYNTAX_ERROR = """
|
||||||
|
def get_text():
|
||||||
|
retun "Hello World"
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
# Create client fixture for FastAPI
|
# Create client fixture for FastAPI
|
||||||
@pytest.fixture(scope="module")
|
@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.interface.listing import CUSTOM_TOOLS
|
||||||
|
from langflow.utils.constants import DEFAULT_PYTHON_FUNCTION
|
||||||
from fastapi.testclient import TestClient
|
from fastapi.testclient import TestClient
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
def test_get_all(client: TestClient):
|
def test_get_all(client: TestClient):
|
||||||
|
|
@ -10,3 +12,67 @@ def test_get_all(client: TestClient):
|
||||||
assert "ZeroShotPrompt" in json_response["prompts"]
|
assert "ZeroShotPrompt" in json_response["prompts"]
|
||||||
# All CUSTOM_TOOLS(dict) should be in the response
|
# All CUSTOM_TOOLS(dict) should be in the response
|
||||||
assert all(tool in json_response["tools"] for tool in CUSTOM_TOOLS.keys())
|
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