diff --git a/config/demo_config.yaml b/config/demo_config.yaml index d36272ba..d4a18a65 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -7,6 +7,10 @@ memory_chat: memory_service: memory_chat_service generation_model: dashscope_generation human_name: 用户 + human_profile_setting: + - name + - gender + - residential location assistant_name: AI memory_service: memory_chat_service: @@ -50,8 +54,8 @@ worker: generation_model_top_k: 1 retrieve_store_worker: class: memory.worker.read.retrieve_store_worker - retrieve_obs_top_k: 100 - retrieve_ins_pf_top_k: 100 + retrieve_obs_top_k: 5 + retrieve_ins_pf_top_k: 5 semantic_rank_worker: class: memory.worker.read.semantic_rank_worker fuse_rerank_worker: @@ -83,8 +87,12 @@ worker: class: memory.worker.write.contra_repeat_worker generation_model: dashscope_generation generation_model_top_k: 1 - today_obs_top_k: 30 + retrieve_top_k: 30 contra_repeat_max_count: 50 + get_reflection_worker: + class: memory.worker.summary.get_reflection_worker + retrieve_top_k: 100 + reflect_obs_cnt_threshold: 32 models: dashscope_generation: class: models.llama_index_generation_model diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 5dc11168..763bad9b 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -1,5 +1,6 @@ import os import time +from typing import List import questionary @@ -27,19 +28,26 @@ class CliMemoryChat(BaseMemoryChat): memory_service: str, generation_model: str, stream: bool = True, - human_name: str = "human", - assistant_name: str = "assistant", + human_name: str = "", + human_profile_setting: List[str] = None, + assistant_name: str = "", **kwargs): self._memory_service: BaseMemoryService | str = memory_service self._generation_model: BaseModel | str = generation_model self.stream: bool = stream self.human_name: str = human_name + self.human_profile_setting: List[str] = human_profile_setting self.assistant_name: str = assistant_name self.kwargs: dict = kwargs self._logo = char_logo("MemoryScope") self._prompt_handler: PromptHandler | None = None + G_CONTEXT.meta_data.update({ + "human_name": human_name, + "assistant_name": assistant_name, + "human_profile_setting": human_profile_setting, + }) self.logger = Logger.get_logger() @@ -66,7 +74,8 @@ class CliMemoryChat(BaseMemoryChat): self._generation_model = G_CONTEXT.model_dict[self._generation_model] return self._generation_model - def get_system_prompt(self) -> Message: + @property + def system_prompt_with_memory(self) -> Message: system_prompt = self.prompt_handler.system_prompt memories: str = self.memory_service.read_memory() @@ -77,13 +86,9 @@ class CliMemoryChat(BaseMemoryChat): return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt) def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen: - new_message: Message = Message(role=MessageRoleEnum.USER.value, - role_name=self.human_name, - content=query) - + new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=self.human_name, content=query) self.memory_service.add_messages(new_message) - system_message: Message = self.get_system_prompt() - return self.generation_model.call(messages=[system_message, new_message], stream=self.stream) + return self.generation_model.call(messages=[self.system_prompt_with_memory, new_message], stream=self.stream) @staticmethod def parse_query_command(query: str): diff --git a/memory_scope/constants/common_constants.py b/memory_scope/constants/common_constants.py index b14d0520..b0000899 100644 --- a/memory_scope/constants/common_constants.py +++ b/memory_scope/constants/common_constants.py @@ -20,6 +20,8 @@ MEMORY = "memory" DEFAULT_SYSTEM_PROMPT = "default_system_prompt" +NOT_REFLECTED_NODES = "not_reflected_nodes" + MODIFIED_MEMORIES = "modified_memories" diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py index bb8c6317..2f91b383 100644 --- a/memory_scope/memory/worker/base_worker.py +++ b/memory_scope/memory/worker/base_worker.py @@ -27,7 +27,7 @@ class BaseWorker(metaclass=ABCMeta): self.logger: Logger = Logger.get_logger() @staticmethod - def _async_run(fn_list, *args, **kwargs): + def async_run(fn_list, *args, **kwargs): async def async_gather(): return await asyncio.gather(*[fn(*args, **kwargs) for fn in fn_list]) diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py index 57b7e616..15656ed8 100644 --- a/memory_scope/memory/worker/memory_base_worker.py +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -2,7 +2,6 @@ from abc import ABCMeta from typing import List, Dict from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS -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 @@ -77,15 +76,13 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): @property def user_name(self) -> str: if self._user_name is None: - message = [x for x in self.messages if x.role == MessageRoleEnum.ASSISTANT.value][-1] - self._user_name = message.role_name + self._user_name = G_CONTEXT.meta_data["human_name"] return self._user_name @property def target_name(self) -> str: if self._target_name is None: - message = [x for x in self.messages if x.role == MessageRoleEnum.USER.value][-1] - self._target_name = message.role_name + self._target_name = G_CONTEXT.meta_data["assistant_name"] return self._target_name @property diff --git a/memory_scope/memory/worker/read/retrieve_store_worker.py b/memory_scope/memory/worker/read/retrieve_store_worker.py index 20e21dd5..d685d826 100644 --- a/memory_scope/memory/worker/read/retrieve_store_worker.py +++ b/memory_scope/memory/worker/read/retrieve_store_worker.py @@ -35,7 +35,7 @@ class RetrieveStoreWorker(MemoryBaseWorker): query, _ = self.get_context(QUERY_WITH_TS) memory_node_list: List[MemoryNode] = [] fn_list = [self.retrieve_from_observation, self.retrieve_from_insight_and_profile] - for result in self._async_run(fn_list=fn_list, query=query): + for result in self.async_run(fn_list=fn_list, query=query): if result: memory_node_list.extend(result) self.logger.info(f"memory_node_list.size={len(memory_node_list)}") diff --git a/memory_scope/memory/worker/read/set_query_worker.py b/memory_scope/memory/worker/read/set_query_worker.py index e0a6c84c..48addbc4 100644 --- a/memory_scope/memory/worker/read/set_query_worker.py +++ b/memory_scope/memory/worker/read/set_query_worker.py @@ -11,7 +11,7 @@ class SetQueryWorker(MemoryBaseWorker): query = self.chat_kwargs["query"] query_timestamp = int(datetime.datetime.now().timestamp()) else: - query = self.messages[-1].content - query_timestamp = self.messages[-1].time_created + query = self.chat_messages[-1].content + query_timestamp = self.chat_messages[-1].time_created self.set_context(QUERY_WITH_TS, (query, query_timestamp)) diff --git a/memory_scope/memory/worker/summary/get_reflection_worker.py b/memory_scope/memory/worker/summary/get_reflection_worker.py index ccc0c6ad..6e37a20a 100644 --- a/memory_scope/memory/worker/summary/get_reflection_worker.py +++ b/memory_scope/memory/worker/summary/get_reflection_worker.py @@ -1,47 +1,39 @@ from typing import List -from memory_scope.utils.response_text_parser import ResponseTextParser -from memory_scope.constants.common_constants import ( - NEW_OBS_NODES, - NOT_REFLECTED_OBS_NODES, - INSIGHT_NODES, - NEW_INSIGHT_KEYS, - NOT_REFLECTED_MERGE_NODES, -) -from memory_scope.constants.language_constants import COLON_WORD, COMMA_WORD -from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.constants.common_constants import NOT_REFLECTED_NODES +from memory_scope.constants.language_constants import COMMA_WORD +from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus +from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker +from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.utils.response_text_parser import ResponseTextParser +from memory_scope.utils.timer import timer from memory_scope.utils.tool_functions import prompt_to_msg class GetReflectionWorker(MemoryBaseWorker): + @timer + def retrieve_not_reflected_memory(self) -> List[MemoryNode]: + filter_dict = { + "user_name": self.user_name, + "target_name": self.target_name, + "status": MemoryNodeStatus.ACTIVE.value, + "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], + "obs_reflected": False, + } + return self.vector_store.retrieve(query=" ", top_k=self.retrieve_top_k, filter_dict=filter_dict) + def _run(self): - # 过滤得到 not_reflected_merge_nodes - new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES) - not_reflected_nodes: List[MemoryNode] = self.get_context( - NOT_REFLECTED_OBS_NODES - ) - not_reflected_merge_nodes: List[MemoryNode] = [] - if new_obs_nodes: - not_reflected_merge_nodes.extend(new_obs_nodes) - if not_reflected_nodes: - not_reflected_merge_nodes.extend(not_reflected_nodes) - not_reflected_merge_nodes = [ - node - for node in not_reflected_merge_nodes - if not node.obs_reflected - ] + not_reflected_nodes: List[MemoryNode] = self.retrieve_not_reflected_memory() # count - not_reflected_count = len(not_reflected_merge_nodes) + not_reflected_count = len(not_reflected_nodes) if not_reflected_count <= self.reflect_obs_cnt_threshold: - self.logger.info( - f"not_reflected_count={not_reflected_count} is not enough, stop reflect." - ) + self.logger.info(f"not_reflected_count={not_reflected_count} is not enough, stop.") return # save context - self.set_context(NOT_REFLECTED_MERGE_NODES, not_reflected_merge_nodes) + self.set_context(NOT_REFLECTED_NODES, not_reflected_nodes) # get profile_keys exist_keys: List[str] = [] diff --git a/memory_scope/memory/worker/summary/load_memory_worker.py b/memory_scope/memory/worker/summary/load_memory_worker.py new file mode 100644 index 00000000..23248b04 --- /dev/null +++ b/memory_scope/memory/worker/summary/load_memory_worker.py @@ -0,0 +1,68 @@ +from typing import List, Dict + +from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus +from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum +from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker +from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.utils.timer import timer +from memory_scope.utils.global_context import G_CONTEXT + +class LoadMemoryWorker(MemoryBaseWorker): + + @timer + async def retrieve_not_reflected_memory(self, query: str) -> List[MemoryNode]: + filter_dict = { + "user_name": self.user_name, + "target_name": self.target_name, + "status": MemoryNodeStatus.ACTIVE.value, + "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], + "obs_reflected": False, + } + return await self.vector_store.async_retrieve(query=query, + top_k=self.retrieve_not_reflected_top_k, + filter_dict=filter_dict) + + @timer + async def retrieve_not_updated_memory(self, query: str) -> List[MemoryNode]: + filter_dict = { + "user_name": self.user_name, + "target_name": self.target_name, + "status": MemoryNodeStatus.ACTIVE.value, + "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], + "obs_updated": False, + } + return await self.vector_store.async_retrieve(query=query, + top_k=self.retrieve_not_updated_top_k, + filter_dict=filter_dict) + + @timer + async def retrieve_profiles(self, query: str) -> List[MemoryNode]: + filter_dict = { + "user_name": self.user_name, + "target_name": self.target_name, + "status": MemoryNodeStatus.ACTIVE.value, + "memory_type": [MemoryTypeEnum.PROFILE.value, MemoryTypeEnum.PROFILE_CUSTOMIZED.value], + } + retrieve_nodes = await self.vector_store.async_retrieve(query=query, + top_k=self.retrieve_profiles_top_k, + filter_dict=filter_dict) + nodes: List[MemoryNode] = [] + human_profile_setting = G_CONTEXT.meta_data.get("human_profile_setting", []) + for attr_key in human_profile_setting: + + + return nodes + + def _run(self): + mock_query = "_" + fn_list = [ + self.retrieve_not_reflected_memory, + self.retrieve_not_updated_memory, + self.retrieve_profiles, + ] + memory_node_dict: Dict[str, MemoryNode] = {} + for nodes in self.async_run(fn_list, query=mock_query): + assert isinstance(nodes[0], MemoryNode) + memory_node_dict.update({n.memory_id: n for n in nodes}) + + memory_node_list = sorted(memory_node_dict.values(), key=lambda x: x.memory_id) diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.py b/memory_scope/memory/worker/write/contra_repeat_worker.py index abcf8945..c3c9f21d 100644 --- a/memory_scope/memory/worker/write/contra_repeat_worker.py +++ b/memory_scope/memory/worker/write/contra_repeat_worker.py @@ -27,7 +27,7 @@ class ContraRepeatWorker(MemoryBaseWorker): "target_name": self.target_name, "status": MemoryNodeStatus.ACTIVE.value, "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], - "obs_dt": dt_handler.datetime_format(), + "dt": dt_handler.datetime_format(), } return self.vector_store.retrieve(query=message.content, top_k=self.today_obs_top_k, filter_dict=filter_dict) @@ -55,11 +55,11 @@ class ContraRepeatWorker(MemoryBaseWorker): for i, n in enumerate(all_obs_nodes): user_query_list.append(f"{i + 1} {n.content}") - system_prompt = self.prompt_config.contra_repeat_system.format(num_obs=len(user_query_list), - user_name=self.user_id) - few_shot = self.prompt_config.contra_repeat_few_shot.format(user_name=self.user_id) - user_query = self.prompt_config.contra_repeat_user_query.format(user_query="\n".join(user_query_list), + system_prompt = self.prompt_handler.contra_repeat_system.format(num_obs=len(user_query_list), user_name=self.user_id) + few_shot = self.prompt_handler.contra_repeat_few_shot.format(user_name=self.user_id) + user_query = self.prompt_handler.contra_repeat_user_query.format(user_query="\n".join(user_query_list), + user_name=self.user_id) contra_repeat_message = self.prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) diff --git a/memory_scope/memory/worker/write/get_observation_with_time_worker.py b/memory_scope/memory/worker/write/get_observation_with_time_worker.py index 71c3692a..eb9defea 100644 --- a/memory_scope/memory/worker/write/get_observation_with_time_worker.py +++ b/memory_scope/memory/worker/write/get_observation_with_time_worker.py @@ -33,8 +33,8 @@ class GetObservationWithTimeWorker(GetObservationWorker): system_prompt = self.prompt_handler.get_observation_with_time_system.format(num_obs=len(user_query_list), user_name=self.target_name) - few_shot = self.prompt_config.get_observation_with_time_few_shot.format(user_name=self.target_name) - user_query = self.prompt_config.get_observation_with_time_user_query.format( + few_shot = self.prompt_handler.get_observation_with_time_few_shot.format(user_name=self.target_name) + user_query = self.prompt_handler.get_observation_with_time_user_query.format( user_query="\n".join(user_query_list), user_name=self.target_name) diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py index d42dcf11..50d30aaf 100644 --- a/memory_scope/memory/worker/write/get_observation_worker.py +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -20,6 +20,7 @@ class GetObservationWorker(MemoryBaseWorker): meta_data = { MemoryTypeEnum.CONVERSATION.value: message.content, TIME_INFER: time_infer, + "keywords": keywords, **dt_handler.dt_info_dict, } @@ -34,11 +35,8 @@ class GetObservationWorker(MemoryBaseWorker): memory_type=MemoryTypeEnum.OBSERVATION.value, status=MemoryNodeStatus.ACTIVE.value, timestamp=message.time_created, - obs_dt=dt_handler.datetime_format(), obs_reflected=False, - obs_updated=False, - obs_keyword=keywords) - node.gen_memory_id() + obs_updated=False) return node def build_prompt(self) -> List[Message]: @@ -61,9 +59,9 @@ class GetObservationWorker(MemoryBaseWorker): system_prompt = self.prompt_handler.get_observation_system.format(num_obs=len(user_query_list), user_name=self.target_name) - few_shot = self.prompt_config.get_observation_few_shot.format(user_name=self.target_name) - user_query = self.prompt_config.get_observation_user_query.format(user_query="\n".join(user_query_list), - user_name=self.target_name) + few_shot = self.prompt_handler.get_observation_few_shot.format(user_name=self.target_name) + user_query = self.prompt_handler.get_observation_user_query.format(user_query="\n".join(user_query_list), + user_name=self.target_name) obtain_obs_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) self.logger.info(f"obtain_obs_message={obtain_obs_message}") diff --git a/memory_scope/memory/worker/write/store_memory_worker.py b/memory_scope/memory/worker/write/store_memory_worker.py index 1499c030..92e4ecbf 100644 --- a/memory_scope/memory/worker/write/store_memory_worker.py +++ b/memory_scope/memory/worker/write/store_memory_worker.py @@ -23,14 +23,12 @@ class StoreMemoryWorker(MemoryBaseWorker): return dt_handler = DatetimeHandler() - node = MemoryNode( - user_name=self.user_name, - target_name=self.target_name, - content=query, - memory_type=MemoryTypeEnum.OBSERVATION.value, - status=MemoryNodeStatus.ACTIVE.value, - timestamp=dt_handler.timestamp, - obs_dt=dt_handler.datetime_format(), - obs_reflected=False, - obs_updated=False) + node = MemoryNode(user_name=self.user_name, + target_name=self.target_name, + content=query, + memory_type=MemoryTypeEnum.OBS_CUSTOMIZED.value, + status=MemoryNodeStatus.ACTIVE.value, + timestamp=dt_handler.timestamp, + obs_reflected=False, + obs_updated=False) self.vector_store.update(node) diff --git a/memory_scope/models/base_model.py b/memory_scope/models/base_model.py index 6971e027..1ea6315f 100644 --- a/memory_scope/models/base_model.py +++ b/memory_scope/models/base_model.py @@ -1,6 +1,5 @@ import inspect import time -import dashscope from abc import abstractmethod, ABCMeta from typing import Any @@ -23,6 +22,7 @@ class BaseModel(metaclass=ABCMeta): max_retries: int = 3, retry_interval: float = 1.0, kwargs_filter: bool = True, + raise_exception: bool = True, **kwargs): self.model_name: str = model_name @@ -31,6 +31,7 @@ class BaseModel(metaclass=ABCMeta): self.max_retries: int = max_retries self.retry_interval: float = retry_interval self.kwargs_filter: bool = kwargs_filter + self.raise_exception: bool = raise_exception self.kwargs: dict = kwargs self.data = {} @@ -85,12 +86,13 @@ class BaseModel(metaclass=ABCMeta): with Timer(self.__class__.__name__, log_time=False) as t: self.before_call(stream=stream, **kwargs) for i in range(self.max_retries): - try: + if self.raise_exception: model_response = self._call(stream=stream, **kwargs) - except dashscope.common.error.AuthenticationError as e: - raise e - except Exception as e: - model_response = ModelResponse(m_type=self.m_type, status=False, details=e.args) + else: + try: + model_response = self._call(stream=stream, **kwargs) + except Exception as e: + model_response = ModelResponse(m_type=self.m_type, status=False, details=e.args) if isinstance(model_response, ModelResponse) and not model_response.status: self.logger.warning(f"call model={self.model_name} failed! cost={t.cost_str} retry_cnt={i} " @@ -114,10 +116,13 @@ class BaseModel(metaclass=ABCMeta): with Timer(self.__class__.__name__, log_time=False) as t: self.before_call(**kwargs) for i in range(self.max_retries): - try: - model_response = await self._async_call(**kwargs) - except Exception as e: - model_response = ModelResponse(status=False, details=e.args) + if self.raise_exception: + model_response = self._async_call(**kwargs) + else: + try: + model_response = self._async_call(**kwargs) + except Exception as e: + model_response = ModelResponse(m_type=self.m_type, status=False, details=e.args) if not model_response.status: self.logger.warning(f"async_call model={self.model_name} failed! cost={t.cost_str} retry_cnt={i} " diff --git a/memory_scope/scheme/memory_node.py b/memory_scope/scheme/memory_node.py index b4f42307..2bc6b010 100644 --- a/memory_scope/scheme/memory_node.py +++ b/memory_scope/scheme/memory_node.py @@ -31,21 +31,16 @@ class MemoryNode(BaseModel): timestamp: int = Field(int(datetime.datetime.now().timestamp()), description="timestamp of the memory node") - obs_dt: str = Field("", description="dt of the observation") + dt: str = Field("", description="dt of the memory node") obs_reflected: bool = Field(False, description="if the observation is reflected") obs_updated: bool = Field(False, description="if the observation has updated user profile or insight") - obs_keyword: str = Field("", description="keywords of the content") - - insight_key: str = Field("", description="insight_key") - - insight_value: str = Field("", description="insight_value") - def __init__(self, **kwargs): super().__init__(**kwargs) - self.gen_memory_id() + self.memory_id = f"{self.user_name}_{self.target_name}_{self.timestamp}_{md5_hash(self.content)[:8]}" + self.dt = datetime.datetime.fromtimestamp(self.timestamp).strftime("%Y%m%d") @property def node_keys(self): @@ -54,5 +49,3 @@ class MemoryNode(BaseModel): def __getitem__(self, key: str): return self.model_dump().get(key) - def gen_memory_id(self): - self.memory_id = f"{self.user_name}_{self.target_name}_{self.timestamp}_{md5_hash(self.content)[:8]}" diff --git a/memory_scope/storage/llama_index_elastic_search_store.py b/memory_scope/storage/llama_index_elastic_search_store.py index f428eb79..d897c2f8 100644 --- a/memory_scope/storage/llama_index_elastic_search_store.py +++ b/memory_scope/storage/llama_index_elastic_search_store.py @@ -7,6 +7,7 @@ from llama_index.vector_stores.elasticsearch import ElasticsearchStore, AsyncDen from memory_scope.models.base_model import BaseModel from memory_scope.scheme.memory_node import MemoryNode from memory_scope.storage.base_vector_store import BaseVectorStore +from memory_scope.utils.logger import Logger class _ElasticsearchStore(ElasticsearchStore): @@ -83,6 +84,7 @@ class LlamaIndexElasticSearchStore(BaseVectorStore): **kwargs) self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store, embed_model=self.embedding_model.model) + self.logger = Logger.get_logger() def retrieve(self, query: str, @@ -100,6 +102,8 @@ class LlamaIndexElasticSearchStore(BaseVectorStore): query: str, top_k: int, filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]: + self.logger.info(f"query={query} top_k={top_k} filter_dict={filter_dict}") + if filter_dict is None: filter_dict = {} diff --git a/memory_scope/utils/global_context.py b/memory_scope/utils/global_context.py index f9ec7a57..5bfc53ca 100644 --- a/memory_scope/utils/global_context.py +++ b/memory_scope/utils/global_context.py @@ -23,5 +23,7 @@ class GlobalContext(object): self.thread_pool: ThreadPoolExecutor | None = None self.language: LanguageEnum = LanguageEnum.EN + self.meta_data: Dict[str, Any] = {} + G_CONTEXT = GlobalContext() diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index 11bfd727..a7bf4426 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -6,10 +6,13 @@ from copy import deepcopy from importlib import import_module import pyfiglet -from termcolor import colored, COLORS +from termcolor import colored from memory_scope.enumeration.message_role_enum import MessageRoleEnum +ALL_COLORS = ["red", "green", "yellow", "blue", "magenta", "cyan", "light_grey", "light_red", "light_green", + "light_yellow", "light_blue", "light_magenta", "light_cyan", "white"] + def underscore_to_camelcase(name: str, is_first_title: bool = True): name_split = name.split("_") @@ -69,7 +72,7 @@ def char_logo(words: str, seed: int = time.time_ns(), color=None): font = pyfiglet.Figlet() rendered_text = font.renderText(words) colored_lines = [] - all_colors = list(COLORS.keys()) + all_colors = ALL_COLORS.copy() random.seed = seed for line in rendered_text.splitlines(): line_color = color