feat: Resolve conflict, auto committed by CodeFlow

This commit is contained in:
qintiancheng.qtc 2024-06-28 16:45:32 +08:00
commit 509f11be84
83 changed files with 2787 additions and 1140 deletions

View file

@ -1,27 +0,0 @@
{
"global_configs": {
"thread_pool_max_count": 5,
"dash_scope_apikey": "",
"open_ai_apikey": "",
"language": "en",
"chat_list": [
"memory_chat"
]
},
"memory_chat": {
"clazz": "chat.memory_chat",
"retrieve": "parse_params,load_profile,extract_time,es_similar,es_keyword,semantic_rank,fuse_rerank",
"generation_model": "dashscope_generation",
"history_msg_count": 3
},
"vector_store": {
"clazz": "storage.base_vector_store",
"index_name": "memory_test",
"password": ""
},
"monitor": {
"clazz": "storage.base_monitor",
"index_name": "memory_test"
},
"workers": "workers"
}

View file

@ -1,57 +1,68 @@
global_config:
thread_pool_max_count: 5
language: en
max_workers: 5
dash_scope_apikey:
open_ai_apikey:
language: en
chat_list:
- memory_chat
memory_chat:
memory_service: memory_chat_service
memory_chat_service:
class: memory.base_memory_service
memory_operations:
- name: read_memory
class: memory.workflow.base_workflow
workflow: parse_params,load_profile,extract_time,es_similar,es_keyword,semantic_rank,fuse_rerank
work_type: frontend
- name: list_memory
class: memory.workflow.base_workflow
workflow: dummy
work_type: frontend
- name: extract_memory
class: memory.workflow.base_workflow
workflow: dummy
work_type: backend
interval_time: 60
min_count: 5
- name: reflect_memory
class: memory.workflow.base_workflow
workflow: dummy
work_type: backend
interval_time: 300
cli_memory_chat:
class: chat.cli_memory_chat
memory_service: memory_chat_service
generation_model: dashscope_generation
human_name: human
assistant_name: assistant
memory_service:
memory_chat_service:
class: memory.service.chat_memory_service
history_msg_count: 32
contextual_msg_count: 6
read_memory_key: read_memory
memory_operations:
read_message:
class: memory.operation.read_memory
workflow: dummy_worker
description: "read session messages of the user"
read_memory:
class: memory.operation.read_memory
workflow: dummy_worker
description: "read related memories of the user"
list_memory:
class: memory.operation.read_memory
workflow: dummy_worker
description: "read all memories of the user"
write_memory:
class: memory.operation.write_memory
workflow: dummy_worker
description: "write observation memories of the user"
interval_time: 60
summary_memory:
class: memory.operation.summary_memory
workflow: dummy_worker
description: "summary observation memories of the user"
interval_time: 300
models:
dashscope_generation:
class: models.llama_index_generation_model
module_name: dashscope_generation
model_name: qwen-max
dashscope_embedding:
class: models.llama_index_embedding_model
module_name: dashscope_embedding
model_name: text-embedding-v2
dashscope_rank:
class: models.llama_index_rank_model
module_name: dashscope_rank
model_name: gte-rerank
vector_store:
clazz: storage.base_vector_store
index_name: memory_test
password: ''
class: storage.llama_index_elastic_search_store
embedding_model: dashscope_embedding
index_name: memory_index
es_url: http://localhost:9200
monitor:
clazz: storage.base_monitor
index_name: memory_test
workers:
- name: update_insight
clazz: worker.summary_long.update_insight
class: storage.dummy_monitor
worker:
dummy_worker:
class: memory.worker.dummy_worker
generation_model: dashscope_generation
embedding_model: dashscope_embedding
rank_model: dashscope_rank
models:
- name: dashscope_generation
clazz: models.llama_index_generation_model
module_name: DashScope
model_name: qwen-max
- name: dashscope_embedding
clazz: models.base_embedding_model
module_name: DashScopeEmbedding
model_name: text-embedding-v2
- name: dashscope_rank
clazz: models.base_rank_model
module_name: DashScopeRerank
model_name: gte-rerank

View file

@ -1,5 +0,0 @@
{
"clazz": "models.base_embedding_model",
"model_name": "text-embedding-v2",
"method_type": "DashScopeEmbedding"
}

View file

@ -1,5 +0,0 @@
{
"clazz": "models.llama_index_generation_model",
"model_name": "qwen-max",
"method_type": "DashScope"
}

View file

@ -1,5 +0,0 @@
{
"clazz": "models.base_rank_model",
"model_name": "gte-rerank",
"method_type": "DashScopeRerank"
}

View file

@ -1,8 +0,0 @@
{
"update_insight": {
"clazz": "worker.summary_long.update_insight",
"generation_model": "dashscope_generation",
"embedding_model": "dashscope_embedding",
"rank_model": "dashscope_rank"
}
}

View file

@ -1,3 +1,3 @@
""" Version of MemoryScope."""
__version__ = "0.1.0-alpha.1"
__version__ = "0.1.0-alpha.1"

View file

@ -2,9 +2,6 @@ from abc import ABCMeta, abstractmethod
class BaseMemoryChat(metaclass=ABCMeta):
def __init__(self, **kwargs):
self.kwargs = kwargs
@abstractmethod
def chat_with_memory(self, query: str):

View file

@ -1,4 +0,0 @@
class BaseMemoryService(object):
def __init__(self, **kwargs):
self.kwargs = kwargs

View file

@ -1,83 +1,177 @@
import datetime
import os
import time
from typing import List
import questionary
from rich.console import Console
from .memory_chat import MemoryChat
from enumeration.message_role_enum import MessageRoleEnum
from scheme.message import Message
from memory_scope.chat.base_memory_chat import BaseMemoryChat
from memory_scope.chat.global_context import G_CONTEXT
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
from memory_scope.memory.service.base_memory_service import BaseMemoryService
from memory_scope.models.base_model import BaseModel
from memory_scope.prompts.memory_chat_prompt import MEMORY_PROMPT, SYSTEM_PROMPT
from memory_scope.scheme.message import Message
from memory_scope.scheme.model_response import ModelResponse, ModelResponseGen
from memory_scope.utils.logger import Logger
from memory_scope.utils.tool_functions import char_logo
class CliMemoryChat(MemoryChat):
class CliMemoryChat(BaseMemoryChat):
USER_COMMANDS = {
"/exit": "exit the CLI",
"/memory": "print the current contents of agent memory",
"/retrieve": "retrieve related memory",
"/log": "log chat progress",
# TODO add more commands
"exit": "exit the CLI",
"help": "get cli commands help",
"stream": "get stream response"
}
def chat_with_memory(self, query): # for testing
def __init__(self,
memory_service: str,
generation_model: str,
stream: bool = True,
human_name: str = "human",
assistant_name: str = "assistant",
**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.assistant_name: str = assistant_name
self.kwargs: dict = kwargs
self._logo = char_logo("MemoryScope")
self.logger = Logger.get_logger()
def print_logo(self):
for line in self._logo:
print(line)
@property
def memory_service(self) -> BaseMemoryService:
if isinstance(self._memory_service, str):
self._memory_service = G_CONTEXT.memory_service_dict[self._memory_service]
self._memory_service.start_service()
return self._memory_service
@property
def generation_model(self) -> BaseModel:
if isinstance(self._generation_model, str):
self._generation_model = G_CONTEXT.model_dict[self._generation_model]
return self._generation_model
def get_system_prompt(self) -> Message:
system_prompt = SYSTEM_PROMPT[G_CONTEXT.language].strip()
memories: str = self.memory_service.read_memory()
if memories:
memory_prompt = MEMORY_PROMPT[G_CONTEXT.language]
system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt, memories]])
return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt)
def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen:
query = query.strip()
if not query:
return
time_created = int(datetime.datetime.now().timestamp())
message = Message(
role=MessageRoleEnum.USER, content=query, time_created=time_created
)
messages = [message]
return self.generation_model.call(messages=messages, stream=True)
new_message: Message = Message(role=MessageRoleEnum.USER.value,
role_name=self.human_name,
content=query)
def retrieve_all(self): # for testing
return "memory 1. 2. 3."
self.memory_service.add_messages(new_message)
system_message: Message = self.get_system_prompt()
model_response = self.generation_model.call(messages=[system_message, new_message], stream=self.stream)
if self.stream:
for _ in model_response:
_.message.role_name = self.assistant_name
yield _
else:
model_response.message.role_name = self.assistant_name
return model_response
def process_commands(self, query: str) -> bool:
continue_run = True
query_split = query.lstrip("/").lower().split(" ")
query = query_split[0]
args = query_split[1:]
if query == "exit":
self.memory_service.stop_service()
continue_run = False
elif query == "help":
questionary.print("CLI commands", "bold")
for cmd, desc in self.USER_COMMANDS.items():
questionary.print(text=f" /{cmd}:", style="bold")
questionary.print(text=f" {desc}")
elif query == "stream":
self.stream = bool(args[0])
questionary.print(f"stream: {self.stream}")
elif query in self.memory_service.op_description_dict:
if not args:
result = self.memory_service.do_operation(op_name=query)
questionary.print(result)
elif args[0].isdigit():
refresh_time = int(args[0])
while True:
time.sleep(refresh_time)
result = self.memory_service.do_operation(op_name=query)
os.system('clear')
self.print_logo()
questionary.print(result)
else:
questionary.print("unknown command received. Please try again!")
else:
questionary.print("unknown command received. Please try again!")
return continue_run
def run(self):
console = Console()
self.print_logo()
self.USER_COMMANDS.update(self.memory_service.op_description_dict)
while True:
query = questionary.text(
"Enter your message or command:",
multiline=False,
qmark=">",
).ask()
try:
query = questionary.text(message=f"{self.human_name}:", multiline=False, qmark=">").unsafe_ask()
if not query:
continue
query = query.rstrip()
query: str = query.strip()
if query == "":
console.print("Empty input received. Try again!")
continue
# handle cli / commands with memory ops
if query.startswith("/"):
if self.process_commands(query=query):
continue
else:
break
# Handle CLI commands
if query.startswith("/"):
if query.lower() == "/exit":
break
elif query.lower() == "/memory":
console.print(self.memory_service.retrieve_all())
elif query.lower() == "/help":
questionary.print("CLI commands", "bold")
for cmd, desc in self.USER_COMMANDS.items():
questionary.print(cmd, "bold")
questionary.print(f" {desc}")
continue
while True:
try:
# with console.status("[bold cyan]Thinking..."):
msg = None
questionary.print("> ", end="", style="fg:yellow")
questionary.print(f"{self.assistant_name}: ", end="", style="bold")
if self.stream:
for msg in self.chat_with_memory(query=query):
console.print(msg.delta, end="")
console.print()
questionary.print(msg.delta, end="")
questionary.print("")
else:
msg = self.chat_with_memory(query=query)
questionary.print(msg.message.content)
self.memory_service.add_messages(msg.message)
except KeyboardInterrupt:
questionary.print("User interrupt occurred.")
is_exit = questionary.confirm("continue exit").unsafe_ask()
if is_exit:
self.memory_service.stop_service()
break
except KeyboardInterrupt:
console.print("User interrupt occurred.")
retry = questionary.confirm("Retry chat_with_memory()?").ask()
if not retry:
break
except Exception as e:
console.print(
f"An exception occurred when running chat_with_memory(): {e}"
)
retry = questionary.confirm("Retry chat_with_memory()?").ask()
if not retry:
break
except Exception as e:
line = f"An exception occurred when running cli memory chat. args={e.args}"
questionary.print(line)
self.logger.exception(line)
continue
questionary.print(f"A memory writing thread is still running, please be patient and wait!")

View file

@ -1,31 +1,27 @@
from concurrent.futures import ThreadPoolExecutor
from typing import Dict, Any
from chat.base_memory_chat import BaseMemoryChat
from enumeration.language_enum import LanguageEnum
from models.base_model import BaseModel
from storage.base_monitor import BaseMonitor
from storage.base_vector_store import BaseVectorStore
from worker.base_worker import BaseWorker
from memory_scope.chat.base_memory_chat import BaseMemoryChat
from memory_scope.enumeration.language_enum import LanguageEnum
from memory_scope.memory.service.base_memory_service import BaseMemoryService
from memory_scope.models.base_model import BaseModel
from memory_scope.storage.base_monitor import BaseMonitor
from memory_scope.storage.base_vector_store import BaseVectorStore
class GlobalContext(object):
def __init__(self):
self.global_configs: Dict[str, Any] = {}
self.worker_dict: Dict[str, Dict[str, BaseWorker]] = {}
self.global_config: Dict[str, Any] = {}
self.worker_config: Dict[str, Dict[str, Any]] = {}
self.memory_service_dict: Dict[str, BaseMemoryService] = {}
self.model_dict: Dict[str, BaseModel] = {}
self.memory_chat_dict: Dict[str, BaseMemoryChat] = {}
self.vector_store: BaseVectorStore | None = None
self.monitor: BaseMonitor | None = None
self.thread_pool: ThreadPoolExecutor | None = None
self.language: LanguageEnum = LanguageEnum.EN
GLOBAL_CONTEXT = GlobalContext()
G_CONTEXT = GlobalContext()

View file

@ -1,67 +0,0 @@
import datetime
from typing import List
from .base_memory_chat import BaseMemoryChat
from .global_context import GLOBAL_CONTEXT
from enumeration.message_role_enum import MessageRoleEnum
from models.base_model import BaseModel
from prompts.memory_chat_prompt import SYSTEM_PROMPT, MEMORY_PROMPT
from scheme.message import Message
from .memory_service import MemoryService
class MemoryChat(BaseMemoryChat):
def __init__(self, generation_model: str, history_msg_count: int, chat_name: str, **kwargs):
super().__init__(**kwargs)
self.memory_service = MemoryService(chat_name=chat_name, **kwargs)
self.generation_model_name: str = generation_model
self.history_msg_count: int = history_msg_count
self._generation_model: BaseModel | None = None
self.history_message_list: List[Message] = []
@property
def generation_model(self):
if self._generation_model is None:
self._generation_model = GLOBAL_CONTEXT.model_dict[
self.generation_model_name
]
return self._generation_model
@staticmethod
def get_system_prompt(related_memories: List[str], time_created: int) -> Message:
system_prompt = SYSTEM_PROMPT[GLOBAL_CONTEXT.language]
if related_memories:
memory_prompt = MEMORY_PROMPT[GLOBAL_CONTEXT.language]
system_prompt = "\n".join([system_prompt, memory_prompt] + related_memories)
return Message(
role=MessageRoleEnum.SYSTEM,
content=system_prompt.strip(),
time_created=time_created,
)
def chat_with_memory(self, query: str):
query = query.strip()
if not query:
return
time_created = int(datetime.datetime.now().timestamp())
new_message: Message = Message(
role=MessageRoleEnum.USER, content=query, time_created=time_created
)
related_memories: List[str] = self.memory_service.retrieve(message=new_message)
system_message = self.get_system_prompt(related_memories, time_created)
self.history_message_list.append(new_message)
self.history_message_list = self.history_message_list[-self.history_msg_count :]
all_messages = [system_message] + self.history_message_list
# TODO at xian zhe
return self.generation_model.call(messages=all_messages, stream=True)
def run(self):
self.memory_service.start_memory_backend()
while True:
query = input("wait for input:")
if query in ["stop", "停止"]:
break
self.chat_with_memory(query=query)

View file

@ -1,70 +0,0 @@
from constants.common_constants import RELATED_MEMORIES
from enumeration.memory_method_enum import MemoryMethodEnum
from scheme.message import Message
from utils.pipeline import Pipeline
from .base_memory_service import BaseMemoryService
class MemoryService(BaseMemoryService):
def __init__(
self,
chat_name: str,
retrieve_pipeline: str,
retrieve_all_pipeline: str,
summary_short_pipeline: str,
summary_long_pipeline: str,
summary_short_interval_time: int = 60,
summary_short_minimum_count: int = 5,
summary_long_interval_time: int = 60 * 5,
summary_long_minimum_count: int = 5 * 5,
**kwargs
):
super().__init__(**kwargs)
self.retrieve_pipeline = Pipeline(
chat_name=chat_name,
memory_method_type=MemoryMethodEnum.RETRIEVE,
pipeline_str=retrieve_pipeline,
)
self.retrieve_all_pipeline = Pipeline(
chat_name=chat_name,
memory_method_type=MemoryMethodEnum.RETRIEVE_ALL,
pipeline_str=retrieve_all_pipeline,
)
self.summary_short_pipeline = Pipeline(
chat_name=chat_name,
memory_method_type=MemoryMethodEnum.SUMMARY_SHORT,
pipeline_str=summary_short_pipeline,
loop_interval_time=summary_short_interval_time,
loop_minimum_count=summary_short_minimum_count,
)
self.summary_long_pipeline = Pipeline(
chat_name=chat_name,
memory_method_type=MemoryMethodEnum.SUMMARY_LONG,
pipeline_str=summary_long_pipeline,
loop_interval_time=summary_long_interval_time,
loop_minimum_count=summary_long_minimum_count,
)
def retrieve(self, message: Message):
self.retrieve_pipeline.submit_message(message, with_lock=False)
self.summary_short_pipeline.submit_message(message)
self.summary_long_pipeline.submit_message(message)
return self.retrieve_pipeline.run(RELATED_MEMORIES)
def retrieve_all(self):
return self.retrieve_all_pipeline.run(RELATED_MEMORIES)
def start_memory_backend(self):
self.summary_short_pipeline.start_loop_run()
self.summary_long_pipeline.start_loop_run()
def get_worker_list(self) -> list:
worker_set = set()
worker_set.update(self.retrieve_pipeline.worker_set)
worker_set.update(self.retrieve_all_pipeline.worker_set)
worker_set.update(self.summary_short_pipeline.worker_set)
worker_set.update(self.summary_long_pipeline.worker_set)
return sorted(worker_set)

