[dev] fix active status to new

This commit is contained in:
jinli.yl 2024-07-08 00:00:47 +08:00
parent 4d6be46a1e
commit e758cf41e6
58 changed files with 36 additions and 4423 deletions

View file

@ -2,7 +2,7 @@
exclude =
scripts/*
src/agentscope/rpc/*
max-line-length = 79
max-line-length = 120
inline-quotes = "
avoid-escape = no
ignore =

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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> <AAA>
<2> <BBB> ddd
<4> <CCC> 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> <AAA> <222>
<2> <BBB>
<4,5> <CCC>
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> <AAA>
<2c> <BBB>
<41> <CCC> <BBB>
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]])
},
]

View file

@ -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 1value可变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="占位符字典")

View file

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

View file

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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