From 2237a6abafc4d36c1379d42bc367c183d285d92a Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sun, 28 Jul 2024 20:12:02 +0800 Subject: [PATCH] fix test file import problem --- examples/api/__init__.py | 0 .../api}/test_interface.py | 0 memoryscope/cli.py | 1 - memoryscope/core/chat/api_memory_chat.py | 5 +- memoryscope/core/chat/cli_memory_chat.py | 5 +- memoryscope/core/config/arguments.py | 4 +- memoryscope/core/config/config_manager.py | 18 +-- memoryscope/core/config/demo_config.yaml | 78 ++++++------- memoryscope/core/memoryscope.py | 2 + memoryscope/core/utils/prompt_handler.py | 6 +- memoryscope/core/worker/memory_base_worker.py | 2 +- tests/models/test_models_lli_embedding.py | 4 +- tests/models/test_models_lli_generation.py | 4 +- tests/models/test_models_lli_rank.py | 2 +- tests/{operations => other}/init_test.py | 0 tests/other/read_yaml.py | 4 +- tests/storages/test_storages_lli_es.py | 4 +- tests/storages/test_storages_lli_synces.py | 4 +- tests/worker/test_workers_cn.py | 106 +++++++++--------- tests/worker/test_workers_en.py | 104 +++++++++-------- 20 files changed, 175 insertions(+), 178 deletions(-) create mode 100644 examples/api/__init__.py rename {tests/operations => examples/api}/test_interface.py (100%) rename tests/{operations => other}/init_test.py (100%) diff --git a/examples/api/__init__.py b/examples/api/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/operations/test_interface.py b/examples/api/test_interface.py similarity index 100% rename from tests/operations/test_interface.py rename to examples/api/test_interface.py diff --git a/memoryscope/cli.py b/memoryscope/cli.py index 1d890358..85d5bec5 100644 --- a/memoryscope/cli.py +++ b/memoryscope/cli.py @@ -8,7 +8,6 @@ from memoryscope.core.memoryscope import MemoryScope def cli_job(**kwargs): - kwargs["memory_chat_type"] = "cli_chat" MemoryScope(**kwargs).default_memory_chat.run() diff --git a/memoryscope/core/chat/api_memory_chat.py b/memoryscope/core/chat/api_memory_chat.py index d1992d73..6aee7c75 100644 --- a/memoryscope/core/chat/api_memory_chat.py +++ b/memoryscope/core/chat/api_memory_chat.py @@ -53,7 +53,10 @@ class ApiMemoryChat(BaseMemoryChat): PromptHandler: An instance of the PromptHandler configured for this CLI session. """ if self._prompt_handler is None: - self._prompt_handler = PromptHandler(__file__, prompt_file="memory_chat_prompt", **self.kwargs) + self._prompt_handler = PromptHandler(__file__, + language=self.context.language, + prompt_file="memory_chat_prompt", + **self.kwargs) return self._prompt_handler @property diff --git a/memoryscope/core/chat/cli_memory_chat.py b/memoryscope/core/chat/cli_memory_chat.py index f2fe8a0a..653b92ec 100644 --- a/memoryscope/core/chat/cli_memory_chat.py +++ b/memoryscope/core/chat/cli_memory_chat.py @@ -68,7 +68,10 @@ class CliMemoryChat(BaseMemoryChat): PromptHandler: An instance of the PromptHandler configured for this CLI session. """ if self._prompt_handler is None: - self._prompt_handler = PromptHandler(__file__, prompt_file="memory_chat_prompt", **self.kwargs) + self._prompt_handler = PromptHandler(__file__, + language=self.context.language, + prompt_file="memory_chat_prompt", + **self.kwargs) return self._prompt_handler def print_logo(self): diff --git a/memoryscope/core/config/arguments.py b/memoryscope/core/config/arguments.py index 1925e31e..f6542d3e 100644 --- a/memoryscope/core/config/arguments.py +++ b/memoryscope/core/config/arguments.py @@ -12,8 +12,8 @@ class Arguments(object): logger_name_time_suffix: str = field(default="%Y%m%d_%H%M%S") - memory_chat_type: str = field(default="cli_chat", metadata={ - "help": "cli_chat(Command-line interaction), api_chat(API interface interaction), etc."}) + memory_chat_class: str = field(default="cli_memory_chat", metadata={ + "help": "cli_memory_chat(Command-line interaction), api_memory_chat(API interface interaction), etc."}) consolidate_memory_interval_time: int = field(default=1, metadata={ "help": "If you feel that the token consumption is relatively high, please increase the time interval."}) diff --git a/memoryscope/core/config/config_manager.py b/memoryscope/core/config/config_manager.py index 7ac8095c..b23eed5c 100644 --- a/memoryscope/core/config/config_manager.py +++ b/memoryscope/core/config/config_manager.py @@ -65,14 +65,9 @@ class ConfigManager(object): @staticmethod def update_memory_chat_by_arguments(config: dict, arguments: Arguments): - if arguments.memory_chat_type == "cli_chat": - memory_chat_class = "chat.cli_memory_chat" - elif arguments.memory_chat_type == "api_chat": - memory_chat_class = "chat.api_memory_chat" - else: - raise NotImplementedError(f"known memory_chat_type={arguments.memory_chat_type}") + memory_chat_class_split = config["class"].split(".") config.update({ - "class": memory_chat_class, + "class": ".".join(memory_chat_class_split[:-1] + [arguments.memory_chat_class]), "human_name": DEFAULT_HUMAN_NAME[LanguageEnum(arguments.language)], "assistant_name": "AI", }) @@ -154,7 +149,7 @@ class ConfigManager(object): def clear_node_all(self, node: str): self.config[node].clear() - def dump_config(self, file_type: Literal["json", "yaml"], to_stream: bool = True, file_path: Optional[str] = None): + def dump_config(self, file_type: Literal["json", "yaml"] = "yaml", file_path: Optional[str] = None) -> str: if file_type == "json": content = json.dumps(self.config, indent=2, ensure_ascii=False) elif file_type == "yaml": @@ -162,9 +157,8 @@ class ConfigManager(object): else: raise NotImplementedError - if to_stream: - print(content) - - if file_type: + if file_path: with open(file_path, "w") as f: f.write(content) + + return content diff --git a/memoryscope/core/config/demo_config.yaml b/memoryscope/core/config/demo_config.yaml index e81ba428..88680e7a 100644 --- a/memoryscope/core/config/demo_config.yaml +++ b/memoryscope/core/config/demo_config.yaml @@ -7,78 +7,78 @@ global: memory_chat: cli_memory_chat: - class: chat.cli_memory_chat + class: core.chat.cli_memory_chat memory_service: memoryscope_service generation_model: generation_model memory_service: memoryscope_service: - class: memory.service.memory_scope_service + class: core.service.memory_scope_service memory_operations: read_message: - class: memory.operation.frontend_operation + class: core.operation.frontend_operation workflow: read_message description: "read short memory" retrieve_memory: - class: memory.operation.frontend_operation + class: core.operation.frontend_operation workflow: set_query,[extract_time|retrieve_obs_ins,semantic_rank],fuse_rerank description: "retrieve long-term memory" list_memory: - class: memory.operation.frontend_operation + class: core.operation.frontend_operation workflow: set_query,retrieve_top_memory,print_memory description: "read all long-term memory of the user" delete_memory: - class: memory.operation.frontend_operation + class: core.operation.frontend_operation workflow: set_query,retrieve_all_memory,delete_memory description: "delete a single long-term memory" delete_all: - class: memory.operation.frontend_operation + class: core.operation.frontend_operation workflow: set_query,retrieve_all_memory,delete_all description: "delete all long-term memory" add_memory: - class: memory.operation.frontend_operation + class: core.operation.frontend_operation workflow: add_memory description: "add a single observation" consolidate_memory: - class: memory.operation.consolidate_memory_op + class: core.operation.consolidate_memory_op workflow: info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat,store_memory description: "summary user's observation memory" interval_time: 1 reflect_and_reconsolidate: - class: memory.operation.backend_operation + class: core.operation.backend_operation workflow: load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat,store_memory description: "summary user's insight memory" interval_time: 15 worker: dummy: - class: memory.worker.dummy_worker + class: core.worker.dummy_worker generation_model: generation_model embedding_model: embedding_model rank_model: rank_model read_message: - class: memory.worker.frontend.read_message_worker + class: core.worker.frontend.read_message_worker set_query: - class: memory.worker.frontend.set_query_worker + class: core.worker.frontend.set_query_worker retrieve_obs_ins: - class: memory.worker.frontend.retrieve_memory_worker + class: core.worker.frontend.retrieve_memory_worker retrieve_obs_top_k: 100 retrieve_ins_top_k: 100 extract_time: - class: memory.worker.frontend.extract_time_worker + class: core.worker.frontend.extract_time_worker generation_model: generation_model semantic_rank: - class: memory.worker.frontend.semantic_rank_worker + class: core.worker.frontend.semantic_rank_worker rank_model: rank_model fuse_rerank: - class: memory.worker.frontend.fuse_rerank_worker + class: core.worker.frontend.fuse_rerank_worker fuse_score_threshold: 0.01 fuse_ratio_dict: conversation: 0.5 @@ -88,88 +88,88 @@ worker: fuse_time_ratio: 2.0 fuse_rerank_top_k: 10 retrieve_top_memory: - class: memory.worker.frontend.retrieve_memory_worker + class: core.worker.frontend.retrieve_memory_worker retrieve_obs_top_k: 100 retrieve_ins_top_k: 100 retrieve_expired_top_k: 100 print_memory: - class: memory.worker.frontend.print_memory_worker + class: core.worker.frontend.print_memory_worker retrieve_all_memory: - class: memory.worker.frontend.retrieve_memory_worker + class: core.worker.frontend.retrieve_memory_worker retrieve_obs_top_k: 1000 retrieve_ins_top_k: 1000 retrieve_expired_top_k: 1000 delete_memory: - class: memory.worker.backend.update_memory_worker + class: core.worker.backend.update_memory_worker method: delete_memory delete_all: - class: memory.worker.backend.update_memory_worker + class: core.worker.backend.update_memory_worker method: delete_all add_memory: - class: memory.worker.backend.update_memory_worker + class: core.worker.backend.update_memory_worker method: from_query info_filter: - class: memory.worker.backend.info_filter_worker + class: core.worker.backend.info_filter_worker generation_model: generation_model load_today_memory: - class: memory.worker.backend.load_memory_worker + class: core.worker.backend.load_memory_worker retrieve_today_top_k: 100 get_observation: - class: memory.worker.backend.get_observation_worker + class: core.worker.backend.get_observation_worker generation_model: generation_model get_observation_with_time: - class: memory.worker.backend.get_observation_with_time_worker + class: core.worker.backend.get_observation_with_time_worker generation_model: generation_model contra_repeat: - class: memory.worker.backend.contra_repeat_worker + class: core.worker.backend.contra_repeat_worker generation_model: generation_model store_memory: - class: memory.worker.backend.update_memory_worker + class: core.worker.backend.update_memory_worker method: from_memory_key memory_key: all load_obs_and_insight: - class: memory.worker.backend.load_memory_worker + class: core.worker.backend.load_memory_worker retrieve_not_reflected_top_k: 100 retrieve_not_updated_top_k: 100 retrieve_insight_top_k: 100 get_reflection_subject: - class: memory.worker.backend.get_reflection_subject_worker + class: core.worker.backend.get_reflection_subject_worker generation_model: generation_model reflect_obs_cnt_threshold: 10 update_insight: - class: memory.worker.backend.update_insight_worker + class: core.worker.backend.update_insight_worker generation_model: generation_model rank_model: rank_model long_contra_repeat: - class: memory.worker.backend.long_contra_repeat_worker + class: core.worker.backend.long_contra_repeat_worker generation_model: generation_model model: generation_model: - class: models.llama_index_generation_model + class: core.models.llama_index_generation_model module_name: dashscope_generation model_name: qwen-max max_tokens: 2000 embedding_model: - class: models.llama_index_embedding_model + class: core.models.llama_index_embedding_model module_name: dashscope_embedding model_name: text-embedding-v2 rank_model: - class: models.llama_index_rank_model + class: core.models.llama_index_rank_model module_name: dashscope_rank model_name: gte-rerank top_n: 500 dummy_generation: - class: models.dummy_generation_model + class: core.models.dummy_generation_model module_name: dummy_generation model_name: dummy_generation_model memory_store: - class: storage.llama_index_es_memory_store + class: core.storage.llama_index_es_memory_store embedding_model: embedding_model index_name: memory_index es_url: http://localhost:9200 retrieve_mode: dense monitor: - class: storage.dummy_monitor \ No newline at end of file + class: core.storage.dummy_monitor \ No newline at end of file diff --git a/memoryscope/core/memoryscope.py b/memoryscope/core/memoryscope.py index 3e0c7978..54b429b9 100644 --- a/memoryscope/core/memoryscope.py +++ b/memoryscope/core/memoryscope.py @@ -82,6 +82,8 @@ class MemoryScope(ConfigManager): if self.context.monitor: self.context.monitor.close() + self.logger.close() + def __enter__(self): self.init_context_by_config() diff --git a/memoryscope/core/utils/prompt_handler.py b/memoryscope/core/utils/prompt_handler.py index d44559f3..a60497e4 100644 --- a/memoryscope/core/utils/prompt_handler.py +++ b/memoryscope/core/utils/prompt_handler.py @@ -16,9 +16,9 @@ class PromptHandler(object): def __init__(self, class_path: str, + language: LanguageEnum | str, prompt_file: str = "", prompt_dict: dict = None, - language_enum: LanguageEnum = LanguageEnum.EN, **kwargs): """ Initializes the PromptHandler with paths to prompt sources and additional keyword arguments. @@ -27,13 +27,13 @@ class PromptHandler(object): class_path (str): The path to the class where prompts are utilized. prompt_file (str, optional): The path to an external file containing prompts. Defaults to "". prompt_dict (dict, optional): A dictionary directly containing prompt definitions. Defaults to None. - language_enum (LanguageEnum): context language. + language (LanguageEnum, str): context language. **kwargs: Additional keyword arguments that might be used in prompt handling. """ class_path: Path = Path(class_path) self._class_dir: Path = class_path.parent self._class_name: str = class_path.stem - self._language_enum: LanguageEnum = language_enum + self._language_enum: LanguageEnum = LanguageEnum(language) self.kwargs = kwargs self._prompt_dict: Dict[str, str] = {} diff --git a/memoryscope/core/worker/memory_base_worker.py b/memoryscope/core/worker/memory_base_worker.py index f4de9044..85b1ad41 100644 --- a/memoryscope/core/worker/memory_base_worker.py +++ b/memoryscope/core/worker/memory_base_worker.py @@ -190,7 +190,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): PromptHandler: An instance of PromptHandler initialized with specific file path and keyword arguments. """ if self._prompt_handler is None: - self._prompt_handler = PromptHandler(self.FILE_PATH, **self.kwargs) + self._prompt_handler = PromptHandler(self.FILE_PATH, language=self.language, **self.kwargs) return self._prompt_handler @property diff --git a/tests/models/test_models_lli_embedding.py b/tests/models/test_models_lli_embedding.py index f21d0bc7..67d3bb13 100644 --- a/tests/models/test_models_lli_embedding.py +++ b/tests/models/test_models_lli_embedding.py @@ -5,8 +5,8 @@ sys.path.append(".") # noqa: E402 import asyncio import unittest -from memoryscope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel -from memoryscope.utils.logger import Logger +from memoryscope.core.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.utils.logger import Logger class TestLLIEmbedding(unittest.TestCase): diff --git a/tests/models/test_models_lli_generation.py b/tests/models/test_models_lli_generation.py index 474851b4..1fbe41b5 100644 --- a/tests/models/test_models_lli_generation.py +++ b/tests/models/test_models_lli_generation.py @@ -6,8 +6,8 @@ import unittest import time import asyncio from memoryscope.scheme.message import Message -from memoryscope.models.llama_index_generation_model import LlamaIndexGenerationModel -from memoryscope.utils.logger import Logger +from memoryscope.core.models.llama_index_generation_model import LlamaIndexGenerationModel +from memoryscope.core.utils.logger import Logger class TestLLILLM(unittest.TestCase): diff --git a/tests/models/test_models_lli_rank.py b/tests/models/test_models_lli_rank.py index 25c59fde..e0e90c7c 100644 --- a/tests/models/test_models_lli_rank.py +++ b/tests/models/test_models_lli_rank.py @@ -1,7 +1,7 @@ import asyncio import unittest -from memoryscope.models.llama_index_rank_model import LlamaIndexRankModel +from memoryscope.core.models.llama_index_rank_model import LlamaIndexRankModel class TestLLIReRank(unittest.TestCase): diff --git a/tests/operations/init_test.py b/tests/other/init_test.py similarity index 100% rename from tests/operations/init_test.py rename to tests/other/init_test.py diff --git a/tests/other/read_yaml.py b/tests/other/read_yaml.py index c54bd6ad..ac0513f4 100644 --- a/tests/other/read_yaml.py +++ b/tests/other/read_yaml.py @@ -2,10 +2,10 @@ import sys sys.path.append(".") # noqa: E402 -from memoryscope.utils.prompt_handler import PromptHandler +from memoryscope.core.utils.prompt_handler import PromptHandler if __name__ == "__main__": file_path: str = __file__ print(file_path) - handler = PromptHandler(__file__, "read_prompt") + handler = PromptHandler(__file__, language="cn", prompt_file="read_prompt", ) print(handler.prompt_dict) diff --git a/tests/storages/test_storages_lli_es.py b/tests/storages/test_storages_lli_es.py index 791286aa..a72d1aab 100644 --- a/tests/storages/test_storages_lli_es.py +++ b/tests/storages/test_storages_lli_es.py @@ -1,8 +1,8 @@ import unittest -from memoryscope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore class TestLlamaIndexElasticSearchStore(unittest.TestCase): diff --git a/tests/storages/test_storages_lli_synces.py b/tests/storages/test_storages_lli_synces.py index bf32bc73..f407a389 100644 --- a/tests/storages/test_storages_lli_synces.py +++ b/tests/storages/test_storages_lli_synces.py @@ -1,8 +1,8 @@ import unittest -from memoryscope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore class TestLlamaIndexElasticSearchStore(unittest.TestCase): diff --git a/tests/worker/test_workers_cn.py b/tests/worker/test_workers_cn.py index e7a65fae..77cad3ff 100644 --- a/tests/worker/test_workers_cn.py +++ b/tests/worker/test_workers_cn.py @@ -1,44 +1,51 @@ import datetime import unittest -from memoryscope.cli import MemoryScope from memoryscope.constants.common_constants import CHAT_MESSAGES, NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, \ - MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES + MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES, \ + MEMORYSCOPE_CONTEXT +from memoryscope.core.config.arguments import Arguments +from memoryscope.core.memoryscope import MemoryScope +from memoryscope.core.utils.tool_functions import init_instance_by_config +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode from memoryscope.scheme.message import Message -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger -from memoryscope.utils.tool_functions import init_instance_by_config class TestWorkersCn(unittest.TestCase): """Tests for LLIEmbedding""" def setUp(self): - datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S') - self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True) - - ms = MemoryScope() - ms.read_config("config/demo_config_cn.yaml") - ms.init_global_content_by_config() + arguments = Arguments( + language="cn", + memory_chat_class="api_memory_chat", + generation_backend="dashscope_generation", + generation_model="qwen-max", + embedding_backend="dashscope_embedding", + embedding_model="text-embedding-v2", + use_dummy_ranker=False, + rank_backend="dashscope_rank", + rank_model="gte-rerank", + ) + self.ms = MemoryScope(arguments=arguments) + config = self.ms.dump_config() + self.ms.logger.info(f"config=\n{config}") def tearDown(self): - self.logger.close() + self.ms.close() @unittest.skip def test_extract_time(self): name = "extract_time" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) query = "明天我去上海出差" query_timestamp = int(datetime.datetime.now().timestamp()) @@ -53,13 +60,12 @@ class TestWorkersCn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜"), @@ -79,13 +85,12 @@ class TestWorkersCn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="你知道北京哪里的海鲜最新鲜吗"), @@ -115,13 +120,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜"), @@ -149,13 +153,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="有没有推荐的策略游戏?最近想找新的挑战。"), @@ -181,13 +184,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation_with_time" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="去年我们一起合作了因果推断技术"), @@ -211,13 +213,12 @@ class TestWorkersCn(unittest.TestCase): name = "contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"), @@ -272,13 +273,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_reflection_subject" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。"), @@ -304,18 +304,17 @@ class TestWorkersCn(unittest.TestCase): @unittest.skip def test_update_insight_worker(self): - reflection_worker = self.test_get_reflection_subject.__wrapped__(self) + reflection_worker: MemoryBaseWorker = self.test_get_reflection_subject.__wrapped__(self) name = "update_insight" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, context=reflection_worker.context, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="用户喜欢打王者荣耀"), @@ -327,18 +326,17 @@ class TestWorkersCn(unittest.TestCase): result = "\n".join(result) worker.logger.info(f"result.update_insight={result}") - @unittest.skip + # @unittest.skip def test_long_contra_repeat_worker(self): name = "long_contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。"), diff --git a/tests/worker/test_workers_en.py b/tests/worker/test_workers_en.py index 0123bc3a..ff421253 100644 --- a/tests/worker/test_workers_en.py +++ b/tests/worker/test_workers_en.py @@ -1,44 +1,51 @@ import datetime import unittest -from memoryscope.cli import MemoryScope from memoryscope.constants.common_constants import CHAT_MESSAGES, NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, \ - MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES + MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES, \ + MEMORYSCOPE_CONTEXT +from memoryscope.core.config.arguments import Arguments +from memoryscope.core.memoryscope import MemoryScope +from memoryscope.core.utils.tool_functions import init_instance_by_config +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode from memoryscope.scheme.message import Message -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger -from memoryscope.utils.tool_functions import init_instance_by_config class TestWorkersEn(unittest.TestCase): """Tests for LLIEmbedding""" def setUp(self): - datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S') - self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True) - - ms = MemoryScope() - ms.read_config("config/demo_config_en.yaml") - ms.init_global_content_by_config() + arguments = Arguments( + language="en", + memory_chat_class="api_memory_chat", + generation_backend="dashscope_generation", + generation_model="qwen-max", + embedding_backend="dashscope_embedding", + embedding_model="text-embedding-v2", + use_dummy_ranker=False, + rank_backend="dashscope_rank", + rank_model="gte-rerank", + ) + self.ms = MemoryScope(arguments=arguments) + config = self.ms.dump_config() + self.ms.logger.info(f"config=\n{config}") def tearDown(self): - self.logger.close() + self.ms.close() @unittest.skip def test_extract_time(self): name = "extract_time" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) query = "I will be on a business trip to Shanghai tomorrow." query_timestamp = int(datetime.datetime.now().timestamp()) @@ -53,13 +60,12 @@ class TestWorkersEn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="I love to eat Sichuan cuisine."), @@ -80,13 +86,12 @@ class TestWorkersEn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="Do you know where the freshest seafood is in Beijing?"), @@ -129,13 +134,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) # FIXME Does the appearance of 'am' indicate the presence of a time keyword? chat_messages = [ @@ -159,13 +163,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, @@ -202,13 +205,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation_with_time" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, @@ -235,13 +237,12 @@ class TestWorkersEn(unittest.TestCase): name = "contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="User is working in Meituan"), @@ -294,13 +295,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_reflection_subject" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="Users are interested in strategy games and looking for new challenges."), @@ -327,18 +327,17 @@ class TestWorkersEn(unittest.TestCase): @unittest.skip def test_update_insight_worker(self): - reflection_worker = self.test_get_reflection_subject.__wrapped__(self) + reflection_worker: MemoryBaseWorker = self.test_get_reflection_subject.__wrapped__(self) name = "update_insight" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, context=reflection_worker.context, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="Users like to play King of Glory"), @@ -355,13 +354,12 @@ class TestWorkersEn(unittest.TestCase): name = "long_contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="Users are interested in strategy games and looking for new challenges."),