memory stream & some workers

This commit is contained in:
hs 2024-06-25 19:20:10 +08:00
parent 55ca839e95
commit e84733f810
62 changed files with 730 additions and 617 deletions

16
.vscode/launch.json vendored Normal file
View file

@ -0,0 +1,16 @@
{
// 使 IntelliSense
//
// 访: https://go.microsoft.com/fwlink/?linkid=830387
"version": "0.2.0",
"configurations": [
{
"name": "Python 调试程序: 当前文件",
"type": "debugpy",
"request": "launch",
"program": "${file}",
"console": "integratedTerminal",
"justMyCode": false
}
]
}

View file

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

View file

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

View file

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

View file

@ -1,12 +1,10 @@
from abc import ABCMeta, abstractmethod
from memory_scope.chat.memory_service import MemoryService
class BaseMemoryChat(metaclass=ABCMeta):
def __init__(self, **kwargs):
self.kwargs = kwargs
def __init__(self, chat_name: str, **kwargs):
self.memory_service = MemoryService(chat_name=chat_name, **kwargs)
@abstractmethod
def chat_with_memory(self, query: str):

View file

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

View file

@ -0,0 +1,83 @@
import datetime
import questionary
from rich.console import Console
from .memory_chat import MemoryChat
from enumeration.message_role_enum import MessageRoleEnum
from scheme.message import Message
class CliMemoryChat(MemoryChat):
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
}
def chat_with_memory(self, query): # for testing
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)
def retrieve_all(self): # for testing
return "memory 1. 2. 3."
def run(self):
console = Console()
while True:
query = questionary.text(
"Enter your message or command:",
multiline=False,
qmark=">",
).ask()
query = query.rstrip()
if query == "":
console.print("Empty input received. Try again!")
continue
# 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..."):
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:
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

View file

@ -1,12 +1,12 @@
from concurrent.futures import ThreadPoolExecutor
from typing import Dict, Any
from memory_scope.chat.base_memory_chat import BaseMemoryChat
from memory_scope.enumeration.language_enum import LanguageEnum
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
from memory_scope.worker.base_worker import BaseWorker
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
class GlobalContext(object):

View file

@ -1,21 +1,20 @@
import datetime
from typing import List
from memory_scope.chat.base_memory_chat import BaseMemoryChat
from memory_scope.chat.global_context import GLOBAL_CONTEXT
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
from memory_scope.models.base_model import BaseModel
from memory_scope.prompts.memory_chat_prompt import SYSTEM_PROMPT, MEMORY_PROMPT
from memory_scope.scheme.message import Message
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,
**kwargs):
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
@ -25,7 +24,9 @@ class MemoryChat(BaseMemoryChat):
@property
def generation_model(self):
if self._generation_model is None:
self._generation_model = GLOBAL_CONTEXT.model_dict[self.generation_model_name]
self._generation_model = GLOBAL_CONTEXT.model_dict[
self.generation_model_name
]
return self._generation_model
@staticmethod
@ -34,7 +35,11 @@ class MemoryChat(BaseMemoryChat):
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)
return Message(
role=MessageRoleEnum.SYSTEM,
content=system_prompt.strip(),
time_created=time_created,
)
def chat_with_memory(self, query: str):
query = query.strip()
@ -42,11 +47,13 @@ class MemoryChat(BaseMemoryChat):
return
time_created = int(datetime.datetime.now().timestamp())
new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created)
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:]
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,43 +1,52 @@
from memory_scope.constants.common_constants import RELATED_MEMORIES
from memory_scope.enumeration.memory_method_enum import MemoryMethodEnum
from memory_scope.scheme.message import Message
from memory_scope.utils.pipeline import Pipeline
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(object):
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,
)
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):
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.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_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)
self.kwargs = kwargs
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)

View file

@ -2,88 +2,27 @@ import json
import os
from concurrent.futures import ThreadPoolExecutor
from typing import Dict, Any, List
import questionary
from rich.console import Console
import sys
import time
import fire
from datetime import datetime
from memory_scope.chat.base_memory_chat import BaseMemoryChat
from memory_scope.chat.global_context import GLOBAL_CONTEXT
from memory_scope.enumeration.language_enum import LanguageEnum
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.utils.logger import Logger
from memory_scope.utils.tool_functions import (
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 memory_scope.chat.memory_chat import MemoryChat
class CliMemoryChat(object): # object -> MemoryChat
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
}
def chat_with_memory(self, query): # for testing
return query
def retrieve_all(self): # for testing
return "memory 1. 2. 3."
def run(self):
console = Console()
while True:
query = questionary.text(
"Enter your message or command:",
multiline=False,
qmark=">",
).ask()
query = query.rstrip()
if query == "":
console.print("Empty input received. Try again!")
continue
# 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
continue
while True:
try:
with console.status("[bold cyan]Thinking..."):
messages = self.chat_with_memory(query=query)
console.print(messages)
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
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
class CliJob(object):
@ -97,7 +36,7 @@ class CliJob(object):
self.logger: Logger = Logger.get_logger("memory_chat")
def init_memory_chat(self):
for chat_name in self.config["chat_list"]:
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
@ -172,10 +111,12 @@ class CliJob(object):
self.init_memory_chat()
self.init_workers()
GLOBAL_CONTEXT.vector_store = init_instance_by_config(
self.config["vector_store"]
)
GLOBAL_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
## 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"])
@staticmethod
def run():
@ -191,4 +132,4 @@ def main(config_path: str):
if __name__ == "__main__":
fire.Fire(main)
fire.Fire(main)

View file

@ -3,3 +3,110 @@ RELATED_MEMORIES = "related_memories"
MESSAGES = "messages"
CHAT_NAME = "chat_name"
PIPELINE = "pipeline"
WORKER = "worker"
MEMORY = "memory"
DEFAULT_SYSTEM_PROMPT = "default_system_prompt"
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"
NEW_OBS_NODES = "new_obs_nodes"
NEW_OBS_WITH_TIME_NODES = "new_obs_with_time_nodes"
INSIGHT_NODES = "insight_nodes"
MERGE_OBS_NODES = "merge_obs_nodes"
NEW_INSIGHT_NODES = "new_insight_nodes"
TODAY_OBS_NODES = "today_obs_nodes"
SIMILAR_OBS_NODES = "similar_obs_nodes"
KEYWORD_OBS_NODES = "keyword_obs_nodes"
NOT_REFLECTED_OBS_NODES = "not_reflected_obs_nodes"
NOT_REFLECTED_MERGE_NODES = "not_reflected_merge_nodes"
NEW_INSIGHT_KEYS = "new_insight_keys"
INSIGHT_KEY = "insight_key"
INSIGHT_VALUE = "insight_value"
DT = "dt"
MSG_TIME = "msg_time"
NEW = "new"
TIME_INFER = "time_infer"
KEY_WORD = "key_word"
REFLECTED = "reflected"
NEW_USER_PROFILE = "new_user_profile"
RECALL_TYPE = "recall_type"
ALL_ONLINE_NODES = "all_online_nodes"
MAX_WORKERS = "max_workers"
TIME_MATCHED = "time_matched"
QUERY_KEYWORDS = "query_keywords"
WEEKDAYS = ["周一", "周二", "周三", "周四", "周五", "周六", "周日"]
DATATIME_WORD_LIST = [
"",
"",
"",
"",
"星期",
"",
"分钟",
"小时",
"",
"上午",
"下午",
"早上",
"早晨",
"晚上",
"中午",
"",
"",
"清晨",
"傍晚",
"凌晨",
"",
]
TIME_FORMAT_V1 = "{year}{month}{day}{weekday}{hour}"
DATATIME_KEY_MAP = {
"": "year",
"": "month",
"": "day",
"": "week",
"星期几": "weekday",
}

