refactor: Update prompt validation process in validate.py

Refactor the prompt validation process in validate.py to improve code readability and maintainability. The changes include:

- Update the import statement for process_prompt_template in api_utils.py.
- Remove unused imports from api_utils.py.
- Simplify the post_validate_prompt function by removing unnecessary conditional statements.
- Replace the validate_prompt function call with process_prompt_template in post_validate_prompt.

These changes enhance the efficiency and organization of the prompt validation process.
This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-06-22 16:53:37 -03:00
commit 8f051d0c21
2 changed files with 55 additions and 44 deletions

View file

@ -1,16 +1,8 @@
from collections import defaultdict
from fastapi import APIRouter, HTTPException from fastapi import APIRouter, HTTPException
from loguru import logger from loguru import logger
from langflow.api.v1.base import Code, CodeValidationResponse, PromptValidationResponse, ValidatePromptRequest from langflow.api.v1.base import Code, CodeValidationResponse, PromptValidationResponse, ValidatePromptRequest
from langflow.base.prompts.api_utils import ( from langflow.base.prompts.api_utils import process_prompt_template
add_new_variables_to_template,
get_old_custom_fields,
remove_old_variables_from_template,
update_input_variables_field,
validate_prompt,
)
from langflow.utils.validate import validate_code from langflow.utils.validate import validate_code
# build router # build router
@ -32,46 +24,20 @@ def post_validate_code(code: Code):
@router.post("/prompt", status_code=200, response_model=PromptValidationResponse) @router.post("/prompt", status_code=200, response_model=PromptValidationResponse)
def post_validate_prompt(prompt_request: ValidatePromptRequest): def post_validate_prompt(prompt_request: ValidatePromptRequest):
try: try:
input_variables = validate_prompt(prompt_request.template) if not prompt_request.frontend_node:
# Check if frontend_node is None before proceeding to avoid attempting to update a non-existent node.
if prompt_request.frontend_node is None:
return PromptValidationResponse( return PromptValidationResponse(
input_variables=input_variables, input_variables=[],
frontend_node=None, frontend_node=None,
) )
if not prompt_request.frontend_node.custom_fields:
prompt_request.frontend_node.custom_fields = defaultdict(list)
old_custom_fields = get_old_custom_fields(prompt_request.frontend_node.custom_fields, prompt_request.name)
add_new_variables_to_template( # Process the prompt template using direct attributes
input_variables, input_variables = process_prompt_template(
prompt_request.frontend_node.custom_fields, template=prompt_request.template,
prompt_request.frontend_node.template, name=prompt_request.name,
prompt_request.name, custom_fields=prompt_request.frontend_node.custom_fields,
frontend_node_template=prompt_request.frontend_node.template,
) )
remove_old_variables_from_template(
old_custom_fields,
input_variables,
prompt_request.frontend_node.custom_fields,
prompt_request.frontend_node.template,
prompt_request.name,
)
update_input_variables_field(input_variables, prompt_request.frontend_node.template)
# If frontend_node.template contains only one field that is type == 'prompt', then we can remove all fields that are not
# 'code', and not in the input_variables list.
prompt_fields = [
key
for key, field in prompt_request.frontend_node.template.items()
if isinstance(field, dict) and field["type"] == "prompt"
]
if len(prompt_fields) == 1:
for key, field in prompt_request.frontend_node.template.copy().items():
if isinstance(field, dict) and field["type"] != "code" and key not in input_variables + prompt_fields:
del prompt_request.frontend_node.template[key]
return PromptValidationResponse( return PromptValidationResponse(
input_variables=input_variables, input_variables=input_variables,
frontend_node=prompt_request.frontend_node, frontend_node=prompt_request.frontend_node,

View file

@ -1,10 +1,13 @@
from collections import defaultdict
from typing import Any, Dict, List, Optional
from fastapi import HTTPException from fastapi import HTTPException
from langchain_core.prompts import PromptTemplate
from loguru import logger from loguru import logger
from langflow.api.v1.base import INVALID_NAMES, check_input_variables from langflow.api.v1.base import INVALID_NAMES, check_input_variables
from langflow.interface.utils import extract_input_variables_from_prompt from langflow.interface.utils import extract_input_variables_from_prompt
from langflow.template.field.prompt import DefaultPromptField from langflow.template.field.prompt import DefaultPromptField
from langchain_core.prompts import PromptTemplate
def validate_prompt(prompt_template: str, silent_errors: bool = False) -> list[str]: def validate_prompt(prompt_template: str, silent_errors: bool = False) -> list[str]:
@ -81,3 +84,45 @@ def remove_old_variables_from_template(old_custom_fields, input_variables, custo
def update_input_variables_field(input_variables, template): def update_input_variables_field(input_variables, template):
if "input_variables" in template: if "input_variables" in template:
template["input_variables"]["value"] = input_variables template["input_variables"]["value"] = input_variables
def process_prompt_template(
template: str, name: str, custom_fields: Optional[Dict[str, List[str]]], frontend_node_template: Dict[str, Any]
):
"""Process and validate prompt template, update template and custom fields."""
# Validate the prompt template and extract input variables
input_variables = validate_prompt(template)
# Initialize custom_fields if None
if custom_fields is None:
custom_fields = defaultdict(list)
# Retrieve old custom fields
old_custom_fields = get_old_custom_fields(custom_fields, name)
# Add new variables to the template
add_new_variables_to_template(input_variables, custom_fields, frontend_node_template, name)
# Remove old variables from the template
remove_old_variables_from_template(old_custom_fields, input_variables, custom_fields, frontend_node_template, name)
# Update the input variables field in the template
update_input_variables_field(input_variables, frontend_node_template)
# Optional: cleanup fields based on specific conditions
cleanup_prompt_template_fields(input_variables, frontend_node_template)
return input_variables
def cleanup_prompt_template_fields(input_variables, template):
"""Removes unused fields if the conditions are met in the template."""
prompt_fields = [
key for key, field in template.items() if isinstance(field, dict) and field.get("type") == "prompt"
]
if len(prompt_fields) == 1:
for key in list(template.keys()): # Use list to copy keys
field = template.get(key, {})
if isinstance(field, dict) and field.get("type") != "code" and key not in input_variables + prompt_fields:
del template[key]