View file

@ -1,135 +1,82 @@
import json
import os
from concurrent.futures import ThreadPoolExecutor
from typing import Dict, Any, List
import sys
import time
import fire
from datetime import datetime
from chat.global_context import GLOBAL_CONTEXT
from enumeration.language_enum import LanguageEnum
from enumeration.model_enum import ModelEnum
from utils.logger import Logger
from utils.tool_functions import (
complete_config_name,
init_instance_by_config,
under_line_to_hump,
)
from chat.memory_chat import MemoryChat
from enumeration.message_role_enum import MessageRoleEnum
from scheme.message import Message
from chat.base_memory_chat import BaseMemoryChat
from models.llama_index_generation_model import LlamaIndexGenerationModel
from models.llama_index_embedding_model import LlamaIndexEmbeddingModel
from models.llama_index_rerank_model import LlamaIndexRerankModel
sys.path.append(".") # noqa: E402
import json
from concurrent.futures import ThreadPoolExecutor
from typing import Dict, Any
import fire
import yaml
from memory_scope.chat.global_context import G_CONTEXT
from memory_scope.enumeration.language_enum import LanguageEnum
from memory_scope.utils.logger import Logger
from memory_scope.utils.tool_functions import init_instance_by_config
from memory_scope.utils.timer import timer
from memory_scope.enumeration.model_enum import ModelEnum
class CliJob(object):
def __init__(self, config_path: str):
self.config_path: str = config_path
self.config_base_dir: str = os.path.dirname(config_path)
def __init__(self):
self.config: Dict[str, Any] = {}
self.logger: Logger = Logger.get_logger("cli_job", to_stream=False)
self.worker_chat_dict: Dict[str, List[str]] = {}
self.logger: Logger = Logger.get_logger("memory_chat")
def load_config(self, path: str):
with open(path) as f:
if path.endswith("yaml"):
self.config = yaml.load(f, yaml.FullLoader)
elif path.endswith("json"):
self.config = json.load(f)
else:
raise RuntimeError("not supported config file type!")
def init_memory_chat(self):
for chat_name in GLOBAL_CONTEXT.global_configs["chat_list"]:
memory_chat_config = self.config[chat_name]
memory_chat: BaseMemoryChat = init_instance_by_config(
memory_chat_config, chat_name=chat_name
)
GLOBAL_CONTEXT.memory_chat_dict[chat_name] = memory_chat
for worker_name in memory_chat.memory_service.get_worker_list():
if worker_name not in self.worker_chat_dict:
self.worker_chat_dict[worker_name] = []
self.worker_chat_dict[worker_name].append(chat_name)
generation_model = memory_chat_config[ModelEnum.GENERATION_MODEL.value]
self.init_model(generation_model)
def init_model(self, model_name: str):
if not model_name or model_name in GLOBAL_CONTEXT.model_dict:
return
with open(
os.path.join(
self.config_base_dir, "model", complete_config_name(model_name)
)
) as f:
model_config = json.load(f)
GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config(model_config)
def init_workers(self):
"""load worker config & init workers"""
worker_config_name: str = self.config["workers"]
with open(
os.path.join(self.config_base_dir, complete_config_name(worker_config_name))
) as f:
worker_config_dict = json.load(f)
for worker_name, worker_config in worker_config_dict.items():
if worker_name not in self.worker_chat_dict:
continue
chat_name_list = self.worker_chat_dict[worker_name]
for chat_name in chat_name_list:
if chat_name not in GLOBAL_CONTEXT.worker_dict:
GLOBAL_CONTEXT.worker_dict[chat_name] = {}
GLOBAL_CONTEXT.worker_dict[chat_name][worker_name] = (
init_instance_by_config(
worker_config,
suffix_name="worker",
**GLOBAL_CONTEXT.global_configs,
)
)
self.init_model(worker_config.get(ModelEnum.EMBEDDING_MODEL.value))
self.init_model(worker_config.get(ModelEnum.GENERATION_MODEL.value))
self.init_model(worker_config.get(ModelEnum.RANK_MODEL.value))
@staticmethod
def set_global_config():
"""TODO set global_configs & set apikey into env"""
GLOBAL_CONTEXT.language = LanguageEnum(
GLOBAL_CONTEXT.global_configs["language"]
)
GLOBAL_CONTEXT.thread_pool = ThreadPoolExecutor(
max_workers=int(GLOBAL_CONTEXT.global_configs["max_workers"])
)
def set_global_config(self):
G_CONTEXT.global_config = global_config = self.config["global_config"]
G_CONTEXT.language = LanguageEnum(global_config["language"])
G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(global_config["max_workers"]))
@timer
def init_global_content_by_config(self):
with open(complete_config_name(self.config_path)) as f:
self.config = json.load(f)
GLOBAL_CONTEXT.global_configs = self.config["global_configs"]
# set global config
self.set_global_config()
self.init_memory_chat()
# init memory_chat
for name, conf in self.config["memory_chat"].items():
G_CONTEXT.memory_chat_dict[name] = init_instance_by_config(conf, name=name)
self.init_workers()
# set memory_service
for name, conf in self.config["memory_service"].items():
G_CONTEXT.memory_service_dict[name] = init_instance_by_config(conf, name=name)
## TODO no db and monitor now
# GLOBAL_CONTEXT.vector_store = init_instance_by_config(
# self.config["vector_store"]
# )
# GLOBAL_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
# init models
for name, conf in self.config["models"].items():
G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name)
@staticmethod
def run():
with GLOBAL_CONTEXT.thread_pool:
memory_chat = list(GLOBAL_CONTEXT.memory_chat_dict.values())[0]
# init vector_store
vector_store_config = self.config["vector_store"]
embedding_model = G_CONTEXT.model_dict[vector_store_config[ModelEnum.EMBEDDING_MODEL.value]]
G_CONTEXT.vector_store = init_instance_by_config(vector_store_config, embedding_model=embedding_model)
# init monitor
G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
# set worker config
G_CONTEXT.worker_config = self.config["worker"]
def run(self, config: str):
self.load_config(config)
self.init_global_content_by_config()
with G_CONTEXT.thread_pool:
memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0]
memory_chat.run()
def main(config_path: str):
job = CliJob(config_path=config_path)
job.init_global_content_by_config()
job.run()
G_CONTEXT.vector_store.close()
G_CONTEXT.monitor.close()
if __name__ == "__main__":
fire.Fire(main)
cli_job = CliJob()
fire.Fire(cli_job.run)

View file

@ -1,3 +1,9 @@
WORKFLOW_NAME = "workflow_name"
RESULT = "result"
CHAT_MESSAGES = "chat_messages"
RELATED_MEMORIES = "related_memories"
MESSAGES = "messages"
@ -16,10 +22,6 @@ RELATED_MEMORIES = "related_memories"
MODIFIED_MEMORIES = "modified_memories"
RESPONSE_EXT_INFO = "response_ext_info"
PROMPT_CONFIG = "prompt_config"
MESSAGES = "messages"
EXTRACT_TIME_DICT = "extract_time_dict"
@ -110,3 +112,5 @@ DATATIME_KEY_MAP = {
"": "week",
"星期几": "weekday",
}
CONTENT_MODIFIED = "content_modified"

View file

@ -0,0 +1,25 @@
from abc import ABCMeta, abstractmethod
from typing import Literal
OPERATION_TYPE = Literal["frontend", "backend"]
class BaseOperation(metaclass=ABCMeta):
operation_type: OPERATION_TYPE = "frontend"
def __init__(self, name: str, description: str = "", **kwargs):
self.name: str = name
self.description: str = description
def init_workflow(self):
pass
@abstractmethod
def run_operation(self):
raise NotImplementedError
def run_operation_backend(self):
pass
def stop_operation_backend(self):
pass

View file

@ -0,0 +1,124 @@
import re
import threading
from concurrent.futures import ThreadPoolExecutor, as_completed
from itertools import zip_longest
from typing import Dict, Any, List
from memory_scope.chat.global_context import G_CONTEXT
from memory_scope.constants.common_constants import WORKFLOW_NAME
from memory_scope.memory.worker.base_worker import BaseWorker
from memory_scope.utils.logger import Logger
from memory_scope.utils.timer import Timer
from memory_scope.utils.tool_functions import init_instance_by_config
class BaseWorkflow(object):
def __init__(self,
name: str,
workflow: str,
thread_pool: ThreadPoolExecutor = G_CONTEXT.thread_pool,
**kwargs):
self.name: str = name
self.workflow: str = workflow
self.thread_pool: ThreadPoolExecutor = thread_pool
self.kwargs = kwargs
self.workflow_worker_list: List[List[List[str]]] = []
self.worker_dict: Dict[str, BaseWorker | bool] = {}
self.context: Dict[str, Any] = {}
self.context_lock = threading.Lock()
self.logger: Logger = Logger.get_logger()
if self.workflow:
self._parse_workflow()
self._print_workflow()
def _parse_workflow(self):
# re-match e.g., [a|b],c,[d,e,f|g,h],j
pattern = r'(\[[^\]]*\]|[^,]+)'
workflow_split = re.findall(pattern, self.workflow)
for workflow_part in workflow_split:
# e.g., [d,e,f|g,h]
workflow_part = workflow_part.strip()
if '[' in workflow_part or ']' in workflow_part:
workflow_part = workflow_part.replace('[', '').replace(']', '')
# e.g., ["d,e,f", "g,h"]
line_split = [x.strip() for x in workflow_part.split("|") if x]
if len(line_split) <= 0:
continue
# is under multi thread cond
is_multi_thread: bool = len(line_split) > 1
# e.g., ["d","e","f"]
line_split_split: List[List[str]] = []
for sub_line_split in line_split:
sub_split = [x.strip() for x in sub_line_split.split(",")]
line_split_split.append(sub_split)
# add workers
for sub_item in sub_split:
self.worker_dict[sub_item] = is_multi_thread
self.workflow_worker_list.append(line_split_split)
def _print_workflow(self):
self.logger.info(f"----- print_workflow_{self.name}_begin -----")
i: int = 0
for workflow_part in self.workflow_worker_list:
if len(workflow_part) == 1:
for w in workflow_part[0]:
self.logger.info(f"stage{i}: {w}")
i += 1
else:
for w_zip in zip_longest(*workflow_part, fillvalue="-"):
self.logger.info(f"stage{i}: {' | '.join(w_zip)}")
i += 1
for w in w_zip:
if w == "-":
continue
self.logger.info(f"----- print_workflow_{self.name}_end -----")
def init_workers(self):
for name in list(self.worker_dict.keys()):
if name not in G_CONTEXT.worker_config:
raise RuntimeError(f"worker={name} is not exists in worker_config!")
self.worker_dict[name] = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=self.worker_dict[name],
context=self.context,
context_lock=self.context_lock)
def _run_sub_workflow(self, worker_list: List[str]) -> bool:
for name in worker_list:
worker = self.worker_dict[name]
worker.run()
if not worker.continue_run:
return False
return True
def run_workflow(self):
with Timer(f"run_workflow_{self.name}"):
self.context[WORKFLOW_NAME] = self.name
for workflow_part in self.workflow_worker_list:
if len(workflow_part) == 1:
if not self._run_sub_workflow(workflow_part[0]):
break
else:
t_list = []
for sub_workflow in workflow_part:
t_list.append(G_CONTEXT.thread_pool.submit(
self._run_sub_workflow, sub_workflow))
flag = True
for future in as_completed(t_list):
if not future.result():
flag = False
break
if not flag:
break

View file

