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 loguru import logger
from langflow.api.v1.base import Code, CodeValidationResponse, PromptValidationResponse, ValidatePromptRequest
from langflow.base.prompts.api_utils import (
add_new_variables_to_template,
get_old_custom_fields,
remove_old_variables_from_template,
update_input_variables_field,
validate_prompt,
)
from langflow.base.prompts.api_utils import process_prompt_template
from langflow.utils.validate import validate_code
# build router
@ -32,46 +24,20 @@ def post_validate_code(code: Code):
@router.post("/prompt", status_code=200, response_model=PromptValidationResponse)
def post_validate_prompt(prompt_request: ValidatePromptRequest):
try:
input_variables = validate_prompt(prompt_request.template)
# Check if frontend_node is None before proceeding to avoid attempting to update a non-existent node.
if prompt_request.frontend_node is None:
if not prompt_request.frontend_node:
return PromptValidationResponse(
input_variables=input_variables,
input_variables=[],
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(
input_variables,
prompt_request.frontend_node.custom_fields,
prompt_request.frontend_node.template,
prompt_request.name,
# Process the prompt template using direct attributes
input_variables = process_prompt_template(
template=prompt_request.template,
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(
input_variables=input_variables,
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 langchain_core.prompts import PromptTemplate
from loguru import logger
from langflow.api.v1.base import INVALID_NAMES, check_input_variables
from langflow.interface.utils import extract_input_variables_from_prompt
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]:
@ -81,3 +84,45 @@ def remove_old_variables_from_template(old_custom_fields, input_variables, custo
def update_input_variables_field(input_variables, template):
if "input_variables" in template:
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]