From 880fa5034f1ef2f658888c709c0b6c42aab67531 Mon Sep 17 00:00:00 2001 From: Gabriel Luiz Freitas Almeida Date: Tue, 27 Jun 2023 18:23:12 -0300 Subject: [PATCH] =?UTF-8?q?=F0=9F=9A=80=20feat(api):=20add=20support=20for?= =?UTF-8?q?=20frontend=5Fnode=20in=20validate=5Fprompt=20endpoint=20?= =?UTF-8?q?=E2=9C=A8=20feat(base.py):=20add=20frontend=5Fnode=20parameter?= =?UTF-8?q?=20to=20ValidatePromptRequest=20and=20PromptValidationResponse?= =?UTF-8?q?=20models=20The=20validate=5Fprompt=20endpoint=20now=20accepts?= =?UTF-8?q?=20a=20frontend=5Fnode=20parameter=20in=20the=20ValidatePromptR?= =?UTF-8?q?equest=20model.=20This=20parameter=20is=20used=20to=20add=20inp?= =?UTF-8?q?ut=20variables=20to=20the=20frontend=5Fnode's=20template=20fiel?= =?UTF-8?q?ds=20and=20custom=20fields.=20The=20PromptValidationResponse=20?= =?UTF-8?q?model=20now=20includes=20the=20frontend=5Fnode=20parameter=20to?= =?UTF-8?q?=20return=20the=20updated=20frontend=5Fnode=20object.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/backend/langflow/api/v1/base.py | 7 +++++-- src/backend/langflow/api/v1/validate.py | 22 +++++++++++++++++++--- 2 files changed, 24 insertions(+), 5 deletions(-) diff --git a/src/backend/langflow/api/v1/base.py b/src/backend/langflow/api/v1/base.py index d595210bb..1a4936a2f 100644 --- a/src/backend/langflow/api/v1/base.py +++ b/src/backend/langflow/api/v1/base.py @@ -1,3 +1,4 @@ +from langflow.template.frontend_node.base import FrontendNode from pydantic import BaseModel, validator from langflow.interface.utils import extract_input_variables_from_prompt @@ -12,8 +13,9 @@ class Code(BaseModel): code: str -class Prompt(BaseModel): +class ValidatePromptRequest(BaseModel): template: str + frontend_node: FrontendNode # Build ValidationResponse class for {"imports": {"errors": []}, "function": {"errors": []}} @@ -32,6 +34,7 @@ class CodeValidationResponse(BaseModel): class PromptValidationResponse(BaseModel): input_variables: list + frontend_node: FrontendNode INVALID_CHARACTERS = { @@ -66,7 +69,7 @@ def validate_prompt(template: str): # if len(input_variables) > 1: # # If there's more than one input variable - return PromptValidationResponse(input_variables=input_variables) + return input_variables def check_input_variables(input_variables: list): diff --git a/src/backend/langflow/api/v1/validate.py b/src/backend/langflow/api/v1/validate.py index 959273a00..3046ad9ba 100644 --- a/src/backend/langflow/api/v1/validate.py +++ b/src/backend/langflow/api/v1/validate.py @@ -3,10 +3,11 @@ from fastapi import APIRouter, HTTPException from langflow.api.v1.base import ( Code, CodeValidationResponse, - Prompt, + ValidatePromptRequest, PromptValidationResponse, validate_prompt, ) +from langflow.template.field.base import TemplateField from langflow.utils.logger import logger from langflow.utils.validate import validate_code @@ -27,9 +28,24 @@ def post_validate_code(code: Code): @router.post("/prompt", status_code=200, response_model=PromptValidationResponse) -def post_validate_prompt(prompt: Prompt): +def post_validate_prompt(prompt: ValidatePromptRequest): try: - return validate_prompt(prompt.template) + input_variables = validate_prompt(prompt.template) + for variable in input_variables: + try: + template_field = TemplateField( + name=variable, field_type="str", show=True, advanced=False + ) + prompt.frontend_node.template.fields.append(template_field) + prompt.frontend_node.custom_fields.append(variable) + except Exception as exc: + logger.exception(exc) + raise HTTPException(status_code=500, detail=str(exc)) from exc + + return PromptValidationResponse( + input_variables=input_variables, + frontend_node=prompt.frontend_node, + ) except Exception as e: logger.exception(e) raise HTTPException(status_code=500, detail=str(e)) from e