[dev] add memoryscope to class path

This commit is contained in:
hs 2024-06-27 12:22:00 +08:00
parent 1b088fcca2
commit 6d30e4b18a
49 changed files with 1911 additions and 540 deletions

View file

@ -5,23 +5,23 @@ global_config:
open_ai_apikey:
memory_chat:
cli_memory_chat:
class: chat_v2.cli_memory_chat
class: memory_scope.chat_v2.cli_memory_chat
memory_service: memory_chat_service
generation_model: dashscope_generation
memory_service:
memory_chat_service:
class: memory.service.chat_memory_service
class: memory_scope.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
class: memory_scope.memory.operation.read_memory
workflow: dummy_worker
description: "read session messages of the user"
contextual_msg_count: 0
read_memory:
class: memory.operation.read_memory
class: memory_scope.memory.operation.read_memory
workflow: dummy_worker
description: "read related memories of the user"
list_memory:
@ -34,7 +34,7 @@ memory_service:
description: "write observation memories of the user"
interval_time: 60
summary_memory:
class: memory.operation.summary_memory
class: memory_scope.memory.operation.summary_memory
workflow: dummy_worker
description: "summary observation memories of the user"
interval_time: 300
@ -61,4 +61,5 @@ workers:
clazz: memory.worker.dummy_worker
generation_model: dashscope_generation
embedding_model: dashscope_embedding
rank_model: dashscope_rank
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

