@@ -18,8 +18,8 @@ from django.core.files.uploadedfile import UploadedFile from langchain import hub from langchain.agents import AgentExecutor, create_structured_chat_agent from langchain.chains import ConversationChain -from langchain.memory import ConversationTokenBufferMemory from langchain_community.tools.google_serper import GoogleSerperResults +from langchain_core.chat_history import InMemoryChatMessageHistory from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, SystemMessage from langchain_core.prompts.prompt import PromptTemplate from langchain_core.runnables import RunnableWithMessageHistory @@ -180,25 +180,25 @@ class Chatgpt(SimpleService): chat_history = self.get_chat_history() conversation = RunnableWithMessageHistory( runnable=self.llm, - get_session_history=lambda _: self.get_chat_history().chat_memory, + get_session_history=lambda _: self.get_chat_history(), ) llm_input = HumanMessage(content=input_content) if file and not image: input_tokens = self.count_text_tokens( - [*chat_history.buffer_as_messages, llm_input, *chunks] + [*chat_history.messages, llm_input, *chunks] ) elif image: input_tokens = self.count_text_tokens([llm_input]) else: input_tokens = self.count_text_tokens( - [*chat_history.buffer_as_messages, llm_input] + [*chat_history.messages, llm_input] ) self.assert_enough_balance( input_tokens, image_size, model=self.llm.model_name ) if image: response = self.llm.invoke([llm_input]) - chat_history.chat_memory.add_ai_message(response) + chat_history.add_ai_message(response) elif file: human_messages = [] chunk_responses = ['Содержание файла: '] @@ -248,7 +248,7 @@ class Chatgpt(SimpleService): content=agent_executor.invoke( { 'input': [llm_input], - 'chat_history': chat_history.buffer_as_messages + 'chat_history': chat_history.messages + [ SystemMessage( content='Учитывай язык диалога перед выдачей ответа' @@ -295,7 +295,7 @@ class Chatgpt(SimpleService): {'input': llm_input.content[0]['text']}, config={'configurable': {'session_id': 'default'}}, ) - chat_history.chat_memory.add_ai_message(response) + chat_history.add_ai_message(response) output_tokens = self.count_text_tokens([response]) if file and not image: @@ -331,10 +331,10 @@ class Chatgpt(SimpleService): raise GenerationException def get_chat_history( - self, - message_limit: int = 10, - token_limit: int = 580, # ~ 1 AIR Token with GPT 3.5 - ) -> ConversationTokenBufferMemory: + self, + message_limit: int = 10, + token_limit: int = 580, # ~ 1 AIR Token with GPT 3.5 + ) -> InMemoryChatMessageHistory: if isinstance(self.store, Chat): air_messages = list( reversed( @@ -355,22 +355,21 @@ class Chatgpt(SimpleService): ).order_by('-created_at')[:message_limit] ) ) - memory = ConversationTokenBufferMemory(llm=self.llm, max_token_limit=token_limit) + memory = InMemoryChatMessageHistory() for msg in air_messages: content = msg.content or '' if msg.from_model: - memory.chat_memory.add_ai_message(content) + memory.add_message(AIMessage(content=content)) else: - memory.chat_memory.add_user_message(HumanMessage(content=content)) + memory.add_message(HumanMessage(content=content)) + + messages = memory.messages + tokens = self.llm.get_num_tokens_from_messages(messages) - buffer = memory.chat_memory.messages - curr_buffer_length = memory.llm.get_num_tokens_from_messages(buffer) + while tokens > token_limit: + messages.pop(0) + tokens = self.llm.get_num_tokens_from_messages(messages) - if curr_buffer_length > memory.max_token_limit: - pruned_memory = [] - while curr_buffer_length > memory.max_token_limit: - pruned_memory.append(buffer.pop(0)) - curr_buffer_length = memory.llm.get_num_tokens_from_messages(buffer) return memory def assert_enough_balance( @@ -547,7 +546,7 @@ class Chatgpt(SimpleService): 'AI:', ), ) - self.assert_enough_balance(chat_history.buffer_as_messages) + self.assert_enough_balance(chat_history.messages) for chunk in conversation.stream(input=input_message.content): if chunk: @@ -556,11 +555,11 @@ class Chatgpt(SimpleService): process_time = timedelta(seconds=time.time() - start_time) self.handle_invoice( self.neuron_model, - self.llm.get_num_tokens_from_messages(chat_history.chat_memory.messages), + self.llm.get_num_tokens_from_messages(chat_history.messages), self.llm.model_name, ) msgs = self.save_results( - [chat_history.chat_memory.messages[-1]], process_time, save + [chat_history.messages[-1]], process_time, save ) return msgs @@ -594,4 +593,4 @@ class Chatgpt(SimpleService): # ... for chunk in itertools.batched(TEMPORARY_TEST_TEXT, 10): - yield ''.join(chunk) + yield ''.join(chunk) \ No newline at end of file