@ -0,0 +1,34 @@
from typing import List
from memory_scope.constants.common_constants import RESULT, CHAT_MESSAGES
from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE
from memory_scope.memory.operation.base_workflow import BaseWorkflow
from memory_scope.scheme.message import Message
class ReadMemory(BaseWorkflow, BaseOperation):
operation_type: OPERATION_TYPE = "frontend"
def __init__(self,
name: str,
description: str,
chat_messages: List[Message],
his_msg_count: int = 0, # supplement to the current query
contextual_msg_count: int = 0, # for the current context dialogue
**kwargs):
super().__init__(name=name, **kwargs)
BaseOperation.__init__(self, name=name, description=description)
self.chat_messages: List[Message] = chat_messages
self.his_msg_count: int = his_msg_count
self.contextual_msg_count: int = contextual_msg_count
def init_workflow(self):
self.init_workers()
def run_operation(self):
max_count = 1 + max(self.his_msg_count, self.contextual_msg_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()
return result

View file

@ -0,0 +1,56 @@
import time
from memory_scope.chat.global_context import G_CONTEXT
from memory_scope.constants.common_constants import RESULT
from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE
from memory_scope.memory.operation.base_workflow import BaseWorkflow
class SummaryMemory(BaseWorkflow, BaseOperation):
operation_type: OPERATION_TYPE = "backend"
def __init__(self,
name: str,
description: str,
interval_time: int = 300,
**kwargs):
super().__init__(name=name, **kwargs)
BaseOperation.__init__(self, name=name, description=description)
self.interval_time: int = interval_time
self._operation_status_run: bool = False
self._loop_switch: bool = False
self._run_thread = None
def init_workflow(self):
self.init_workers()
def run_operation(self):
if self._operation_status_run:
return
self._operation_status_run = True
self.run_workflow()
result = self.context.get(RESULT)
self.context.clear()
self._operation_status_run = False
return result
def _loop_operation(self):
while self._loop_switch:
for _ in range(self.interval_time):
if self._loop_switch:
time.sleep(1)
else:
break
if self._loop_switch:
self.run_operation()
def run_operation_backend(self):
if not self._loop_switch:
self._loop_switch = True
self._run_thread = G_CONTEXT.thread_pool.submit(self._loop_operation)
def stop_operation_backend(self):
self._loop_switch = False

View file

@ -0,0 +1,83 @@
import time
from typing import List
from memory_scope.chat.global_context import G_CONTEXT
from memory_scope.constants.common_constants import CHAT_MESSAGES, RESULT
from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE
from memory_scope.memory.operation.base_workflow import BaseWorkflow
from memory_scope.scheme.message import Message
class WriteMemory(BaseWorkflow, BaseOperation):
operation_type: OPERATION_TYPE = "backend"
def __init__(self,
name: str,
description: str,
chat_messages: List[Message],
his_msg_count: int = 0,
message_lock=None,
interval_time: int = 60,
contextual_msg_count: int = 6,
**kwargs):
super().__init__(name=name, **kwargs)
BaseOperation.__init__(self, name=name, description=description)
self.chat_messages: List[Message] = chat_messages
self.his_msg_count: int = his_msg_count
self.message_lock = message_lock
self.interval_time: int = interval_time
self.contextual_msg_count: int = contextual_msg_count
self._operation_status_run: bool = False
self._loop_switch: bool = False
@property
def not_memorized_size(self):
return sum([not x.memorized for x in self.chat_messages])
def set_memorized(self):
if self.message_lock:
with self.message_lock:
for msg in self.chat_messages:
msg.memorized = True
def init_workflow(self):
self.init_workers()
def run_operation(self):
if self._operation_status_run:
return
self._operation_status_run = True
not_memorized_size = self.not_memorized_size
if not_memorized_size < self.contextual_msg_count:
return
max_count = not_memorized_size + self.his_msg_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()
self.set_memorized()
self._operation_status_run = False
return result
def _loop_operation(self):
while self._loop_switch:
for _ in range(self.interval_time):
if self._loop_switch:
time.sleep(1)
else:
break
if self._loop_switch:
self.run_operation()
def run_operation_backend(self):
if not self._loop_switch:
self._loop_switch = True
return G_CONTEXT.thread_pool.submit(self._loop_operation)
def stop_operation_backend(self):
self._loop_switch = False

View file

@ -0,0 +1,52 @@
import threading
from abc import ABCMeta, abstractmethod
from typing import List, Dict
from memory_scope.memory.operation.base_operation import BaseOperation
from memory_scope.scheme.message import Message
from memory_scope.utils.logger import Logger
class BaseMemoryService(metaclass=ABCMeta):
def __init__(self,
memory_operations: Dict[str, dict],
read_memory_key: str = "read_memory",
**kwargs):
self.memory_operations: Dict[str, dict] = memory_operations
self.read_memory_key: str = read_memory_key
self._operation_dict: Dict[str, BaseOperation] = {}
self._op_description_dict: Dict[str, str] = {}
self.chat_messages: List[Message] = []
self.message_lock = threading.Lock
self.logger = Logger.get_logger()
self.kwargs = kwargs
@abstractmethod
def _init_operation(self, memory_operations: Dict[str, dict]):
raise NotImplementedError
@abstractmethod
def add_messages(self, messages: List[Message] | Message):
raise NotImplementedError
def start_service(self):
pass
@abstractmethod
def do_operation(self, op_name: str):
raise NotImplementedError
@property
def op_description_dict(self) -> Dict[str, str]:
if not self._op_description_dict:
self._op_description_dict = {k: v.description for k, v in self._operation_dict.items()}
return self._op_description_dict
def read_memory(self):
assert self.read_memory_key in self._operation_dict, f"op={self.read_memory_key} is not inited!"
return self.do_operation(self.read_memory_key)
def stop_service(self):
pass

View file

@ -0,0 +1,55 @@
from typing import List, Dict
from memory_scope.memory.service.base_memory_service import BaseMemoryService
from memory_scope.scheme.message import Message
from memory_scope.utils.tool_functions import init_instance_by_config
class ChatMemoryService(BaseMemoryService):
def __init__(self, history_msg_count: int = 32, contextual_msg_count: int = 6, **kwargs):
super().__init__(**kwargs)
self.history_msg_count: int = history_msg_count
self.contextual_msg_count: int = contextual_msg_count
assert self.history_msg_count >= self.contextual_msg_count
self._init_operation(self.memory_operations)
def _init_operation(self, memory_operations: Dict[str, dict]):
for name, operation_config in memory_operations.items():
if name in self._operation_dict:
self.logger.warning(f"memory operation={name} is repeated!")
continue
self._operation_dict[name] = init_instance_by_config(
config=operation_config,
name=name,
chat_messages=self.chat_messages,
message_lock=self.message_lock,
contextual_msg_count=self.contextual_msg_count)
def add_messages(self, messages: List[Message] | Message):
if isinstance(messages, Message):
messages = [messages]
messages = sorted(messages, key=lambda x: x.time_created)
self.chat_messages.extend(messages)
if len(self.chat_messages) > self.history_msg_count:
gap_size = len(self.chat_messages) - self.history_msg_count
for _ in range(gap_size):
self.chat_messages.pop(0)
def start_service(self):
for _, operation in self._operation_dict.items():
operation.init_workflow()
if operation.operation_type == "backend":
operation.run_operation_backend()
def do_operation(self, op_name: str):
if op_name not in self._operation_dict:
self.logger.warning(f"op_name={op_name} is not inited!")
return
return self._operation_dict[op_name].run_operation()
def stop_service(self):
for _, operation in self._operation_dict.items():
if operation.operation_type == "backend":
operation.stop_operation_backend()

View file

@ -0,0 +1,56 @@
from abc import ABCMeta, abstractmethod
from typing import Any, Dict
from memory_scope.utils.logger import Logger
from memory_scope.utils.timer import Timer
class BaseWorker(metaclass=ABCMeta):
def __init__(self,
name: str,
context: Dict[str, Any],
context_lock=None,
raise_exception: bool = True,
is_multi_thread: bool = False,
**kwargs):
self.name: str = name
self.context: Dict[str, Any] = context
self.context_lock = context_lock
self.raise_exception: bool = raise_exception
self.is_multi_thread: bool = is_multi_thread
self.kwargs: dict = kwargs
self.continue_run: bool = True
self.logger: Logger = Logger.get_logger()
@abstractmethod
def _run(self):
raise NotImplementedError
def run(self):
self.logger.info(f"----- worker_{self.name}_begin -----")
with Timer(self.name, log_time=False) as t:
if self.raise_exception:
self._run()
else:
try:
self._run()
except Exception as e:
self.logger.exception(f"run {self.name} failed! args={e.args}")
self.logger.info(f"----- worker_{self.name}_end cost={t.cost_str}-----")
def get_context(self, key: str, default=None):
return self.context.get(key, default)
def set_context(self, key: str, value: Any):
if self.is_multi_thread:
with self.context_lock:
self.context_dict[key] = value
else:
self.context[key] = value
def __getattr__(self, key: str):
return self.kwargs[key]

View file

@ -0,0 +1,12 @@
import datetime
from memory_scope.constants.common_constants import RESULT, WORKFLOW_NAME
from memory_scope.memory.worker.base_worker import BaseWorker
class DummyWorker(BaseWorker):
def _run(self):
workflow_name = self.get_context(WORKFLOW_NAME)
self.logger.info(f"enter workflow={workflow_name}.dummy_worker!")
ts = int(datetime.datetime.now().timestamp())
self.set_context(RESULT, f"test {workflow_name} \nts={ts}")

View file

@ -0,0 +1,82 @@
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
class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
def __init__(self,
embedding_model: str = "",
generation_model: str = "",
rank_model: str = "",
**kwargs):
super(MemoryBaseWorker, self).__init__(**kwargs)
self._embedding_model: BaseModel | str = embedding_model
self._generation_model: BaseModel | str = generation_model
self._rank_model: BaseModel | str = rank_model
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(CHAT_MESSAGES)
@messages.setter
def messages(self, value):
self.set_context(CHAT_MESSAGES, value)
@property
def embedding_model(self) -> BaseModel:
if isinstance(self._embedding_model, str):
self._embedding_model = G_CONTEXT.model_dict[self._embedding_model]
return self._embedding_model
@property
def generation_model(self) -> BaseModel:
if isinstance(self._generation_model, str):
self._generation_model = G_CONTEXT.model_dict[self._generation_model]
return self._generation_model
@property
def rank_model(self) -> BaseModel:
if isinstance(self._rank_model, str):
self._rank_model = G_CONTEXT.model_dict[self._rank_model]
return self._rank_model
@property
def vector_store(self) -> BaseVectorStore:
if self._vector_store is None:
self._vector_store = G_CONTEXT.vector_store
return self._vector_store
@property
def monitor(self):
if self._monitor is None:
self._monitor = G_CONTEXT.monitor
return self._monitor
@property
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]
@staticmethod
def get_prompt(prompt: dict) -> str:
return prompt[G_CONTEXT.global_configs["language"]]

View file

@ -1,4 +1 @@
from utils.registry import Registry
# __all__ = ["LlamaIndexEmbeddingModel", "LlamaIndexGenerationModel", "LlamaIndexRerankModel"]
MODEL_REGISTRY = Registry("models")

View file

@ -1,12 +1,15 @@
import inspect
import time
from abc import abstractmethod, ABCMeta
from typing import Any
from enumeration.model_enum import ModelEnum
from . import MODEL_REGISTRY
from .response import ModelResponse, ModelResponseGen
from utils.logger import Logger
from utils.timer import Timer
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.scheme.model_response import ModelResponse, ModelResponseGen
from memory_scope.utils.logger import Logger
from memory_scope.utils.registry import Registry
from memory_scope.utils.timer import Timer
MODEL_REGISTRY = Registry("models")
class BaseModel(metaclass=ABCMeta):
@ -14,7 +17,7 @@ class BaseModel(metaclass=ABCMeta):
def __init__(self,
model_name: str,
method_type: str,
module_name: str,
timeout: int = None,
max_retries: int = 3,
retry_interval: float = 1.0,
@ -22,24 +25,32 @@ class BaseModel(metaclass=ABCMeta):
**kwargs):
self.model_name: str = model_name
self.method_type: str = method_type
self.module_name: str = module_name
self.timeout: int = timeout
self.max_retries: int = max_retries
self.retry_interval: float = retry_interval
self.kwargs_filter: bool = kwargs_filter
self.kwargs: dict = kwargs
self.data = {}
self._model: Any = None
self.logger = Logger.get_logger()
obj_cls = MODEL_REGISTRY.get(self.method_type)
if not obj_cls:
raise RuntimeError(f"method_type={self.method_type} is not supported!")
@property
def model(self):
if self._model is None:
if self.module_name not in MODEL_REGISTRY.module_dict:
raise RuntimeError(f"method_type={self.module_name} is not supported!")
obj_cls = MODEL_REGISTRY[self.module_name]
if kwargs_filter:
allowed_kwargs = list(inspect.signature(obj_cls.__init__).parameters.keys())
kwargs = {key: value for key, value in kwargs.items() if key in allowed_kwargs}
self.model = obj_cls(**kwargs)
if self.kwargs_filter:
allowed_kwargs = list(inspect.signature(obj_cls.__init__).parameters.keys())
kwargs = {key: value for key, value in self.kwargs.items() if key in allowed_kwargs}
else:
kwargs = self.kwargs
self._model = obj_cls(**kwargs)
return self._model
@abstractmethod
def before_call(self, **kwargs) -> None:
@ -70,8 +81,8 @@ class BaseModel(metaclass=ABCMeta):
:param kwargs:
:return:
"""
self.before_call(stream=stream, **kwargs)
with Timer(self.__class__.__name__, log_time=False) as t:
self.before_call(stream=stream, **kwargs)
for i in range(self.max_retries):
try:
model_response = self._call(stream=stream, **kwargs)
@ -97,8 +108,8 @@ class BaseModel(metaclass=ABCMeta):
:param kwargs:
:return:
"""
self.before_call(**kwargs)
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)

View file

@ -2,20 +2,15 @@ from typing import List
from llama_index.embeddings.dashscope import DashScopeEmbedding
from models import MODEL_REGISTRY
from models.base_model import BaseModel
from models.response import ModelResponse, ModelResponseGen
from enumeration.model_enum import ModelEnum
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY
from memory_scope.scheme.model_response import ModelResponse
class LlamaIndexEmbeddingModel(BaseModel):
m_type: ModelEnum = ModelEnum.EMBEDDING_MODEL
MODEL_REGISTRY.batch_register(
[
DashScopeEmbedding,
]
)
MODEL_REGISTRY.register("dashscope_embedding", DashScopeEmbedding)
def before_call(self, **kwargs):
text: str | List[str] = kwargs.pop("text", "")
@ -42,16 +37,11 @@ class LlamaIndexEmbeddingModel(BaseModel):
:param kwargs:
:return:
"""
return ModelResponse(
m_type=self.m_type, raw=self.model.get_text_embedding_batch(**self.data)
)
return ModelResponse(m_type=self.m_type, raw=self.model.get_text_embedding_batch(**self.data))
async def _async_call(self, **kwargs) -> ModelResponse:
"""
:param kwargs:
:return:
"""
return ModelResponse(
m_type=self.m_type,
raw=await self.model.aget_text_embedding_batch(**self.data),
)
return ModelResponse(m_type=self.m_type, raw=await self.model.aget_text_embedding_batch(**self.data))

View file

@ -1,73 +1,62 @@
from typing import List, Dict
from llama_index.core.base.llms.types import (
ChatMessage,
ChatResponse,
CompletionResponse,
)
import datetime
from typing import List
from llama_index.core.base.llms.types import ChatMessage, ChatResponse, CompletionResponse
from llama_index.llms.dashscope import DashScope
from enumeration.model_enum import ModelEnum
from . import MODEL_REGISTRY
from .base_model import BaseModel
from .response import ModelResponse, ModelResponseGen
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY
from memory_scope.scheme.model_response import ModelResponse, ModelResponseGen
from memory_scope.scheme.message import Message
class LlamaIndexGenerationModel(BaseModel):
m_type: ModelEnum = ModelEnum.GENERATION_MODEL
# TODO rename module name at xianzhe
MODEL_REGISTRY.batch_register([
DashScope,
])
MODEL_REGISTRY.register("dashscope_generation", DashScope)
def before_call(self, **kwargs) -> None:
def before_call(self, **kwargs):
prompt: str = kwargs.pop("prompt", "")
messages: List[Dict[str, str]] = kwargs.pop("messages", [])
messages: List[Message] = kwargs.pop("messages", [])
if prompt:
input_text = prompt
input_type = 'prompt'
llama_input = input_text
self.data = {"prompt": prompt}
elif messages:
input_text = messages
input_type = 'messages'
llama_input = [ChatMessage(role=x.role, content=x.content) for x in input_text]
self.data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]}
else:
raise RuntimeError("prompt and messages is both empty!")
self.data = {input_type: llama_input}
def after_call(self,
model_response: ModelResponse,
stream: bool = False,
**kwargs) -> ModelResponse | ModelResponseGen:
model_response.message = Message(role=MessageRoleEnum.ASSISTANT, content="")
call_result = model_response.raw
if stream:
def gen() -> ModelResponseGen:
text = ""
for response in call_result:
delta = response.delta
text += delta
model_response.text = text
model_response.delta = delta
model_response.message.content += response.delta
model_response.delta = response.delta
yield model_response
return gen()
else:
if isinstance(call_result, CompletionResponse):
content = call_result.text
model_response.message.content = call_result.text
elif isinstance(call_result, ChatResponse):
content = call_result.message.content
model_response.message.content = call_result.message.content
else:
raise NotImplementedError
model_response.text = content
return model_response
def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen:
assert "prompt" in self.data or "messages" in self.data
results = ModelResponse(m_type=self.m_type)
if 'prompt' in self.data:
if "prompt" in self.data:
if stream:
response = self.model.stream_complete(**self.data)
else:

View file

@ -4,24 +4,21 @@ from llama_index.core.data_structs import Node
from llama_index.core.schema import NodeWithScore
from llama_index.postprocessor.dashscope_rerank import DashScopeRerank
from models import MODEL_REGISTRY
from models.base_model import BaseModel
from models.response import ModelResponse, ModelResponseGen
from enumeration.model_enum import ModelEnum
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY
from memory_scope.scheme.model_response import ModelResponse
class LlamaIndexRerankModel(BaseModel):
class LlamaIndexRankModel(BaseModel):
m_type: ModelEnum = ModelEnum.RANK_MODEL
MODEL_REGISTRY.batch_register([
DashScopeRerank
])
MODEL_REGISTRY.register("dashscope_rank", DashScopeRerank)
def before_call(self, **kwargs) -> None:
assert "query" in kwargs or "documents" in kwargs
query: str = kwargs.pop("query", "")
documents: List[str] = kwargs.pop("documents", [])
if isinstance(documents, str):
documents = [documents]
assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}"
@ -29,10 +26,7 @@ class LlamaIndexRerankModel(BaseModel):
nodes = [NodeWithScore(node=Node(text=doc), score=-1.0) for doc in documents]
self._get_documents_mapping(documents)
self.data = {
"nodes": nodes,
"query_str": query,
}
self.data = {"nodes": nodes, "query_str": query}
def after_call(self, model_response: ModelResponse, **kwargs) -> ModelResponse:
if not model_response.rank_scores:

View file

@ -0,0 +1,8 @@
from ..enumeration.language_enum import LanguageEnum
GET_INSIGHT_SYSTEM_PROMPT = {}
GET_INSIGHT_FEW_SHOT_PROMPT = {}
GET_INSIGHT_USER_QUERY_PROMPT = {}

View file

@ -0,0 +1,8 @@
from ..enumeration.language_enum import LanguageEnum
GET_REFLECTION_SYSTEM_PROMPT = {}
GET_REFLECTION_FEW_SHOT_PROMPT = {}
GET_REFLECTION_USER_QUERY_PROMPT = {}

View file

@ -0,0 +1,8 @@
from ..enumeration.language_enum import LanguageEnum
LONG_CONTRA_REPEAT_SYSTEM_PROMPT = {}
LONG_CONTRA_REPEAT_FEW_SHOT_PROMPT = {}
LONG_CONTRA_REPEAT_USER_QUERY_PROMPT = {}

View file

@ -1,4 +1,4 @@
from enumeration.language_enum import LanguageEnum
from ..enumeration.language_enum import LanguageEnum
SYSTEM_PROMPT = {
LanguageEnum.CN: """

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,3 +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,9 +1,19 @@
import datetime
from typing import Dict
from pydantic import Field, BaseModel
class Message(BaseModel):
role: str = Field(..., description="The role of the message sender (user, assistant, system)")
role_name: str = Field("", description="role name")
content: str = Field(..., description="The body of the message")
time_created: int = Field("", description="Timestamp when the message was created")
time_created: int = Field(int(datetime.datetime.now().timestamp()),
description="Timestamp when the message was created")
memorized: bool = Field(False, description="indicate whether message is memorized")
meta_data: Dict[str, str] = Field({}, description="meta data for msg")

View file

@ -3,11 +3,12 @@ from typing import Generator, List, Dict, Any
from pydantic import BaseModel, Field
from enumeration.model_enum import ModelEnum
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.scheme.message import Message
class ModelResponse(BaseModel):
text: str = Field("", description="generation model result")
message: Message | None = Field(None, description="generation model result")
delta: str = Field("", description="New text that just streamed in (only used when streaming)")
@ -27,12 +28,7 @@ class ModelResponse(BaseModel):
def __str__(self, max_size=100, **kwargs):
result = {}
try:
all_dict = self.model_dump()
except Exception:
all_dict = self.dict()
for key, value in all_dict.items():
for key, value in self.model_dump().items():
if key == "raw" or not value:
continue

View file

