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:
parent
55c6168bc7
commit
8f051d0c21
2 changed files with 55 additions and 44 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue