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
|
||||
Loading…
Add table
Add a link
Reference in a new issue