🐛 fix(validate.py): rename parameter 'prompt' to 'prompt_request' in post_validate_prompt function for clarity

✨ feat(validate.py): refactor post_validate_prompt function to improve code readability and maintainability
The parameter 'prompt' in the 'post_validate_prompt' function has been renamed to 'prompt_request' to improve clarity and avoid confusion with the 'prompt' variable used within the function. The function has also been refactored to improve code readability and maintainability by extracting logic into separate helper functions. The helper functions 'get_old_custom_fields', 'add_new_variables_to_template', 'remove_old_variables_from_template', and 'update_input_variables_field' have been added to handle specific tasks within the 'post_validate_prompt' function. This refactoring improves the overall structure and organization of the code.
This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-07-03 23:10:41 -03:00
commit d5c7fb9dc5

View file

@ -28,13 +28,41 @@ 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: ValidatePromptRequest): def post_validate_prompt(prompt_request: ValidatePromptRequest):
try: try:
input_variables = validate_prompt(prompt.template) input_variables = validate_prompt(prompt_request.template)
# Reinitialize custom_fields
old_custom_fields = prompt.frontend_node.custom_fields.copy() old_custom_fields = get_old_custom_fields(prompt_request)
prompt.frontend_node.custom_fields = []
# Add new variables to the template add_new_variables_to_template(input_variables, prompt_request)
remove_old_variables_from_template(
old_custom_fields, input_variables, prompt_request
)
update_input_variables_field(input_variables, prompt_request)
return PromptValidationResponse(
input_variables=input_variables,
frontend_node=prompt_request.frontend_node,
)
except Exception as e:
logger.exception(e)
raise HTTPException(status_code=500, detail=str(e)) from e
def get_old_custom_fields(prompt_request):
try:
old_custom_fields = prompt_request.frontend_node.custom_fields[
prompt_request.name
].copy()
except KeyError:
old_custom_fields = []
prompt_request.frontend_node.custom_fields[prompt_request.name] = []
return old_custom_fields
def add_new_variables_to_template(input_variables, prompt_request):
for variable in input_variables: for variable in input_variables:
try: try:
template_field = TemplateField( template_field = TemplateField(
@ -43,34 +71,33 @@ def post_validate_prompt(prompt: ValidatePromptRequest):
field_type="str", field_type="str",
show=True, show=True,
advanced=False, advanced=False,
input_types=["Document", "BaseOutputParser"], input_types=["BaseLoader", "BaseOutputParser"],
) )
prompt.frontend_node.template[variable] = template_field.to_dict() prompt_request.frontend_node.template[variable] = template_field.to_dict()
prompt.frontend_node.custom_fields.append(variable) prompt_request.frontend_node.custom_fields[prompt_request.name].append(
variable
)
except Exception as exc: except Exception as exc:
logger.exception(exc) logger.exception(exc)
raise HTTPException(status_code=500, detail=str(exc)) from exc raise HTTPException(status_code=500, detail=str(exc)) from exc
# Remove variables that are not in the template anymore
def remove_old_variables_from_template(
old_custom_fields, input_variables, prompt_request
):
for variable in old_custom_fields: for variable in old_custom_fields:
if variable not in input_variables: if variable not in input_variables:
try: try:
prompt.frontend_node.template.pop(variable, None) prompt_request.frontend_node.template.pop(variable, None)
except Exception as exc: except Exception as exc:
logger.exception(exc) logger.exception(exc)
raise HTTPException(status_code=500, detail=str(exc)) from exc raise HTTPException(status_code=500, detail=str(exc)) from exc
# Now we will set the field "input_variables" to the new list of variables
# if it exists
if "input_variables" in prompt.frontend_node.template:
prompt.frontend_node.template["input_variables"]["value"] = input_variables
return PromptValidationResponse( def update_input_variables_field(input_variables, prompt_request):
input_variables=input_variables, if "input_variables" in prompt_request.frontend_node.template:
frontend_node=prompt.frontend_node, prompt_request.frontend_node.template["input_variables"][
) "value"
except Exception as e: ] = input_variables
logger.exception(e)
raise HTTPException(status_code=500, detail=str(e)) from e