Underground-Digital/Workflow-Engine
0
1import enum2import json3import os4from typing import TYPE_CHECKING, Optional5 6from core.app.app_config.entities import PromptTemplateEntity7from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity8from core.file import file_manager9from core.memory.token_buffer_memory import TokenBufferMemory10from core.model_runtime.entities.message_entities import (11 PromptMessage,12 PromptMessageContent,13 SystemPromptMessage,14 TextPromptMessageContent,15 UserPromptMessage,16)17from core.prompt.entities.advanced_prompt_entities import MemoryConfig18from core.prompt.prompt_transform import PromptTransform19from core.prompt.utils.prompt_template_parser import PromptTemplateParser20from models.model import AppMode21 22if TYPE_CHECKING:23 from core.file.models import File24 25 26class ModelMode(str, enum.Enum):27 COMPLETION = "completion"28 CHAT = "chat"29 30 @classmethod31 def value_of(cls, value: str) -> "ModelMode":32 """33 Get value of given mode.34 35 :param value: mode value36 :return: mode37 """38 for mode in cls:39 if mode.value == value:40 return mode41 raise ValueError(f"invalid mode value {value}")42 43 44prompt_file_contents = {}45 46 47class SimplePromptTransform(PromptTransform):48 """49 Simple Prompt Transform for Chatbot App Basic Mode.50 """51 52 def get_prompt(53 self,54 app_mode: AppMode,55 prompt_template_entity: PromptTemplateEntity,56 inputs: dict,57 query: str,58 files: list["File"],59 context: Optional[str],60 memory: Optional[TokenBufferMemory],61 model_config: ModelConfigWithCredentialsEntity,62 ) -> tuple[list[PromptMessage], Optional[list[str]]]:63 inputs = {key: str(value) for key, value in inputs.items()}64 65 model_mode = ModelMode.value_of(model_config.mode)66 if model_mode == ModelMode.CHAT:67 prompt_messages, stops = self._get_chat_model_prompt_messages(68 app_mode=app_mode,69 pre_prompt=prompt_template_entity.simple_prompt_template,70 inputs=inputs,71 query=query,72 files=files,73 context=context,74 memory=memory,75 model_config=model_config,76 )77 else:78 prompt_messages, stops = self._get_completion_model_prompt_messages(79 app_mode=app_mode,80 pre_prompt=prompt_template_entity.simple_prompt_template,81 inputs=inputs,82 query=query,83 files=files,84 context=context,85 memory=memory,86 model_config=model_config,87 )88 89 return prompt_messages, stops90 91 def get_prompt_str_and_rules(92 self,93 app_mode: AppMode,94 model_config: ModelConfigWithCredentialsEntity,95 pre_prompt: str,96 inputs: dict,97 query: Optional[str] = None,98 context: Optional[str] = None,99 histories: Optional[str] = None,100 ) -> tuple[str, dict]:101 # get prompt template102 prompt_template_config = self.get_prompt_template(103 app_mode=app_mode,104 provider=model_config.provider,105 model=model_config.model,106 pre_prompt=pre_prompt,107 has_context=context is not None,108 query_in_prompt=query is not None,109 with_memory_prompt=histories is not None,110 )111 112 variables = {k: inputs[k] for k in prompt_template_config["custom_variable_keys"] if k in inputs}113 114 for v in prompt_template_config["special_variable_keys"]:115 # support #context#, #query# and #histories#116 if v == "#context#":117 variables["#context#"] = context or ""118 elif v == "#query#":119 variables["#query#"] = query or ""120 elif v == "#histories#":121 variables["#histories#"] = histories or ""122 123 prompt_template = prompt_template_config["prompt_template"]124 prompt = prompt_template.format(variables)125 126 return prompt, prompt_template_config["prompt_rules"]127 128 def get_prompt_template(129 self,130 app_mode: AppMode,131 provider: str,132 model: str,133 pre_prompt: str,134 has_context: bool,135 query_in_prompt: bool,136 with_memory_prompt: bool = False,137 ) -> dict:138 prompt_rules = self._get_prompt_rule(app_mode=app_mode, provider=provider, model=model)139 140 custom_variable_keys = []141 special_variable_keys = []142 143 prompt = ""144 for order in prompt_rules["system_prompt_orders"]:145 if order == "context_prompt" and has_context:146 prompt += prompt_rules["context_prompt"]147 special_variable_keys.append("#context#")148 elif order == "pre_prompt" and pre_prompt:149 prompt += pre_prompt + "\n"150 pre_prompt_template = PromptTemplateParser(template=pre_prompt)151 custom_variable_keys = pre_prompt_template.variable_keys152 elif order == "histories_prompt" and with_memory_prompt:153 prompt += prompt_rules["histories_prompt"]154 special_variable_keys.append("#histories#")155 156 if query_in_prompt:157 prompt += prompt_rules.get("query_prompt", "{{#query#}}")158 special_variable_keys.append("#query#")159 160 return {161 "prompt_template": PromptTemplateParser(template=prompt),162 "custom_variable_keys": custom_variable_keys,163 "special_variable_keys": special_variable_keys,164 "prompt_rules": prompt_rules,165 }166 167 def _get_chat_model_prompt_messages(168 self,169 app_mode: AppMode,170 pre_prompt: str,171 inputs: dict,172 query: str,173 context: Optional[str],174 files: list["File"],175 memory: Optional[TokenBufferMemory],176 model_config: ModelConfigWithCredentialsEntity,177 ) -> tuple[list[PromptMessage], Optional[list[str]]]:178 prompt_messages = []179 180 # get prompt181 prompt, _ = self.get_prompt_str_and_rules(182 app_mode=app_mode,183 model_config=model_config,184 pre_prompt=pre_prompt,185 inputs=inputs,186 query=None,187 context=context,188 )189 190 if prompt and query:191 prompt_messages.append(SystemPromptMessage(content=prompt))192 193 if memory:194 prompt_messages = self._append_chat_histories(195 memory=memory,196 memory_config=MemoryConfig(197 window=MemoryConfig.WindowConfig(198 enabled=False,199 )200 ),201 prompt_messages=prompt_messages,202 model_config=model_config,203 )204 205 if query:206 prompt_messages.append(self.get_last_user_message(query, files))207 else:208 prompt_messages.append(self.get_last_user_message(prompt, files))209 210 return prompt_messages, None211 212 def _get_completion_model_prompt_messages(213 self,214 app_mode: AppMode,215 pre_prompt: str,216 inputs: dict,217 query: str,218 context: Optional[str],219 files: list["File"],220 memory: Optional[TokenBufferMemory],221 model_config: ModelConfigWithCredentialsEntity,222 ) -> tuple[list[PromptMessage], Optional[list[str]]]:223 # get prompt224 prompt, prompt_rules = self.get_prompt_str_and_rules(225 app_mode=app_mode,226 model_config=model_config,227 pre_prompt=pre_prompt,228 inputs=inputs,229 query=query,230 context=context,231 )232 233 if memory:234 tmp_human_message = UserPromptMessage(content=prompt)235 236 rest_tokens = self._calculate_rest_token([tmp_human_message], model_config)237 histories = self._get_history_messages_from_memory(238 memory=memory,239 memory_config=MemoryConfig(240 window=MemoryConfig.WindowConfig(241 enabled=False,242 )243 ),244 max_token_limit=rest_tokens,245 human_prefix=prompt_rules.get("human_prefix", "Human"),246 ai_prefix=prompt_rules.get("assistant_prefix", "Assistant"),247 )248 249 # get prompt250 prompt, prompt_rules = self.get_prompt_str_and_rules(251 app_mode=app_mode,252 model_config=model_config,253 pre_prompt=pre_prompt,254 inputs=inputs,255 query=query,256 context=context,257 histories=histories,258 )259 260 stops = prompt_rules.get("stops")261 if stops is not None and len(stops) == 0:262 stops = None263 264 return [self.get_last_user_message(prompt, files)], stops265 266 def get_last_user_message(self, prompt: str, files: list["File"]) -> UserPromptMessage:267 if files:268 prompt_message_contents: list[PromptMessageContent] = []269 prompt_message_contents.append(TextPromptMessageContent(data=prompt))270 for file in files:271 prompt_message_contents.append(file_manager.to_prompt_message_content(file))272 273 prompt_message = UserPromptMessage(content=prompt_message_contents)274 else:275 prompt_message = UserPromptMessage(content=prompt)276 277 return prompt_message278 279 def _get_prompt_rule(self, app_mode: AppMode, provider: str, model: str) -> dict:280 """281 Get simple prompt rule.282 :param app_mode: app mode283 :param provider: model provider284 :param model: model name285 :return:286 """287 prompt_file_name = self._prompt_file_name(app_mode=app_mode, provider=provider, model=model)288 289 # Check if the prompt file is already loaded290 if prompt_file_name in prompt_file_contents:291 return prompt_file_contents[prompt_file_name]292 293 # Get the absolute path of the subdirectory294 prompt_path = os.path.join(os.path.dirname(os.path.realpath(__file__)), "prompt_templates")295 json_file_path = os.path.join(prompt_path, f"{prompt_file_name}.json")296 297 # Open the JSON file and read its content298 with open(json_file_path, encoding="utf-8") as json_file:299 content = json.load(json_file)300 301 # Store the content of the prompt file302 prompt_file_contents[prompt_file_name] = content303 304 return content305 306 def _prompt_file_name(self, app_mode: AppMode, provider: str, model: str) -> str:307 # baichuan308 is_baichuan = False309 if provider == "baichuan":310 is_baichuan = True311 else:312 baichuan_supported_providers = ["huggingface_hub", "openllm", "xinference"]313 if provider in baichuan_supported_providers and "baichuan" in model.lower():314 is_baichuan = True315 316 if is_baichuan:317 if app_mode == AppMode.COMPLETION:318 return "baichuan_completion"319 else:320 return "baichuan_chat"321 322 # common323 if app_mode == AppMode.COMPLETION:324 return "common_completion"325 else:326 return "common_chat"327 