@ -18,8 +18,8 @@ class BaseMonitor(metaclass=ABCMeta):
:return:
"""
@abstractmethod
def flush(self):
"""
:return:
"""
pass
def close(self):
pass

View file

@ -1,61 +1,34 @@
from abc import ABCMeta, abstractmethod
from typing import Dict, List
from models.base_model import BaseModel
from memory_scope.scheme.memory_node import MemoryNode
class BaseVectorStore(metaclass=ABCMeta):
def __init__(self, index_name: str, embedding_model: BaseModel, content_key: str = "text", **kwargs):
self.index_name: str = index_name
self.embedding_model: BaseModel = embedding_model
self.content_key: str = content_key
self.kwargs: dict = kwargs
@abstractmethod
def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]):
"""
:param text:
:param limit_size:
:param filter_dict:
:return:
"""
pass
@abstractmethod
async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]):
"""
:param text:
:param limit_size:
:param filter_dict:
:return:
"""
pass
@abstractmethod
def insert(self, node: MemoryNode):
""" TODO 是否overwrite
:return:
"""
pass
@abstractmethod
def insert_batch(self):
"""
:return:
"""
def insert_batch(self, nodes: List[MemoryNode]):
pass
@abstractmethod
def delete(self):
"""
:return:
"""
def delete(self, node: MemoryNode):
pass
def update(self, node: MemoryNode):
pass
@abstractmethod
def flush(self):
"""
:return:
"""
pass
def close(self):
pass

View file

@ -0,0 +1,12 @@
from memory_scope.storage.base_monitor import BaseMonitor
class DummyMonitor(BaseMonitor):
def add(self):
pass
def add_token(self):
pass
def close(self):
pass

View file

@ -1,15 +1,28 @@
from typing import Dict, List, Any
from llama_index.core.schema import TextNode
from llama_index.core.vector_stores import VectorStoreQuery
from llama_index.core import VectorStoreIndex, StorageContext, ServiceContext
from llama_index.core import VectorStoreIndex
from llama_index.core.schema import TextNode, NodeWithScore
from llama_index.vector_stores.elasticsearch import ElasticsearchStore, AsyncDenseVectorStrategy
from llama_index.core.vector_stores.types import MetadataFilters, ExactMatchFilter, VectorStoreQueryMode
from memory_scope.models.base_model import BaseModel
from memory_scope.storage.base_vector_store import BaseVectorStore
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.storage.base_vector_store import BaseVectorStore
class _ElasticsearchStore(ElasticsearchStore):
async def adelete(self, ref_doc_id: str, **delete_kwargs: Any) -> None:
"""
Async delete node from Elasticsearch index.
Args:
ref_doc_id: ID of the node to delete.
delete_kwargs: Optional. Additional arguments to
pass to AsyncElasticsearch delete_by_query.
Raises:
Exception: If AsyncElasticsearch delete_by_query fails.
"""
return await self._store.delete(query={"term": {"_id": ref_doc_id}}, **delete_kwargs)
def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str, Any]:
@ -24,7 +37,7 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str
"""
result = {
"bool" : {}
"bool": {}
}
for key, value in standard_filters.items():
if isinstance(value, list):
@ -32,10 +45,10 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str
for v in value:
operands.append(
{
"term":
{
f"metadata.{key}.keyword": {"value": v}
}
"term":
{
f"metadata.{key}.keyword": {"value": v}
}
}
)
result['bool'].update({"should": operands})
@ -56,67 +69,69 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str
class LlamaIndexElasticSearchStore(BaseVectorStore):
def __init__(self,
index_name: str,
embedding_model: BaseModel,
content_key: str = "text",
def __init__(self,
embedding_model: BaseModel,
index_name: str,
es_url: str,
use_hybrid: bool = True,
**kwargs):
self.index_name: str = index_name
self.embedding_model: BaseModel = embedding_model
self.es_store = ElasticsearchStore(index_name=self.index_name,
retrieval_strategy=AsyncDenseVectorStrategy(hybrid=True),
**kwargs)
self.service_context = ServiceContext.from_defaults(embed_model=self.embedding_model, llm=None)
self.es_store = _ElasticsearchStore(index_name=index_name,
es_url=es_url,
retrieval_strategy=AsyncDenseVectorStrategy(hybrid=use_hybrid),
**kwargs)
self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store,
service_context=self.service_context)
def retrieve(self, query: str, top_k: int = 3, filter_dict: Dict[str, List[str]] = {}) -> MemoryNode:
filter = _to_elasticsearch_filter(filter_dict)
embed_model=self.embedding_model.model)
def retrieve(self,
query: str,
top_k: int,
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
if filter_dict is None:
filter_dict = {}
es_filter = _to_elasticsearch_filter(filter_dict)
retriever = self.index.as_retriever(
vector_store_kwargs={
"es_filter": filter
},
similarity_top_k=top_k
)
textnodes = retriever.retrieve(query)
results = self._textnodes2memorynodes(textnodes)
vector_store_kwargs={"es_filter": es_filter},
similarity_top_k=top_k)
text_nodes = retriever.retrieve(query)
return [self._text_node_2_memory_node(n) for n in text_nodes]
async def async_retrieve(self,
query: str,
top_k: int,
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
if filter_dict is None:
filter_dict = {}
es_filter = _to_elasticsearch_filter(filter_dict)
retriever = self.index.as_retriever(
vector_store_kwargs={"es_filter": es_filter},
similarity_top_k=top_k)
text_nodes: List[NodeWithScore] = await retriever.aretrieve(query)
return [self._text_node_2_memory_node(n) for n in text_nodes]
return results
async def async_retrieve(self, query: str, top_k: int = 3, filter_dict: Dict[str, List[str]] = {}) -> MemoryNode:
raise NotImplementedError
## return await super().async_retrieve(text, limit_size, filter_dict)
def insert(self, node: MemoryNode):
node = self._memorynode2textnode(node)
self.index.insert_nodes([node])
self.index.insert_nodes([self._memory_node_2_text_node(node)])
def insert_batch(self, node: MemoryNode) -> None:
raise NotImplementedError
def delete(self, node: MemoryNode):
memory_id = node.memory_id
return self.es_store.delete(memory_id)
def delete(self):
raise NotImplementedError
def flush(self):
raise NotImplementedError
def _memorynode2textnode(self, memory_node: MemoryNode) -> TextNode:
content = memory_node.content
meta = memory_node.model_dump(exclude={"content"})
return TextNode(text=content, metadata=meta)
def update(self, node: MemoryNode):
self.delete(node)
self.insert(node)
def _textnode2memorynode(self, text_node: TextNode) -> MemoryNode:
content = text_node.text
meta = text_node.metadata
mem_node = MemoryNode(content=content, **meta)
return mem_node
def close(self):
self.es_store.close()
def _textnodes2memorynodes(self, text_nodes: TextNode) -> MemoryNode:
mem_nodes = [self._textnode2memorynode(node) for node in text_nodes]
return mem_nodes
@staticmethod
def _memory_node_2_text_node(memory_node: MemoryNode) -> TextNode:
return TextNode(id_=memory_node.memory_id,
text=memory_node.content,
metadata=memory_node.model_dump(exclude={"content"}))
@staticmethod
def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode:
return MemoryNode(content=text_node.text, **text_node.metadata)

View file

@ -1,185 +0,0 @@
import re
import threading
import time
from concurrent.futures import as_completed
from itertools import zip_longest
from typing import Dict, Any, List
from chat.global_context import GLOBAL_CONTEXT
from constants.common_constants import MESSAGES, CHAT_NAME
from enumeration.memory_method_enum import MemoryMethodEnum
from scheme.message import Message
from utils.logger import Logger
from utils.timer import Timer
from worker.base_worker import BaseWorker
class Pipeline(object):
def __init__(self,
chat_name: str,
memory_method_type: MemoryMethodEnum,
pipeline_str: str,
history_msg_count: int = 3,
loop_interval_time: int = 300,
loop_minimum_count: int = 20):
self.chat_name: str = chat_name
self.memory_method_type: MemoryMethodEnum = memory_method_type
self.pipeline_str: str = pipeline_str
self.history_msg_count: int = history_msg_count
self.loop_interval_time: int = loop_interval_time
self.loop_minimum_count: int = loop_minimum_count
# pipeline上下文和锁
self.context: Dict[str, Any] = {}
self.context_lock = threading.Lock()
# pipeline run config
self.loop_switch: bool = False
self.pipeline_list: list[list] = []
self.worker_set: set[str] = set()
self.worker_dict: Dict[str, BaseWorker] = {}
self.injected: bool = False
# message list
self.history_message_list: List[Message] = []
self.current_message_list: List[Message] = []
self.message_lock = threading.Lock()
# 日志
self.logger: Logger = Logger.get_logger()
self._parse_pipeline()
def _parse_pipeline(self):
if not self.pipeline_str:
return
# re-match e.g., [a|b],c,[d,e,f|g,h],j
pattern = r'(\[[^\]]*\]|[^,]+)'
pipeline_split = re.findall(pattern, self.pipeline_str)
self.pipeline_list = []
for pipeline_part in pipeline_split:
# e.g., [d,e,f|g,h]
pipeline_part = pipeline_part.strip()
if '[' in pipeline_part or ']' in pipeline_part:
pipeline_part = pipeline_part.replace('[', '').replace(']', '')
# e.g., ["d,e,f", "g,h"]
line_split = [x.strip() for x in pipeline_part.split("|") if x]
if len(line_split) <= 0:
continue
# e.g., ["d","e","f"]
line_split_split = []
for sub_line_split in line_split:
sub_split = [x.strip() for x in sub_line_split.split(",")]
line_split_split.append(sub_split)
# add to workers
self.worker_set.update(sub_split)
self.pipeline_list.append(line_split_split)
def _visit_and_inject_workers(self):
if self.injected:
return
self.worker_dict = GLOBAL_CONTEXT.worker_dict[self.chat_name]
self.logger.info(f"----- {self.chat_name} {self.memory_method_type.value} Pipeline Begin -----")
i: int = 0
for pipeline_part in self.pipeline_list:
if len(pipeline_part) == 1:
for w in pipeline_part[0]:
self.logger.info(f"stage{i}: {w}")
i += 1
if w not in self.worker_dict:
raise RuntimeError(f"worker={w} is not inited.")
# 注入context
self.worker_dict[w].set_context_dict(self.context)
else:
for w_zip in zip_longest(*pipeline_part, fillvalue="-"):
self.logger.info(f"stage{i}: {' | '.join(w_zip)}")
i += 1
for w in w_zip:
if w == "-":
continue
if w not in self.worker_dict:
raise RuntimeError(f"worker={w} is not inited.")
# 注入context & lock
self.worker_dict[w].set_context_dict(self.context, self.context_lock)
self.logger.info(f"----- {self.chat_name} {self.memory_method_type.value} Pipeline End -----")
self.injected = True
def _worker_run(self, worker_list: list[str]) -> bool:
for worker_name in worker_list:
worker = self.worker_dict[worker_name]
worker.run()
if not worker.continue_run:
return False
return True
def _run(self):
self._visit_and_inject_workers()
with Timer(f"pipeline_{self.chat_name}_{self.memory_method_type.value}"):
self.context[MESSAGES] = self.history_message_list + self.current_message_list
self.context[CHAT_NAME] = self.chat_name
for pipeline_part in self.pipeline_list:
if len(pipeline_part) == 1:
if not self._worker_run(pipeline_part[0]):
break
else:
t_list = []
for worker_list in pipeline_part:
t_list.append(GLOBAL_CONTEXT.thread_pool.submit(self._worker_run, worker_list))
flag = True
for future in as_completed(t_list):
if not future.result():
flag = False
break
if not flag:
break
def _thread_loop(self):
while self.loop_switch:
time.sleep(self.loop_interval_time)
if len(self.current_message_list) < self.loop_minimum_count:
continue
self._run()
self.context.clear()
self.history_message_list = self.history_message_list.extend(self.current_message_list)[
-self.history_msg_count:]
with self.message_lock:
self.current_message_list.clear()
def start_loop_run(self):
if not self.loop_switch:
self.loop_switch = True
return GLOBAL_CONTEXT.thread_pool.submit(self._thread_loop)
def run(self, result_key: str = None):
self._run()
# 获取result
result = None
if result_key:
result = self.context.get(result_key)
self.context.clear()
# 清理 msg
self.history_message_list = self.history_message_list.extend(self.current_message_list)[
-self.history_msg_count:]
self.current_message_list.clear()
return result
def submit_message(self, message: Message, with_lock=True):
if with_lock:
with self.message_lock:
self.current_message_list.append(message)
else:
self.current_message_list.append(message)

View file

@ -10,7 +10,8 @@ class Registry(object):
self.name: str = name
self.module_dict: Dict[str, Any] = {}
def register(self, module: Any, module_name: str = None):
def register(self, module_name: str = None, module: Any = None):
assert module is not None
if module_name is None:
module_name = module.__name__
@ -27,6 +28,6 @@ class Registry(object):
raise NotImplementedError
self.module_dict.update(module_name_dict)
def get(self, module_name: str):
assert module_name in self.module_dict, f'{module_name} not found in {self.name}'
def __getitem__(self, module_name: str):
assert module_name in self.module_dict, f"{module_name} not found in {self.name}"
return self.module_dict[module_name]

View file

@ -1,10 +1,10 @@
import re
from utils.logger import Logger
from memory_scope.utils.logger import Logger
class ResponseTextParser(object):
pattern_v1 = re.compile(r'<(.*?)>')
pattern_v1 = re.compile(r"<(.*?)>")
def __init__(self, response_text: str):
self.response_text: str = response_text.strip()
@ -12,22 +12,26 @@ class ResponseTextParser(object):
def parse_v1(self, prefix: str = ""):
result = []
for line in self.response_text.split('\n'):
for line in self.response_text.split("\n"):
line = line.strip()
if not line:
continue
matches = [match.group(1) for match in self.pattern_v1.finditer(line)]
if matches:
result.append(matches)
self.logger.info(f"{prefix} response_text={self.response_text} result={result}", stacklevel=2)
self.logger.info(
f"{prefix} response_text={self.response_text} result={result}", stacklevel=2
)
return result
def parse_v2(self, prefix: str = ""):
result = []
for line in self.response_text.split('\n'):
for line in self.response_text.split("\n"):
line = line.strip()
if not line or line == "":
continue
result.append(line)
self.logger.info(f"{prefix} response_text={self.response_text} result={result}", stacklevel=2)
self.logger.info(
f"{prefix} response_text={self.response_text} result={result}", stacklevel=2
)
return result

View file

@ -1,6 +1,6 @@
import time
from .logger import Logger
from memory_scope.utils.logger import Logger
class Timer(object):

View file

@ -1,8 +1,16 @@
import hashlib
import random
import re
from importlib import import_module
import time
from copy import deepcopy
from datetime import datetime
from importlib import import_module
from enumeration.message_role_enum import MessageRoleEnum
import pyfiglet
from termcolor import colored, COLORS
from memory_scope.constants.common_constants import WEEKDAYS
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
def under_line_to_hump(underline_str):
@ -10,27 +18,27 @@ def under_line_to_hump(underline_str):
return sub[0:1].upper() + sub[1:]
def init_instance_by_config(
config: dict, default_clazz_path: str = "", suffix_name: str = "", **kwargs
):
clazz_path = config.pop("clazz")
if not clazz_path:
raise RuntimeError("empty clazz_path!")
clazz_name_split = clazz_path.split(".")
clazz_name: str = clazz_name_split[-1]
if suffix_name and not clazz_name.endswith(suffix_name):
clazz_name = f"{clazz_name}_{suffix_name}"
def init_instance_by_config(config: dict, default_class_path: str = "memory_scope", suffix_name: str = "", **kwargs):
config_copy = deepcopy(config)
origin_class_path: str = config_copy.pop("class")
if not origin_class_path:
raise RuntimeError("empty class path!")
# 构造path
clazz_paths = []
if default_clazz_path:
clazz_paths.append(default_clazz_path)
clazz_paths.extend(clazz_name_split[:-1])
clazz_paths.append(clazz_name)
module = import_module(".".join(clazz_paths))
class_name_split = origin_class_path.split(".")
class_name: str = class_name_split[-1]
if suffix_name and not class_name.lower().endswith(suffix_name.lower()):
class_name = f"{class_name}_{suffix_name}"
class_name_split[-1] = class_name
cls_name = under_line_to_hump(clazz_name)
return getattr(module, cls_name)(**config, **kwargs)
class_paths = []
if default_class_path and not origin_class_path.startswith(default_class_path):
class_paths.append(default_class_path)
class_paths.extend(class_name_split)
module = import_module(".".join(class_paths))
cls_name = under_line_to_hump(class_name)
config_copy.update(kwargs)
return getattr(module, cls_name)(**config_copy)
def complete_config_name(config_name: str, suffix: str = ".json"):
@ -67,17 +75,16 @@ def get_datetime_info_dict(parse_dt: datetime):
}
def time_to_formatted_str(
time: datetime | str | int | float = None,
date_format: str = "%Y%m%d", # e.g. %Y%m%d -> "20240528", add %H:%M:%S
string_format: str = "",
) -> str:
if isinstance(time, str | int | float):
if isinstance(time, str):
time = float(time)
current_dt = datetime.fromtimestamp(time)
elif isinstance(time, datetime):
current_dt = time
def time_to_formatted_str(dt: datetime | str | int | float = None,
date_format: str = "%Y%m%d", # e.g. %Y%m%d -> "20240528", add %H:%M:%S
string_format: str = "") -> str:
if isinstance(dt, str | int | float):
if isinstance(dt, str):
dt = float(dt)
current_dt = datetime.fromtimestamp(dt)
elif isinstance(dt, datetime):
current_dt = dt
else:
current_dt = datetime.now()
@ -88,3 +95,28 @@ def time_to_formatted_str(
return_str = string_format.format(**get_datetime_info_dict(current_dt))
return return_str
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())
random.seed = seed
for line in rendered_text.splitlines():
line_color = color
if line_color is None:
random.shuffle(all_colors)
line_color = all_colors[0]
colored_line = ""
for char in line:
colored_char = colored(char, line_color, attrs=['bold'])
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()