View file

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

View file

@ -2,11 +2,11 @@ import inspect
import time
from abc import abstractmethod, ABCMeta
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.models import MODEL_REGISTRY
from memory_scope.models.response import ModelResponse, ModelResponseGen
from memory_scope.utils.logger import Logger
from memory_scope.utils.timer import Timer
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
class BaseModel(metaclass=ABCMeta):

View file

@ -2,18 +2,20 @@ from typing import List
from llama_index.embeddings.dashscope import DashScopeEmbedding
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.models import MODEL_REGISTRY
from memory_scope.models.base_model import BaseModel
from memory_scope.models.response import ModelResponse
from models import MODEL_REGISTRY
from models.base_model import BaseModel
from models.response import ModelResponse, ModelResponseGen
from enumeration.model_enum import ModelEnum
class LlamaIndexEmbeddingModel(BaseModel):
m_type: ModelEnum = ModelEnum.EMBEDDING_MODEL
MODEL_REGISTRY.batch_register([
DashScopeEmbedding,
])
MODEL_REGISTRY.batch_register(
[
DashScopeEmbedding,
]
)
def before_call(self, **kwargs):
text: str | List[str] = kwargs.pop("text", "")
@ -40,11 +42,16 @@ 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,16 +1,15 @@
from typing import List, Dict
from llama_index.core.base.llms.types import ChatMessage
from llama_index.core.base.llms.types import (
ChatMessage,
ChatResponse,
CompletionResponse,
)
from llama_index.llms.dashscope import DashScope
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.models import MODEL_REGISTRY
from memory_scope.models.base_model import BaseModel
from memory_scope.models.response import ModelResponse, ModelResponseGen
from enumeration.model_enum import ModelEnum
from . import MODEL_REGISTRY
from .base_model import BaseModel
from .response import ModelResponse, ModelResponseGen
class LlamaIndexGenerationModel(BaseModel):
@ -32,7 +31,7 @@ class LlamaIndexGenerationModel(BaseModel):
elif messages:
input_text = messages
input_type = 'messages'
llama_input = [ChatMessage(role=x['role'], content=x['content']) for x in input_text]
llama_input = [ChatMessage(role=x.role, content=x.content) for x in input_text]
else:
raise RuntimeError("prompt and messages is both empty!")

View file

@ -4,10 +4,11 @@ from llama_index.core.data_structs import Node
from llama_index.core.schema import NodeWithScore
from llama_index.postprocessor.dashscope_rerank import DashScopeRerank
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.models import MODEL_REGISTRY
from memory_scope.models.base_model import BaseModel
from memory_scope.models.response import ModelResponse
from models import MODEL_REGISTRY
from models.base_model import BaseModel
from models.response import ModelResponse, ModelResponseGen
from enumeration.model_enum import ModelEnum
class LlamaIndexRerankModel(BaseModel):

View file

@ -3,7 +3,7 @@ from typing import Generator, List, Dict, Any
from pydantic import BaseModel, Field
from memory_scope.enumeration.model_enum import ModelEnum
from enumeration.model_enum import ModelEnum
class ModelResponse(BaseModel):

View file

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

View file

@ -1,8 +1,7 @@
from abc import ABCMeta, abstractmethod
from typing import Dict, List
from memory_scope.models.base_model import BaseModel
from memory_scope.scheme.memory_node import MemoryNode
from models.base_model import BaseModel
class BaseVectorStore(metaclass=ABCMeta):
@ -21,6 +20,7 @@ class BaseVectorStore(metaclass=ABCMeta):
:param filter_dict:
:return:
"""
pass
@abstractmethod
async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]):
@ -30,27 +30,32 @@ class BaseVectorStore(metaclass=ABCMeta):
:param filter_dict:
:return:
"""
pass
@abstractmethod
def insert(self, node: MemoryNode):
""" TODO 是否overwrite
:return:
"""
pass
@abstractmethod
def insert_batch(self):
"""
:return:
"""
pass
@abstractmethod
def delete(self):
"""
:return:
"""
pass
@abstractmethod
def flush(self):
"""
:return:
"""
pass

View file

@ -5,13 +5,13 @@ from concurrent.futures import as_completed
from itertools import zip_longest
from typing import Dict, Any, List
from memory_scope.chat.global_context import GLOBAL_CONTEXT
from memory_scope.constants.common_constants import MESSAGES, CHAT_NAME
from memory_scope.enumeration.memory_method_enum import MemoryMethodEnum
from memory_scope.scheme.message import Message
from memory_scope.utils.logger import Logger
from memory_scope.utils.timer import Timer
from memory_scope.worker.base_worker import BaseWorker
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):

View file

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

View file

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

View file

@ -1,15 +1,18 @@
import re
from importlib import import_module
from datetime import datetime
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
from enumeration.message_role_enum import MessageRoleEnum
def under_line_to_hump(underline_str):
sub = re.sub(r'(_\w)', lambda x: x.group(1)[1].upper(), underline_str)
sub = re.sub(r"(_\w)", lambda x: x.group(1)[1].upper(), underline_str)
return sub[0:1].upper() + sub[1:]
def init_instance_by_config(config: dict, default_clazz_path: str = "", suffix_name: str = "", **kwargs):
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!")
@ -44,6 +47,44 @@ def prompt_to_msg(system_prompt: str, few_shot: str, user_query: str):
},
{
"role": MessageRoleEnum.USER.value,
"content": "\n".join([x.strip() for x in [few_shot, system_prompt, user_query]])
"content": "\n".join(
[x.strip() for x in [few_shot, system_prompt, user_query]]
),
},
]
def get_datetime_info_dict(parse_dt: datetime):
return {
"year": parse_dt.year,
"month": parse_dt.month,
"day": parse_dt.day,
"hour": parse_dt.hour,
"minute": parse_dt.minute,
"second": parse_dt.second,
"week": parse_dt.isocalendar().week,
"weekday": WEEKDAYS[parse_dt.isocalendar().weekday - 1],
}
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
else:
current_dt = datetime.now()
return_str = ""
if date_format:
return_str = current_dt.strftime(date_format)
elif string_format:
return_str = string_format.format(**get_datetime_info_dict(current_dt))
return return_str

View file

@ -1,7 +1,7 @@
from typing import Any, Dict
from memory_scope.utils.logger import Logger
from memory_scope.utils.timer import Timer
from utils.logger import Logger
from utils.timer import Timer
class BaseWorker(object):
@ -63,6 +63,9 @@ class BaseWorker(object):
else:
self.context_dict[key] = value
def __getattr__(self, key):
return self.kwargs[key]
@property
def name_simple(self) -> str:
if not self._name_simple:

View file

@ -0,0 +1,6 @@
from memory_base_worker import MemoryBaseWorker
class DummyWorker(MemoryBaseWorker):
def _run(self):
pass

