Underground-Digital/Workflow-Engine
0
1from typing import Optional2 3from core.app.app_config.features.file_upload.manager import FileUploadConfigManager4from core.file import file_manager5from core.file.models import FileType6from core.model_manager import ModelInstance7from core.model_runtime.entities import (8 AssistantPromptMessage,9 ImagePromptMessageContent,10 PromptMessage,11 PromptMessageContent,12 PromptMessageRole,13 TextPromptMessageContent,14 UserPromptMessage,15)16from core.prompt.utils.extract_thread_messages import extract_thread_messages17from extensions.ext_database import db18from factories import file_factory19from models.model import AppMode, Conversation, Message, MessageFile20from models.workflow import WorkflowRun21 22 23class TokenBufferMemory:24 def __init__(self, conversation: Conversation, model_instance: ModelInstance) -> None:25 self.conversation = conversation26 self.model_instance = model_instance27 28 def get_history_prompt_messages(29 self, max_token_limit: int = 2000, message_limit: Optional[int] = None30 ) -> list[PromptMessage]:31 """32 Get history prompt messages.33 :param max_token_limit: max token limit34 :param message_limit: message limit35 """36 app_record = self.conversation.app37 38 # fetch limited messages, and return reversed39 query = (40 db.session.query(41 Message.id,42 Message.query,43 Message.answer,44 Message.created_at,45 Message.workflow_run_id,46 Message.parent_message_id,47 )48 .filter(49 Message.conversation_id == self.conversation.id,50 )51 .order_by(Message.created_at.desc())52 )53 54 if message_limit and message_limit > 0:55 message_limit = min(message_limit, 500)56 else:57 message_limit = 50058 59 messages = query.limit(message_limit).all()60 61 # instead of all messages from the conversation, we only need to extract messages62 # that belong to the thread of last message63 thread_messages = extract_thread_messages(messages)64 65 # for newly created message, its answer is temporarily empty, we don't need to add it to memory66 if thread_messages and not thread_messages[0].answer:67 thread_messages.pop(0)68 69 messages = list(reversed(thread_messages))70 71 prompt_messages = []72 for message in messages:73 files = db.session.query(MessageFile).filter(MessageFile.message_id == message.id).all()74 if files:75 file_extra_config = None76 if self.conversation.mode not in {AppMode.ADVANCED_CHAT, AppMode.WORKFLOW}:77 file_extra_config = FileUploadConfigManager.convert(self.conversation.model_config)78 else:79 if message.workflow_run_id:80 workflow_run = (81 db.session.query(WorkflowRun).filter(WorkflowRun.id == message.workflow_run_id).first()82 )83 84 if workflow_run:85 file_extra_config = FileUploadConfigManager.convert(86 workflow_run.workflow.features_dict, is_vision=False87 )88 89 if file_extra_config and app_record:90 file_objs = file_factory.build_from_message_files(91 message_files=files, tenant_id=app_record.tenant_id, config=file_extra_config92 )93 else:94 file_objs = []95 96 if not file_objs:97 prompt_messages.append(UserPromptMessage(content=message.query))98 else:99 prompt_message_contents: list[PromptMessageContent] = []100 prompt_message_contents.append(TextPromptMessageContent(data=message.query))101 for file_obj in file_objs:102 if file_obj.type in {FileType.IMAGE, FileType.AUDIO}:103 prompt_message = file_manager.to_prompt_message_content(file_obj)104 prompt_message_contents.append(prompt_message)105 106 prompt_messages.append(UserPromptMessage(content=prompt_message_contents))107 else:108 prompt_messages.append(UserPromptMessage(content=message.query))109 110 prompt_messages.append(AssistantPromptMessage(content=message.answer))111 112 if not prompt_messages:113 return []114 115 # prune the chat message if it exceeds the max token limit116 curr_message_tokens = self.model_instance.get_llm_num_tokens(prompt_messages)117 118 if curr_message_tokens > max_token_limit:119 pruned_memory = []120 while curr_message_tokens > max_token_limit and len(prompt_messages) > 1:121 pruned_memory.append(prompt_messages.pop(0))122 curr_message_tokens = self.model_instance.get_llm_num_tokens(prompt_messages)123 124 return prompt_messages125 126 def get_history_prompt_text(127 self,128 human_prefix: str = "Human",129 ai_prefix: str = "Assistant",130 max_token_limit: int = 2000,131 message_limit: Optional[int] = None,132 ) -> str:133 """134 Get history prompt text.135 :param human_prefix: human prefix136 :param ai_prefix: ai prefix137 :param max_token_limit: max token limit138 :param message_limit: message limit139 :return:140 """141 prompt_messages = self.get_history_prompt_messages(max_token_limit=max_token_limit, message_limit=message_limit)142 143 string_messages = []144 for m in prompt_messages:145 if m.role == PromptMessageRole.USER:146 role = human_prefix147 elif m.role == PromptMessageRole.ASSISTANT:148 role = ai_prefix149 else:150 continue151 152 if isinstance(m.content, list):153 inner_msg = ""154 for content in m.content:155 if isinstance(content, TextPromptMessageContent):156 inner_msg += f"{content.data}\n"157 elif isinstance(content, ImagePromptMessageContent):158 inner_msg += "[image]\n"159 160 string_messages.append(f"{role}: {inner_msg.strip()}")161 else:162 message = f"{role}: {m.content}"163 string_messages.append(message)164 165 return "\n".join(string_messages)166 