Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
token_buffer_memory.py166 linesDownload Raw Back to memory
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