View file

@ -0,0 +1,22 @@
from typing import List
from constants.common_constants import INSIGHT_NODES
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
from cli import GLOBAL_CONTEXT
class EsInsightWorker(MemoryBaseWorker):
def _run(self):
insight_nodes = self.vector_store.retrieve(
size=self.kwargs.es_insight_top_k,
filter_dict={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memoryType": MemoryTypeEnum.INSIGHT.value,
},
)
self.logger.info(f"insight_nodes.size={len(insight_nodes)}")
self.set_context(INSIGHT_NODES, insight_nodes)

View file

@ -0,0 +1,22 @@
from typing import List
from constants.common_constants import NEW, NEW_OBS_NODES
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class EsNewObsWorker(MemoryBaseWorker):
def _run(self):
new_obs_nodes = self.vector_store.retrieve(
size=self.kwargs.es_new_obs_top_k,
filter_dict={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memoryType": MemoryTypeEnum.OBSERVATION.value,
f"metaData.{NEW}": "1",
},
)
self.logger.info(f"es new obs, size={len(new_obs_nodes)}")
self.set_context(NEW_OBS_NODES, new_obs_nodes)

View file

@ -0,0 +1,29 @@
from typing import List
from constants.common_constants import REFLECTED, NOT_REFLECTED_OBS_NODES
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class EsNotReflectedWorker(MemoryBaseWorker):
def _run(self):
not_reflected_obs_nodes = self.vector_store.retrieve(
size=self.kwargs.es_new_obs_top_k,
filter_dict={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memoryType": [
MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value,
],
f"metaData.{REFLECTED}": "0",
},
)
self.logger.info(
f"retrieve_not_reflected_obs.size={len(not_reflected_obs_nodes)}"
)
self.set_context(NOT_REFLECTED_OBS_NODES, not_reflected_obs_nodes)

View file

@ -0,0 +1,37 @@
from typing import List
from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_recall_type import MemoryRecallType
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class EsSimilarWorker(MemoryBaseWorker):
def __init__(self, es_similar_top_k, *args, **kwargs):
super(EsSimilarWorker, self).__init__(*args, **kwargs)
self.es_similar_top_k = es_similar_top_k
def _run(self):
query = self.messages[-1].content
similar_obs_nodes = self.vector_store.retrieve(
text=query,
size=self.es_similar_top_k,
filter_dict={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memoryType": [
MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.INSIGHT.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value,
],
},
)
for node in similar_obs_nodes:
node.metaData[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}")
self.set_context(SIMILAR_OBS_NODES, similar_obs_nodes)

View file

@ -0,0 +1,32 @@
from typing import List
from utils.tool_functions import time_to_formatted_str
from constants.common_constants import TODAY_OBS_NODES, DT
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class EsTodayObsWorker(MemoryBaseWorker):
def __init__(self, es_today_obs_top_k, *args, **kwargs):
super(EsTodayObsWorker, self).__init__(*args, **kwargs)
self.es_today_obs_top_k = es_today_obs_top_k
def _run(self):
if not self.messages:
self.logger.warning("messages is empty!")
return
msg_time_created = self.messages[-1].time_created
today_obs_nodes = self.vector_store.retrieve(
size=self.es_today_obs_top_k,
filter_dict={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memoryType": MemoryTypeEnum.OBSERVATION.value,
f"metaData.{DT}": time_to_formatted_str(msg_time_created),
},
)
self.logger.info(f"retrieve_today_obs.size={len(today_obs_nodes)}")
self.set_context(TODAY_OBS_NODES, today_obs_nodes)

View file

@ -0,0 +1,25 @@
from typing import List, Dict
from constants import common_constants
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class LoadProfileWorker(MemoryBaseWorker):
def _run(self):
user_profile_node = self.vector_store(
size=10000,
filter_dict={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memoryType": [
MemoryTypeEnum.PROFILE.value,
MemoryTypeEnum.PROFILE_CUSTOMIZED.value,
],
},
)
self.set_context(common_constants.USER_PROFILE, user_profile_node)
self.logger.info(f"retrieve_user_profile.size={len(user_profile_node)}")

View file

@ -1,12 +1,12 @@
from typing import List
from memory_scope.chat.global_context import GLOBAL_CONTEXT
from memory_scope.constants.common_constants import MESSAGES, CHAT_NAME
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
from memory_scope.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
class MemoryBaseWorker(BaseWorker):

View file

