fix test file import problem

This commit is contained in:
jinli.yl 2024-07-28 20:12:02 +08:00
parent 70ecdb1d94
commit 2237a6abaf
20 changed files with 175 additions and 178 deletions

0
examples/api/__init__.py Normal file
View file

View file

@ -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()

View file

@ -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

View file

@ -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):

View file

@ -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."})

View file

@ -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

View file

@ -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
class: core.storage.dummy_monitor

View file

@ -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()

View file

@ -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] = {}

View file

@ -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

View file

@ -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):

View file

@ -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):

View file

@ -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):

View file

@ -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)

View file

@ -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):

View file

@ -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):

View file

@ -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="用户对策略游戏感兴趣,寻找新挑战。"),

View file

@ -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."),