From e758cf41e6c2d3676d52c024ac073fd51aeb9e12 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Mon, 8 Jul 2024 00:00:47 +0800 Subject: [PATCH] [dev] fix active status to new --- .flake8 | 2 +- memory_scope/memory/worker/base_worker.py | 5 +- .../memory/worker/memory_base_worker.py | 2 +- .../summary/get_reflection_subject_worker.py | 2 +- .../worker/summary/summary_collect_worker.py | 51 +-- .../worker/summary/update_insight_worker.py | 12 +- .../worker/write/store_memory_worker.py | 6 +- .../storage/llama_index_es_memory_store.py | 4 + old/__init__.py | 0 old/dash_client.py | 124 ------ old/dash_embedding_client.py | 103 ----- old/dash_generate_client.py | 129 ------ old/dash_rerank_client.py | 119 ----- old/elastic_search_client.py | 419 ------------------ old/memory_node.py | 71 --- old/memory_wrap_node.py | 33 -- old/summary_long/__init__.py | 0 old/summary_long/get_insight_worker.py | 140 ------ old/summary_long/get_reflection_worker.py | 80 ---- old/summary_long/long_contra_repeat_worker.py | 118 ----- old/summary_long/summary_collect_worker.py | 33 -- old/summary_long/update_insight_worker.py | 151 ------- old/summary_long/update_profile_worker.py | 210 --------- old/summary_short/__init__.py | 0 old/summary_short/contra_repeat_worker.py | 98 ---- .../get_observation_with_time_worker.py | 134 ------ old/summary_short/get_observation_worker.py | 122 ----- old/summary_short/info_filter_worker.py | 64 --- old/tool_functions.py | 198 --------- old/user_attribute.py | 33 -- old/user_profile_handler.py | 102 ----- old/worker/__init__.py | 0 old/worker/base_worker.py | 73 --- old/worker/dummy_worker.py | 6 - old/worker/es/__init__.py | 0 old/worker/es/es_insight_worker.py | 22 - old/worker/es/es_new_obs_worker.py | 22 - old/worker/es/es_not_reflected_worker.py | 29 -- old/worker/es/es_similar_worker.py | 37 -- old/worker/es/es_today_obs_worker.py | 32 -- old/worker/es/load_profile_worker.py | 25 -- old/worker/retrieve/__init__.py | 0 old/worker/retrieve/extract_time_worker.py | 74 ---- old/worker/retrieve/fuse_rerank_worker.py | 120 ----- old/worker/retrieve/memory_store_worker.py | 35 -- old/worker/retrieve/semantic_rank_worker.py | 62 --- old/worker/summary_long/__init__.py | 0 old/worker/summary_long/get_insight_worker.py | 166 ------- .../summary_long/get_reflection_worker.py | 99 ----- .../summary_long/long_contra_repeat_worker.py | 129 ------ .../summary_long/summary_collect_worker.py | 47 -- .../summary_long/update_insight_worker.py | 177 -------- .../summary_long/update_profile_worker.py | 241 ---------- old/worker/summary_short/__init__.py | 0 .../summary_short/contra_repeat_worker.py | 117 ----- .../get_observation_with_time_worker.py | 167 ------- .../summary_short/get_observation_worker.py | 144 ------ .../summary_short/info_filter_worker.py | 70 --- 58 files changed, 36 insertions(+), 4423 deletions(-) delete mode 100644 old/__init__.py delete mode 100644 old/dash_client.py delete mode 100644 old/dash_embedding_client.py delete mode 100644 old/dash_generate_client.py delete mode 100644 old/dash_rerank_client.py delete mode 100644 old/elastic_search_client.py delete mode 100644 old/memory_node.py delete mode 100644 old/memory_wrap_node.py delete mode 100644 old/summary_long/__init__.py delete mode 100644 old/summary_long/get_insight_worker.py delete mode 100644 old/summary_long/get_reflection_worker.py delete mode 100644 old/summary_long/long_contra_repeat_worker.py delete mode 100644 old/summary_long/summary_collect_worker.py delete mode 100644 old/summary_long/update_insight_worker.py delete mode 100644 old/summary_long/update_profile_worker.py delete mode 100644 old/summary_short/__init__.py delete mode 100644 old/summary_short/contra_repeat_worker.py delete mode 100644 old/summary_short/get_observation_with_time_worker.py delete mode 100644 old/summary_short/get_observation_worker.py delete mode 100644 old/summary_short/info_filter_worker.py delete mode 100644 old/tool_functions.py delete mode 100644 old/user_attribute.py delete mode 100644 old/user_profile_handler.py delete mode 100644 old/worker/__init__.py delete mode 100644 old/worker/base_worker.py delete mode 100644 old/worker/dummy_worker.py delete mode 100644 old/worker/es/__init__.py delete mode 100644 old/worker/es/es_insight_worker.py delete mode 100644 old/worker/es/es_new_obs_worker.py delete mode 100644 old/worker/es/es_not_reflected_worker.py delete mode 100644 old/worker/es/es_similar_worker.py delete mode 100644 old/worker/es/es_today_obs_worker.py delete mode 100644 old/worker/es/load_profile_worker.py delete mode 100644 old/worker/retrieve/__init__.py delete mode 100644 old/worker/retrieve/extract_time_worker.py delete mode 100644 old/worker/retrieve/fuse_rerank_worker.py delete mode 100644 old/worker/retrieve/memory_store_worker.py delete mode 100644 old/worker/retrieve/semantic_rank_worker.py delete mode 100644 old/worker/summary_long/__init__.py delete mode 100644 old/worker/summary_long/get_insight_worker.py delete mode 100644 old/worker/summary_long/get_reflection_worker.py delete mode 100644 old/worker/summary_long/long_contra_repeat_worker.py delete mode 100644 old/worker/summary_long/summary_collect_worker.py delete mode 100644 old/worker/summary_long/update_insight_worker.py delete mode 100644 old/worker/summary_long/update_profile_worker.py delete mode 100644 old/worker/summary_short/__init__.py delete mode 100644 old/worker/summary_short/contra_repeat_worker.py delete mode 100644 old/worker/summary_short/get_observation_with_time_worker.py delete mode 100644 old/worker/summary_short/get_observation_worker.py delete mode 100644 old/worker/summary_short/info_filter_worker.py diff --git a/.flake8 b/.flake8 index b44183ed..b55b82c1 100644 --- a/.flake8 +++ b/.flake8 @@ -2,7 +2,7 @@ exclude = scripts/* src/agentscope/rpc/* -max-line-length = 79 +max-line-length = 120 inline-quotes = " avoid-escape = no ignore = diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py index 70269ec7..d4812ada 100644 --- a/memory_scope/memory/worker/base_worker.py +++ b/memory_scope/memory/worker/base_worker.py @@ -66,12 +66,9 @@ class BaseWorker(metaclass=ABCMeta): def set_context(self, key: str, value: Any): if self.is_multi_thread: with self.context_lock: - self.context_dict[key] = value + self.context[key] = value else: self.context[key] = value def has_content(self, key: str): return key in self.context - - def __getattr__(self, key: str): - return self.kwargs[key] diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py index 05074f91..3898b167 100644 --- a/memory_scope/memory/worker/memory_base_worker.py +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -5,8 +5,8 @@ from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS from memory_scope.memory.worker.base_worker import BaseWorker from memory_scope.models.base_model import BaseModel from memory_scope.scheme.message import Message -from memory_scope.storage.base_monitor import BaseMonitor from memory_scope.storage.base_memory_store import BaseMemoryStore +from memory_scope.storage.base_monitor import BaseMonitor from memory_scope.utils.global_context import G_CONTEXT from memory_scope.utils.prompt_handler import PromptHandler diff --git a/memory_scope/memory/worker/summary/get_reflection_subject_worker.py b/memory_scope/memory/worker/summary/get_reflection_subject_worker.py index 673845df..601fd0da 100644 --- a/memory_scope/memory/worker/summary/get_reflection_subject_worker.py +++ b/memory_scope/memory/worker/summary/get_reflection_subject_worker.py @@ -22,7 +22,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): meta_data=meta_data, key=insight_key, memory_type=MemoryTypeEnum.INSIGHT.value, - status=MemoryNodeStatus.ACTIVE.value) + status=MemoryNodeStatus.NEW.value) def _run(self): not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES) diff --git a/memory_scope/memory/worker/summary/summary_collect_worker.py b/memory_scope/memory/worker/summary/summary_collect_worker.py index a77644ea..ae07a237 100644 --- a/memory_scope/memory/worker/summary/summary_collect_worker.py +++ b/memory_scope/memory/worker/summary/summary_collect_worker.py @@ -1,46 +1,23 @@ -from typing import List, Dict +from typing import List -from memory_scope.constants.common_constants import ( - NEW_INSIGHT_NODES, - MODIFIED_MEMORIES, - INSIGHT_NODES, - NEW_OBS_NODES, - NOT_REFLECTED_OBS_NODES, - NEW, - NOT_REFLECTED_MERGE_NODES, -) -from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, INSIGHT_NODES, MERGE_OBS_NODES, \ + NOT_UPDATED_NODES from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker +from memory_scope.scheme.memory_node import MemoryNode 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 - ) + keys = [ + INSIGHT_NODES, + MERGE_OBS_NODES, + NOT_UPDATED_NODES, + NOT_REFLECTED_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.obs_updated} - ) - 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.obs_updated = "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}) + memory_nodes: List[MemoryNode] = [] + for key in keys: + memory_nodes.extend(self.get_context(key)) - self.set_context(MODIFIED_MEMORIES, list(all_node_dict.values())) + self.memory_store.update_memories(memory_nodes) diff --git a/memory_scope/memory/worker/summary/update_insight_worker.py b/memory_scope/memory/worker/summary/update_insight_worker.py index 75348207..b235cd25 100644 --- a/memory_scope/memory/worker/summary/update_insight_worker.py +++ b/memory_scope/memory/worker/summary/update_insight_worker.py @@ -2,6 +2,7 @@ from typing import List from memory_scope.constants.common_constants import INSIGHT_NODES, NOT_UPDATED_NODES, NOT_REFLECTED_NODES from memory_scope.constants.language_constants import COLON_WORD, NONE_WORD, REPEATED_WORD +from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker from memory_scope.scheme.memory_node import MemoryNode from memory_scope.utils.datetime_handler import DatetimeHandler @@ -48,6 +49,8 @@ class UpdateInsightWorker(MemoryBaseWorker): insight_node.meta_data.update({k: str(v) for k, v in dt_handler.dt_info_dict.items()}) insight_node.timestamp = dt_handler.timestamp insight_node.dt = dt_handler.datetime_format() + if insight_node.status == MemoryNodeStatus.ACTIVE.value: + insight_node.status = MemoryNodeStatus.CONTENT_MODIFIED.value self.logger.info(f"after_update_{insight_node.key} value={insight_value}") return insight_node @@ -89,6 +92,10 @@ class UpdateInsightWorker(MemoryBaseWorker): self.logger.info(f"update_{insight_node.key} insight_value={insight_value} is invalid.") return insight_node + if insight_node.value == insight_value: + self.logger.info(f"value={insight_value} is same!") + return insight_node + self.update_insight_node(insight_node=insight_node, insight_value=insight_value) return insight_node @@ -102,7 +109,7 @@ class UpdateInsightWorker(MemoryBaseWorker): return for node in insight_nodes: - if node.content: + if node.status == MemoryNodeStatus.ACTIVE.value: self.submit_async_task(fn=self.filter_obs_nodes, insight_node=node, not_updated_nodes=not_updated_nodes) @@ -126,3 +133,6 @@ class UpdateInsightWorker(MemoryBaseWorker): # get result self.gather_async_result() + + for node in not_updated_nodes: + node.obs_updated = True diff --git a/memory_scope/memory/worker/write/store_memory_worker.py b/memory_scope/memory/worker/write/store_memory_worker.py index 95f8ac05..13bf8783 100644 --- a/memory_scope/memory/worker/write/store_memory_worker.py +++ b/memory_scope/memory/worker/write/store_memory_worker.py @@ -14,7 +14,7 @@ class StoreMemoryWorker(MemoryBaseWorker): if self.has_content(store_key): memory_nodes: List[MemoryNode] = self.get_context(store_key) - self.memory_store.update_batch(memory_nodes) + self.memory_store.update_memories(memory_nodes) elif store_key in self.chat_kwargs: query = self.chat_kwargs[store_key] @@ -27,8 +27,8 @@ class StoreMemoryWorker(MemoryBaseWorker): target_name=self.target_name, content=query, memory_type=MemoryTypeEnum.OBS_CUSTOMIZED.value, - status=MemoryNodeStatus.ACTIVE.value, + status=MemoryNodeStatus.NEW.value, timestamp=dt_handler.timestamp, obs_reflected=False, obs_updated=False) - self.memory_store.update(node) + self.memory_store.update_memories(node) diff --git a/memory_scope/storage/llama_index_es_memory_store.py b/memory_scope/storage/llama_index_es_memory_store.py index 68885615..d5a83581 100644 --- a/memory_scope/storage/llama_index_es_memory_store.py +++ b/memory_scope/storage/llama_index_es_memory_store.py @@ -134,6 +134,10 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): self.es_store.close() def update_memories(self, nodes: MemoryNode | List[MemoryNode]): + if not nodes: + self.logger.warning("empty nodes!") + return + if isinstance(nodes, MemoryNode): nodes = [nodes] diff --git a/old/__init__.py b/old/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/old/dash_client.py b/old/dash_client.py deleted file mode 100644 index 57b2cd22..00000000 --- a/old/dash_client.py +++ /dev/null @@ -1,124 +0,0 @@ -import json -import time -from http import HTTPStatus - -import requests - -from utils.logger import Logger -from utils.timer import Timer -from enumeration.env_type import EnvType - - -class DashClient(object): - - def __init__(self, - request_id: str, - dash_scope_uid: str, - authorization: str, - workspace: str, - model_name: str, - env_type: EnvType | str = EnvType.DAILY, - timeout: int = None, - max_retry_count: int = 2, - retry_sleep_time: float = 1.0, - **kwargs): - - self.model_name: str = model_name - self.env_type: EnvType = EnvType(env_type) - self.timeout: int = timeout - self.max_retry_count: int = max_retry_count - self.retry_sleep_time: float = retry_sleep_time - self.kwargs: dict = kwargs - - # 20240506 update by 泉雨 - # if authorization: - # workspace = "" - # dash_scope_uid = "" - - self.headers = { - 'Content-Type': 'application/json', - 'Authorization': authorization, - 'X-Request-Id': request_id, - 'X-DashScope-Uid': dash_scope_uid, - 'X-DashScope-WorkSpace': workspace, - } - - self.url: str = "" - self.data = {} - - self.logger = Logger.get_logger() - - def before_call(self, model_name: str = None, **kwargs): - pass - - def after_call(self, response_obj, **kwargs): - pass - - def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs): - if model_name is None: - model_name = self.model_name - - self.before_call(model_name=model_name, **kwargs) - - with Timer(self.__class__.__name__, log_time=False) as t: - self.logger.debug(f"url={self.url} header={self.headers} data={self.data} timeout={self.timeout}") - response = requests.post(url=self.url, - headers=self.headers, - data=json.dumps(self.data), - timeout=self.timeout) - - if response.status_code == HTTPStatus.OK: - response_obj = json.loads(response.text) - self.logger.info(f"{self.__class__.__name__} env={self.env_type.value} {t.get_cost_info()}, " - f"call model={model_name} success! retry_cnt={retry_cnt}", - stacklevel=3) - return self.after_call(response_obj, **kwargs), True - - else: - self.logger.warning(f"{self.__class__.__name__} env={self.env_type.value} {t.get_cost_info()}, " - f"call model={model_name} failed! retry_cnt={retry_cnt} details={response.text}", - stacklevel=3) - return None, False - - def call(self, model_name: str = None, **kwargs): - for i in range(self.max_retry_count): - result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs) - if flag: - return result - else: - time.sleep(self.retry_sleep_time) - - return None - - -class LLIClient(object): - - def __init__(self, - model_name: str, - timeout: int = None, - max_retry_count: int = 2, - retry_sleep_time: float = 1.0, - **kwargs): - - self.model_name: str = model_name - self.timeout: int = timeout - self.max_retry_count: int = max_retry_count - self.retry_sleep_time: float = retry_sleep_time - self.kwargs: dict = kwargs - - self.data = {} - self.logger = Logger.get_logger() - - - def before_call(self, **kwargs): - pass - - def after_call(self, **kwargs): - pass - - def call_once(self, **kwargs): - pass - - def call(self, **kwargs): - pass - diff --git a/old/dash_embedding_client.py b/old/dash_embedding_client.py deleted file mode 100644 index 91ca82fd..00000000 --- a/old/dash_embedding_client.py +++ /dev/null @@ -1,103 +0,0 @@ -from typing import List, Dict - -import dashscope -import time - -from models import EMB -from models.dash_client import DashClient, LLIClient - -from typing import List, Dict -from utils.registry import build_from_cfg -from utils.timer import Timer - - -from constants.common_constants import DASH_ENV_URL_DICT, DASH_API_URL_DICT -from enumeration.dash_api_enum import DashApiEnum - - -class DashEmbeddingClient(DashClient): - """ - url: https://help.aliyun.com/document_detail/2782232.html?spm=a2c4g.2782227.0.0.76195b1d9UeBAk#a6a39590fegqx - """ - - def __init__(self, model_name: str = dashscope.TextEmbedding.Models.text_embedding_v2, **kwargs): - super(DashEmbeddingClient, self).__init__(model_name=model_name, **kwargs) - self.url = DASH_ENV_URL_DICT.get(self.env_type) + DASH_API_URL_DICT.get(DashApiEnum.EMBEDDING) - - def before_call(self, model_name: str = None, **kwargs): - text: str | List[str] = kwargs.pop("text", "") - # text_type: query or document - text_type: str = kwargs.pop("text_type", "query") - - if isinstance(text, str): - text = [text] - - self.kwargs["text_type"] = text_type - self.data = { - "model": model_name, - "input": { - "texts": text, - }, - "parameters": {**kwargs, **self.kwargs}, - } - - def after_call(self, response_obj, **kwargs) -> Dict[int, List[float]] | List[float]: - embedding_results = {} - for emb in response_obj["output"]["embeddings"]: - embedding_results[emb["text_index"]] = emb["embedding"] - - if len(embedding_results) == 1: - embedding_results = list(embedding_results.values())[0] - return embedding_results - -class LLIEmbedding(LLIClient): - - def __init__(self, method, model_name, **kwargs): - super(LLIEmbedding, self).__init__(model_name, **kwargs) - self.config = { - "method": method, - "model_name": model_name, - **kwargs} - self.embedder = build_from_cfg(self.config, EMB) - - def before_call(self, **kwargs): - text: str | List[str] = kwargs.pop("text", "") - - if isinstance(text, str): - text = [text] - self.data = dict(texts=text) - - def after_call(self, emb: Dict[int, List[float]], **kwargs) -> Dict[int, List[float]] | List[float]: - embedding_results = {} - for idx, e in enumerate(emb): - embedding_results[idx] = e - - if len(embedding_results) == 1: - embedding_results = list(embedding_results.values())[0] - return embedding_results - - - def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs): - if model_name is None: - model_name = self.model_name - - self.before_call(model_name=model_name, **kwargs) - with Timer(self.__class__.__name__, log_time=False) as t: - self.logger.debug(f"data={self.data} timeout={self.timeout}") - try: - results = self.embedder.get_text_embedding_batch(**self.data) - results = self.after_call(results) - return results, True - except Exception as e: - self.logger.debug(f"Get Error in Embedding: {e}") - return None, False - - - def call(self, model_name: str = None, **kwargs): - for i in range(self.max_retry_count): - result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs) - if flag: - return result - else: - time.sleep(self.retry_sleep_time) - return None \ No newline at end of file diff --git a/old/dash_generate_client.py b/old/dash_generate_client.py deleted file mode 100644 index aefbbe5f..00000000 --- a/old/dash_generate_client.py +++ /dev/null @@ -1,129 +0,0 @@ -from typing import List, Dict - -import dashscope - -from models.dash_client import DashClient, LLIClient -from constants.common_constants import DASH_ENV_URL_DICT, DASH_API_URL_DICT -from enumeration.dash_api_enum import DashApiEnum - -import time -from typing import List, Dict -from utils.timer import Timer -from models import LLM -from utils.registry import build_from_cfg -from llama_index.core.base.llms.types import ChatMessage -from llama_index.core.base.llms.types import ( - ChatResponse, - CompletionResponse, -) - - -class DashGenerateClient(DashClient): - """ - url: https://help.aliyun.com/document_detail/2712576.html - """ - - def __init__(self, model_name: str = dashscope.Generation.Models.qwen_max, **kwargs): - super(DashGenerateClient, self).__init__(model_name=model_name, **kwargs) - self.url = DASH_ENV_URL_DICT.get(self.env_type) + DASH_API_URL_DICT.get(DashApiEnum.GENERATION) - - def before_call(self, model_name: str = None, **kwargs): - prompt: str = kwargs.pop("prompt", "") - messages: List[Dict[str, str]] = kwargs.pop("messages", []) - - input_text = {} - if prompt: - input_text["prompt"] = prompt - elif messages: - input_text["messages"] = messages - else: - raise RuntimeError("prompt and messages is both empty!") - - self.data = { - "model": model_name, - "input": input_text, - "parameters": {**kwargs, **self.kwargs}, - } - - def after_call(self, response_obj, **kwargs): - self.logger.debug(f"response_obj={response_obj}") - output = response_obj["output"] - if "text" in output: - return output["text"] - elif "choices" in output: - return output["choices"][0]["message"]["content"] - else: - raise NotImplementedError - - -class LLILLM(LLIClient): - - def __init__(self, method, model_name: str, **kwargs): - super(LLILLM, self).__init__(model_name, **kwargs) - self.config = { - "method": method, - "model_name": model_name, - **kwargs} - self.llm = build_from_cfg(self.config, LLM) - - - def before_call(self, model_name: str = None, **kwargs): - prompt: str = kwargs.pop("prompt", "") - messages: List[Dict[str, str]] = kwargs.pop("messages", []) - - if prompt: - input_text = prompt - input_type = 'prompt' - llama_input = input_text - elif messages: - input_text = messages - input_type = 'messages' - llama_input = [ChatMessage( - role=x['role'], content=x['content'] - ) for x in input_text] - else: - raise RuntimeError("prompt and messages is both empty!") - - self.data = { - input_type: llama_input, - } - - def after_call(self, response_obj: ChatResponse | CompletionResponse, **kwargs) -> str: - self.logger.debug(f"response_obj={response_obj}") - if isinstance(response_obj, CompletionResponse): - return response_obj.text - elif isinstance(response_obj, ChatResponse): - return response_obj.message.content - else: - raise NotImplementedError - - - def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs): - if model_name is None: - model_name = self.model_name - - self.before_call(model_name=model_name, **kwargs) - - with Timer(self.__class__.__name__, log_time=False) as t: - self.logger.debug(f"data={self.data} timeout={self.timeout}") - if True: - # try: - if 'prompt' in self.data: - results = self.llm.complete(**self.data) - else: - results = self.llm.chat(**self.data) - results = self.after_call(results) - return results, True - # except: - # return None, False - - - def call(self, model_name: str = None, **kwargs): - for i in range(self.max_retry_count): - result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs) - print("dashscope llm results:",result) - if flag: - return result - else: - time.sleep(self.retry_sleep_time) - return None \ No newline at end of file diff --git a/old/dash_rerank_client.py b/old/dash_rerank_client.py deleted file mode 100644 index 66ae7ef0..00000000 --- a/old/dash_rerank_client.py +++ /dev/null @@ -1,119 +0,0 @@ -from typing import List - -import dashscope - -from models.dash_client import DashClient, LLIClient -from constants.common_constants import DASH_ENV_URL_DICT, DASH_API_URL_DICT -from enumeration.dash_api_enum import DashApiEnum - - -import time -from typing import List -from models import RERANKER -from utils.timer import Timer - -from utils.registry import build_from_cfg -from llama_index.core.data_structs import Node -from llama_index.core.schema import NodeWithScore # type: ignore - - -class DashReRankClient(DashClient): - """ - url: https://help.aliyun.com/document_detail/2780059.html - """ - - def __init__(self, model_name: str = dashscope.TextReRank.Models.gte_rerank, **kwargs): - super(DashReRankClient, self).__init__(model_name=model_name, **kwargs) - self.url = DASH_ENV_URL_DICT.get(self.env_type) + DASH_API_URL_DICT.get(DashApiEnum.RERANK) - - def before_call(self, model_name: str = None, **kwargs): - query: str = kwargs.pop("query", "") - documents: List[str] = kwargs.pop("documents", []) - top_n: int | None = kwargs.pop("top_n", None) - return_documents: bool = kwargs.pop("return_documents", False) - - assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}" - if top_n is None: - top_n = len(documents) - - self.kwargs.update({ - "top_n": top_n, - "return_documents": return_documents, - }) - self.data = { - "model": model_name, - "input": { - "query": query, - "documents": documents, - }, - "parameters": {**kwargs, **self.kwargs}, - } - - def after_call(self, response_obj, **kwargs): - return response_obj["output"]["results"] - -class LLIReRank(LLIClient): - - def __init__(self, method, model_name, **kwargs): - super(LLIReRank, self).__init__(model_name, **kwargs) - - self.config = { - "method": method, - "model_name": model_name, - **kwargs} - self.reranker = build_from_cfg(self.config, RERANKER) - - - def before_call(self, model_name: str = None, **kwargs): - query: str = kwargs.pop("query", "") - documents: List[str] = kwargs.pop("documents", []) - top_n: int | None = kwargs.pop("top_n", None) - return_documents: bool = kwargs.pop("return_documents", False) - - assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}" - if top_n is None: - top_n = len(documents) - - nodes = [NodeWithScore(node=Node(text=text), score=1.0) for text in documents] - - self.data = { - "nodes": nodes, - "query_str": query, - } - - - def after_call(self, nodes, **kwargs): - results = [] - for idx, node in enumerate(nodes): - results.append(dict(index=idx, - relevance_score=node.score, - document=node.node.text)) - return results - - - def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs): - if model_name is None: - model_name = self.model_name - - self.before_call(model_name=model_name, **kwargs) - - with Timer(self.__class__.__name__, log_time=False) as t: - self.logger.debug(f"data={self.data} timeout={self.timeout}") - try: - results = self.reranker.postprocess_nodes(**self.data) - results = self.after_call(results) - return results, True - except Exception as e: - self.logger.debug(f"Rerank falls, data={self.data}") - # return None, False - - - def call(self, model_name: str = None, **kwargs): - for i in range(self.max_retry_count): - result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs) - if flag: - return result - else: - time.sleep(self.retry_sleep_time) - return None - \ No newline at end of file diff --git a/old/elastic_search_client.py b/old/elastic_search_client.py deleted file mode 100644 index aefe518b..00000000 --- a/old/elastic_search_client.py +++ /dev/null @@ -1,419 +0,0 @@ -from elasticsearch import Elasticsearch -from elasticsearch.helpers import bulk - - -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 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 -from llama_index.vector_stores.elasticsearch import AsyncDenseVectorStrategy - -class ElasticSearchClient(object): - def __init__(self, - es_user_name: str, - es_password: str, - es_index_name: str, - embedding_client: DashEmbeddingClient | None = None, - env_type: EnvType | str = EnvType.DAILY, - content_key: str = "content", - vector_key: str = "vector", - **kwargs): - - self.es_index_name: str = es_index_name - self.embedding_client: DashEmbeddingClient = embedding_client - self.content_key: str = content_key - self.vector_key: str = vector_key - - self.es_client = Elasticsearch( - hosts=[ES_ENV_URL_DICT.get(EnvType(env_type))], - basic_auth=(es_user_name, es_password), - **kwargs) - - self.logger = Logger.get_logger() - self.logger.debug(f"connect es_client info={self.es_client.info()}") - - def log_index_info(self): - index_info = self.es_client.indices.get(index=self.es_index_name) - self.logger.info(f"index={self.es_index_name} exists. index_info={index_info}") - - def insert(self, _id: str, body: dict): - assert body and self.content_key in body, f"body={body} is illegal!" - - # text_type: document - content = body[self.content_key] - vector = self.embedding_client.call(text=content, text_type="document") - if not vector: - self.logger.warning(f"embedding_client call failed, stop es insert!") - return - - body[self.vector_key] = vector - response = self.es_client.index(id=_id, index=self.es_index_name, body=body) - self.logger.info(f"insert response={response}") - - def insert_batch(self, doc_list: list): - """ - doc_list = [ - { - "_id": 2, - "_source": { - "author": "john", - "text": "Elasticsearch: cool.", - "timestamp": "2023-03-23T10:00:00" - } - }, - { - "_id": 3, - "_source": { - "author": "jane", - "text": "Elasticsearch: very cool.", - "timestamp": "2023-03-23T11:00:00" - } - } - ] - """ - text_list = [] - for doc in doc_list: - assert "_id" in doc and "_source" in doc - content = doc["_source"][self.content_key] - text_list.append(content) - - vector_dict = self.embedding_client.call(text=text_list, text_type="document") - if not vector_dict: - self.logger.warning(f"embedding_client call failed, stop es insert!") - return - - # add _index - for i, doc in enumerate(doc_list): - doc["_index"] = self.es_index_name - vector = vector_dict[i] - doc["_source"][self.vector_key] = vector - - # 执行批量插入 - responses = bulk(self.es_client, doc_list) - - # 输出批量插入的响应 - for response in responses[1]: - self.logger.info(f"insert_batch response={response}") - - def print_hits(self, hits: list): - for hit in hits: - print_kwargs = { - "_id": hit['_id'], - "_score": hit['_score'], - } - for k, v in hit['_source'].items(): - # 不打印vector - if k == self.vector_key: - v = len(v) - print_kwargs[k] = v - self.logger.info(" ".join([f"{k}={v}" for k, v in print_kwargs.items()])) - - def exact_search(self, - size: int, - exact_filters: dict = None, - wildcard_filters: dict = None, - print_hits: bool = False, - exclude_vector: bool = True): - """ - { - "match": { - "category": "electronics" # 一级字段过滤 - } - }, - { - "match": { - "product.name": "laptop" # 二级字段过滤 - } - } - { - "terms": { - "product.keyA": ["a", "b", "c"] # 二级字段keyA的精确值必须为a、b、c中的一 - } - } - """ - must_list = [] - for key, value in exact_filters.items(): - if not key: - continue - if isinstance(value, str): - must_list.append({"match": {key: value}}) - elif isinstance(value, list): - must_list.append({"terms": {key: value}}) - - query = { - "size": size, - "query": { - "bool": { - "must": must_list - } - }, - # 添加_source配置以排除vector字段 - "_source": { - "excludes": [self.vector_key] if exclude_vector else [] - } - } - - if wildcard_filters: - should_list = [] - for key, value in wildcard_filters.items(): - if not key: - continue - if isinstance(value, str): - should_list.append({"wildcard": {key: f"*{value}*"}}) - elif isinstance(value, list): - for v in value: - should_list.append({"wildcard": {key: f"*{v}*"}}) - - query["query"]["bool"].update({ - "should": should_list, - "minimum_should_match": 1, - }) - self.logger.info(f"query={query}") - - response = self.es_client.search(index=self.es_index_name, body=query) - hits = response['hits']['hits'] - - # 耗时log - self.logger.info(f"exact_search cost={response['took']}ms " - f"size={len(hits)} " - f"timed_out={response['timed_out']} " - f"shards={response['_shards']} " - f"exact_filters={exact_filters}", stacklevel=2) - - # 每一条结果log一次 - if print_hits: - self.print_hits(hits) - - return hits - - def exact_search_v2(self, - size: int, - term_filters: dict = None, - match_filters: dict = None, - print_hits: bool = False, - exclude_vector: bool = True): - - """ -"bool": { - "must": [ - {"term": {"field1": "固定值"}}, # 一级目录关键字过滤(等于某个值) - {"terms": {"field2": ["a", "b", "c"]}} # 二级目录关键字过滤(等于三个中的任意一个) - ], - "should": [ # 至少匹配其中之一 - {"match": {"key": "ccc"}}, # key包含"ccc" - {"match": {"key": "bbb"}} # 或者key包含"bbb" - ], - "minimum_should_match": 1 # 至少有一个`should`条件匹配 -} - """ - - query = { - "size": size, - "query": { - "bool": { - - } - }, - # 添加_source配置以排除vector字段 - "_source": { - "excludes": [self.vector_key] if exclude_vector else [] - } - } - - if term_filters: - must_list = [] - for k, v in term_filters.items(): - if isinstance(v, list): - must_list.append({"terms": {k: v}}) - elif isinstance(v, str): - must_list.append({"term": {k: v}}) - else: - raise NotImplemented - query["query"]["bool"]["must"] = must_list - - if match_filters: - match_list = [] - for k, v in match_filters.items(): - if isinstance(v, list): - for v_sub in v: - match_list.append({"match": {k: v_sub}}) - elif isinstance(v, str): - match_list.append({"match": {k: v}}) - else: - raise NotImplemented - query["query"]["bool"]["should"] = match_list - query["query"]["bool"]["minimum_should_match"] = 1 - - self.logger.info(query) - response = self.es_client.search(index=self.es_index_name, body=query) - hits = response['hits']['hits'] - - # 耗时log - self.logger.info(f"exact_search cost={response['took']}ms " - f"size={len(hits)} " - f"timed_out={response['timed_out']} " - f"shards={response['_shards']}", stacklevel=2) - - # 每一条结果log一次 - if print_hits: - self.print_hits(hits) - - return hits - - def similar_search(self, - text: str, - size: int, - exact_filters: dict = None, - print_hits: bool = False, - exclude_vector: bool = True): - - if exact_filters is None: - exact_filters = {} - - # 过滤or - or_filters = {} - for k in list(exact_filters.keys()): - v = exact_filters[k] - if isinstance(v, list): - exact_filters.pop(k) - or_filters[k] = v - - vector = self.embedding_client.call(text=text) - if not vector: - self.logger.warning(f"embedding_client call failed, stop select from es!") - return - - query = { - # 返回最相似的top_k个文档 - "size": size, - "query": { - "bool": { - "must": { - "script_score": { - # 对所有文档执行 - "query": { - "match_all": {} - }, - "script": { - # 使用余弦相似度+1,es不能返回负数 - "source": f"cosineSimilarity(params.query_vector, '{self.vector_key}') + 1.0", - "params": {"query_vector": vector} - } - } - }, - "filter": [ - {"term": {k: v}} for k, v in exact_filters.items() - ], - } - }, - # 添加_source配置以排除vector字段 - "_source": { - "excludes": [self.vector_key] if exclude_vector else [] - } - } - - if or_filters: - k_v_pair = [] - for k, v_list in or_filters.items(): - for v in v_list: - k_v_pair.append((k, v)) - query["query"]["bool"]["should"] = [{"term": {k: v}} for k, v in k_v_pair] - query["query"]["bool"]["minimum_should_match"] = 1 - - response = self.es_client.search(index=self.es_index_name, body=query) - hits = response['hits']['hits'] - - # 耗时log - self.logger.info(f"similar_search cost={response['took']}ms " - f"size={len(hits)} " - f"timed_out={response['timed_out']} " - f"shards={response['_shards']} " - f"text={text} " - f"exact_filters={exact_filters}", stacklevel=2) - - # 还原打分 - for hit in hits: - hit['_score'] -= 1 - - # 每一条结果log一次 - if print_hits: - self.print_hits(hits) - - return hits - - -class LLIElasticSearch(object): - def __init__(self, - es_index_name: str, - embedding_client: LLIEmbedding | None = None, - retrieve_topk: int = 3, - content_key: str = "text", - ): - self.es_index_name = es_index_name - self.content_key = content_key - self.embedding_client: LLIEmbedding = embedding_client - # using local es for debug convenient - self.es_client = ElasticsearchStore(index_name="my_index", - es_url="http://localhost:9200", - retrieval_strategy=AsyncDenseVectorStrategy(hybrid=True)) - - self.service_context = ServiceContext.from_defaults(embed_model=self.embedding_client, llm=None) - self.storage_context = StorageContext.from_defaults(vector_store=self.es_client) - self.index = VectorStoreIndex(storage_context=self.storage_context, - service_context=self.service_context) - - self.retriever = self.index.as_retriever(similarity_top_k=retrieve_topk) - self.logger = Logger.get_logger() - - def log_index_info(self, ): - pass - - def print_hits(self, hits: list): - for hit in hits: - print_kwargs = { - "_id": hit['_id'], - "_score": hit['_score'], - } - for k, v in hit['_source'].items(): - # 不打印vector - if k == self.vector_key: - v = len(v) - print_kwargs[k] = v - self.logger.info(" ".join([f"{k}={v}" for k, v in print_kwargs.items()])) - - def similar_search(self, - text: str, - size: int, ): - - ret_nodes = self.retriever.retrieve(text) - return ret_nodes - - def insert_batch(self, doc_list:list[str]): - node_list = [] - for doc in doc_list: - assert "_id" in doc and "_source" in doc - content = doc["_source"]["text"] - doc["_source"].pop("text") - meta = doc["_source"] - node = TextNode(text=content, metadata=meta) - node.node_id(doc['_id']) - node_list.append(node) - self.index.insert_nodes(node_list) - - - def insert(self, _id: str, body: dict): - assert body and self.content_key in body, f"body={body} is illegal!" - content = body[self.content_key] - body.pop(self.content_key) - meta = body - node = TextNode(text=content, metadata=meta) - self.index.insert_nodes([node]) - - - \ No newline at end of file diff --git a/old/memory_node.py b/old/memory_node.py deleted file mode 100644 index 4a682fd9..00000000 --- a/old/memory_node.py +++ /dev/null @@ -1,71 +0,0 @@ -import re -from typing import Dict, List - -from pydantic import Field, BaseModel - - -class MemoryNode(BaseModel): - """ - 除了 content_modified,其他均和数据库字段保持统一 - 根据code判断,如果code是空,则为新增的memoryNode,如果有值,则为更新 - if content_modified is true,则需要调用embedding服务 - """ - id: str = Field("", description="唯一主键 uuid64") - - code: str = Field("", description="和id保持一致,为空则是新增") - - # 0520新增 - timeCreated: str = Field("", description="Memory创建时间(算法不关注)") - - # 0520新增 - timeModified: str = Field("", description="Memory更新时间(算法不关注)") - - content: str = Field("", description="记忆内容") - - memoryId: str = Field("", description="记忆 id,检索区分字段") - - # 0520新增 - scene: str = Field("", description="source: TONGYI_MAIN_CHAT/TONGYI_CHAR_CHAT/BAILIAN/ASSISTANT") - - # 0520新增 - # NOTE 百炼服务端只召回observation, insight, profile, obs_customized, profile_customized - memoryType: str = Field("", description="conversation, observation, insight, " - "profile, obs_customized, profile_customized") - - # 0520新增,但不是数据库字段 - content_modified: bool = Field(False, description="content是否被更新,if true,则需要调用embedding服务") - - # reflected: 1 is reflected before, 0 has not reflected, 如果是用户自定义,写入空值"". - metaData: Dict[str, str] = Field({}, description="元信息: infoScore, algoVersion, datetime, reflected") - - status: str = Field("active", description="active or expired") - - tenantId: str = Field("", description="request id") - - vector: List[float] = Field([], description="content embedding result, return empty") - - def get_time_info(self, time_format: str): - pattern = re.compile(r'\{([^}]*)}') - keys = pattern.findall(time_format) - - match_flag = True - kv_dict = {} - for k in keys: - if k not in self.metaData: - match_flag = False - break - v = self.metaData[k] - if not v: - match_flag = False - break - - kv_dict[k] = v - - if match_flag: - return time_format.format(**kv_dict) - return "" - - def to_dict(self): - res = {"content": self.content, "memoryId": self.memoryId, "memoryType": self.memoryType, - "status": self.status, "metaData": self.metaData} - return res diff --git a/old/memory_wrap_node.py b/old/memory_wrap_node.py deleted file mode 100644 index 28abd65e..00000000 --- a/old/memory_wrap_node.py +++ /dev/null @@ -1,33 +0,0 @@ -from pydantic import Field, BaseModel - -from scheme.memory_node import MemoryNode - - -class MemoryNode(BaseModel): - id: str = Field("", description="uuid64") - - score_similar: float = Field(0, description="相似度打分") - - score_rank: float = Field(0, description="排序打分") - - score_rerank: float = Field(0, description="重排打分") - - memory_node: MemoryNode = Field(None, description="memory node 核心,返回给上游的结构") - - @classmethod - def init_from_es(cls, hit: dict): - memory_node = MemoryNode(**hit['_source']) - return cls(id=hit['_id'], score_similar=hit['_score'], memory_node=memory_node) - - @classmethod - def init_from_attrs(cls, **kwargs): - _id: str = kwargs.get("_id", "") - score_similar: float = kwargs.pop("score_similar", 0) - score_rank: float = kwargs.pop("score_rank", 0) - score_rerank: float = kwargs.pop("score_rerank", 0) - memory_node = MemoryNode(**kwargs) - return cls(id=_id, - score_similar=score_similar, - score_rank=score_rank, - score_rerank=score_rerank, - memory_node=memory_node) diff --git a/old/summary_long/__init__.py b/old/summary_long/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/old/summary_long/get_insight_worker.py b/old/summary_long/get_insight_worker.py deleted file mode 100644 index a118cfe9..00000000 --- a/old/summary_long/get_insight_worker.py +++ /dev/null @@ -1,140 +0,0 @@ -from datetime import datetime -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_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 GetInsightWorker(MemoryBaseWorker): - def __init__(self, insight_obs_max_cnt, es_insight_similar_top_k, get_insight_model, get_insight_max_token, get_insight_temperature, get_insight_top_k, **kwargs): - super(GetInsightWorker,self).__init__(*args,**kwargs) - self.insight_obs_max_cnt = insight_obs_max_cnt - self.get_insight_model = get_insight_model - self.get_insight_max_token = get_insight_max_token - self.get_insight_temperature = get_insight_temperature - 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) -> 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, - } - 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.init_from_attrs(content=content, - memoryId=self.memory_id, - scene=self.scene, - memoryType=MemoryTypeEnum.INSIGHT.value, - content_modified=True, # 新增的insight需要置为true - metaData=meta_data, - status=MemoryNodeStatus.ACTIVE.value, - tenantId=self.tenant_id) - - def reflect_new_insight_key(self, - insight_key: str, - not_reflected_merge_nodes: List[MemoryNode]) -> MemoryNode | None: - - # 检索历史memory - hits = self.es_client.similar_search(text=insight_key, - size=self.es_insight_similar_top_k, - exact_filters={ - "memoryId": self.memory_id, - "status": MemoryNodeStatus.ACTIVE.value, - "scene": self.scene.lower(), - "memoryType": [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.rerank_client.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.prompt_config.get_insight_system, - few_shot=self.prompt_config.get_insight_few_shot, - user_query=self.prompt_config.get_insight_user_query.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.gene_client.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.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}") - - self.set_context(NEW_INSIGHT_NODES, new_insight_nodes) - - # set REFLECTED - for node in not_reflected_merge_nodes: - scheme.memory_node.metaData[REFLECTED] = "1" diff --git a/old/summary_long/get_reflection_worker.py b/old/summary_long/get_reflection_worker.py deleted file mode 100644 index 9385f15d..00000000 --- a/old/summary_long/get_reflection_worker.py +++ /dev/null @@ -1,80 +0,0 @@ -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 scheme.memory_node import MemoryNode -from worker.memory_base_worker import MemoryBaseWorker - - -class GetReflectionWorker(MemoryBaseWorker): - def __init__(self, reflect_obs_cnt_threshold, reflect_num_questions, reflect_obs_model, reflect_obs_max_token, reflect_obs_temperature, reflect_obs_top_k, *args, **kwargs): - super(GetReflectionWorker,self).__init__(*args, **kwargs) - self.reflect_obs_cnt_threshold = reflect_obs_cnt_threshold - self.reflect_num_questions = reflect_num_questions - self.reflect_obs_model = reflect_obs_model - self.reflect_obs_max_token = reflect_obs_max_token - self.reflect_obs_temperature = reflect_obs_temperature - self.reflect_obs_top_k = reflect_obs_top_k - - 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 scheme.memory_node.metaData.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.memory_node.metaData.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.memory_node.content for n in not_reflected_merge_nodes] - reflect_message = self.prompt_to_msg( - system_prompt=self.prompt_config.get_reflect_system.format( - num_questions=self.reflect_num_questions), - few_shot=self.prompt_config.get_reflect_few_shot, - user_query=self.prompt_config.get_reflect_user_query.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.gene_client.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) diff --git a/old/summary_long/long_contra_repeat_worker.py b/old/summary_long/long_contra_repeat_worker.py deleted file mode 100644 index 0d18adb2..00000000 --- a/old/summary_long/long_contra_repeat_worker.py +++ /dev/null @@ -1,118 +0,0 @@ -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_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 LongContraRepeatWorker(MemoryBaseWorker): - def __init__(es_contra_repeat_similar_top_k, long_contra_repeat_threshold, merge_obs_model, merge_obs_max_token, merge_obs_temperature, merge_obs_top_k, *args, **kwargs): - super(LongContraRepeatWorker, self).__init__(*args, **kwargs) - self.es_contra_repeat_similar_top_k = es_contra_repeat_similar_top_k - self.merge_obs_model = merge_obs_model - self.merge_obs_max_token = merge_obs_max_token - self.merge_obs_temperature = merge_obs_temperature - self.merge_obs_top_k = merge_obs_top_k - - 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) - # 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_scheme.memory_node.content - hits = self.es_client.similar_search(text=text, - size=self.es_contra_repeat_similar_top_k, - exact_filters={ - "memoryId": self.memory_id, - "status": MemoryNodeStatus.ACTIVE.value, - "scene": self.scene.lower(), - "memoryType": [MemoryTypeEnum.OBSERVATION.value, - MemoryTypeEnum.OBS_CUSTOMIZED.value], - }) - - related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits] - - 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.memory_node.metaData.get(MSG_TIME, ""), reverse=True) - for i, n in enumerate(all_obs_nodes): - user_query_list.append(f"{i + 1} {n.memory_node.content}") - merge_obs_message = self.prompt_to_msg( - system_prompt=self.prompt_config.long_contra_repeat_system.format(num_obs=len(user_query_list)), - few_shot=self.prompt_config.long_contra_repeat_few_shot, - user_query=self.prompt_config.long_contra_repeat_user_query.format(user_query="\n".join(user_query_list))) - self.logger.info(f"merge_obs_message={merge_obs_message}") - - # call LLM - response_text = self.gene_client.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 != "无": - scheme.memory_node.status = MemoryNodeStatus.EXPIRED.value - merge_obs_nodes.append(node) - 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) diff --git a/old/summary_long/summary_collect_worker.py b/old/summary_long/summary_collect_worker.py deleted file mode 100644 index c65477f4..00000000 --- a/old/summary_long/summary_collect_worker.py +++ /dev/null @@ -1,33 +0,0 @@ -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 scheme.memory_node import MemoryNode -from worker.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.memory_node.content_modified}) - if new_insight_nodes: - all_node_dict.update({n.memory_node.content: n for n in new_insight_nodes}) - if new_obs_nodes: - # 设置为非新 - for n in new_obs_nodes: - n.memory_node.metaData[NEW] = "0" - all_node_dict.update({n.memory_node.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())) diff --git a/old/summary_long/update_insight_worker.py b/old/summary_long/update_insight_worker.py deleted file mode 100644 index 02725069..00000000 --- a/old/summary_long/update_insight_worker.py +++ /dev/null @@ -1,151 +0,0 @@ -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 scheme.memory_node import MemoryNode -from worker.memory_base_worker import MemoryBaseWorker - - -class UpdateInsightWorker(MemoryBaseWorker): - def __init__(update_insight_threshold, update_insight_max_thread, update_insight_model, update_insight_max_token, update_insight_temperature, update_insight_top_k,*args, **kwargs): - super(UpdateInsightWorker, self).__init__(*args, **kwargs) - self.update_insight_threshold = update_insight_threshold - self.update_insight_max_thread = update_insight_max_thread - self.update_insight_model = update_insight_model - self.update_insight_max_token = update_insight_max_token - self.update_insight_temperature = update_insight_temperature - self.update_insight_top_k = update_insight_top_k - - 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_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 - - result = self.rerank_client.call(query=insight_key, - documents=[x.memory_node.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 rank_node in result: - index = rank_node["index"] - score = rank_node["relevance_score"] - 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_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"句子:{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, - user_query=self.prompt_config.update_insight_user_query.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.gene_client.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_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[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[: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_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}") diff --git a/old/summary_long/update_profile_worker.py b/old/summary_long/update_profile_worker.py deleted file mode 100644 index a02496a3..00000000 --- a/old/summary_long/update_profile_worker.py +++ /dev/null @@ -1,210 +0,0 @@ -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 scheme.memory_node import MemoryNode -from node.user_attribute import UserAttribute -from worker.memory_base_worker import MemoryBaseWorker - - -class UpdateProfileWorker(MemoryBaseWorker): - def __init__(update_profile_max_thread, update_profile_threshold, extra_user_attrs, update_profile_model, update_profile_max_token, update_profile_temperature, update_profile_top_k, *args, **kwargs): - super(UpdateProfileWorker,self).__init__(*args, **kwargs) - self.update_profile_max_thread = update_profile_max_thread - self.extra_user_attrs = extra_user_attrs - self.update_profile_threshold = update_profile_threshold - self.update_profile_model = update_profile_model - self.update_profile_max_token = update_profile_max_token - self.update_profile_temperature = update_profile_temperature - self.update_profile_top_k = update_profile_top_k - - # @property - # def extra_user_attrs(self): - # return self.request.extra_user_attrs - - def filter_obs_nodes(self, - user_attr: UserAttribute, - new_obs_nodes: List[MemoryNode]) -> (UserAttribute, List[MemoryNode], float): - max_score: float = 0 - filtered_nodes: List[MemoryNode] = [] - result = self.rerank_client.call(query=user_attr.description, - documents=[x.memory_node.content for x in new_obs_nodes]) - - if not result: - self.add_run_info(f"update_user_attr={user_attr.memory_key} call rerank failed!") - return user_attr, filtered_nodes, max_score - - # 找到大于阈值的obs node - filtered_nodes: List[MemoryNode] = [] - for rank_node in result: - index = rank_node["index"] - score = rank_node["relevance_score"] - 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.memory_key} desc={user_attr.description} " - 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[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"句子:{scheme.memory_node.content}") - update_profile = f"{user_attr.memory_key}({user_attr.description})" - update_profile_value = update_profile + ":" + ",".join(user_attr.value) - - if user_attr.is_unique == 1: - update_profile_message = self.prompt_to_msg( - system_prompt=self.prompt_config.update_unique_profile_system, - few_shot=self.prompt_config.update_unique_profile_few_shot, - user_query=self.prompt_config.update_unique_profile_user_query.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.prompt_config.update_plural_profile_system, - few_shot=self.prompt_config.update_plural_profile_few_shot, - user_query=self.prompt_config.update_plural_profile_user_query.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.gene_client.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.memory_key} call llm failed!") - return user_attr - - profile_list = ResponseTextParser(response_text).parse_v1(f"update_attr {user_attr.memory_key}") - if not profile_list: - self.add_run_info(f"update_one_user_attr key={user_attr.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.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.is_unique == 1: - user_attr.value = [profile.strip()] - else: - attr_value_list = profile.replace(",", ",").split(",") - user_attr.value = [x.strip() for x in sorted(list(set(user_attr.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.description: - user_attr.description = user_attr_desc - continue - - # 增加新属性 - new_attr = UserAttribute(memory_id=self.config.memory_id, - scene=self.scene, - memory_key=user_attr_key, - is_unique=int(user_attr_unique), - is_mutable=1, - memory_type=MemoryTypeEnum.PROFILE, - description=user_attr_desc, - 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[UserAttribute] = [] - self.set_context(NEW_USER_PROFILE, new_user_profile) - - for user_attr_key, user_attr in self.user_profile_dict.items(): - # 不可修改直接跳过 - if user_attr.is_mutable != 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: UserAttribute = result - self.logger.info(f"after_update_profile memory_key={user_attribute.memory_key} " - f"desc={user_attribute.description} value={user_attribute.value}") - new_user_profile.append(user_attribute) diff --git a/old/summary_short/__init__.py b/old/summary_short/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/old/summary_short/contra_repeat_worker.py b/old/summary_short/contra_repeat_worker.py deleted file mode 100644 index 1a413e78..00000000 --- a/old/summary_short/contra_repeat_worker.py +++ /dev/null @@ -1,98 +0,0 @@ -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_status_enum import MemoryNodeStatus -from scheme.memory_node import MemoryNode -from worker.memory_base_worker import MemoryBaseWorker - - -class ContraRepeatWorker(MemoryBaseWorker): - def __init__(self, merge_obs_model, merge_obs_max_token, merge_obs_temperature, merge_obs_top_k, *args, **kwargs): - super(ContraRepeatWorker, self).__init__(*args, **kwargs) - self.merge_obs_model = merge_obs_model - self.merge_obs_max_token = merge_obs_max_token - self.merge_obs_temperature = merge_obs_temperature - self.merge_obs_top_k = merge_obs_top_k - - 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.memory_node.metaData.get(MSG_TIME, ""), reverse=True) - for i, n in enumerate(all_obs_nodes): - user_query_list.append(f"{i + 1} {n.memory_node.content}") - merge_obs_message = self.prompt_to_msg( - system_prompt=self.prompt_config.contra_repeat_system.format(num_obs=len(user_query_list)), - few_shot=self.prompt_config.contra_repeat_few_shot, - user_query=self.prompt_config.contra_repeat_user_query.format(user_query="\n".join(user_query_list))) - self.logger.info(f"merge_obs_message={merge_obs_message}") - - # call LLM - response_text = self.gene_client.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 != "无": - scheme.memory_node.status = MemoryNodeStatus.EXPIRED.value - merge_obs_nodes.append(node) - 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) diff --git a/old/summary_short/get_observation_with_time_worker.py b/old/summary_short/get_observation_with_time_worker.py deleted file mode 100644 index 845c0e24..00000000 --- a/old/summary_short/get_observation_with_time_worker.py +++ /dev/null @@ -1,134 +0,0 @@ -from datetime import datetime -from typing import List - -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_status_enum import MemoryNodeStatus -from enumeration.memory_type_enum import MemoryTypeEnum -from scheme.memory_node import MemoryNode -from node.message import Message -from worker.memory_base_worker import MemoryBaseWorker - - -class GetObservationWithTimeWorker(MemoryBaseWorker): - def __init__(self, summary_messages_model, summary_messages_max_token, summary_messages_temperature, summary_messages_top_k, *args, **kwargs): - super(GetObservationWithTimeWorker, self).__init__(*args, **kwargs) - self.summary_messages_model = summary_messages_model - self.summary_messages_max_token = summary_messages_max_token - self.summary_messages_temperature = summary_messages_temperature - self.summary_messages_top_k = summary_messages_top_k - - 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, # 关键词 - } - - # 事件时间 - 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, - memoryId=self.memory_id, - timeCreated=message.time_created, - scene=self.scene, - memoryType=MemoryTypeEnum.OBSERVATION.value, - content_modified=True, # 新增的obs需要置为true - metaData=meta_data, - status=MemoryNodeStatus.ACTIVE.value, - tenantId=self.tenant_id) - - 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.prompt_config.get_observation_with_time_system.format(num_obs=len(user_query_list)), - few_shot=self.prompt_config.get_observation_with_time_few_shot, - user_query=self.prompt_config.get_observation_with_time_user_query.format( - user_query="\n".join(user_query_list))) - self.logger.info(f"obtain_obs_message={obtain_obs_message}") - - # call LLM - response_text: str = self.gene_client.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) diff --git a/old/summary_short/get_observation_worker.py b/old/summary_short/get_observation_worker.py deleted file mode 100644 index 0ab6b939..00000000 --- a/old/summary_short/get_observation_worker.py +++ /dev/null @@ -1,122 +0,0 @@ -from datetime import datetime -from typing import List - -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_status_enum import MemoryNodeStatus -from enumeration.memory_type_enum import MemoryTypeEnum -from scheme.memory_node import MemoryNode -from node.message import Message -from worker.memory_base_worker import MemoryBaseWorker - - -class GetObservationWorker(MemoryBaseWorker): - def __init__(self, summary_messages_model, summary_messages_max_token, summary_messages_temperature, summary_messages_top_k, *args, **kwargs): - super(GetObservationWorker, self).__init__(*args,**kwargs) - self.summary_messages_model = summary_messages_model - self.summary_messages_max_token = summary_messages_max_token - self.summary_messages_temperature = summary_messages_temperature - self.summary_messages_top_k = summary_messages_top_k - - 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, # 关键词 - } - meta_data.update({k: str(v) for k, v in get_datetime_info_dict(created_dt).items()}) - - return MemoryNode.init_from_attrs(content=obs_content, - memoryId=self.memory_id, - timeCreated=message.time_created, - scene=self.scene, - memoryType=MemoryTypeEnum.OBSERVATION.value, - content_modified=True, # 新增的obs需要置为true - metaData=meta_data, - status=MemoryNodeStatus.ACTIVE.value, - tenantId=self.tenant_id) - - 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.prompt_config.get_observation_system.format(num_obs=len(user_query_list)), - few_shot=self.prompt_config.get_observation_few_shot, - user_query=self.prompt_config.get_observation_user_query.format(user_query="\n".join(user_query_list))) - self.logger.info(f"obtain_obs_message={obtain_obs_message}") - - # call LLM - response_text: str = self.gene_client.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) diff --git a/old/summary_short/info_filter_worker.py b/old/summary_short/info_filter_worker.py deleted file mode 100644 index 0723ea54..00000000 --- a/old/summary_short/info_filter_worker.py +++ /dev/null @@ -1,64 +0,0 @@ -from common.response_text_parser import ResponseTextParser -from enumeration.message_role_enum import MessageRoleEnum -from worker.memory_base_worker import MemoryBaseWorker - - -class InfoFilterWorker(MemoryBaseWorker): - def __init__(self, info_filter_msg_max_size, info_filter_model, info_filter_max_token, info_filter_temperature, info_filter_top_k, *args, **kwargs): - super(InfoFilterWorker,self).__init__(*args,**kwargs) - self.info_filter_msg_max_size - self.info_filter_model = info_filter_model - self.info_filter_max_token = info_filter_max_token - self.info_filter_temperature = info_filter_temperature - self.info_filter_top_k = info_filter_top_k - - 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.prompt_config.info_filter_system.format(batch_size=len(info_messages)), - few_shot=self.prompt_config.info_filter_few_shot, - user_query=self.prompt_config.info_filter_user_query.format(user_query=user_query)) - self.logger.info(f"info_filter_message={info_filter_message}") - - # call llm - response_text = self.gene_client.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 diff --git a/old/tool_functions.py b/old/tool_functions.py deleted file mode 100644 index d32778b6..00000000 --- a/old/tool_functions.py +++ /dev/null @@ -1,198 +0,0 @@ -import re -from datetime import datetime -from importlib import import_module -from typing import Dict, List - -from constants.common_constants import WEEKDAYS - -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) - return sub[0:1].upper() + sub[1:] - - -def parse_response_text_v1(response_text: str) -> dict: - """ - parse text like: - <1> - <2> ddd - <4> dddd<555> - - result = {1: "AAA", 2: "BBB", 4: "CCC"} - """ - result_dict: Dict[int, str] = {} - - # 确保第一个数字,后面是string - matches = re.findall(r'<(\d+)>\s*<([^>]+)>', response_text.strip()) - - # matches 为空返回 - for key, value in matches: - result_dict[int(key)] = value - - return result_dict - - -def parse_response_text_v2(response_text: str) -> Dict[int, List[str]]: - """ - parse text like: - XXX - <1> <222> - <2> - <4,5> - - result = {1: ["AAA", "222"], 2: "BBB", 4: "CCC"} - """ - result_dict: Dict[int, List[str]] = {} - for line in response_text.strip().split("\n"): - if "> <" not in line: - continue - ll = [x.removeprefix("<").removesuffix(">") for x in line.strip().split("> <")] - idx: str = ll[0] - values: List[str] = ll[1:] - if idx.isdigit(): - idx_int = int(idx) - else: - idx_split = idx.split(",") - if len(idx_split) == 0: - continue - idx = idx_split[0] - if idx.isdigit(): - idx_int = int(idx) - else: - continue - if values: - result_dict[idx_int] = values - - return result_dict - - -def parse_response_text_v3(response_text: str) -> List[List[str]]: - """ - parse text like: - XXX - <1> - <2c> - <41> - - result = [["1", "AAA"], ["2c", "BBB"], ["41", "CCC", "BBB"]] - """ - result_list: List[List[str]] = [] - for line in response_text.strip().split("\n"): - if "> <" not in line: - continue - ll = [x.removeprefix("<").removesuffix(">") for x in line.strip().split("> <")] - result_list.append(ll) - return result_list - - -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 extract_date_parts(input_string: str): - # Extending our pattern to handle "每" (every) as a possible value. - patterns = { - 'year': r'(\d+|每)年', - 'month': r'(\d+|每)月', - 'day': r'(\d+|每)日', - 'weekday': r'周([一二三四五六日])?', - 'hour': r'(\d+)点' - } - weekday_dict = {"一": 1, "二": 2, "三": 3, "四": 4, "五": 5, "六": 6, "日": 7} - extracted_data = {} - - # Search for patterns in the input string and populate the dictionary - for key, pattern in patterns.items(): - match = re.search(pattern, input_string) - if match: # If there is a match, include it in the output dictionary - if match.group(1) == "每": - extracted_data[key] = -1 - elif match.group(1) in weekday_dict.keys(): - extracted_data[key] = weekday_dict[match.group(1)] - else: - extracted_data[key] = int(match.group(1)) - return extracted_data - - -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 - - -def init_instance_by_config(config: dict|object, default_module_path: str = None, try_kwargs: dict = {}, accept_types: type = None): - if isinstance(config, accept_types): - return config - - import_module(config.pop("path", default_module_path)) - clazz = getattr(module, config.pop("name")) - try: - return clazz(**config, **try_kwargs) - except: - return clazz(**config) - - -def init_instance_by_config_v2(config: dict, default_clazz_path: str = "", suffix_name: str = "", **kwargs): - clazz_path = config.pop("clazz") - if not clazz_path: - raise RuntimeError("empty clazz_path!") - clazz_name_split = clazz_path.split(".") - clazz_name: str = clazz_name_split[-1] - if suffix_name and not clazz_name.endswith(suffix_name): - clazz_name = f"{clazz_name}_{suffix_name}" - - # 构造path - clazz_paths = [] - if default_clazz_path: - clazz_paths.append(default_clazz_path) - clazz_paths.extend(clazz_name_split[:-1]) - clazz_paths.append(clazz_name) - module = import_module(".".join(clazz_paths)) - - cls_name = under_line_to_hump(clazz_name) - return getattr(module, cls_name)(**config, **kwargs) - - -def complete_config_name(config_name: str, suffix: str = ".json"): - if not config_name.endswith(suffix): - config_name += suffix - return config_name - - -def prompt_to_msg(system_prompt: str, few_shot: str, user_query: str): - return [ - { - "role": MessageRoleEnum.SYSTEM.value, - "content": system_prompt.strip(), - }, - { - "role": MessageRoleEnum.USER.value, - "content": "\n".join([x.strip() for x in [few_shot, system_prompt, user_query]]) - }, - ] diff --git a/old/user_attribute.py b/old/user_attribute.py deleted file mode 100644 index dab4f4d6..00000000 --- a/old/user_attribute.py +++ /dev/null @@ -1,33 +0,0 @@ -from typing import Dict, List - -from pydantic import Field, BaseModel - - -class UserAttribute(BaseModel): - """ - 用户画像的一条属性,和数据库保持一致,只会选择status为1的属性透传过来。 - status会透传过来。 - 如果code为空,则为新增,否则是更新。 - 确保请求是10条,返回是原始10条+加上新增的条数(如果可以新增)。只会对正确的请求操作数据库。 - """ - id: str = Field("", description="唯一主键") - - memory_id: str = Field("", description="memory id") - - # 从key改成memory_key - memory_key: str = Field("", description="memory key") - - value: List[str] = Field([], description="value") - - is_unique: int = Field(1, description="属性是否唯一,if 1 value只有一个,if 0, value 可以很多个") - - is_mutable: int = Field(1, description="是否可变,if 1,value可变,if 1,不可变(用户定义)") - - memory_type: str = Field("", description="profile, profile_customized") - - description: str = Field("", description="memory id") - - status: int = Field(1, - description="0为删除,1为active,状态,算法不感知,只为了保存用户删除的画像,给算法传status为valid的用户画像") - - ext_info: Dict[str, str] = Field({}, description="占位符字典") diff --git a/old/user_profile_handler.py b/old/user_profile_handler.py deleted file mode 100644 index 4d61bec4..00000000 --- a/old/user_profile_handler.py +++ /dev/null @@ -1,102 +0,0 @@ -import json -from typing import List, Dict - -from enumeration.memory_status_enum import MemoryNodeStatus -from scheme.memory_node import MemoryNode -from node.user_attribute import UserAttribute - - -class UserProfileHandler(object): - @classmethod - def format_content(cls, key: str, description: str, value: str | List[str] = None): - if not key.startswith("用户"): - key = f"用户的{key}" - - if not description.startswith("用户"): - description = f"用户{description}" - - content = f"{key}({description})" - - if value: - if isinstance(value, list): - value = ",".join(value) - content = f"{content}:{value}" - - return content - - """ - 提供UserAttribute 和 MemoryNode 的相互转化 - """ - - @classmethod - def to_nodes(cls, - user_profile: List[UserAttribute] | Dict[str, UserAttribute] | None = None, - split_value: bool = False) -> List[MemoryNode]: - - user_profile_dict: Dict[str, UserAttribute] = {} - if user_profile: - if isinstance(user_profile, list): - for user_attr in user_profile: - user_profile_dict[user_attr.memory_key] = user_attr - elif isinstance(user_profile, dict): - user_profile_dict.update(user_profile) - - 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 = MemoryNode.init_from_attrs(id=_id, - code=_id, - content="", - memoryId=user_attr.memory_id, - scene=user_attr.scene, - memoryType=user_attr.memory_type, - content_modified=True, - metaData={ - "memory_key": user_attr.memory_key, - "value": json.dumps(user_attr.value, ensure_ascii=False), - "is_unique": str(user_attr.is_unique), - "is_mutable": str(user_attr.is_mutable), - "description": user_attr.description, - "status": MemoryNodeStatus.ACTIVE.value, - "ext_info": json.dumps(user_attr.ext_info, - ensure_ascii=False), - }, - status=MemoryNodeStatus.ACTIVE.value) - - if split_value: - for value in user_attr.value: - content = cls.format_content(user_attr.memory_key, user_attr.description, value) - attr_node_copy = attr_node.copy(deep=True) - attr_node_copy.memory_node.content = content - user_profile_nodes.append(attr_node_copy) - else: - content = cls.format_content(user_attr.memory_key, user_attr.description, user_attr.value) - 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[MemoryNode]) -> Dict[str, UserAttribute]: - user_profile_dict: Dict[str, UserAttribute] = {} - - for node in user_profile_nodes: - user_attr = UserAttribute( - code=node.id, - 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 diff --git a/old/worker/__init__.py b/old/worker/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/old/worker/base_worker.py b/old/worker/base_worker.py deleted file mode 100644 index 7f0dd62f..00000000 --- a/old/worker/base_worker.py +++ /dev/null @@ -1,73 +0,0 @@ -from typing import Any, Dict - -from ..utils.logger import Logger -from ..utils.timer import Timer - - -class BaseWorker(object): - - def __init__(self, raise_exception: bool = True, **kwargs): - super(BaseWorker, self).__init__(**kwargs) - # 异常是否继续执行 - self.raise_exception: bool = raise_exception - - # True 为正常运行,False会结束整个pipeline - self.continue_run: bool = True - - # 短name - self._name_simple: str = "" - - # 是否多线程环境 - self.is_multi_thread: bool = False - - # pipeline 上下文 - self.context_dict: Dict[str, Any] | None = None - self.context_lock = None - - # 日志 - self.logger: Logger = Logger.get_logger() - - # worker 参数 - self.kwargs: dict = kwargs - - def _run(self): - raise NotImplementedError - - def run(self): - self.logger.info(f"----- Begin {self.name_simple} -----") - with Timer(self.name_simple, log_time=False) as t: - if self.raise_exception: - self._run() - else: - try: - self._run() - except Exception as e: - self.logger.exception(f"run {self.name_simple} failed! args={e.args}") - - self.logger.info(f"----- End {self.name_simple} cost={t.cost_str}-----") - - def set_context_dict(self, context_dict: dict, context_lock=None): - self.context_dict = context_dict - if context_lock is not None: - self.context_lock = context_lock - self.is_multi_thread = True - - def get_context(self, key: str, default=None): - return self.context_dict.get(key, default) - - def set_context(self, key: str, value: Any): - if self.is_multi_thread: - # add lock to multi thread - with self.context_lock: - self.context_dict[key] = value - 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: - self._name_simple = self.__class__.__name__.replace("Worker", "") - return self._name_simple diff --git a/old/worker/dummy_worker.py b/old/worker/dummy_worker.py deleted file mode 100644 index 87ba5dbe..00000000 --- a/old/worker/dummy_worker.py +++ /dev/null @@ -1,6 +0,0 @@ -from memory_base_worker import MemoryBaseWorker - - -class DummyWorker(MemoryBaseWorker): - def _run(self): - pass \ No newline at end of file diff --git a/old/worker/es/__init__.py b/old/worker/es/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/old/worker/es/es_insight_worker.py b/old/worker/es/es_insight_worker.py deleted file mode 100644 index cd658119..00000000 --- a/old/worker/es/es_insight_worker.py +++ /dev/null @@ -1,22 +0,0 @@ -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_memories( - size=self.kwargs.es_insight_top_k, - filter_dict={ - "memory_id": self.memory_id, - "status": MemoryNodeStatus.ACTIVE.value, - "memory_type": MemoryTypeEnum.INSIGHT.value, - }, - ) - self.logger.info(f"insight_nodes.size={len(insight_nodes)}") - self.set_context(INSIGHT_NODES, insight_nodes) diff --git a/old/worker/es/es_new_obs_worker.py b/old/worker/es/es_new_obs_worker.py deleted file mode 100644 index 021e963f..00000000 --- a/old/worker/es/es_new_obs_worker.py +++ /dev/null @@ -1,22 +0,0 @@ -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_memories( - size=self.kwargs.es_new_obs_top_k, - filter_dict={ - "memory_id": self.memory_id, - "status": MemoryNodeStatus.ACTIVE.value, - "memory_type": MemoryTypeEnum.OBSERVATION.value, - f"meta_data.{NEW}": "1", - }, - ) - self.logger.info(f"es new obs, size={len(new_obs_nodes)}") - self.set_context(NEW_OBS_NODES, new_obs_nodes) diff --git a/old/worker/es/es_not_reflected_worker.py b/old/worker/es/es_not_reflected_worker.py deleted file mode 100644 index 10c96d7f..00000000 --- a/old/worker/es/es_not_reflected_worker.py +++ /dev/null @@ -1,29 +0,0 @@ -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_memories( - size=self.kwargs.es_new_obs_top_k, - filter_dict={ - "memory_id": self.memory_id, - "status": MemoryNodeStatus.ACTIVE.value, - "memory_type": [ - MemoryTypeEnum.OBSERVATION.value, - MemoryTypeEnum.OBS_CUSTOMIZED.value, - ], - f"meta_data.{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) diff --git a/old/worker/es/es_similar_worker.py b/old/worker/es/es_similar_worker.py deleted file mode 100644 index 2b419398..00000000 --- a/old/worker/es/es_similar_worker.py +++ /dev/null @@ -1,37 +0,0 @@ -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_memories( - text=query, - size=self.es_similar_top_k, - filter_dict={ - "memory_id": self.memory_id, - "status": MemoryNodeStatus.ACTIVE.value, - "memory_type": [ - MemoryTypeEnum.OBSERVATION.value, - MemoryTypeEnum.INSIGHT.value, - MemoryTypeEnum.OBS_CUSTOMIZED.value, - ], - }, - ) - - for node in similar_obs_nodes: - 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}") - self.set_context(SIMILAR_OBS_NODES, similar_obs_nodes) diff --git a/old/worker/es/es_today_obs_worker.py b/old/worker/es/es_today_obs_worker.py deleted file mode 100644 index 34f65f02..00000000 --- a/old/worker/es/es_today_obs_worker.py +++ /dev/null @@ -1,32 +0,0 @@ -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_memories( - size=self.es_today_obs_top_k, - filter_dict={ - "memory_id": self.memory_id, - "status": MemoryNodeStatus.ACTIVE.value, - "memory_type": MemoryTypeEnum.OBSERVATION.value, - f"meta_Data.{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) diff --git a/old/worker/es/load_profile_worker.py b/old/worker/es/load_profile_worker.py deleted file mode 100644 index 90ad2dc9..00000000 --- a/old/worker/es/load_profile_worker.py +++ /dev/null @@ -1,25 +0,0 @@ -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={ - "memory_id": self.memory_id, - "status": MemoryNodeStatus.ACTIVE.value, - "memory_type": [ - 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)}") diff --git a/old/worker/retrieve/__init__.py b/old/worker/retrieve/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/old/worker/retrieve/extract_time_worker.py b/old/worker/retrieve/extract_time_worker.py deleted file mode 100644 index 4770360f..00000000 --- a/old/worker/retrieve/extract_time_worker.py +++ /dev/null @@ -1,74 +0,0 @@ -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 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""" -任务指令:从语句与语句发生的时间,推断并提取语句内容中指向的时间段。回答尽可能完整的时间段。 -语句:{query} -时间:{query_time_str} -回答: - """.strip() - - def _run(self): - # save to context - extract_time_dict = {} - self.set_context(EXTRACT_TIME_DICT, extract_time_dict) - - # get query & time_created_dt - query = self.messages[-1].content - time_created = self.messages[-1].time_created - - # find datetime keyword - contain_datetime = False - for datetime_word in DATATIME_WORD_LIST: - if datetime_word in query: - contain_datetime = True - break - if not contain_datetime: - self.logger.info(f"contain_datetime={contain_datetime}") - 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 - ) - self.logger.info(f"extract_time_prompt={extract_time_prompt}") - - # call sft model - - 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+)" - matches = re.findall(pattern, response_text) - for key, value in matches: - if key in DATATIME_KEY_MAP.keys(): - extract_time_dict[DATATIME_KEY_MAP[key]] = value - - self.logger.info(f"response_text={response_text} filters={extract_time_dict}") diff --git a/old/worker/retrieve/fuse_rerank_worker.py b/old/worker/retrieve/fuse_rerank_worker.py deleted file mode 100644 index 9c3f9854..00000000 --- a/old/worker/retrieve/fuse_rerank_worker.py +++ /dev/null @@ -1,120 +0,0 @@ -from typing import Dict, List - -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): - @staticmethod - def format_time_infer( - time_infer: str, extract_time_dict: Dict[str, str], meta_data: Dict[str, str] - ): - if time_infer: - return time_infer - - time_infer = "" - if "year" in extract_time_dict: - value = meta_data.get("msg_year") - if value: - time_infer += f"{value}年" - elif value == "-1": - 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 += "每月" - - if "day" in extract_time_dict: - value = meta_data.get("msg_day") - if value: - time_infer += f"{value}日" - elif value == "-1": - time_infer += "每日" - - if "weekday" in extract_time_dict: - value = meta_data.get("msg_weekday") - if value: - time_infer += value - - return time_infer - - def _run(self): - # 解析时间meta信息 - extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT) - 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") - return - - filtered_nodes = [] - for node in all_online_nodes: - if node.score_rank < self.fuse_score_threshold: - continue - - # 根据类型给ratio - type_ratio: float = self.fuse_ratio_dict.get(node.memory_type, 0.1) - - # 时间系数,完全匹配才行 - fuse_time_ratio: float = 1.0 - match_event_flag = False - match_msg_flag = False - if extract_time_dict: - match_event_flag = True - for k, v in extract_time_dict.items(): - event_value = node.meta_data.get(f"event_{k}", "") - if event_value in ["-1", v]: - continue - else: - match_event_flag = False - break - - match_msg_flag = True - for k, v in extract_time_dict.items(): - msg_value = node.meta_data.get(f"msg_{k}", "") - if msg_value == v: - continue - else: - match_msg_flag = False - break - - if match_event_flag or match_msg_flag: - fuse_time_ratio = self.fuse_time_ratio - node.meta_data[TIME_MATCHED] = "1" - - node.score_rerank = node.score_rank * type_ratio * fuse_time_ratio - 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 = filtered_nodes[: self.output_max_count] - related_memories: List[str] = [] - for node in filtered_nodes: - content = node.content - - # 如果命中时间逻辑 - 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) diff --git a/old/worker/retrieve/memory_store_worker.py b/old/worker/retrieve/memory_store_worker.py deleted file mode 100644 index 4efea52d..00000000 --- a/old/worker/retrieve/memory_store_worker.py +++ /dev/null @@ -1,35 +0,0 @@ -from typing import List - -from utils.user_profile_handler import UserProfileHandler -from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE, CONTENT_MODIFIED -from scheme.memory_node import MemoryNode -from worker.memory_base_worker import MemoryBaseWorker - - -class MemoryStoreWorker(MemoryBaseWorker): - - def _run(self): - 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.memory_id}_content_{n.content}" - n.code = n.id - # TODO add batch insert - n.meta_data.pop(CONTENT_MODIFIED) - self.vector_store.insert(n) - else: - self.logger.warning("modified_memories is empty!") - - new_user_profile: List[MemoryNode] = self.get_context(NEW_USER_PROFILE) - if new_user_profile: - 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!") diff --git a/old/worker/retrieve/semantic_rank_worker.py b/old/worker/retrieve/semantic_rank_worker.py deleted file mode 100644 index 0fc48bc9..00000000 --- a/old/worker/retrieve/semantic_rank_worker.py +++ /dev/null @@ -1,62 +0,0 @@ -from typing import List, Dict - -from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE, KEYWORD_OBS_NODES, ALL_ONLINE_NODES, \ - QUERY_KEYWORDS -from enumeration.memory_recall_enum import MemoryRecallType -from scheme.memory_node import MemoryNode -from worker.memory_base_worker import MemoryBaseWorker - - -class SemanticRankWorker(MemoryBaseWorker): - - def user_profile_to_nodes(self) -> List[MemoryNode]: - user_profile_nodes: List[MemoryNode] = self.user_profile_dict - for node in user_profile_nodes: - # 从画像侧召回 - 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: List[MemoryNode] = self.get_context(SIMILAR_OBS_NODES) - if similar_obs_nodes: - for node in similar_obs_nodes: - 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[node.content] = node - - if not all_node_dict: - self.add_run_info("all_node_dict is empty!", continue_run=False) - return - - # call recall model - query_keywords = self.get_context(QUERY_KEYWORDS) - # TODO 根据效果更改 - # query: str = "用户:" + self.messages[-1].content - query: str = self.messages[-1].content - if query_keywords: - query_keyword_join = ",".join(query_keywords) - query = f"{query} 用户的{query_keyword_join}。" - documents = list(all_node_dict.keys()) - result = self.rank_model.call(query=query, documents=documents) - - if not result: - self.add_run_info("semantic call recall model failed!") - return - - # set score - for index, score in result.rank_scores.items(): - content = documents[index] - node = all_node_dict[content] - 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()) - self.set_context(ALL_ONLINE_NODES, all_online_nodes) diff --git a/old/worker/summary_long/__init__.py b/old/worker/summary_long/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/old/worker/summary_long/get_insight_worker.py b/old/worker/summary_long/get_insight_worker.py deleted file mode 100644 index c7be74e8..00000000 --- a/old/worker/summary_long/get_insight_worker.py +++ /dev/null @@ -1,166 +0,0 @@ -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" diff --git a/old/worker/summary_long/get_reflection_worker.py b/old/worker/summary_long/get_reflection_worker.py deleted file mode 100644 index 1647c9b6..00000000 --- a/old/worker/summary_long/get_reflection_worker.py +++ /dev/null @@ -1,99 +0,0 @@ -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) diff --git a/old/worker/summary_long/long_contra_repeat_worker.py b/old/worker/summary_long/long_contra_repeat_worker.py deleted file mode 100644 index 42f2705e..00000000 --- a/old/worker/summary_long/long_contra_repeat_worker.py +++ /dev/null @@ -1,129 +0,0 @@ -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) diff --git a/old/worker/summary_long/summary_collect_worker.py b/old/worker/summary_long/summary_collect_worker.py deleted file mode 100644 index 62e0699b..00000000 --- a/old/worker/summary_long/summary_collect_worker.py +++ /dev/null @@ -1,47 +0,0 @@ -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())) diff --git a/old/worker/summary_long/update_insight_worker.py b/old/worker/summary_long/update_insight_worker.py deleted file mode 100644 index 93e246a0..00000000 --- a/old/worker/summary_long/update_insight_worker.py +++ /dev/null @@ -1,177 +0,0 @@ -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}" - ) diff --git a/old/worker/summary_long/update_profile_worker.py b/old/worker/summary_long/update_profile_worker.py deleted file mode 100644 index de2b9061..00000000 --- a/old/worker/summary_long/update_profile_worker.py +++ /dev/null @@ -1,241 +0,0 @@ -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) diff --git a/old/worker/summary_short/__init__.py b/old/worker/summary_short/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/old/worker/summary_short/contra_repeat_worker.py b/old/worker/summary_short/contra_repeat_worker.py deleted file mode 100644 index c1168d58..00000000 --- a/old/worker/summary_short/contra_repeat_worker.py +++ /dev/null @@ -1,117 +0,0 @@ -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) diff --git a/old/worker/summary_short/get_observation_with_time_worker.py b/old/worker/summary_short/get_observation_with_time_worker.py deleted file mode 100644 index 0ff9aeb4..00000000 --- a/old/worker/summary_short/get_observation_with_time_worker.py +++ /dev/null @@ -1,167 +0,0 @@ -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) diff --git a/old/worker/summary_short/get_observation_worker.py b/old/worker/summary_short/get_observation_worker.py deleted file mode 100644 index 63c74339..00000000 --- a/old/worker/summary_short/get_observation_worker.py +++ /dev/null @@ -1,144 +0,0 @@ -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) diff --git a/old/worker/summary_short/info_filter_worker.py b/old/worker/summary_short/info_filter_worker.py deleted file mode 100644 index f030642e..00000000 --- a/old/worker/summary_short/info_filter_worker.py +++ /dev/null @@ -1,70 +0,0 @@ -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