@ -1,19 +1,11 @@
import re
from common.tool_functions import time_to_formatted_str
from constants.common_constants import DATATIME_WORD_LIST, DATATIME_KEY_MAP
from constants.common_constants import EXTRACT_TIME_DICT
from utils.tool_functions import time_to_formatted_str
from constants.common_constants import DATATIME_WORD_LIST, DATATIME_KEY_MAP, EXTRACT_TIME_DICT
from worker.memory_base_worker import MemoryBaseWorker
class ExtractTimeWorker(MemoryBaseWorker):
def __init__(self, parse_time_model, parse_time_max_token, parse_time_temperature, parse_time_top_k, *args, **kwargs):
super(ExtractTimeWorker, self).__init__(*args, **kwargs)
self.parse_time_model = parse_time_model
self.parse_time_max_token = parse_time_max_token
self.parse_time_temperature = parse_time_temperature
self.parse_time_top_k = parse_time_top_k
@staticmethod
def get_parse_time_prompt(query: str, query_time_str: str):
return f"""
@ -51,7 +43,8 @@ class ExtractTimeWorker(MemoryBaseWorker):
self.logger.info(f"extract_time_prompt={extract_time_prompt}")
# call sft model
response_text = self.gene_client.call(prompt=extract_time_prompt,
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,

View file

@ -2,7 +2,7 @@ from typing import Dict, List
from constants.common_constants import RELATED_MEMORIES, EXTRACT_TIME_DICT, ALL_ONLINE_NODES, \
TIME_MATCHED
from node.memory_wrap_node import MemoryWrapNode
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
@ -55,7 +55,7 @@ class FuseRerankWorker(MemoryBaseWorker):
def _run(self):
# 解析时间meta信息
extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT)
all_online_nodes: List[MemoryWrapNode] = self.get_context(ALL_ONLINE_NODES)
all_online_nodes: List[MemoryNode] = self.get_context(ALL_ONLINE_NODES)
if not all_online_nodes:
self.add_run_info("all_online_nodes is empty, stop")
@ -67,7 +67,7 @@ class FuseRerankWorker(MemoryBaseWorker):
continue
# 根据类型给ratio
type_ratio: float = self.fuse_ratio_dict.get(node.memory_node.memoryType, 0.1)
type_ratio: float = self.fuse_ratio_dict.get(scheme.memory_node.memoryType, 0.1)
# 时间系数,完全匹配才行
fuse_time_ratio: float = 1.0
@ -76,7 +76,7 @@ class FuseRerankWorker(MemoryBaseWorker):
if extract_time_dict:
match_event_flag = True
for k, v in extract_time_dict.items():
event_value = node.memory_node.metaData.get(f"event_{k}", "")
event_value = scheme.memory_node.metaData.get(f"event_{k}", "")
if event_value in ["-1", v]:
continue
else:
@ -85,7 +85,7 @@ class FuseRerankWorker(MemoryBaseWorker):
match_msg_flag = True
for k, v in extract_time_dict.items():
msg_value = node.memory_node.metaData.get(f"msg_{k}", "")
msg_value = scheme.memory_node.metaData.get(f"msg_{k}", "")
if msg_value == v:
continue
else:
@ -94,10 +94,10 @@ class FuseRerankWorker(MemoryBaseWorker):
if match_event_flag or match_msg_flag:
fuse_time_ratio = self.fuse_time_ratio
node.memory_node.metaData[TIME_MATCHED] = "1"
scheme.memory_node.metaData[TIME_MATCHED] = "1"
node.score_rerank = node.score_rank * type_ratio * fuse_time_ratio
self.logger.info(f"content={node.memory_node.content} f_event={int(match_event_flag)} "
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}")
filtered_nodes.append(node)
@ -106,18 +106,18 @@ class FuseRerankWorker(MemoryBaseWorker):
filtered_nodes = filtered_nodes[: self.output_max_count]
related_memories: List[str] = []
for node in filtered_nodes:
content = node.memory_node.content
content = scheme.memory_node.content
# 如果命中时间逻辑
if node.memory_node.metaData.get(TIME_MATCHED, "") == "1":
# time_infer = node.memory_node.metaData.get(TIME_INFER)
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=node.memory_node.metaData)
# meta_data=scheme.memory_node.metaData)
time_infer = self.format_time_infer(time_infer="",
extract_time_dict=extract_time_dict,
meta_data=node.memory_node.metaData)
meta_data=scheme.memory_node.metaData)
content = f"{time_infer}: {content}"
related_memories.append(content)

View file

@ -2,17 +2,17 @@ from typing import List
from utils.user_profile_handler import UserProfileHandler
from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE
from node.memory_node import MemoryNode
from node.memory_wrap_node import MemoryWrapNode
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[MemoryWrapNode] | 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], MemoryWrapNode):
if isinstance(modified_memories[0], MemoryNode):
modified_memories = [n.memory_node for n in modified_memories]
for n in modified_memories:

View file

@ -4,38 +4,33 @@ 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 node.memory_wrap_node import MemoryWrapNode
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class SemanticRankWorker(MemoryBaseWorker):
def user_profile_to_nodes(self) -> List[MemoryWrapNode]:
user_profile_nodes: List[MemoryWrapNode] = UserProfileHandler.to_nodes(self.user_profile_dict, split_value=True)
def user_profile_to_nodes(self) -> List[MemoryNode]:
user_profile_nodes: List[MemoryNode] = UserProfileHandler.to_nodes(self.user_profile_dict, split_value=True)
for node in user_profile_nodes:
# 从画像侧召回
node.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.PROFILE
self.logger.info(f"user profile node={node.memory_node.content}")
scheme.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.PROFILE
self.logger.info(f"user profile node={scheme.memory_node.content}")
return user_profile_nodes
def _run(self):
all_node_dict: Dict[str, MemoryWrapNode] = {}
all_node_dict: Dict[str, MemoryNode] = {}
# 优先级: similar_obs_nodes < keyword_obs_nodes < profile_nodes
similar_obs_nodes: List[MemoryWrapNode] = self.get_context(SIMILAR_OBS_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[node.memory_node.content] = node
all_node_dict[scheme.memory_node.content] = node
keyword_obs_nodes: List[MemoryWrapNode] = self.get_context(KEYWORD_OBS_NODES)
if keyword_obs_nodes:
for node in keyword_obs_nodes:
all_node_dict[node.memory_node.content] = node
profile_nodes: List[MemoryWrapNode] = self.user_profile_to_nodes()
profile_nodes: List[MemoryNode] = self.user_profile_to_nodes()
if profile_nodes:
for node in profile_nodes:
all_node_dict[node.memory_node.content] = node
all_node_dict[scheme.memory_node.content] = node
if not all_node_dict:
self.add_run_info(f"all_node_dict is empty!", continue_run=False)
@ -61,8 +56,8 @@ class SemanticRankWorker(MemoryBaseWorker):
content = documents[rank_node["index"]]
node = all_node_dict[content]
node.score_rank = rank_node["relevance_score"]
self.logger.info(f"query={query} content={node.memory_node.content} score_rank={node.score_rank}")
self.logger.info(f"query={query} content={scheme.memory_node.content} score_rank={node.score_rank}")
# save context
all_online_nodes: List[MemoryWrapNode] = list(all_node_dict.values())
all_online_nodes: List[MemoryNode] = list(all_node_dict.values())
self.set_context(ALL_ONLINE_NODES, all_online_nodes)

View file

@ -2,13 +2,13 @@ from elasticsearch import Elasticsearch
from elasticsearch.helpers import bulk
from memory_scope.models.dash_embedding_client import DashEmbeddingClient, LLIEmbedding
from models.dash_embedding_client import DashEmbeddingClient, LLIEmbedding
from common.dash_embedding_client import DashEmbeddingClient
from common.logger import Logger
from constants.common_constants import ES_ENV_URL_DICT
from enumeration.env_type import EnvType
from memory_scope.utils.logger import Logger
from utils.logger import Logger
from llama_index.core import VectorStoreIndex, StorageContext, ServiceContext
from llama_index.vector_stores.elasticsearch import ElasticsearchStore
from llama_index.core.schema import TextNode

View file

@ -1,26 +0,0 @@
from typing import List
from constants.common_constants import INSIGHT_NODES
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class EsInsightWorker(MemoryBaseWorker):
def __init__(self, es_insight_top_k, *args, **kwargs):
super(EsInsightWorker, self).__init__(*args, **kwargs)
self.es_insight_top_k = es_insight_top_k
def _run(self):
hits = self.es_client.exact_search_v2(size=self.es_insight_top_k,
term_filters={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"scene": self.scene.lower(),
"memoryType": MemoryTypeEnum.INSIGHT.value,
})
insight_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(hit) for hit in hits]
self.logger.info(f"insight_nodes.size={len(insight_nodes)}")
self.set_context(INSIGHT_NODES, insight_nodes)

View file

@ -1,53 +0,0 @@
from typing import List
from constants.common_constants import KEY_WORD, KEYWORD_OBS_NODES, RECALL_TYPE, QUERY_KEYWORDS
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_recall_type import MemoryRecallType
from enumeration.memory_type_enum import MemoryTypeEnum
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class EsKeywordWorker(MemoryBaseWorker):
def _run(self):
query = self.messages[-1].content
# keywords = jieba.analyse.extract_tags(query, topK=3, withWeight=False, allowPOS=()) # 分解关键词
# keywords = jieba.cut(query, cut_all=False) # 使用精确模式分词
# 查询相关关键词
keywords = set()
query_keywords = set()
for key, values in self.config.key_word_relate_dict.items():
if key in query:
keywords.add(key)
keywords.update(values)
query_keywords.add(values[0])
keywords = list(keywords)
self.set_context(QUERY_KEYWORDS, query_keywords)
# 任意一个匹配都算
hits = self.es_client.exact_search_v2(size=self.config.es_keyword_top_k,
term_filters={
"memoryId": self.config.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"scene": self.scene.lower(),
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.INSIGHT.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value],
},
match_filters={
f"metaData.{KEY_WORD}": keywords,
})
# 初始化成MemoryWrapNode并加入召回源的参数
keyword_obs_nodes: List[MemoryWrapNode] = []
for hit in hits:
node = MemoryWrapNode.init_from_es(hit)
node.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.KEYWORD.value
keyword_obs_nodes.append(node)
self.logger.info(f"keyword_obs_nodes size={len(keyword_obs_nodes)}")
for node in keyword_obs_nodes:
self.logger.info(f"node={node.memory_node.content} score_similar={node.score_similar}")
self.set_context(KEYWORD_OBS_NODES, keyword_obs_nodes)

View file

@ -1,27 +0,0 @@
from typing import List
from constants.common_constants import NEW, NEW_OBS_NODES
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class EsNewObsWorker(MemoryBaseWorker):
def __init__(self, es_new_obs_top_k, *args, **kwargs):
super(EsNewObsWorker, self).__init__(*args, **kwargs)
self.es_new_obs_top_k = es_new_obs_top_k
def _run(self):
hits = self.es_client.exact_search_v2(size=self.es_new_obs_top_k,
term_filters={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"scene": self.scene.lower(),
"memoryType": MemoryTypeEnum.OBSERVATION.value,
f"metaData.{NEW}": "1",
})
new_obs_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(hit) for hit in hits]
self.logger.info(f"es new obs, size={len(new_obs_nodes)}")
self.set_context(NEW_OBS_NODES, new_obs_nodes)

View file

@ -1,28 +0,0 @@
from typing import List
from constants.common_constants import REFLECTED, NOT_REFLECTED_OBS_NODES
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class EsNotReflectedWorker(MemoryBaseWorker):
def __init__(self, es_not_reflected_top_k, *args, **kwargs):
super(EsNotReflectedWorker, self).__init__(*args, **kwargs)
self.es_new_obs_top_k = es_new_obs_top_k
def _run(self):
hits = self.es_client.exact_search_v2(size=self.es_not_reflected_top_k,
term_filters={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"scene": self.scene.lower(),
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value],
f"metaData.{REFLECTED}": "0",
})
not_reflected_obs_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(hit) for hit in hits]
self.logger.info(f"retrieve_not_reflected_obs.size={len(not_reflected_obs_nodes)}")
self.set_context(NOT_REFLECTED_OBS_NODES, not_reflected_obs_nodes)

View file

@ -1,30 +0,0 @@
from typing import List
from constants.common_constants import ALL_NODES, ALL_MEMORIES
from enumeration.memory_node_status import MemoryNodeStatus
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class EsRetrieveAllWorker(MemoryBaseWorker):
def _run(self):
# msg_time_created = self.messages[-1].time_created
hits = self.es_client.exact_search_v2(size=1000,
term_filters={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"scene": self.scene.lower(),
# "memoryType": MemoryTypeEnum.OBSERVATION.value,
# f"metaData.{DT}": time_to_formatted_str(msg_time_created),
})
all_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(hit) for hit in hits]
self.logger.info(f"retrieve_all.size={len(all_nodes)}")
self.set_context(ALL_NODES, all_nodes)
all_memories = []
if all_nodes:
for node in all_nodes:
all_memories.append(node.memory_node.to_dict())
self.set_context(ALL_MEMORIES, all_memories)

View file

@ -1,38 +0,0 @@
from typing import List
from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_recall_type import MemoryRecallType
from enumeration.memory_type_enum import MemoryTypeEnum
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class EsSimilarWorker(MemoryBaseWorker):
def __init__(self, es_similar_top_k, *args, **kwargs):
super(EsSimilarWorker, self).__init__(*args, **kwargs)
self.es_similar_top_k = es_similar_top_k
def _run(self):
query = self.messages[-1].content
hits = self.es_client.similar_search(text=query,
size=self.es_similar_top_k,
exact_filters={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"scene": self.scene.lower(),
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.INSIGHT.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value],
})
# 初始化成MemoryWrapNode并加入召回源的参数
similar_obs_nodes: List[MemoryWrapNode] = []
for hit in hits:
node = MemoryWrapNode.init_from_es(hit)
node.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.SIMILAR.value
similar_obs_nodes.append(node)
self.logger.info(f"similar_obs_nodes.size={len(similar_obs_nodes)}")
for node in similar_obs_nodes:
self.logger.info(f"node={node.memory_node.content} score_similar={node.score_similar}")
self.set_context(SIMILAR_OBS_NODES, similar_obs_nodes)

View file

@ -1,31 +0,0 @@
from typing import List
from common.tool_functions import time_to_formatted_str
from constants.common_constants import TODAY_OBS_NODES, DT
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class EsTodayObsWorker(MemoryBaseWorker):
def __init__(self, es_today_obs_top_k, *args, **kwargs):
super(EsTodayObsWorker, self).__init__(*args, **kwargs)
self.es_today_obs_top_k = es_today_obs_top_k
def _run(self):
if not self.messages:
self.logger.warning("messages is empty!")
return
msg_time_created = self.messages[-1].time_created
hits = self.es_client.exact_search_v2(size=self.es_today_obs_top_k,
term_filters={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"scene": self.scene.lower(),
"memoryType": MemoryTypeEnum.OBSERVATION.value,
f"metaData.{DT}": time_to_formatted_str(msg_time_created),
})
today_obs_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(hit) for hit in hits]
self.logger.info(f"retrieve_today_obs.size={len(today_obs_nodes)}")
self.set_context(TODAY_OBS_NODES, today_obs_nodes)

View file

@ -1,34 +0,0 @@
from typing import List, Dict
from utils.user_profile_handler import UserProfileHandler
from constants import common_constants
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from node.memory_wrap_node import MemoryWrapNode
from node.user_attribute import UserAttribute
from pipeline.memory import MemoryServiceRequestModel
from worker.memory_base_worker import MemoryBaseWorker
class LoadProfileWorker(MemoryBaseWorker):
def _run(self):
hits = self.es_client.exact_search_v2(size=10000,
term_filters={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"scene": self.scene.lower(),
"memoryType": [MemoryTypeEnum.PROFILE.value,
MemoryTypeEnum.PROFILE_CUSTOMIZED.value],
})
user_profile_node: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(hit) for hit in hits]
user_profile_dict: Dict[str, UserAttribute] = UserProfileHandler.to_user_attr(user_profile_node)
request: MemoryServiceRequestModel = self.get_context(common_constants.REQUEST)
for user_attr in request.user_profile:
user_profile_dict[user_attr.memory_key] = user_attr
request.user_profile = list(user_profile_dict.values())
self.logger.info(f"retrieve_user_profile.size={len(user_profile_dict)}")
for key, user_attr in user_profile_dict.items():
self.logger.info(f"{key}: {user_attr.description}: {user_attr.value}")

View file

@ -1,9 +1,9 @@
from pydantic import Field, BaseModel
from node.memory_node import MemoryNode
from scheme.memory_node import MemoryNode
class MemoryWrapNode(BaseModel):
class MemoryNode(BaseModel):
id: str = Field("", description="uuid64")
score_similar: float = Field(0, description="相似度打分")

View file

@ -4,9 +4,9 @@ from typing import List
from common.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
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from node.memory_wrap_node import MemoryWrapNode
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
@ -20,7 +20,7 @@ class GetInsightWorker(MemoryBaseWorker):
self.get_insight_top_k = get_insight_top_k
self.es_insight_similar_top_k = es_insight_similar_top_k
def new_insight_node(self, insight_key: str, insight_value: str) -> MemoryWrapNode:
def new_insight_node(self, insight_key: str, insight_value: str) -> MemoryNode:
created_dt = datetime.now()
dt = time_to_formatted_str(time=created_dt)
@ -33,7 +33,7 @@ class GetInsightWorker(MemoryBaseWorker):
meta_data.update({k: str(v) for k, v in get_datetime_info_dict(created_dt).items()})
content = f"用户的{insight_key}{insight_value}"
return MemoryWrapNode.init_from_attrs(content=content,
return MemoryNode.init_from_attrs(content=content,
memoryId=self.memory_id,
scene=self.scene,
memoryType=MemoryTypeEnum.INSIGHT.value,
@ -44,7 +44,7 @@ class GetInsightWorker(MemoryBaseWorker):
def reflect_new_insight_key(self,
insight_key: str,
not_reflected_merge_nodes: List[MemoryWrapNode]) -> MemoryWrapNode | None:
not_reflected_merge_nodes: List[MemoryNode]) -> MemoryNode | None:
# 检索历史memory
hits = self.es_client.similar_search(text=insight_key,
@ -58,7 +58,7 @@ class GetInsightWorker(MemoryBaseWorker):
})
# 转化成 MemoryNodeWrap 合并新增nodes
related_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(x) for x in hits]
related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits]
related_nodes.extend(not_reflected_merge_nodes)
# content去重
@ -106,12 +106,12 @@ class GetInsightWorker(MemoryBaseWorker):
return self.new_insight_node(insight_key=insight_key, insight_value=response_text)
def _run(self):
new_insight_keys: List[MemoryWrapNode] = self.get_context(NEW_INSIGHT_KEYS)
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[MemoryWrapNode] = self.get_context(NOT_REFLECTED_MERGE_NODES)
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
@ -124,11 +124,11 @@ class GetInsightWorker(MemoryBaseWorker):
not_reflected_merge_nodes=not_reflected_merge_nodes)
# save output
new_insight_nodes: List[MemoryWrapNode] = []
new_insight_nodes: List[MemoryNode] = []
for result in self.join_threads():
if result:
new_insight_nodes.append(result)
assert isinstance(result, MemoryWrapNode)
assert isinstance(result, MemoryNode)
insight_key = result.memory_node.metaData.get(INSIGHT_KEY, "")
insight_value = result.memory_node.metaData.get(INSIGHT_VALUE, "")
self.logger.info(f"after_get_insight insight_key={insight_key} insight_value={insight_value}")
@ -137,4 +137,4 @@ class GetInsightWorker(MemoryBaseWorker):
# set REFLECTED
for node in not_reflected_merge_nodes:
node.memory_node.metaData[REFLECTED] = "1"
scheme.memory_node.metaData[REFLECTED] = "1"

View file

@ -3,7 +3,7 @@ from typing import List
from common.response_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 node.memory_wrap_node import MemoryWrapNode
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
@ -19,15 +19,15 @@ class GetReflectionWorker(MemoryBaseWorker):
def _run(self):
# 过滤得到 not_reflected_merge_nodes
new_obs_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_NODES)
not_reflected_nodes: List[MemoryWrapNode] = self.get_context(NOT_REFLECTED_OBS_NODES)
not_reflected_merge_nodes: List[MemoryWrapNode] = []
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.memory_node.metaData.get(REFLECTED, "") == "0"]
if scheme.memory_node.metaData.get(REFLECTED, "") == "0"]
# count
not_reflected_count = len(not_reflected_merge_nodes)
@ -45,7 +45,7 @@ class GetReflectionWorker(MemoryBaseWorker):
self.logger.info(f"profile_keys={profile_keys}")
# get insight_keys
insight_nodes: List[MemoryWrapNode] = self.get_context(INSIGHT_NODES)
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
if insight_nodes:
insight_keys = [n.memory_node.metaData.get(INSIGHT_KEY) for n in insight_nodes]
insight_keys = [x.strip() for x in insight_keys if x]

View file

@ -3,9 +3,9 @@ from typing import List
from common.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_node_status import MemoryNodeStatus
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from node.memory_wrap_node import MemoryWrapNode
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
@ -20,12 +20,12 @@ class LongContraRepeatWorker(MemoryBaseWorker):
def _run(self):
# 合并当前的obs和今日的obs
new_obs_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_NODES)
# new_obs_with_time_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_WITH_TIME_NODES)
# oday_obs_nodes: List[MemoryWrapNode] = self.get_context(TODAY_OBS_NODES)
all_obs_nodes: List[MemoryWrapNode] = []
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)
# oday_obs_nodes: List[MemoryNode] = self.get_context(TODAY_OBS_NODES)
all_obs_nodes: List[MemoryNode] = []
for new_obs_node in new_obs_nodes:
text = new_obs_node.memory_node.content
text = new_obs_scheme.memory_node.content
hits = self.es_client.similar_search(text=text,
size=self.es_contra_repeat_similar_top_k,
exact_filters={
@ -36,7 +36,7 @@ class LongContraRepeatWorker(MemoryBaseWorker):
MemoryTypeEnum.OBS_CUSTOMIZED.value],
})
related_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(x) for x in hits]
related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits]
has_match = False
for related_node in related_nodes:
@ -82,7 +82,7 @@ class LongContraRepeatWorker(MemoryBaseWorker):
return
# add merged obs
merge_obs_nodes: List[MemoryWrapNode] = []
merge_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_merge_obs_list:
if not obs_content_list:
continue
@ -108,11 +108,11 @@ class LongContraRepeatWorker(MemoryBaseWorker):
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
continue
node: MemoryWrapNode = all_obs_nodes[idx]
node: MemoryNode = all_obs_nodes[idx]
if keep_flag != "":
node.memory_node.status = MemoryNodeStatus.EXPIRED.value
scheme.memory_node.status = MemoryNodeStatus.EXPIRED.value
merge_obs_nodes.append(node)
self.logger.info(f"after contra repeat: {node.memory_node.content} {node.memory_node.status}")
self.logger.info(f"after contra repeat: {scheme.memory_node.content} {scheme.memory_node.status}")
# save context
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)

View file

@ -2,21 +2,21 @@ 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
from node.memory_wrap_node import MemoryWrapNode
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class SummaryCollectWorker(MemoryBaseWorker):
def _run(self):
insight_nodes: List[MemoryWrapNode] = self.get_context(INSIGHT_NODES)
new_insight_nodes: List[MemoryWrapNode] = self.get_context(NEW_INSIGHT_NODES)
new_obs_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_NODES)
not_reflected_nodes: List[MemoryWrapNode] = self.get_context(NOT_REFLECTED_OBS_NODES)
not_reflected_merge_nodes: List[MemoryWrapNode] = self.get_context(NOT_REFLECTED_MERGE_NODES)
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, MemoryWrapNode] = {}
all_node_dict: Dict[str, MemoryNode] = {}
if insight_nodes:
all_node_dict.update({n.id: n for n in insight_nodes if n.memory_node.content_modified})
if new_insight_nodes:

View file

@ -2,7 +2,7 @@ from typing import List
from common.response_text_parser import ResponseTextParser
from constants.common_constants import INSIGHT_NODES, NEW_OBS_NODES, INSIGHT_KEY, INSIGHT_VALUE
from node.memory_wrap_node import MemoryWrapNode
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
@ -17,13 +17,13 @@ class UpdateInsightWorker(MemoryBaseWorker):
self.update_insight_top_k = update_insight_top_k
def filter_obs_nodes(self,
insight_node: MemoryWrapNode,
new_obs_nodes: List[MemoryWrapNode]) -> (MemoryWrapNode, List[MemoryWrapNode], float):
insight_node: MemoryNode,
new_obs_nodes: List[MemoryNode]) -> (MemoryNode, List[MemoryNode], float):
max_score: float = 0
filtered_nodes: List[MemoryWrapNode] = []
filtered_nodes: List[MemoryNode] = []
insight_key = insight_node.memory_node.metaData.get(INSIGHT_KEY, "")
insight_value = insight_node.memory_node.metaData.get(INSIGHT_VALUE, "")
insight_key = insight_scheme.memory_node.metaData.get(INSIGHT_KEY, "")
insight_value = insight_scheme.memory_node.metaData.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
@ -55,18 +55,18 @@ class UpdateInsightWorker(MemoryBaseWorker):
return insight_node, filtered_nodes, max_score
def update_insight(self,
insight_node: MemoryWrapNode,
filtered_nodes: List[MemoryWrapNode]) -> MemoryWrapNode:
insight_node: MemoryNode,
filtered_nodes: List[MemoryNode]) -> MemoryNode:
insight_key = insight_node.memory_node.metaData.get(INSIGHT_KEY, "")
insight_value = insight_node.memory_node.metaData.get(INSIGHT_VALUE, "")
insight_key = insight_scheme.memory_node.metaData.get(INSIGHT_KEY, "")
insight_value = insight_scheme.memory_node.metaData.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.memory_node.content}")
user_query_list.append(f"句子:{scheme.memory_node.content}")
update_insight_message = self.prompt_to_msg(
system_prompt=self.prompt_config.update_insight_system,
few_shot=self.prompt_config.update_insight_few_shot,
@ -102,14 +102,14 @@ class UpdateInsightWorker(MemoryBaseWorker):
self.logger.info(f"insight_value={insight_value}, skip.")
return insight_node
insight_node.memory_node.metaData[INSIGHT_VALUE] = insight_value
insight_node.memory_node.content_modified = True
insight_scheme.memory_node.metaData[INSIGHT_VALUE] = insight_value
insight_scheme.memory_node.content_modified = True
return insight_node
def _run(self):
# 获取新的obs和insight
new_obs_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_NODES)
insight_nodes: List[MemoryWrapNode] = self.get_context(INSIGHT_NODES)
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
@ -145,7 +145,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
# 等待结果
for result in self.join_threads():
if result:
insight_node: MemoryWrapNode = result
insight_key = insight_node.memory_node.metaData.get(INSIGHT_KEY, "")
insight_value = insight_node.memory_node.metaData.get(INSIGHT_VALUE, "")
insight_node: MemoryNode = result
insight_key = insight_scheme.memory_node.metaData.get(INSIGHT_KEY, "")
insight_value = insight_scheme.memory_node.metaData.get(INSIGHT_VALUE, "")
self.logger.info(f"after_update_insight insight_key={insight_key} insight_value={insight_value}")

View file

@ -3,7 +3,7 @@ from typing import List
from common.response_text_parser import ResponseTextParser
from constants.common_constants import NEW_OBS_NODES, NEW_USER_PROFILE
from enumeration.memory_type_enum import MemoryTypeEnum
from node.memory_wrap_node import MemoryWrapNode
from scheme.memory_node import MemoryNode
from node.user_attribute import UserAttribute
from worker.memory_base_worker import MemoryBaseWorker
@ -25,9 +25,9 @@ class UpdateProfileWorker(MemoryBaseWorker):
def filter_obs_nodes(self,
user_attr: UserAttribute,
new_obs_nodes: List[MemoryWrapNode]) -> (UserAttribute, List[MemoryWrapNode], float):
new_obs_nodes: List[MemoryNode]) -> (UserAttribute, List[MemoryNode], float):
max_score: float = 0
filtered_nodes: List[MemoryWrapNode] = []
filtered_nodes: List[MemoryNode] = []
result = self.rerank_client.call(query=user_attr.description,
documents=[x.memory_node.content for x in new_obs_nodes])
@ -36,7 +36,7 @@ class UpdateProfileWorker(MemoryBaseWorker):
return user_attr, filtered_nodes, max_score
# 找到大于阈值的obs node
filtered_nodes: List[MemoryWrapNode] = []
filtered_nodes: List[MemoryNode] = []
for rank_node in result:
index = rank_node["index"]
score = rank_node["relevance_score"]
@ -47,20 +47,20 @@ class UpdateProfileWorker(MemoryBaseWorker):
keep_flag = "keep"
max_score = max(max_score, score)
self.logger.info(f"key={user_attr.memory_key} desc={user_attr.description} "
f"content={node.memory_node.content} score={score} keep_flag={keep_flag}")
f"content={scheme.memory_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: UserAttribute, filtered_nodes: List[MemoryWrapNode]) -> UserAttribute:
def update_user_attr(self, user_attr: UserAttribute, filtered_nodes: List[MemoryNode]) -> UserAttribute:
self.logger.info(f"update_user_attr memory_key={user_attr.memory_key} desc={user_attr.description} "
f"value={user_attr.value} doc.size={len(filtered_nodes)}")
# 根据不同的参数类型是否多值分别给出prompt
user_query_list = []
for node in filtered_nodes:
user_query_list.append(f"句子:{node.memory_node.content}")
user_query_list.append(f"句子:{scheme.memory_node.content}")
update_profile = f"{user_attr.memory_key}{user_attr.description}"
update_profile_value = update_profile + "" + "".join(user_attr.value)
@ -158,7 +158,7 @@ class UpdateProfileWorker(MemoryBaseWorker):
self.user_profile_dict[user_attr_key] = new_attr
def _run(self):
new_obs_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_NODES)
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()))

View file

@ -3,8 +3,8 @@ from typing import List
from common.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_node_status import MemoryNodeStatus
from node.memory_wrap_node import MemoryWrapNode
from enumeration.memory_status_enum import MemoryNodeStatus
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
@ -18,10 +18,10 @@ class ContraRepeatWorker(MemoryBaseWorker):
def _run(self):
# 合并当前的obs和今日的obs
new_obs_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_NODES)
new_obs_with_time_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_WITH_TIME_NODES)
today_obs_nodes: List[MemoryWrapNode] = self.get_context(TODAY_OBS_NODES)
all_obs_nodes: List[MemoryWrapNode] = []
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:
@ -62,7 +62,7 @@ class ContraRepeatWorker(MemoryBaseWorker):
return
# add merged obs
merge_obs_nodes: List[MemoryWrapNode] = []
merge_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_merge_obs_list:
if not obs_content_list:
continue
@ -88,11 +88,11 @@ class ContraRepeatWorker(MemoryBaseWorker):
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
continue
node: MemoryWrapNode = all_obs_nodes[idx]
node: MemoryNode = all_obs_nodes[idx]
if keep_flag != "":
node.memory_node.status = MemoryNodeStatus.EXPIRED.value
scheme.memory_node.status = MemoryNodeStatus.EXPIRED.value
merge_obs_nodes.append(node)
self.logger.info(f"after contra repeat: {node.memory_node.content} {node.memory_node.status}")
self.logger.info(f"after contra repeat: {scheme.memory_node.content} {scheme.memory_node.status}")
# save context
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)

View file

@ -5,9 +5,9 @@ from common.response_text_parser import ResponseTextParser
from common.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
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from node.memory_wrap_node import MemoryWrapNode
from scheme.memory_node import MemoryNode
from node.message import Message
from worker.memory_base_worker import MemoryBaseWorker
@ -40,7 +40,7 @@ class GetObservationWithTimeWorker(MemoryBaseWorker):
# 对话时间
meta_data.update({f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()})
return MemoryWrapNode.init_from_attrs(content=obs_content,
return MemoryNode.init_from_attrs(content=obs_content,
memoryId=self.memory_id,
timeCreated=message.time_created,
scene=self.scene,
@ -97,7 +97,7 @@ class GetObservationWithTimeWorker(MemoryBaseWorker):
return
# gene new obs nodes
new_obs_nodes: List[MemoryWrapNode] = []
new_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_obs_list:
if not obs_content_list:
continue

View file

@ -5,9 +5,9 @@ from common.response_text_parser import ResponseTextParser
from common.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
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from node.memory_wrap_node import MemoryWrapNode
from scheme.memory_node import MemoryNode
from node.message import Message
from worker.memory_base_worker import MemoryBaseWorker
@ -36,7 +36,7 @@ class GetObservationWorker(MemoryBaseWorker):
}
meta_data.update({k: str(v) for k, v in get_datetime_info_dict(created_dt).items()})
return MemoryWrapNode.init_from_attrs(content=obs_content,
return MemoryNode.init_from_attrs(content=obs_content,
memoryId=self.memory_id,
timeCreated=message.time_created,
scene=self.scene,
@ -89,7 +89,7 @@ class GetObservationWorker(MemoryBaseWorker):
return
# gene new obs nodes
new_obs_nodes: List[MemoryWrapNode] = []
new_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_obs_list:
if not obs_content_list:
continue

View file

@ -5,7 +5,7 @@ from typing import Dict, List
from constants.common_constants import WEEKDAYS
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
from enumeration.message_role_enum import MessageRoleEnum
def under_line_to_hump(underline_str):
@ -150,7 +150,7 @@ def init_instance_by_config(config: dict|object, default_module_path: str = None
if isinstance(config, accept_types):
return config
import_module(config.pop("path", default_module_path))
import_module(config.pop("path", default_module_path))
clazz = getattr(module, config.pop("name"))
try:
return clazz(**config, **try_kwargs)

View file

@ -10,13 +10,10 @@ class UserAttribute(BaseModel):
如果code为空则为新增否则是更新
确保请求是10条返回是原始10条+加上新增的条数如果可以新增只会对正确的请求操作数据库
"""
code: str = Field("", description="唯一主键 code")
id: str = Field("", description="唯一主键")
memory_id: str = Field("", description="memory id")
# 上游可能没有传这个参数可能隐藏在memory_id做区分
scene: str = Field("", description="source: TONGYI_MAIN_CHAT/TONGYI_CHAR_CHAT/BAILIAN/ASSISTANT")
# 从key改成memory_key
memory_key: str = Field("", description="memory key")

