[bugfix] deep copy in messages

This commit is contained in:
jinli.yl 2024-06-28 16:32:25 +08:00
parent 5f1bcf50d8
commit c5386ceaaf
5 changed files with 35 additions and 14 deletions

View file

@ -27,7 +27,7 @@ class ReadMemory(BaseWorkflow, BaseOperation):
def run_operation(self):
max_count = 1 + max(self.his_msg_count, self.contextual_msg_count)
self.context[CHAT_MESSAGES] = [x.copy() for x in self.chat_messages[-max_count:]]
self.context[CHAT_MESSAGES] = [x.copy(deep=True) for x in self.chat_messages[-max_count:]]
self.run_workflow()
result = self.context.get(RESULT)
self.context.clear()

View file

@ -56,7 +56,7 @@ class WriteMemory(BaseWorkflow, BaseOperation):
return
max_count = not_memorized_size + self.his_msg_count
self.context[CHAT_MESSAGES] = [x.copy() for x in self.chat_messages[-max_count:]]
self.context[CHAT_MESSAGES] = [x.copy(deep=True) for x in self.chat_messages[-max_count:]]
self.run_workflow()
result = self.context.get(RESULT)
self.context.clear()

View file

@ -2,8 +2,11 @@ from abc import ABCMeta
from typing import List
from memory_scope.chat.global_context import G_CONTEXT
from memory_scope.constants.common_constants import CHAT_MESSAGES
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
from memory_scope.memory.worker.base_worker import BaseWorker
from memory_scope.models.base_model import BaseModel
from memory_scope.scheme.message import Message
from memory_scope.storage.base_monitor import BaseMonitor
from memory_scope.storage.base_vector_store import BaseVectorStore
@ -24,17 +27,15 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
self._vector_store: BaseVectorStore | None = None
self._monitor: BaseMonitor | None = None
self._user_id: str | None = None
@property
def messages(self) -> List[Message]:
return self.get_context(MESSAGES)
return self.get_context(CHAT_MESSAGES)
@messages.setter
def messages(self, value):
self.set_context(MESSAGES, value)
@property
def chat_name(self):
return self.get_context(CHAT_NAME)
self.set_context(CHAT_MESSAGES, value)
@property
def embedding_model(self) -> BaseModel:
@ -67,11 +68,15 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
return self._monitor
@property
def memory_id(self) -> str:
pass
def user_id(self) -> str:
if self._user_id is None:
message = [x for x in self.messages if x.role == MessageRoleEnum.USER.value][-1]
self._user_id = message.role_name
return self._user_id
def __getattr__(self, key: str):
return self.kwargs[key]
def get_prompt(self, x):
return x[GLOBAL_CONTEXT.global_configs["language"]]
@staticmethod
def get_prompt(prompt: dict) -> str:
return prompt[G_CONTEXT.global_configs["language"]]

View file

@ -1,13 +1,17 @@
import datetime
from typing import Dict, List
import
from pydantic import Field, BaseModel
from memory_scope.utils.tool_functions import md5_hash
class MemoryNode(BaseModel):
user_id: str = Field("", description="unique memory id for user")
memory_id: str = Field("", description="unique id for memory item")
user_id: str = Field("", description="unique memory id for user")
content: str = Field("", description="memory content")
score_similar: float = Field(0, description="es similar score")
@ -24,9 +28,14 @@ class MemoryNode(BaseModel):
vector: List[float] = Field([], description="content embedding result, return empty")
timestamp: int = Field(int(datetime.datetime.now().timestamp()), description="timestamp of the memory node")
@property
def node_keys(self):
return list(self.model_json_schema()["properties"].keys())
def __getitem__(self, key: str):
return self.model_dump().get(key)
def gen_memory_id(self):
self.memory_id = f"{self.user_id}_{self.timestamp}_{md5_hash(self.content)[:8]}"

View file

@ -1,3 +1,4 @@
import hashlib
import random
import re
import time
@ -113,3 +114,9 @@ def char_logo(words: str, seed: int = time.time_ns(), color=None):
colored_line += colored_char
colored_lines.append(colored_line)
return colored_lines
def md5_hash(input_string: str):
m = hashlib.md5()
m.update(input_string.encode('utf-8'))
return m.hexdigest()