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