View file

@ -1,8 +1,8 @@
import json
from typing import List, Dict
from enumeration.memory_node_status import MemoryNodeStatus
from node.memory_wrap_node import MemoryWrapNode
from enumeration.memory_status_enum import MemoryNodeStatus
from scheme.memory_node import MemoryNode
from node.user_attribute import UserAttribute
@ -25,13 +25,13 @@ class UserProfileHandler(object):
return content
"""
提供UserAttribute MemoryWrapNode 的相互转化
提供UserAttribute MemoryNode 的相互转化
"""
@classmethod
def to_nodes(cls,
user_profile: List[UserAttribute] | Dict[str, UserAttribute] | None = None,
split_value: bool = False) -> List[MemoryWrapNode]:
split_value: bool = False) -> List[MemoryNode]:
user_profile_dict: Dict[str, UserAttribute] = {}
if user_profile:
@ -41,14 +41,14 @@ class UserProfileHandler(object):
elif isinstance(user_profile, dict):
user_profile_dict.update(user_profile)
user_profile_nodes: List[MemoryWrapNode] = []
user_profile_nodes: List[MemoryNode] = []
for _, user_attr in user_profile_dict.items():
# 获取id
_id = user_attr.code
if not _id:
_id = f"{user_attr.memory_id}_{user_attr.scene}_profile_{user_attr.memory_key}"
attr_node = MemoryWrapNode.init_from_attrs(id=_id,
attr_node = MemoryNode.init_from_attrs(id=_id,
code=_id,
content="",
memoryId=user_attr.memory_id,
@ -75,28 +75,28 @@ class UserProfileHandler(object):
user_profile_nodes.append(attr_node_copy)
else:
content = cls.format_content(user_attr.memory_key, user_attr.description, user_attr.value)
attr_node.memory_node.content = content
attr_scheme.memory_node.content = content
user_profile_nodes.append(attr_node)
return user_profile_nodes
@classmethod
def to_user_attr(cls, user_profile_nodes: List[MemoryWrapNode]) -> Dict[str, UserAttribute]:
def to_user_attr(cls, user_profile_nodes: List[MemoryNode]) -> Dict[str, UserAttribute]:
user_profile_dict: Dict[str, UserAttribute] = {}
for node in user_profile_nodes:
user_attr = UserAttribute(
code=node.id,
memory_id=node.memory_node.memoryId,
scene=node.memory_node.scene,
memory_key=node.memory_node.metaData["memory_key"],
value=json.loads(node.memory_node.metaData["value"]),
is_unique=int(node.memory_node.metaData["is_unique"]),
is_mutable=int(node.memory_node.metaData["is_mutable"]),
memory_type=node.memory_node.memoryType,
description=node.memory_node.metaData["description"],
status=1 if node.memory_node.metaData["status"] == MemoryNodeStatus.ACTIVE.value else 0,
ext_info=json.loads(node.memory_node.metaData["ext_info"]),
memory_id=scheme.memory_node.memoryId,
scene=scheme.memory_node.scene,
memory_key=scheme.memory_node.metaData["memory_key"],
value=json.loads(scheme.memory_node.metaData["value"]),
is_unique=int(scheme.memory_node.metaData["is_unique"]),
is_mutable=int(scheme.memory_node.metaData["is_mutable"]),
memory_type=scheme.memory_node.memoryType,
description=scheme.memory_node.metaData["description"],
status=1 if scheme.memory_node.metaData["status"] == MemoryNodeStatus.ACTIVE.value else 0,
ext_info=json.loads(scheme.memory_node.metaData["ext_info"]),
)
user_profile_dict[user_attr.memory_key] = user_attr
return user_profile_dict