@ -5,7 +5,6 @@ class BaseMemoryChat(metaclass=ABCMeta):
def __init__(self, **kwargs):
self.kwargs = kwargs
@abstractmethod
def chat_with_memory(self, query: str):
"""

View file

@ -1,4 +1,54 @@
class BaseMemoryService(object):
def __init__(self, **kwargs):
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
self._init_operation(memory_operations)
@abstractmethod
def _init_operation(self, memory_operations: Dict[str, dict]):
raise NotImplementedError
@abstractmethod
def add_messages(self, messages: List[Message] | Message):
raise NotImplementedError
def prepare_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.operate(self.read_memory_key)
# def __getattr__(self, key):
# return self.kwargs[key]

View file

@ -0,0 +1,50 @@
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
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 prepare_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()

View file

@ -1,83 +1,139 @@
import datetime
import time
from typing import Dict, 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 GlobalContext
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 ..models.model_response import ModelResponse, ModelResponseGen
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, **kwargs):
super().__init__(**kwargs)
self._memory_service: BaseMemoryService | str = memory_service
self._generation_model: BaseModel | str = generation_model
self.stream: bool = True
@property
def memory_service(self) -> BaseMemoryService:
if isinstance(self._memory_service, str):
self._memory_service = GlobalContext.memory_service_dict[self._memory_service]
self._memory_service.prepare_service()
return self._memory_service
@property
def generation_model(self) -> BaseModel:
if isinstance(self._generation_model, str):
self._generation_model = GlobalContext.model_dict[self._generation_model]
return self._generation_model
@staticmethod
def get_system_prompt(related_memories: List[str], time_created: int) -> Message:
system_prompt = SYSTEM_PROMPT[GlobalContext.language]
if related_memories:
memory_prompt = MEMORY_PROMPT[GlobalContext.language]
all_prompt_list = [system_prompt, memory_prompt]
all_prompt_list.extend(related_memories)
system_prompt = "\n".join([x.strip() for x in all_prompt_list])
return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt, time_created=time_created)
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, content=query, time_created=time_created)
self.submit_messages(new_message)
related_memories: List[str] = self.memory_service.read_memory()
system_message: Message = self.get_system_prompt(related_memories, time_created)
if self.stream:
for result in self.generation_model.call(messages=[system_message, new_message], stream=self.stream):
yield result
def retrieve_all(self): # for testing
return "memory 1. 2. 3."
self.submit_messages(result.text)
def run(self):
console = Console()
op_description_dict: Dict[str, str] = self.memory_service.get_op_description_dict()
self.USER_COMMANDS.update({f"/{k}": v for k, v in op_description_dict.items()})
while True:
query = questionary.text(
"Enter your message or command:",
"Please enter your message or command:",
multiline=False,
qmark=">",
).ask()
query = query.rstrip()
query: str = query.rstrip()
if query == "":
console.print("Empty input received. Try again!")
print("Empty input received. Please try again!")
continue
# Handle CLI commands
# handle cli / commands with memory ops
if query.startswith("/"):
if query.lower() == "/exit":
query_split = query.lstrip("/").lower().split(" ")
query = query_split[0]
args = query_split[1:]
if query == "exit":
break
elif query.lower() == "/memory":
console.print(self.memory_service.retrieve_all())
elif query.lower() == "/help":
elif query == "help":
questionary.print("CLI commands", "bold")
for cmd, desc in self.USER_COMMANDS.items():
questionary.print(cmd, "bold")
questionary.print(f" {desc}")
print(f" {desc}")
elif query == "stream":
questionary.print(f"stream: {self.stream}")
self.stream = ~self.stream
elif query in op_description_dict:
if not args:
result = self.memory_service.do_operation(op_name=query)
print(result)
elif args[0].isdigit():
refresh_time = int(args[0])
try:
while True:
time.sleep(refresh_time)
result = self.memory_service.do_operation(op_name=query)
print(result, flush=True)
except KeyboardInterrupt:
print("stop refresh!")
else:
print("unknown command received. Please try again!")
else:
print("unknown command received. Please try again!")
continue
while True:
try:
# with console.status("[bold cyan]Thinking..."):
for msg in self.chat_with_memory(query=query):
console.print(msg.delta, end="")
console.print()
if self.stream:
for msg in self.chat_with_memory(query=query):
print(msg.text, flush=True)
print()
else:
msg = self.chat_with_memory(query=query)
print(msg.text)
break
except KeyboardInterrupt:
console.print("User interrupt occurred.")
questionary.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}"
)
questionary.print(f"An exception occurred when running chat_with_memory(): {e}")
retry = questionary.confirm("Retry chat_with_memory()?").ask()
if not retry:
break

View file

@ -1,19 +1,19 @@
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 .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
class GlobalContext(object):
def __init__(self):
self.global_configs: Dict[str, Any] = {}
self.worker_dict: Dict[str, Dict[str, BaseWorker]] = {}
self.worker_config: Dict[str, Dict[str, BaseWorker]] = {}
self.model_dict: Dict[str, BaseModel] = {}

View file

@ -43,7 +43,7 @@ class MemoryChat(BaseMemoryChat):
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 :]
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)

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

@ -3,8 +3,6 @@ import time
from typing import Dict, List
import questionary
from rich.console import Console
from memory_scope.chat_v2.base_memory_chat import BaseMemoryChat
from memory_scope.chat_v2.global_context import G_CONTEXT
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
@ -12,18 +10,21 @@ 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 ..models.model_response import ModelResponse, ModelResponseGen
class CliMemoryChat(BaseMemoryChat):
USER_COMMANDS = {
"exit": "exit the CLI",
"help": "get cli commands help",
"stream": "get stream response"
}
def __init__(self, memory_service: str, generation_model: str, **kwargs):
super().__init__(**kwargs)
self._memory_service: BaseMemoryService | str = memory_service
self._generation_model: BaseModel | str = generation_model
self.stream: bool = True
@property
def memory_service(self) -> BaseMemoryService:
@ -48,79 +49,91 @@ class CliMemoryChat(BaseMemoryChat):
system_prompt = "\n".join([x.strip() for x in all_prompt_list])
return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt, time_created=time_created)
def chat_with_memory(self, query: str):
def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen:
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)
self.submit_messages(new_message)
related_memories: List[str] = self.memory_service.read_memory()
system_message: Message = self.get_system_prompt(related_memories, time_created)
return self.generation_model.call(messages=[system_message, new_message], stream=True)
if self.stream:
for result in self.generation_model.call(messages=[system_message, new_message], stream=self.stream):
yield result
self.submit_messages(result.text)
def run(self):
op_description_dict: Dict[str, str] = self.memory_service.op_description_dict
self.USER_COMMANDS.update({f"/{k}": v for k, v in op_description_dict.items()})
console = Console()
while True:
query = questionary.text(
"Please enter your message or command:",
multiline=False,
qmark=">",
).ask()
query: str = query.rstrip()
if query == "":
console.print("Empty input received. Please try again!")
continue
# handle cli / commands with memory ops
if query.startswith("/"):
query_split = query.lstrip("/").lower().split(" ")
query = query_split[0]
args = query_split[1:]
if query == "exit":
break
elif query == "help":
questionary.print("CLI commands", "bold")
for cmd, desc in self.USER_COMMANDS.items():
questionary.print(cmd, "bold")
questionary.print(f" {desc}")
elif query in op_description_dict:
if not args:
result = self.memory_service.operate(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.operate(op_name=query)
questionary.print(result)
else:
console.print("unknown command received. Please try again!")
else:
console.print("unknown command received. Please try again!")
continue
def run(self):
op_description_dict: Dict[str, str] = self.memory_service.get_op_description_dict()
self.USER_COMMANDS.update({f"/{k}": v for k, v in op_description_dict.items()})
while True:
try:
# with console.status("[bold cyan]Thinking..."):
for msg in self.chat_with_memory(query=query):
console.print(msg.delta, end="")
console.print()
break
except KeyboardInterrupt:
console.print("User interrupt occurred.")
retry = questionary.confirm("Retry chat_with_memory()?").ask()
if not retry:
query = questionary.text(
"Please enter your message or command:",
multiline=False,
qmark=">",
).ask()
query: str = query.rstrip()
if query == "":
print("Empty input received. Please try again!")
continue
# handle cli / commands with memory ops
if query.startswith("/"):
query_split = query.lstrip("/").lower().split(" ")
query = query_split[0]
args = query_split[1:]
if query == "exit":
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:
elif query == "help":
questionary.print("CLI commands", "bold")
for cmd, desc in self.USER_COMMANDS.items():
questionary.print(cmd, "bold")
print(f" {desc}")
elif query == "stream":
questionary.print(f"stream: {self.stream}")
self.stream = ~self.stream
elif query in op_description_dict:
if not args:
result = self.memory_service.do_operation(op_name=query)
print(result)
elif args[0].isdigit():
refresh_time = int(args[0])
try:
while True:
time.sleep(refresh_time)
result = self.memory_service.do_operation(op_name=query)
print(result, flush=True)
except KeyboardInterrupt:
print("stop refresh!")
else:
print("unknown command received. Please try again!")
else:
print("unknown command received. Please try again!")
continue
while True:
try:
if self.stream:
for msg in self.chat_with_memory(query=query):
print(msg.text, flush=True)
print()
else:
msg = self.chat_with_memory(query=query)
print(msg.text)
break
except KeyboardInterrupt:
questionary.print("User interrupt occurred.")
retry = questionary.confirm("Retry chat_with_memory()?").ask()
if not retry:
break
except Exception as e:
questionary.print(f"An exception occurred when running chat_with_memory(): {e}")
retry = questionary.confirm("Retry chat_with_memory()?").ask()
if not retry:
break

View file

@ -24,5 +24,3 @@ class GlobalContext(pydantic.BaseModel):
thread_pool: ThreadPoolExecutor | None = pydantic.Field(None, description="global thread_pool")
language: LanguageEnum = pydantic.Field(LanguageEnum.CN, description="language: cn / en")
G_CONTEXT = GlobalContext()

View file

@ -1,135 +1,76 @@
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 typing import Dict, Any
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
import yaml
import fire
from .chat_v2.global_context import G_CONTEXT
from .enumeration.language_enum import LanguageEnum
from .utils.logger import Logger
from .utils.tool_functions import init_instance_by_config
class CliJob(object):
def __init__(self, config_path: str):
def __init__(self, config_path: str, config_suffix: str = ".yaml"):
self.config_path: str = config_path
self.config_base_dir: str = os.path.dirname(config_path)
self.config_suffix: str = config_suffix
self.config: Dict[str, Any] = {}
self.worker_chat_dict: Dict[str, List[str]] = {}
self.logger: Logger = Logger.get_logger("memory_chat")
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))
self.logger: Logger = Logger.get_logger("cli_job")
@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(global_config: Dict[str, Any]):
"""set global_configs & set apikey into env
:return:
TODO at sen
"""
G_CONTEXT.global_config = global_config
G_CONTEXT.language = LanguageEnum(global_config["language"])
G_CONTEXT.thread_pool = ThreadPoolExecutor(
max_workers=int(global_config["max_workers"])
)
def init_global_content_by_config(self):
with open(complete_config_name(self.config_path)) as f:
self.config = json.load(f)
# load config
config_path = self.config_path
if not self.config_path.endswith(self.config_suffix):
config_path += self.config_suffix
with open(config_path) as f:
self.config = yaml.load(f, yaml.FullLoader)
GLOBAL_CONTEXT.global_configs = self.config["global_configs"]
self.set_global_config()
# set global_config
self.set_global_config(self.config["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)
# init vector_store
G_CONTEXT.vector_store = init_instance_by_config(
self.config["vector_store"]
)
# init monitor
G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
# set worker config
G_CONTEXT.worker_config = self.config["workers"]
@staticmethod
def run():
with GLOBAL_CONTEXT.thread_pool:
memory_chat = list(GLOBAL_CONTEXT.memory_chat_dict.values())[0]
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()
if __name__ == "__main__":
fire.Fire(main)

View file

@ -1,67 +0,0 @@
from concurrent.futures import ThreadPoolExecutor
from typing import Dict, Any
import yaml
from chat_v2.global_context import G_CONTEXT
from enumeration.language_enum import LanguageEnum
from utils.logger import Logger
from utils.tool_functions import init_instance_by_config
class CliJob(object):
def __init__(self, config_path: str, config_suffix: str = ".yaml"):
self.config_path: str = config_path
self.config_suffix: str = config_suffix
self.config: Dict[str, Any] = {}
self.logger: Logger = Logger.get_logger("cli_job")
@staticmethod
def set_global_config(global_config: Dict[str, Any]):
""" set global_configs & set apikey into env
:return:
TODO at sen
"""
G_CONTEXT.global_config = global_config
G_CONTEXT.language = LanguageEnum(global_config["language"])
G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(global_config["max_workers"]))
def init_global_content_by_config(self):
# load config
config_path = self.config_path
if not self.config_path.endswith(self.config_suffix):
config_path += self.config_suffix
with open(config_path) as f:
self.config = yaml.load(f, yaml.FullLoader)
# set global_config
self.set_global_config(self.config["global_config"])
# 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)
# 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)
# init models
for name, conf in self.config["models"].items():
G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name)
# init vector_store
G_CONTEXT.vector_store = init_instance_by_config(self.config["vector_store"])
# init monitor
G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
# set worker config
G_CONTEXT.worker_config = self.config["workers"]
@staticmethod
def run():
with G_CONTEXT.thread_pool:
memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0]
memory_chat.run()

View file

@ -22,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"
@ -116,3 +112,5 @@ DATATIME_KEY_MAP = {
"": "week",
"星期几": "weekday",
}
CONTENT_MODIFIED = "content_modified"

View file

@ -23,8 +23,6 @@ class BaseMemoryService(metaclass=ABCMeta):
self.logger = Logger.get_logger()
self.kwargs = kwargs
self._init_operation(memory_operations)
@abstractmethod
def _init_operation(self, memory_operations: Dict[str, dict]):
raise NotImplementedError

View file

@ -6,26 +6,28 @@ 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):
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)
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):

View file

@ -19,7 +19,9 @@ class LlamaIndexRerankModel(BaseModel):
query: str = kwargs.pop("query", "")
documents: List[str] = kwargs.pop("documents", [])
assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}"
assert (
query and documents
), f"query or documents is empty! query={query}, documents={len(documents)}"
# using -1.0 as dummy scores
nodes = [NodeWithScore(node=Node(text=doc), score=-1.0) for doc in documents]
@ -41,7 +43,9 @@ class LlamaIndexRerankModel(BaseModel):
return model_response
def _call(self, **kwargs) -> ModelResponse:
return ModelResponse(m_type=self.m_type, raw=self.model.postprocess_nodes(**self.data))
return ModelResponse(
m_type=self.m_type, raw=self.model.postprocess_nodes(**self.data)
)
async def _async_call(self, **kwargs) -> ModelResponse:
raise NotImplementedError

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

@ -0,0 +1,7 @@
from ..enumeration.language_enum import LanguageEnum
UPDATE_INSIGHT_SYSTEM_PROMPT = {}
UPDATE_INSIGHT_FEW_SHOT_PROMPT = {}
UPDATE_INSIGHT_USER_QUERY_PROMPT = {}

View file

@ -4,7 +4,7 @@ from 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,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

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

@ -1,20 +1,20 @@
from typing import List
from typing import List, Dict
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
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
from ..scheme.memory_node import MemoryNode
from ..constants import common_constants
class MemoryBaseWorker(BaseWorker):
def __init__(self,
embedding_model: str,
generation_model: str,
rank_model: str,
**kwargs):
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
@ -40,25 +40,29 @@ class MemoryBaseWorker(BaseWorker):
return self.get_context(CHAT_NAME)
@property
def embedding_model(self):
def embedding_model(self) -> BaseModel:
if self._embedding_model is None:
self._embedding_model = GLOBAL_CONTEXT.model_dict.get(self.embedding_model_name)
self._embedding_model = GLOBAL_CONTEXT.model_dict.get(
self.embedding_model_name
)
return self._embedding_model
@property
def generation_model(self):
def generation_model(self) -> BaseModel:
if self._generation_model is None:
self._generation_model = GLOBAL_CONTEXT.model_dict.get(self.generation_model_name)
self._generation_model = GLOBAL_CONTEXT.model_dict.get(
self.generation_model_name
)
return self._generation_model
@property
def rank_model(self):
def rank_model(self) -> BaseModel:
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):
def vector_store(self) -> BaseVectorStore:
if self._vector_store is None:
self._vector_store = GLOBAL_CONTEXT.vector_store
return self._vector_store
@ -68,3 +72,22 @@ class MemoryBaseWorker(BaseWorker):
if self._monitor is None:
self._monitor = GLOBAL_CONTEXT.monitor
return self._monitor
@property
def user_profile_dict(self) -> Dict[str, MemoryNode]:
if not self._user_profile_dict:
self._user_profile_dict = {
user_attr.meta_data.get("memory_key", ""): user_attr
for user_attr in self.get_context(common_constants.USER_PROFILE)
}
return self._user_profile_dict
@property
def memory_id(self) -> str:
pass
def __getattr__(self, key):
return self.kwargs[key]
def get_prompt(self, x):
return x[GLOBAL_CONTEXT.global_configs["language"]]

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

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

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

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