View file

@ -1,70 +0,0 @@
from typing import List
from chat.global_context import GLOBAL_CONTEXT
from constants.common_constants import MESSAGES, CHAT_NAME
from models.base_model import BaseModel
from scheme.message import Message
from storage.base_monitor import BaseMonitor
from storage.base_vector_store import BaseVectorStore
from worker.base_worker import BaseWorker
class MemoryBaseWorker(BaseWorker):
def __init__(self,
embedding_model: str,
generation_model: str,
rank_model: str,
**kwargs):
super(MemoryBaseWorker, self).__init__(**kwargs)
self.embedding_model_name: str = embedding_model
self.generation_model_name: str = generation_model
self.rank_model_name: str = rank_model
self._embedding_model: BaseModel | None = None
self._generation_model: BaseModel | None = None
self._rank_model: BaseModel | None = None
self._vector_store: BaseVectorStore | None = None
self._monitor: BaseMonitor | None = None
@property
def messages(self) -> List[Message]:
return self.get_context(MESSAGES)
@messages.setter
def messages(self, value):
self.set_context(MESSAGES, value)
@property
def chat_name(self):
return self.get_context(CHAT_NAME)
@property
def embedding_model(self):
if self._embedding_model is None:
self._embedding_model = GLOBAL_CONTEXT.model_dict.get(self.embedding_model_name)
return self._embedding_model
@property
def generation_model(self):
if self._generation_model is None:
self._generation_model = GLOBAL_CONTEXT.model_dict.get(self.generation_model_name)
return self._generation_model
@property
def rank_model(self):
if self._rank_model is None:
self._rank_model = GLOBAL_CONTEXT.model_dict.get(self.rank_model_name)
return self._rank_model
@property
def vector_store(self):
if self._vector_store is None:
self._vector_store = GLOBAL_CONTEXT.vector_store
return self._vector_store
@property
def monitor(self):
if self._monitor is None:
self._monitor = GLOBAL_CONTEXT.monitor
return self._monitor

View file

@ -1,36 +0,0 @@
import json
from config.bailian_memory_config import BailianMemoryConfig
from constants.common_constants import REQUEST, CONFIG
from pipeline.memory import MemoryServiceRequestModel
from worker.base_worker import BaseWorker
class ParseParamsWorker(BaseWorker):
def _run(self):
# 参数合并
memory_config = {}
# 更新环境变量
memory_config.update(self.context_handler.env_configs)
# 更新请求参数
request: MemoryServiceRequestModel = self.context_handler.get_context(REQUEST)
memory_config.update(request.model_dump(exclude=set("ext_info", )))
# 更新ext_info
if request.ext_info:
memory_config.update(request.ext_info)
# 存入上下文
memory_config_model: BailianMemoryConfig = BailianMemoryConfig(**memory_config)
self.context_handler.set_context(CONFIG, memory_config_model)
# 打印
self.logger.info(f"memory_config_model={json.dumps(memory_config_model.model_dump(), ensure_ascii=False)}")
# 上游可能没有传这个参数可能隐藏在memory_id做区分
if request.user_profile:
for user_attr in request.user_profile:
user_attr.scene = request.scene

0
old/worker/__init__.py Normal file
View file

View file

@ -1,7 +1,7 @@
from typing import Any, Dict
from utils.logger import Logger
from utils.timer import Timer
from ..utils.logger import Logger
from ..utils.timer import Timer
class BaseWorker(object):

View file

View file

@ -13,9 +13,9 @@ class EsInsightWorker(MemoryBaseWorker):
insight_nodes = self.vector_store.retrieve(
size=self.kwargs.es_insight_top_k,
filter_dict={
"memoryId": self.memory_id,
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memoryType": MemoryTypeEnum.INSIGHT.value,
"memory_type": MemoryTypeEnum.INSIGHT.value,
},
)
self.logger.info(f"insight_nodes.size={len(insight_nodes)}")

View file

@ -12,10 +12,10 @@ class EsNewObsWorker(MemoryBaseWorker):
new_obs_nodes = self.vector_store.retrieve(
size=self.kwargs.es_new_obs_top_k,
filter_dict={
"memoryId": self.memory_id,
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memoryType": MemoryTypeEnum.OBSERVATION.value,
f"metaData.{NEW}": "1",
"memory_type": MemoryTypeEnum.OBSERVATION.value,
f"meta_data.{NEW}": "1",
},
)
self.logger.info(f"es new obs, size={len(new_obs_nodes)}")

View file

