dify/api/core/prompt/prompt_builder.py

39 lines
1.8 KiB
Python
Raw Normal View History

2023-05-15 08:51:32 +08:00
import re
from langchain.prompts import SystemMessagePromptTemplate, HumanMessagePromptTemplate, AIMessagePromptTemplate
from langchain.schema import BaseMessage
2023-06-27 15:30:38 +08:00
from core.prompt.prompt_template import JinjaPromptTemplate
2023-05-15 08:51:32 +08:00
class PromptBuilder:
@classmethod
def to_system_message(cls, prompt_content: str, inputs: dict) -> BaseMessage:
2023-06-27 15:30:38 +08:00
prompt_template = JinjaPromptTemplate.from_template(prompt_content)
2023-05-15 08:51:32 +08:00
system_prompt_template = SystemMessagePromptTemplate(prompt=prompt_template)
prompt_inputs = {k: inputs[k] for k in system_prompt_template.input_variables if k in inputs}
system_message = system_prompt_template.format(**prompt_inputs)
return system_message
@classmethod
def to_ai_message(cls, prompt_content: str, inputs: dict) -> BaseMessage:
2023-06-27 15:30:38 +08:00
prompt_template = JinjaPromptTemplate.from_template(prompt_content)
2023-05-15 08:51:32 +08:00
ai_prompt_template = AIMessagePromptTemplate(prompt=prompt_template)
prompt_inputs = {k: inputs[k] for k in ai_prompt_template.input_variables if k in inputs}
ai_message = ai_prompt_template.format(**prompt_inputs)
return ai_message
@classmethod
def to_human_message(cls, prompt_content: str, inputs: dict) -> BaseMessage:
2023-06-27 15:30:38 +08:00
prompt_template = JinjaPromptTemplate.from_template(prompt_content)
2023-05-15 08:51:32 +08:00
human_prompt_template = HumanMessagePromptTemplate(prompt=prompt_template)
human_message = human_prompt_template.format(**inputs)
return human_message
@classmethod
def process_template(cls, template: str):
2023-06-27 15:30:38 +08:00
processed_template = re.sub(r'\{{2}(.+)\}{2}', r'{\1}', template)
# processed_template = re.sub(r'\{([a-zA-Z_]\w+?)\}', r'\1', template)
# processed_template = re.sub(r'\{\{([a-zA-Z_]\w+?)\}\}', r'{\1}', processed_template)
2023-05-15 08:51:32 +08:00
return processed_template