@ -14,13 +14,13 @@ class EsNotReflectedWorker(MemoryBaseWorker):
not_reflected_obs_nodes = self.vector_store.retrieve(
size=self.kwargs.es_new_obs_top_k,
filter_dict={
"memoryId": self.memory_id,
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memoryType": [
"memory_type": [
MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value,
],
f"metaData.{REFLECTED}": "0",
f"meta_data.{REFLECTED}": "0",
},
)
self.logger.info(

View file

@ -19,9 +19,9 @@ class EsSimilarWorker(MemoryBaseWorker):
text=query,
size=self.es_similar_top_k,
filter_dict={
"memoryId": self.memory_id,
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memoryType": [
"memory_type": [
MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.INSIGHT.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value,
@ -30,7 +30,7 @@ class EsSimilarWorker(MemoryBaseWorker):
)
for node in similar_obs_nodes:
node.metaData[RECALL_TYPE] = MemoryRecallType.SIMILAR.value
node.meta_data[RECALL_TYPE] = MemoryRecallType.SIMILAR.value
self.logger.info(f"similar_obs_nodes.size={len(similar_obs_nodes)}")
for node in similar_obs_nodes:
self.logger.info(f"node={node.content} score_similar={node.score_similar}")

View file

@ -21,10 +21,10 @@ class EsTodayObsWorker(MemoryBaseWorker):
today_obs_nodes = self.vector_store.retrieve(
size=self.es_today_obs_top_k,
filter_dict={
"memoryId": self.memory_id,
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memoryType": MemoryTypeEnum.OBSERVATION.value,
f"metaData.{DT}": time_to_formatted_str(msg_time_created),
"memory_type": MemoryTypeEnum.OBSERVATION.value,
f"meta_Data.{DT}": time_to_formatted_str(msg_time_created),
},
)

View file

@ -13,9 +13,9 @@ class LoadProfileWorker(MemoryBaseWorker):
user_profile_node = self.vector_store(
size=10000,
filter_dict={
"memoryId": self.memory_id,
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memoryType": [
"memory_type": [
MemoryTypeEnum.PROFILE.value,
MemoryTypeEnum.PROFILE_CUSTOMIZED.value,
],

View file

View file

@ -1,11 +1,16 @@
import re
from utils.tool_functions import time_to_formatted_str
from constants.common_constants import DATATIME_WORD_LIST, DATATIME_KEY_MAP, EXTRACT_TIME_DICT
from constants.common_constants import (
DATATIME_WORD_LIST,
DATATIME_KEY_MAP,
EXTRACT_TIME_DICT,
)
from worker.memory_base_worker import MemoryBaseWorker
class ExtractTimeWorker(MemoryBaseWorker):
# TODO add en version
@staticmethod
def get_parse_time_prompt(query: str, query_time_str: str):
return f"""
@ -35,27 +40,32 @@ class ExtractTimeWorker(MemoryBaseWorker):
return
# prepare prompt
# TODO add en version
time_format = "{year}{month}{day}日,{year}年第{week}周,{weekday}{hour}{minute}{second}秒。"
query_time_str = time_to_formatted_str(time=time_created,
date_format="",
string_format=time_format)
extract_time_prompt = self.get_parse_time_prompt(query=query, query_time_str=query_time_str)
query_time_str = time_to_formatted_str(
time=time_created, date_format="", string_format=time_format
)
extract_time_prompt = self.get_parse_time_prompt(
query=query, query_time_str=query_time_str
)
self.logger.info(f"extract_time_prompt={extract_time_prompt}")
# call sft model
self.generation_model.call(prompt=extract_time_prompt,
model_name=self.parse_time_model,
max_token=self.parse_time_max_token,
temperature=self.parse_time_temperature,
top_k=self.parse_time_top_k)
response_text = self.generation_model.call(
prompt=extract_time_prompt,
model_name=self.parse_time_model,
max_token=self.parse_time_max_token,
temperature=self.parse_time_temperature,
top_k=self.parse_time_top_k,
)
# if empty, return
if not response_text:
return
# re-match time info to dict
pattern = r'-\s*(\S+)(\d+)'
pattern = r"-\s*(\S+)(\d+)"
matches = re.findall(pattern, response_text)
for key, value in matches:
if key in DATATIME_KEY_MAP.keys():

View file

@ -1,25 +1,20 @@
from typing import Dict, List
from constants.common_constants import RELATED_MEMORIES, EXTRACT_TIME_DICT, ALL_ONLINE_NODES, \
TIME_MATCHED
from constants.common_constants import (
RELATED_MEMORIES,
EXTRACT_TIME_DICT,
ALL_ONLINE_NODES,
TIME_MATCHED,
)
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class FuseRerankWorker(MemoryBaseWorker):
def __init__(self, fuse_time_ratio, fuse_score_threshold, fuse_ratio_dict, *args, **kwargs):
super(FuseRerankWorker, self).__init__(*args, **kwargs)
self.fuse_score_threshold = fuse_score_threshold
self.fuse_ratio_dict = fuse_ratio_dict
# self.default_system_prompt = default_system_prompt
self.fuse_time_ratio = fuse_time_ratio
@property
def output_max_count(self):
return self.request.user.output_max_count
@staticmethod
def format_time_infer(time_infer: str, extract_time_dict: Dict[str, str], meta_data: Dict[str, str]):
def format_time_infer(
time_infer: str, extract_time_dict: Dict[str, str], meta_data: Dict[str, str]
):
if time_infer:
return time_infer
@ -29,21 +24,21 @@ class FuseRerankWorker(MemoryBaseWorker):
if value:
time_infer += f"{value}"
elif value == "-1":
time_infer += f"每年"
time_infer += "每年"
if "month" in extract_time_dict:
value = meta_data.get("msg_month")
if value:
time_infer += f"{value}"
elif value == "-1":
time_infer += f"每月"
time_infer += "每月"
if "day" in extract_time_dict:
value = meta_data.get("msg_day")
if value:
time_infer += f"{value}"
elif value == "-1":
time_infer += f"每日"
time_infer += "每日"
if "weekday" in extract_time_dict:
value = meta_data.get("msg_weekday")
@ -67,7 +62,7 @@ class FuseRerankWorker(MemoryBaseWorker):
continue
# 根据类型给ratio
type_ratio: float = self.fuse_ratio_dict.get(scheme.memory_node.memoryType, 0.1)
type_ratio: float = self.fuse_ratio_dict.get(node.memory_type, 0.1)
# 时间系数,完全匹配才行
fuse_time_ratio: float = 1.0
@ -76,7 +71,7 @@ class FuseRerankWorker(MemoryBaseWorker):
if extract_time_dict:
match_event_flag = True
for k, v in extract_time_dict.items():
event_value = scheme.memory_node.metaData.get(f"event_{k}", "")
event_value = node.meta_data.get(f"event_{k}", "")
if event_value in ["-1", v]:
continue
else:
@ -85,7 +80,7 @@ class FuseRerankWorker(MemoryBaseWorker):
match_msg_flag = True
for k, v in extract_time_dict.items():
msg_value = scheme.memory_node.metaData.get(f"msg_{k}", "")
msg_value = node.meta_data.get(f"msg_{k}", "")
if msg_value == v:
continue
else:
@ -94,32 +89,32 @@ class FuseRerankWorker(MemoryBaseWorker):
if match_event_flag or match_msg_flag:
fuse_time_ratio = self.fuse_time_ratio
scheme.memory_node.metaData[TIME_MATCHED] = "1"
node.meta_data[TIME_MATCHED] = "1"
node.score_rerank = node.score_rank * type_ratio * fuse_time_ratio
self.logger.info(f"content={scheme.memory_node.content} f_event={int(match_event_flag)} "
f"f_msg={int(match_msg_flag)} score_rerank={node.score_rerank}")
self.logger.info(
f"content={node.content} f_event={int(match_event_flag)} "
f"f_msg={int(match_msg_flag)} score_rerank={node.score_rerank}"
)
filtered_nodes.append(node)
# get output & save context
filtered_nodes = sorted(filtered_nodes, key=lambda x: x.score_rerank, reverse=True)
filtered_nodes = sorted(
filtered_nodes, key=lambda x: x.score_rerank, reverse=True
)
filtered_nodes = filtered_nodes[: self.output_max_count]
related_memories: List[str] = []
for node in filtered_nodes:
content = scheme.memory_node.content
content = node.content
# 如果命中时间逻辑
if scheme.memory_node.metaData.get(TIME_MATCHED, "") == "1":
# time_infer = scheme.memory_node.metaData.get(TIME_INFER)
# if not time_infer:
# time_infer = self.format_time_infer(time_infer=time_infer,
# extract_time_dict=extract_time_dict,
# meta_data=scheme.memory_node.metaData)
time_infer = self.format_time_infer(time_infer="",
extract_time_dict=extract_time_dict,
meta_data=scheme.memory_node.metaData)
if node.meta_data.get(TIME_MATCHED, "") == "1":
time_infer = self.format_time_infer(
time_infer="",
extract_time_dict=extract_time_dict,
meta_data=node.meta_data,
)
content = f"{time_infer}: {content}"
related_memories.append(content)
self.set_context(RELATED_MEMORIES, related_memories)
# self.set_context(DEFAULT_SYSTEM_PROMPT, self.default_system_prompt)

View file

@ -1,33 +1,35 @@
from typing import List
from utils.user_profile_handler import UserProfileHandler
from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE
from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE, CONTENT_MODIFIED
from scheme.memory_node import MemoryNode
from scheme.memory_node import MemoryNode
from node.user_attribute import UserAttribute
from worker.memory_base_worker import MemoryBaseWorker
class MemoryStoreWorker(MemoryBaseWorker):
def _run(self):
modified_memories: List[MemoryNode] | List[MemoryNode] = self.get_context(MODIFIED_MEMORIES)
modified_memories: List[MemoryNode] | List[MemoryNode] = self.get_context(
MODIFIED_MEMORIES
)
if modified_memories:
if isinstance(modified_memories[0], MemoryNode):
modified_memories = [n.memory_node for n in modified_memories]
for n in modified_memories:
if not n.id:
n.id = f"{n.memoryId}_{n.scene}_content_{n.content}"
n.id = f"{n.memory_id}_content_{n.content}"
n.code = n.id
# TODO add batch insert
self.es_client.insert(n.id, body=n.model_dump(exclude=set("content_modified", )))
n.meta_data.pop(CONTENT_MODIFIED)
self.vector_store.insert(n)
else:
self.logger.warning("modified_memories is empty!")
new_user_profile: List[UserAttribute] = self.get_context(NEW_USER_PROFILE)
new_user_profile: List[MemoryNode] = self.get_context(NEW_USER_PROFILE)
if new_user_profile:
new_user_nodes: List[MemoryNode] = [n.memory_node for n in UserProfileHandler.to_nodes(new_user_profile)]
for n in new_user_nodes:
self.es_client.insert(n.id, body=n.model_dump(exclude=set("content_modified", )))
for n in new_user_profile:
n.meta_data.pop(CONTENT_MODIFIED)
self.vector_store.insert(n)
else:
self.logger.warning("new_user_profile is empty!")

View file

@ -1,9 +1,8 @@
from typing import List, Dict
from utils.user_profile_handler import UserProfileHandler
from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE, KEYWORD_OBS_NODES, ALL_ONLINE_NODES, \
QUERY_KEYWORDS
from enumeration.memory_recall_type import MemoryRecallType
from enumeration.memory_recall_enum import MemoryRecallType
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
@ -11,29 +10,29 @@ from worker.memory_base_worker import MemoryBaseWorker
class SemanticRankWorker(MemoryBaseWorker):
def user_profile_to_nodes(self) -> List[MemoryNode]:
user_profile_nodes: List[MemoryNode] = UserProfileHandler.to_nodes(self.user_profile_dict, split_value=True)
user_profile_nodes: List[MemoryNode] = self.user_profile_dict
for node in user_profile_nodes:
# 从画像侧召回
scheme.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.PROFILE
self.logger.info(f"user profile node={scheme.memory_node.content}")
node.meta_data[RECALL_TYPE] = MemoryRecallType.PROFILE
self.logger.info(f"user profile node={node.content}")
return user_profile_nodes
def _run(self):
all_node_dict: Dict[str, MemoryNode] = {}
# 优先级: similar_obs_nodes < < profile_nodes
# 优先级: similar_obs_nodes < profile_nodes
similar_obs_nodes: List[MemoryNode] = self.get_context(SIMILAR_OBS_NODES)
if similar_obs_nodes:
for node in similar_obs_nodes:
all_node_dict[scheme.memory_node.content] = node
all_node_dict[node.content] = node
profile_nodes: List[MemoryNode] = self.user_profile_to_nodes()
if profile_nodes:
for node in profile_nodes:
all_node_dict[scheme.memory_node.content] = node
all_node_dict[node.content] = node
if not all_node_dict:
self.add_run_info(f"all_node_dict is empty!", continue_run=False)
self.add_run_info("all_node_dict is empty!", continue_run=False)
return
# call recall model
@ -45,18 +44,18 @@ class SemanticRankWorker(MemoryBaseWorker):
query_keyword_join = "".join(query_keywords)
query = f"{query} 用户的{query_keyword_join}"
documents = list(all_node_dict.keys())
result = self.rerank_client.call(query=query, documents=documents)
result = self.rank_model.call(query=query, documents=documents)
if not result:
self.add_run_info(f"semantic call recall model failed!")
self.add_run_info("semantic call recall model failed!")
return
# set score
for rank_node in result:
content = documents[rank_node["index"]]
for index, score in result.rank_scores.items():
content = documents[index]
node = all_node_dict[content]
node.score_rank = rank_node["relevance_score"]
self.logger.info(f"query={query} content={scheme.memory_node.content} score_rank={node.score_rank}")
node.score_rank = score
self.logger.info(f"query={query} content={node.content} score_rank={node.score_rank}")
# save context
all_online_nodes: List[MemoryNode] = list(all_node_dict.values())

View file

View file

@ -0,0 +1,166 @@
from datetime import datetime
from typing import List
from ...utils.tool_functions import time_to_formatted_str, get_datetime_info_dict
from ...constants.common_constants import (
NEW_INSIGHT_NODES,
DT,
NOT_REFLECTED_MERGE_NODES,
NEW_INSIGHT_KEYS,
INSIGHT_KEY,
INSIGHT_VALUE,
REFLECTED,
CONTENT_MODIFIED
)
from ...enumeration.memory_status_enum import MemoryNodeStatus
from ...enumeration.memory_type_enum import MemoryTypeEnum
from ...scheme.memory_node import MemoryNode
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.get_insight_prompt import (
GET_INSIGHT_FEW_SHOT_PROMPT,
GET_INSIGHT_SYSTEM_PROMPT,
GET_INSIGHT_USER_QUERY_PROMPT
)
class GetInsightWorker(MemoryBaseWorker):
def new_insight_node(self, insight_key: str, insight_value: str) -> MemoryNode:
created_dt = datetime.now()
dt = time_to_formatted_str(time=created_dt)
# 组合meta_data
meta_data = {
DT: dt,
INSIGHT_KEY: insight_key,
INSIGHT_VALUE: insight_value,
CONTENT_MODIFIED: True, # 新增的insight需要置为true
}
meta_data.update(
{k: str(v) for k, v in get_datetime_info_dict(created_dt).items()}
)
content = f"用户的{insight_key}{insight_value}"
return MemoryNode(
content=content,
memory_id=self.memory_id,
memory_type=MemoryTypeEnum.INSIGHT.value,
meta_data=meta_data,
status=MemoryNodeStatus.ACTIVE.value,
)
def reflect_new_insight_key(
self, insight_key: str, not_reflected_merge_nodes: List[MemoryNode]
) -> MemoryNode | None:
# 检索历史memory
hits = self.vector_store.similar_search(
text=insight_key,
size=self.es_insight_similar_top_k,
exact_filters={
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": [
MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value,
],
},
)
# 转化成 MemoryNodeWrap 合并新增nodes
related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits]
related_nodes.extend(not_reflected_merge_nodes)
# content去重
related_node_dict = {n.memory_node.content: n for n in related_nodes}
related_nodes = sorted(
list(related_node_dict.values()), key=lambda x: x.memory_node.id
)
documents = [n.memory_node.content for n in related_nodes]
# 重排所有记忆
result = self.rank_model.call(query=insight_key, documents=documents)
if not result:
self.add_run_info(
f"reflect insight_key={insight_key} call rerank client failed!"
)
return
# 根据打分过滤
for rank_node in result:
index = rank_node["index"]
score = rank_node["relevance_score"]
related_nodes[index].score_rank = score
related_nodes_sorted = sorted(
related_nodes, key=lambda x: x.score_rank, reverse=True
)[: self.insight_obs_max_cnt]
# 生成prompt
user_query_list = [x.memory_node.content for x in related_nodes_sorted]
get_insight_message = self.prompt_to_msg(
system_prompt=self.get_prompt(GET_INSIGHT_SYSTEM_PROMPT),
few_shot=self.get_prompt(GET_INSIGHT_FEW_SHOT_PROMPT),
user_query=self.get_prompt(GET_INSIGHT_USER_QUERY_PROMPT).format(
insight_key=insight_key, user_query="\n".join(user_query_list)
),
)
self.logger.info(f"get_insight_message={get_insight_message}")
# call LLM, 提取insight
response_text = self.generation_model.call(
messages=get_insight_message,
model_name=self.get_insight_model,
max_token=self.get_insight_max_token,
temperature=self.get_insight_temperature,
top_k=self.get_insight_top_k,
)
# return if empty
if not response_text:
self.add_run_info("reflect_upon_user_attr call llm failed!")
return
response_text = response_text.strip()
if response_text in [""]:
return
return self.new_insight_node(
insight_key=insight_key, insight_value=response_text
)
def _run(self):
new_insight_keys: List[MemoryNode] = self.get_context(NEW_INSIGHT_KEYS)
if not new_insight_keys:
self.add_run_info("new_insight_keys is empty! stop insight.")
return
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(
NOT_REFLECTED_MERGE_NODES
)
if not not_reflected_merge_nodes:
self.add_run_info("not_reflected_merge_nodes is empty! stop get insight.")
return
# submit insight task
for insight_key in new_insight_keys:
self.submit_thread(
self.reflect_new_insight_key,
sleep_time=1,
insight_key=insight_key,
not_reflected_merge_nodes=not_reflected_merge_nodes,
)
# save output
new_insight_nodes: List[MemoryNode] = []
for result in self.join_threads():
if result:
new_insight_nodes.append(result)
assert isinstance(result, MemoryNode)
insight_key = result.meta_data.get(INSIGHT_KEY, "")
insight_value = result.meta_data.get(INSIGHT_VALUE, "")
self.logger.info(
f"after_get_insight insight_key={insight_key} insight_value={insight_value}"
)
self.set_context(NEW_INSIGHT_NODES, new_insight_nodes)
# set REFLECTED
for node in not_reflected_merge_nodes:
node.meta_data[REFLECTED] = "1"

View file

@ -0,0 +1,99 @@
from typing import List
from ...utilsresponse_text_parser import ResponseTextParser
from ...constants.common_constants import (
NEW_OBS_NODES,
NOT_REFLECTED_OBS_NODES,
REFLECTED,
INSIGHT_NODES,
INSIGHT_KEY,
NEW_INSIGHT_KEYS,
NOT_REFLECTED_MERGE_NODES,
)
from ...scheme.memory_node import MemoryNode
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.get_reflection_prompt import (
GET_REFLECTION_FEW_SHOT_PROMPT,
GET_REFLECTION_SYSTEM_PROMPT,
GET_REFLECTION_USER_QUERY_PROMPT
)
class GetReflectionWorker(MemoryBaseWorker):
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 node.meta_data.get(REFLECTED, "") == "0"
]
# count
not_reflected_count = len(not_reflected_merge_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."
)
return
# save context
self.set_context(NOT_REFLECTED_MERGE_NODES, not_reflected_merge_nodes)
# get profile_keys
exist_keys: List[str] = []
profile_keys: List[str] = list(self.user_profile_dict.keys())
exist_keys.extend(profile_keys)
self.logger.info(f"profile_keys={profile_keys}")
# get insight_keys
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
if insight_nodes:
insight_keys = [
n.meta_data.get(INSIGHT_KEY) for n in insight_nodes
]
insight_keys = [x.strip() for x in insight_keys if x]
exist_keys.extend(insight_keys)
self.logger.info(f"insight_keys={insight_keys}")
# gen reflect prompt
user_query_list = [n.content for n in not_reflected_merge_nodes]
reflect_message = self.prompt_to_msg(
system_prompt=self.get_prompt(GET_REFLECTION_SYSTEM_PROMPT).format(
num_questions=self.reflect_num_questions
),
few_shot=self.get_prompt(GET_REFLECTION_FEW_SHOT_PROMPT),
user_query=self.get_prompt(GET_REFLECTION_USER_QUERY_PROMPT).format(
exist_keys="".join(exist_keys), user_query="\n".join(user_query_list)
),
)
self.logger.info(f"reflect_message={reflect_message}")
# # call LLM
response_text = self.generation_model.call(
messages=reflect_message,
model_name=self.reflect_obs_model,
max_token=self.reflect_obs_max_token,
temperature=self.reflect_obs_temperature,
top_k=self.reflect_obs_top_k,
)
# return if empty
if not response_text:
self.add_run_info("reflect_obs_questions call llm failed!")
return
# parse text & save
new_insight_keys = ResponseTextParser(response_text).parse_v2(
"get_insight_keys"
)
if new_insight_keys:
self.set_context(NEW_INSIGHT_KEYS, new_insight_keys)

View file

@ -0,0 +1,129 @@
from typing import List
from ...utils.response_text_parser import ResponseTextParser
from ...constants.common_constants import (
NEW_OBS_NODES,
MSG_TIME,
MODIFIED_MEMORIES,
)
from ...enumeration.memory_status_enum import MemoryNodeStatus
from ...enumeration.memory_type_enum import MemoryTypeEnum
from ...scheme.memory_node import MemoryNode
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.long_contra_repeat_prompt import (
LONG_CONTRA_REPEAT_FEW_SHOT_PROMPT,
LONG_CONTRA_REPEAT_SYSTEM_PROMPT,
LONG_CONTRA_REPEAT_USER_QUERY_PROMPT,
)
class LongContraRepeatWorker(MemoryBaseWorker):
def _run(self):
# 合并当前的obs和今日的obs
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
all_obs_nodes: List[MemoryNode] = []
for new_obs_node in new_obs_nodes:
text = new_obs_node.content
related_nodes = self.vector_store.similar_search(
text=text,
size=self.es_contra_repeat_similar_top_k,
exact_filters={
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": [
MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value,
],
},
)
has_match = False
for related_node in related_nodes:
if related_node.score_similar < self.long_contra_repeat_threshold:
continue
else:
has_match = True
all_obs_nodes.append(related_node)
if has_match:
all_obs_nodes.append(new_obs_node)
if not all_obs_nodes:
self.add_run_info("all_obs_nodes is empty!")
return
# gene prompt
user_query_list = []
all_obs_nodes = sorted(
all_obs_nodes,
key=lambda x: x.meta_data.get(MSG_TIME, ""),
reverse=True,
)
for i, n in enumerate(all_obs_nodes):
user_query_list.append(f"{i + 1} {n.content}")
merge_obs_message = self.prompt_to_msg(
system_prompt=self.get_prompt(LONG_CONTRA_REPEAT_SYSTEM_PROMPT).format(
num_obs=len(user_query_list)
),
few_shot=self.get_prompt(LONG_CONTRA_REPEAT_FEW_SHOT_PROMPT),
user_query=self.get_prompt(LONG_CONTRA_REPEAT_USER_QUERY_PROMPT).format(
user_query="\n".join(user_query_list)
),
)
self.logger.info(f"merge_obs_message={merge_obs_message}")
# call LLM
response_text = self.generation_model.call(
messages=merge_obs_message,
model_name=self.merge_obs_model,
max_token=self.merge_obs_max_token,
temperature=self.merge_obs_temperature,
top_k=self.merge_obs_top_k,
)
# return if empty
if not response_text:
self.add_run_info("contra repeat call llm failed!")
return
# parse text
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
if len(idx_merge_obs_list) <= 0:
self.add_run_info("idx_merge_obs_list is empty!")
return
# add merged obs
merge_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_merge_obs_list:
if not obs_content_list:
continue
# [6, 逃课]
if len(obs_content_list) != 2:
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
continue
idx, keep_flag = obs_content_list
if not idx.isdigit():
self.logger.warning(f"idx={idx} is invalid!")
continue
# 序号需要修正-1
idx = int(idx) - 1
if idx >= len(all_obs_nodes):
self.logger.warning(f"idx={idx} is invalid!")
continue
if keep_flag not in ["矛盾", "被包含", ""]:
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
continue
node: MemoryNode = all_obs_nodes[idx]
if keep_flag != "":
node.status = MemoryNodeStatus.EXPIRED.value
merge_obs_nodes.append(node)
self.logger.info(f"after contra repeat: {node.content} {node.status}")
# save context
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)

View file

@ -0,0 +1,47 @@
from typing import List, Dict
from ...constants.common_constants import (
NEW_INSIGHT_NODES,
MODIFIED_MEMORIES,
INSIGHT_NODES,
NEW_OBS_NODES,
NOT_REFLECTED_OBS_NODES,
NEW,
NOT_REFLECTED_MERGE_NODES,
CONTENT_MODIFIED,
)
from ...scheme.memory_node import MemoryNode
from ..memory_base_worker import MemoryBaseWorker
class SummaryCollectWorker(MemoryBaseWorker):
def _run(self):
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
new_insight_nodes: List[MemoryNode] = self.get_context(NEW_INSIGHT_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] = self.get_context(
NOT_REFLECTED_MERGE_NODES
)
# 合并逻辑复杂务必check
all_node_dict: Dict[str, MemoryNode] = {}
if insight_nodes:
all_node_dict.update(
{n.id: n for n in insight_nodes if n.meta_data.get(CONTENT_MODIFIED, False)}
)
if new_insight_nodes:
all_node_dict.update({n.content: n for n in new_insight_nodes})
if new_obs_nodes:
# 设置为非新
for n in new_obs_nodes:
n.meta_data[NEW] = "0"
all_node_dict.update({n.content: n for n in new_obs_nodes})
if not_reflected_merge_nodes and not_reflected_nodes:
# 进入reflect阶段
all_node_dict.update({n.id: n for n in not_reflected_nodes})
self.set_context(MODIFIED_MEMORIES, list(all_node_dict.values()))

View file

@ -0,0 +1,177 @@
from typing import List
from ...utils.response_text_parser import ResponseTextParser
from ...constants.common_constants import (
INSIGHT_NODES,
NEW_OBS_NODES,
INSIGHT_KEY,
INSIGHT_VALUE,
CONTENT_MODIFIED,
)
from ...scheme.memory_node import MemoryNode
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.update_insight_prompt import (
UPDATE_INSIGHT_FEW_SHOT_PROMPT,
UPDATE_INSIGHT_SYSTEM_PROMPT,
UPDATE_INSIGHT_USER_QUERY_PROMPT,
)
class UpdateInsightWorker(MemoryBaseWorker):
def filter_obs_nodes(
self, insight_node: MemoryNode, new_obs_nodes: List[MemoryNode]
) -> (MemoryNode, List[MemoryNode], float):
max_score: float = 0
filtered_nodes: List[MemoryNode] = []
insight_key = insight_node.meta_data.get(INSIGHT_KEY, "")
insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "")
if not insight_key or not insight_value:
self.logger.warning(
f"insight_key={insight_key} insight_value={insight_value} is empty!"
)
return insight_node, filtered_nodes, max_score
result = self.rank_model.call(
query=insight_key, documents=[x.content for x in new_obs_nodes]
)
if not result:
self.add_run_info(f"update_insight={insight_key} call rerank failed!")
return insight_node, filtered_nodes, max_score
# 找到大于阈值的obs node
for index, score in result.rank_scores.items():
node = new_obs_nodes[index]
keep_flag = "filtered"
if score >= self.update_insight_threshold:
filtered_nodes.append(node)
keep_flag = "keep"
max_score = max(max_score, score)
self.logger.info(
f"insight_key={insight_key} insight_value={insight_value} "
f"score={score} keep_flag={keep_flag}"
)
if not filtered_nodes:
self.logger.info(f"update_insight={insight_key} filtered_nodes is empty!")
return insight_node, filtered_nodes, max_score
def update_insight(
self, insight_node: MemoryNode, filtered_nodes: List[MemoryNode]
) -> MemoryNode:
insight_key = insight_node.meta_data.get(INSIGHT_KEY, "")
insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "")
self.logger.info(
f"update_insight insight_key={insight_key} insight_value={insight_value} "
f"doc.size={len(filtered_nodes)}"
)
# gen prompt
user_query_list = []
for node in filtered_nodes:
user_query_list.append(f"句子:{node.content}")
update_insight_message = self.prompt_to_msg(
system_prompt=self.get_prompt(UPDATE_INSIGHT_SYSTEM_PROMPT),
few_shot=self.get_prompt(UPDATE_INSIGHT_FEW_SHOT_PROMPT),
user_query=self.get_prompt(UPDATE_INSIGHT_USER_QUERY_PROMPT).format(
user_query="\n".join(user_query_list),
insight_key=insight_key,
insight_key_value=insight_key + "" + insight_value,
),
)
self.logger.info(f"update_insight_message={update_insight_message}")
# call LLM
response_text: str = self.generation_model.call(
messages=update_insight_message,
model_name=self.update_insight_model,
max_token=self.update_insight_max_token,
temperature=self.update_insight_temperature,
top_k=self.update_insight_top_k,
)
# return if empty
if not response_text:
self.add_run_info(
f"update_insight insight_key={insight_key} call llm failed!"
)
return insight_node
profile_list = ResponseTextParser(response_text).parse_v1(
f"update_profile {insight_key}"
)
if not profile_list:
self.add_run_info(
f"update_insight insight_key={insight_key} profile_list empty 1!"
)
return insight_node
profile_list = profile_list[0]
if not profile_list:
self.add_run_info(
f"update_insight insight_key={insight_key} profile_list empty 2"
)
return insight_node
insight_value = profile_list[0]
if not insight_value or insight_value in ["", "重复"]:
self.logger.info(f"insight_value={insight_value}, skip.")
return insight_node
insight_node.meta_data[INSIGHT_VALUE] = insight_value
insight_node.meta_data[CONTENT_MODIFIED] = True
return insight_node
def _run(self):
# 获取新的obs和insight
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
if not new_obs_nodes:
self.logger.info("new_obs_nodes is empty, stop update sights!")
return
if not insight_nodes:
self.logger.info("insight_nodes is empty, stop update sights!")
return
# 提交打分任务
for node in insight_nodes:
self.submit_thread(
self.filter_obs_nodes,
sleep_time=0.1,
insight_node=node,
new_obs_nodes=new_obs_nodes,
)
# 选择topN
result_list = []
for result in self.join_threads():
insight_node, filtered_nodes, max_score = result
if not filtered_nodes:
continue
result_list.append(result)
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
if len(result_sorted) > self.update_insight_max_thread:
result_sorted = result_sorted[: self.update_insight_max_thread]
# 提交LLM update任务
for insight_node, filtered_nodes, _ in result_sorted:
self.submit_thread(
self.update_insight,
sleep_time=1,
insight_node=insight_node,
filtered_nodes=filtered_nodes,
)
# 等待结果
for result in self.join_threads():
if result:
insight_node: MemoryNode = result
insight_key = insight_node.meta_data.get(INSIGHT_KEY, "")
insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "")
self.logger.info(
f"after_update_insight insight_key={insight_key} insight_value={insight_value}"
)

View file

@ -0,0 +1,241 @@
from typing import List
from ...utils.response_text_parser import ResponseTextParser
from ...constants.common_constants import NEW_OBS_NODES, NEW_USER_PROFILE
from ...enumeration.memory_type_enum import MemoryTypeEnum
from ....memory_node import MemoryNode
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.update_profile_prompt import (
UPDATE_PLURAL_PROFILE_FEW_SHOT_PROMPT,
UPDATE_PLURAL_PROFILE_SYSTEM_PROMPT,
UPDATE_PLURAL_PROFILE_USER_QUERY_PROMPT,
UPDATE_UNIQUE_PROFILE_FEW_SHOT_PROMPT,
UPDATE_UNIQUE_PROFILE_SYSTEM_PROMPT,
UPDATE_UNIQUE_PROFILE_USER_QUERY_PROMPT
)
from ...chat.global_context import GlobalContext
class UpdateProfileWorker(MemoryBaseWorker):
@property
def extra_user_attrs(self):
return GlobalContext.global_configs.get("extra_user_attrs", [])
def filter_obs_nodes(
self, user_attr: MemoryNode, new_obs_nodes: List[MemoryNode]
) -> (MemoryNode, List[MemoryNode], float):
max_score: float = 0
filtered_nodes: List[MemoryNode] = []
result = self.rank_model.call(
query=user_attr.meta_data.get("description", ""),
documents=[x.content for x in new_obs_nodes],
)
if not result:
self.add_run_info(
f"update_user_attr={user_attr.meta_data.get("memory_key", "")} call rerank failed!"
)
return user_attr, filtered_nodes, max_score
# 找到大于阈值的obs node
filtered_nodes: List[MemoryNode] = []
for index, score in result.rank_scores.items():
node = new_obs_nodes[index]
keep_flag = "filtered"
if score >= self.update_profile_threshold:
filtered_nodes.append(node)
keep_flag = "keep"
max_score = max(max_score, score)
self.logger.info(
f"key={user_attr.meta_data.get("memory_key", "")} desc={user_attr.meta_data.get("description", "")} "
f"content={node.content} score={score} keep_flag={keep_flag}"
)
if not filtered_nodes:
self.logger.info(f"update_user_attr={user_attr} filtered_nodes is empty!")
return user_attr, filtered_nodes, max_score
def update_user_attr(
self, user_attr: MemoryNode, filtered_nodes: List[MemoryNode]
) -> MemoryNode:
self.logger.info(
f"update_user_attr memory_key={user_attr.meta_data.get("memory_key", "")} desc={user_attr.meta_data.get("description", "")} "
f"value={user_attr.meta_data.get("value", "")} doc.size={len(filtered_nodes)}"
)
# 根据不同的参数类型是否多值分别给出prompt
user_query_list = []
for node in filtered_nodes:
user_query_list.append(f"句子:{node.content}")
update_profile = f"{user_attr.meta_data.get("memory_key", "")}{user_attr.meta_data.get("description", "")}"
update_profile_value = update_profile + "" + "".join(user_attr.meta_data.get("value", ""))
if user_attr.meta_data.get("is_unique", 0) == 1:
update_profile_message = self.prompt_to_msg(
system_prompt=self.get_prompt(UPDATE_UNIQUE_PROFILE_SYSTEM_PROMPT),
few_shot=self.get_prompt(UPDATE_UNIQUE_PROFILE_FEW_SHOT_PROMPT),
user_query=self.get_prompt(UPDATE_UNIQUE_PROFILE_USER_QUERY_PROMPT).format(
user_query="\n".join(user_query_list),
update_profile=update_profile,
update_profile_value=update_profile_value,
),
)
else:
update_profile_message = self.prompt_to_msg(
system_prompt=self.get_prompt(UPDATE_PLURAL_PROFILE_SYSTEM_PROMPT),
few_shot=self.get_prompt(UPDATE_PLURAL_PROFILE_FEW_SHOT_PROMPT),
user_query=self.get_prompt(UPDATE_PLURAL_PROFILE_USER_QUERY_PROMPT).format(
user_query="\n".join(user_query_list),
update_profile=update_profile,
update_profile_value=update_profile_value,
),
)
self.logger.info(f"update_profile_message={update_profile_message}")
# call LLM
response_text: str = self.generation_model.call(
messages=update_profile_message,
model_name=self.update_profile_model,
max_token=self.update_profile_max_token,
temperature=self.update_profile_temperature,
top_k=self.update_profile_top_k,
)
# return if empty
if not response_text:
self.add_run_info(
f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} call llm failed!"
)
return user_attr
profile_list = ResponseTextParser(response_text).parse_v1(
f"update_attr {user_attr.meta_data.get("memory_key", "")}"
)
if not profile_list:
self.add_run_info(
f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} profile_list empty 1!"
)
return user_attr
profile_list = profile_list[0]
if not profile_list:
self.add_run_info(
f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} profile_list empty 2"
)
return user_attr
profile = profile_list[0]
if not profile or profile in ["", "重复"]:
self.logger.info(f"profile={profile}, skip.")
return user_attr
# check 英文中午逗号
if user_attr.meta_data.get("is_unique", 0) == 1:
user_attr.meta_data["value"] = [profile.strip()]
else:
attr_value_list = profile.replace("", ",").split(",")
user_attr.meta_data["value"] = [
x.strip() for x in sorted(list(set(user_attr.meta_data.get("value", "") + attr_value_list)))
]
return user_attr
def add_extra_user_attrs(self):
# 解析为空返回
extra_user_attr_list = [x.strip() for x in self.extra_user_attrs if x.strip()]
if not extra_user_attr_list:
return
for user_attr_info in extra_user_attr_list:
user_attr_split = user_attr_info.split(":")
# 格式不对返回
if len(user_attr_split) < 1:
continue
user_attr_key = user_attr_split[0]
user_attr_desc = ""
if len(user_attr_split) >= 2:
user_attr_desc = user_attr_split[1]
user_attr_unique = 0
if len(user_attr_split) >= 3:
user_attr_unique = int(user_attr_split[2])
# 已经包含返回
if user_attr_key in self.user_profile_dict:
user_attr = self.user_profile_dict[user_attr_key]
# description为空补充description
if not user_attr.meta_data.get("description", ""):
user_attr.meta_data["description"] = user_attr_desc
continue
# 增加新属性
new_attr = MemoryNode(
memory_id=self.memory_id,
meta_data={
"memory_key": user_attr_key,
"is_unique": int(user_attr_unique),
"is_mutable": 1,
"description": user_attr_desc
},
memory_type=MemoryTypeEnum.PROFILE,
status=1,
)
self.user_profile_dict[user_attr_key] = new_attr
def _run(self):
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
if not new_obs_nodes:
self.logger.info("new_obs_nodes is empty, stop user profile!")
self.set_context(NEW_USER_PROFILE, list(self.user_profile_dict.values()))
return
# 增加环境变量配置的属性
if self.extra_user_attrs:
self.add_extra_user_attrs()
new_user_profile: List[MemoryNode] = []
self.set_context(NEW_USER_PROFILE, new_user_profile)
for user_attr_key, user_attr in self.user_profile_dict.items():
# 不可修改直接跳过
if user_attr.meta_data.get("is_mutable", 0) != 1:
new_user_profile.append(user_attr)
self.logger.info(f"{user_attr_key} is not mutable! continue")
continue
self.submit_thread(
self.filter_obs_nodes,
sleep_time=0.1,
user_attr=user_attr,
new_obs_nodes=new_obs_nodes,
)
# 选择topN
result_list = []
for result in self.join_threads():
user_attr, filtered_nodes, max_score = result
if not filtered_nodes:
continue
result_list.append(result)
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
if len(result_sorted) > self.update_profile_max_thread:
result_sorted = result_sorted[: self.update_profile_max_thread]
# 提交LLM update任务
for user_attr, filtered_nodes, _ in result_sorted:
self.submit_thread(
self.update_user_attr,
sleep_time=1,
user_attr=user_attr,
filtered_nodes=filtered_nodes,
)
# collect result & save
for result in self.join_threads():
if result:
user_attribute: MemoryNode = result
self.logger.info(
f"after_update_profile memory_key={user_attribute.meta_data.get("memory_key", "")} "
f"desc={user_attribute.meta_data.get("description", "")} value={user_attribute.meta_data.get("value", "")}"
)
new_user_profile.append(user_attribute)

View file

View file

@ -0,0 +1,117 @@
from typing import List
from ...utils.response_text_parser import ResponseTextParser
from ...constants.common_constants import (
NEW_OBS_NODES,
TODAY_OBS_NODES,
MSG_TIME,
NEW_OBS_WITH_TIME_NODES,
MODIFIED_MEMORIES,
)
from ...enumeration.memory_status_enum import MemoryNodeStatus
from ...scheme.memory_node import MemoryNode
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.contra_repeat_prompt import (
CONTRA_REPEAT_FEW_SHOT_PROMPT,
CONTRA_REPEAT_SYSTEM_PROMPT,
CONTRA_REPEAT_USER_QUERY_PROMPT,
)
class ContraRepeatWorker(MemoryBaseWorker):
def _run(self):
# 合并当前的obs和今日的obs
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
new_obs_with_time_nodes: List[MemoryNode] = self.get_context(
NEW_OBS_WITH_TIME_NODES
)
today_obs_nodes: List[MemoryNode] = self.get_context(TODAY_OBS_NODES)
all_obs_nodes: List[MemoryNode] = []
if new_obs_nodes:
all_obs_nodes.extend(new_obs_nodes)
if new_obs_with_time_nodes:
all_obs_nodes.extend(new_obs_with_time_nodes)
if today_obs_nodes:
all_obs_nodes.extend(today_obs_nodes)
if not all_obs_nodes:
self.add_run_info("all_obs_nodes is empty!")
return
# gene prompt
user_query_list = []
all_obs_nodes = sorted(
all_obs_nodes,
key=lambda x: x.meta_data.get(MSG_TIME, ""),
reverse=True,
)
for i, n in enumerate(all_obs_nodes):
user_query_list.append(f"{i + 1} {n.content}")
merge_obs_message = self.prompt_to_msg(
system_prompt=self.get_prompt(CONTRA_REPEAT_SYSTEM_PROMPT).format(
num_obs=len(user_query_list)
),
few_shot=self.get_prompt(CONTRA_REPEAT_FEW_SHOT_PROMPT),
user_query=self.get_prompt(CONTRA_REPEAT_USER_QUERY_PROMPT).format(
user_query="\n".join(user_query_list)
),
)
self.logger.info(f"merge_obs_message={merge_obs_message}")
# call LLM
response_text = self.generation_model.call(
messages=merge_obs_message,
model_name=self.merge_obs_model,
max_token=self.merge_obs_max_token,
temperature=self.merge_obs_temperature,
top_k=self.merge_obs_top_k,
)
# return if empty
if not response_text:
self.add_run_info("contra repeat call llm failed!")
return
# parse text
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
if len(idx_merge_obs_list) <= 0:
self.add_run_info("idx_merge_obs_list is empty!")
return
# add merged obs
merge_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_merge_obs_list:
if not obs_content_list:
continue
# [6, 逃课]
if len(obs_content_list) != 2:
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
continue
idx, keep_flag = obs_content_list
if not idx.isdigit():
self.logger.warning(f"idx={idx} is invalid!")
continue
# 序号需要修正-1
idx = int(idx) - 1
if idx >= len(all_obs_nodes):
self.logger.warning(f"idx={idx} is invalid!")
continue
if keep_flag not in ["矛盾", "被包含", ""]:
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
continue
node: MemoryNode = all_obs_nodes[idx]
if keep_flag != "":
node.status = MemoryNodeStatus.EXPIRED.value
merge_obs_nodes.append(node)
self.logger.info(
f"after contra repeat: {node.content} {node.status}"
)
# save context
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)

View file

@ -0,0 +1,167 @@
from datetime import datetime
from typing import List
from ...utils.response_text_parser import ResponseTextParser
from ...utils.tool_functions import (
time_to_formatted_str,
get_datetime_info_dict,
extract_date_parts,
)
from ...constants.common_constants import (
REFLECTED,
DT,
TIME_INFER,
NEW,
MSG_TIME,
KEY_WORD,
DATATIME_WORD_LIST,
NEW_OBS_WITH_TIME_NODES,
CONTENT_MODIFIED,
)
from ...enumeration.memory_status_enum import MemoryNodeStatus
from ...enumeration.memory_type_enum import MemoryTypeEnum
from ...scheme.memory_node import MemoryNode
from ...scheme.message import Message
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.get_observation_with_time_prompt import (
GET_OBSERVATION_WITH_TIME_FEW_SHOT_PROMPT,
GET_OBSERVATION_WITH_TIME_SYSTEM_PROMPT,
GET_OBSERVATION_WITH_TIME_USER_QUERY_PROMPT,
)
class GetObservationWithTimeWorker(MemoryBaseWorker):
def add_observation(
self, message: Message, obs_content: str, time_infer: str, keywords: str
):
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
dt = time_to_formatted_str(time=created_dt)
# 组合meta_data
meta_data = {
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
REFLECTED: "0", # reflect标记
DT: dt, # 当天标记
NEW: "1", # summary-long标记
MSG_TIME: message.time_created, # 对话时间
TIME_INFER: time_infer, # 推断的时间
KEY_WORD: keywords, # 关键词
CONTENT_MODIFIED: True, # 新增的obs需要置为true
}
# 事件时间
meta_data.update(
{f"event_{k}": str(v) for k, v in extract_date_parts(time_infer).items()}
)
# 对话时间
meta_data.update(
{f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()}
)
return MemoryNode.init_from_attrs(
content=obs_content,
memory_id=self.memory_id,
memory_type=MemoryTypeEnum.OBSERVATION.value,
meta_data=meta_data,
status=MemoryNodeStatus.ACTIVE.value,
)
def _run(self):
# gene prompt
user_query_list = []
i = 1
for msg in self.messages:
match = False
for time_keyword in DATATIME_WORD_LIST:
if time_keyword in msg.content:
match = True
break
if match:
dt = time_to_formatted_str(
time=msg.time_created,
date_format="",
string_format="{year}{month}{day}{weekday}{hour}",
)
user_query_list.append(f"{i} {dt} 用户:{msg.content}")
i += 1
if not user_query_list:
self.add_run_info(
f"get obs with time user_query_list={user_query_list} is empty"
)
return
obtain_obs_message = self.prompt_to_msg(
system_prompt=self.get_prompt(
GET_OBSERVATION_WITH_TIME_SYSTEM_PROMPT
).format(num_obs=len(user_query_list)),
few_shot=self.get_prompt(GET_OBSERVATION_WITH_TIME_FEW_SHOT_PROMPT),
user_query=self.get_prompt(
GET_OBSERVATION_WITH_TIME_USER_QUERY_PROMPT
).format(user_query="\n".join(user_query_list)),
)
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
# call LLM
response_text: str = self.generation_model.call(
messages=obtain_obs_message,
model_name=self.summary_messages_model,
max_token=self.summary_messages_max_token,
temperature=self.summary_messages_temperature,
top_k=self.summary_messages_top_k,
)
# return if empty
if not response_text:
self.add_run_info("summary call llm failed!", continue_run=False)
return
# parse text
idx_obs_list = ResponseTextParser(response_text).parse_v1("get_obs_time")
if len(idx_obs_list) <= 0:
self.add_run_info("idx_obs_list is empty!", continue_run=False)
return
# gene new obs nodes
new_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_obs_list:
if not obs_content_list:
continue
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
if len(obs_content_list) != 4:
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
continue
idx, time_infer, obs_content, keywords = obs_content_list
if obs_content in ["", "重复"]:
continue
if not idx.isdigit():
self.logger.warning(f"idx={idx} is invalid!")
continue
if time_infer == "":
time_infer = ""
# 序号需要修正-1
idx = int(idx) - 1
if idx >= len(self.messages):
self.logger.warning(
f"idx={idx} is invalid! messages.size={len(self.messages)}"
)
continue
new_obs_nodes.append(
self.add_observation(
message=self.messages[idx],
obs_content=obs_content,
time_infer=time_infer,
keywords=keywords,
)
)
# save context
self.set_context(NEW_OBS_WITH_TIME_NODES, new_obs_nodes)

View file

@ -0,0 +1,144 @@
from datetime import datetime
from typing import List
from ...utils.response_text_parser import ResponseTextParser
from ...utils.tool_functions import time_to_formatted_str, get_datetime_info_dict
from ...constants.common_constants import (
REFLECTED,
DT,
NEW_OBS_NODES,
TIME_INFER,
NEW,
MSG_TIME,
KEY_WORD,
DATATIME_WORD_LIST,
CONTENT_MODIFIED,
)
from ...enumeration.memory_status_enum import MemoryNodeStatus
from ...enumeration.memory_type_enum import MemoryTypeEnum
from ...scheme.memory_node import MemoryNode
from ...scheme.message import Message
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.get_observation_prompt import (
GET_OBSERVATION_FEW_SHOT_PROMPT,
GET_OBSERVATION_SYSTEM_PROMPT,
GET_OBSERVATION_USER_QUERY_PROMPT,
)
class GetObservationWorker(MemoryBaseWorker):
def add_observation(self, message: Message, obs_content: str, keywords: str):
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
dt = time_to_formatted_str(time=created_dt)
# 组合meta_data
meta_data = {
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
REFLECTED: "0", # reflect标记
DT: dt, # 当天标记
NEW: "1", # summary-long标记
MSG_TIME: message.time_created, # 对话时间
TIME_INFER: "", # 推断的时间
KEY_WORD: keywords, # 关键词
CONTENT_MODIFIED: True, # 新增的obs需要置为true
}
meta_data.update(
{k: str(v) for k, v in get_datetime_info_dict(created_dt).items()}
)
return MemoryNode(
content=obs_content,
memory_id=self.memory_id,
memory_type=MemoryTypeEnum.OBSERVATION.value,
meta_data=meta_data,
status=MemoryNodeStatus.ACTIVE.value,
)
def _run(self):
# gene prompt
user_query_list = []
i = 1
for msg in self.messages:
match = False
for time_keyword in DATATIME_WORD_LIST:
if time_keyword in msg.content:
match = True
break
if not match:
user_query_list.append(f"{i} 用户:{msg.content}")
i += 1
if not user_query_list:
self.add_run_info(f"get obs user_query_list={user_query_list} is empty")
return
obtain_obs_message = self.prompt_to_msg(
system_prompt=self.get_prompt(GET_OBSERVATION_SYSTEM_PROMPT).format(
num_obs=len(user_query_list)
),
few_shot=self.get_prompt(GET_OBSERVATION_FEW_SHOT_PROMPT),
user_query=self.get_prompt(GET_OBSERVATION_USER_QUERY_PROMPT).format(
user_query="\n".join(user_query_list)
),
)
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
# call LLM
response_text: str = self.generation_model.call(
messages=obtain_obs_message,
model_name=self.summary_messages_model,
max_token=self.summary_messages_max_token,
temperature=self.summary_messages_temperature,
top_k=self.summary_messages_top_k,
)
# return if empty
if not response_text:
self.add_run_info("summary call llm failed!", continue_run=False)
return
# parse text
idx_obs_list = ResponseTextParser(response_text).parse_v1("get obs")
if len(idx_obs_list) <= 0:
self.add_run_info("idx_obs_list is empty!", continue_run=False)
return
# gene new obs nodes
new_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_obs_list:
if not obs_content_list:
continue
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
if len(obs_content_list) != 4:
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
continue
idx, time_infer, obs_content, keywords = obs_content_list
if obs_content in ["", "重复"]:
continue
if not idx.isdigit():
self.logger.warning(f"idx={idx} is invalid!")
continue
# 序号需要修正-1
idx = int(idx) - 1
if idx >= len(self.messages):
self.logger.warning(
f"idx={idx} is invalid! messages.size={len(self.messages)}"
)
continue
new_obs_nodes.append(
self.add_observation(
message=self.messages[idx],
obs_content=obs_content,
keywords=keywords,
)
)
# save context
self.set_context(NEW_OBS_NODES, new_obs_nodes)

View file

@ -0,0 +1,70 @@
from ...utils.response_text_parser import ResponseTextParser
from enumeration.message_role_enum import MessageRoleEnum
from worker.memory_base_worker import MemoryBaseWorker
from ...chat.global_context import GlobalContext
from ...prompts.info_filter_prompt import INFO_FILTER_FEW_SHOT_PROMPT, INFO_FILTER_SYSTEM_PROMPT, INFO_FILTER_USER_QUERY_PROMPT
class InfoFilterWorker(MemoryBaseWorker):
def _run(self):
# filter user msg
info_messages = []
for msg in self.messages:
if msg.role != MessageRoleEnum.USER.value:
continue
if len(msg.content) >= self.info_filter_msg_max_size:
continue
info_messages.append(msg)
# gene prompt
user_query = "\n".join(
[f"{i + 1} 用户:{msg.content}" for i, msg in enumerate(info_messages)]
)
info_filter_message = self.prompt_to_msg(
system_prompt=self.get_prompt(INFO_FILTER_SYSTEM_PROMPT).format(
batch_size=len(info_messages)
),
few_shot=self.get_prompt(INFO_FILTER_FEW_SHOT_PROMPT),
user_query=self.get_prompt(INFO_FILTER_USER_QUERY_PROMPT).format(
user_query=user_query
),
)
self.logger.info(f"info_filter_message={info_filter_message}")
# call llm
response_text = self.generation_model.call(
messages=info_filter_message,
model_name=self.info_filter_model,
max_token=self.info_filter_max_token,
temperature=self.info_filter_temperature,
top_k=self.info_filter_top_k,
)
# return if empty
if not response_text:
self.add_run_info("info score call llm failed!", continue_run=False)
return
# parse text
info_score_list = ResponseTextParser(response_text).parse_v1("info_filter")
if len(info_score_list) != len(info_messages):
self.add_run_info(
f"info_score_size != info_messages_size, "
f"{len(info_score_list)} vs {len(info_messages)}",
continue_run=False,
)
return
# 过滤value=0的messages
filtered_messages = []
for msg, info_score in zip(info_messages, info_score_list):
if not info_score:
continue
score = info_score[0]
# if score in ("1", "2",):
if score in ("2",):
msg.info_score = score
filtered_messages.append(msg)
# 后续不会关注为0的msg直接丢弃
self.messages = filtered_messages

13
test.py Normal file
View file

@ -0,0 +1,13 @@
from memory_scope.cli import CliJob
import fire
def main(config_path: str):
job = CliJob(config_path=config_path)
job.init_global_content_by_config()
job.run()
if __name__ == "__main__":
# fire.Fire(main)
main("config/config.yaml")

View file

@ -1,7 +1,12 @@
import sys
sys.path.append(".") # noqa: E402
import asyncio
import unittest
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
from memory_scope.utils.logger import Logger
class TestLLIEmbedding(unittest.TestCase):
@ -9,26 +14,31 @@ class TestLLIEmbedding(unittest.TestCase):
def setUp(self):
config = {
"method_type": "DashScopeEmbedding",
"module_name": "dashscope_embedding",
"model_name": "text-embedding-v2",
"clazz": "models.base_embedding_model"
}
self.emb = LlamaIndexEmbeddingModel(**config)
print()
self.logger = Logger.get_logger()
def test_single_embedding(self):
text = "您吃了吗?"
result = self.emb.call(text=text)
print(result)
self.logger.info(result.m_type)
self.logger.info(len(result.embedding_results))
def test_batch_embedding(self):
texts = ["您吃了吗?",
"吃了吗您?"]
result = self.emb.call(text=texts)
print(result)
print()
self.logger.info(result)
def test_async_embedding(self):
texts = ["您吃了吗?",
"吃了吗您?"]
# 调用异步函数并等待其结果
result = asyncio.run(self.emb.async_call(text=texts))
print(result)
print()
self.logger.info(result)

View file

@ -1,6 +1,13 @@
import unittest
import sys
sys.path.append(".") # noqa: E402
import unittest
import time
from memory_scope.scheme.message import Message
from memory_scope.models.llama_index_generation_model import LlamaIndexGenerationModel
from memory_scope.utils.logger import Logger
class TestLLILLM(unittest.TestCase):
@ -8,52 +15,41 @@ class TestLLILLM(unittest.TestCase):
def setUp(self):
config = {
"method_type": "DashScope",
"module_name": "dashscope_generation",
"model_name": "qwen-max",
"clazz": "models.llama_index_generation_model"
}
self.llm = LlamaIndexGenerationModel(**config)
self.logger = Logger.get_logger()
def test_llm_prompt(self):
prompt = "你是谁?"
ans = self.llm.call(
stream=False,
prompt=prompt
)
print(ans.text)
@unittest.skip("tmp")
ans = self.llm.call(stream=False, prompt=prompt)
self.logger.info(ans.message.content)
def test_llm_messages(self):
messages = [{"role": "system", "content": "you are a helpful assistant."},
{"role": "user", "content": "你是谁?"}]
ans = self.llm.call(
stream=False,
messages=messages
)
print(ans.text)
@unittest.skip("tmp")
messages = [Message(role="system", content="you are a helpful assistant."),
Message(role="user", content="你是谁?")]
ans = self.llm.call(stream=False, messages=messages)
self.logger.info(ans.message.content)
def test_llm_prompt_stream(self):
prompt = "你如何看待黄金上涨?"
ans = self.llm.call(
stream=True,
prompt=prompt
)
import sys
import time
ans = self.llm.call(stream=True, prompt=prompt)
self.logger.info("-----start-----")
for a in ans:
sys.stdout.write(a.delta)
sys.stdout.flush()
time.sleep(0.1)
@unittest.skip("tmp")
def test_llm_messages(self):
messages = [{"role": "system", "content": "you are a helpful assistant."},
{"role": "user", "content": "你如何看待黄金上涨?"}]
ans = self.llm.call(
stream=True,
messages=messages
)
import sys
import time
self.logger.info("-----end-----")
def test_llm_messages_stream(self):
messages = [Message(role="system", content="you are a helpful assistant."),
Message(role="user", content="你如何看待黄金上涨?")]
ans = self.llm.call(stream=True, messages=messages)
self.logger.info("-----start-----")
for a in ans:
sys.stdout.write(a.delta)
sys.stdout.flush()
time.sleep(0.1)
time.sleep(0.1)
self.logger.info("-----end-----")

View file

@ -1,7 +1,6 @@
import json
import unittest
from memory_scope.models.llama_index_rerank_model import LlamaIndexRerankModel
from memory_scope.models.llama_index_rank_model import LlamaIndexRankModel
class TestLLIReRank(unittest.TestCase):
@ -9,11 +8,11 @@ class TestLLIReRank(unittest.TestCase):
def setUp(self):
config = {
"method_type": "DashScopeRerank",
"module_name": "dashscope_rank",
"model_name": "gte-rerank",
"clazz": "models.llama_index_rerank_model"
}
self.reranker = LlamaIndexRerankModel(**config)
self.reranker = LlamaIndexRankModel(**config)
def test_rerank(self):
query = "吃啥?"

View file

@ -1,90 +1,147 @@
import unittest
from llama_index.core.vector_stores.types import MetadataFilter, MetadataFilters, FilterCondition, FilterOperator
from llama_index.core.schema import TextNode
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.storage.llama_index_elastic_search_store import LlamaIndexElasticSearchStore
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
class TestLlamaIndexElasticSearchStore(unittest.TestCase):
"""Tests for LLIEmbedding"""
def setUp(self):
config = {
"method_type": "DashScopeEmbedding",
"module_name": "dashscope_embedding",
"model_name": "text-embedding-v2",
"clazz": "models.llama_index_embedding_model"
}
emb = LlamaIndexEmbeddingModel(**config).model
emb = LlamaIndexEmbeddingModel(**config)
config = {
"index_name" : "0625_3",
"es_url" : "http://localhost:9200",
"embedding_model" : emb,
"index_name": "0626_1",
"es_url": "http://localhost:9200",
"embedding_model": emb,
}
self.es_store = LlamaIndexElasticSearchStore(**config)
self.data = [
MemoryNode(
content="The lives of two mob hitmen, a boxer, a gangster and his wife, and a pair of diner bandits intertwine in four tales of violence and redemption.",
content="The lives of two mob hitmen, a boxer, a gangster and his wife, "
"and a pair of diner bandits intertwine in four tales of violence and redemption.",
memory_type="observation",
id="0"
user_id="0",
status="valid",
memory_id="aaa123",
),
MemoryNode(
content="When the menace known as the Joker wreaks havoc and chaos on the people of Gotham, Batman must accept one of the greatest psychological and physical tests of his ability to fight injustice.",
content="When the menace known as the Joker wreaks havoc and chaos on the people of Gotham, "
"Batman must accept one of the greatest psychological and physical tests of his "
"ability to fight injustice.",
memory_type="observation",
id="1"
user_id="1",
status="valid",
memory_id="bbb456",
meta_data={"1": "1"}
),
MemoryNode(
content="An insomniac office worker and a devil-may-care soapmaker form an underground fight club that evolves into something much, much more.",
content="An insomniac office worker and a devil-may-care soapmaker form an underground fight "
"club that evolves into something much, much more.",
memory_type="insights",
id="2"
user_id="2",
status="valid",
memory_id="ccc789",
meta_data={"2": "2"}
),
MemoryNode(
content="A thief who steals corporate secrets through the use of dream-sharing technology is given the inverse task of planting an idea into thed of a C.E.O.",
content="A thief who steals corporate secrets through the use of dream-sharing technology "
"is given the inverse task of planting an idea into thed of a C.E.O.",
memory_type="insights",
id="3"
user_id="3",
status="valid",
memory_id="ddd012",
meta_data={"3": "3"}
),
MemoryNode(
content="A computer hacker learns from mysterious rebels about the true nature of his reality and his role in the war against its controllers.",
content="A computer hacker learns from mysterious rebels about the true nature of his reality "
"and his role in the war against its controllers.",
memory_type="profile",
id="4"
user_id="4",
status="valid",
memory_id="eee345",
meta_data={"4": "4"}
),
MemoryNode(
content="Two detectives, a rookie and a veteran, hunt a serial killer who uses the seven deadly sins as his motives.",
content="Two detectives, a rookie and a veteran, hunt a serial killer who uses the seven "
"deadly sins as his motives.",
memory_type="profile",
id="5"
user_id="5",
status="valid",
memory_id="fff678",
meta_data={"5": "5"},
),
MemoryNode(
content="An organized crime dynasty's aging patriarch transfers control of his clandestine empire to his reluctant son.",
content="An organized crime dynasty's aging patriarch transfers control of his clandestine "
"empire to his reluctant son.",
memory_type="insights",
id="6"),
user_id="6",
status="valid",
memory_id="ggg901",
meta_data={"5": "5"}
),
MemoryNode(
content="ggggggggg",
memory_type="profile",
id="6"),
]
# @unittest.skip("tmp")
def test_insert(self, ):
for node in self.data:
self.es_store.insert(node)
user_id="6",
status="valid",
memory_id="ggg234",
meta_data={"5": "5"}
# @unittest.skip("tmp")
def test_retrieve(self, ):
filter = {
"id": ["1", "2", "3"],
"memory_type": "insights",
),
]
def test_retrieve(self):
filter_dict = {
"user_id": "6",
}
res = self.es_store.retrieve(query="hacker", filter_dict=filter, top_k=10)
for node in self.data:
self.es_store.insert(node)
self.es_store.insert(MemoryNode(
content="xxxxxx",
memory_type="profile",
user_id="6",
status="valid",
memory_id="ggg567",
meta_data={"5": "5"}
))
res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10)
print(len(res))
print(res)
self.es_store.update(MemoryNode(
content="test update",
memory_type="profile",
user_id="6",
status="invalid",
memory_id="ggg567"
))
res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10)
print(len(res))
print(res)
self.es_store.delete(MemoryNode(
content="test update",
memory_type="profile",
user_id="6",
status="invalid",
memory_id="ggg567"
))
res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10)
print(len(res))
print(res)
def tearDown(self):
self.es_store.close()