add cli_config & fix api

This commit is contained in:
hs 2024-06-18 16:50:45 +08:00
parent a154d4c160
commit 9c79fd95d8
37 changed files with 162 additions and 122 deletions

View file

@ -1 +1,66 @@
# runtime params apikey or overwrite config params
import json
from pydantic import BaseModel
from models.dash_embedding_client import DashEmbeddingClient
from models.dash_generate_client import DashGenerateClient
from models.dash_rerank_client import DashReRankClient
from models.elastic_search_client import ElasticSearchClient
from pipeline.memory_service import MemoryService
from enumeration.memory_type_enum import MemoryTypeEnum
from utils.tool_functions import init_instance_by_config
class Wrapper:
"""Wrapper class for anything that needs to set up during init"""
def __init__(self):
self._provider = None
def register(self, provider):
self._provider = provider
def __getattr__(self, key):
if self.__dict__.get("_provider", None) is None:
raise AttributeError("Please run init() first using qlib")
return getattr(self._provider, key)
class Config(BaseModel):
thread_pool_max_count: int = 5
pipeline: dict = {
"retrive": """
parse_params,es.load_profile,[retrieve.extract_time|es.es_similar|es.es_keyword],retrieve.semantic_rank,retrieve.fuse_rerank
""".strip(),
"summary_long": """
parse_params,summary_short.info_filter,[es.es_today_obs|summary_short.get_observation|summary_short.get_observation_with_time],summary_short.contra_repeat,memory_store
""".strip(),
"summary_short": """
parse_params,[es.load_profile|es.es_new_obs|es.es_insight],[summary_long.update_insight|summary_long.get_reflection,summary_long.get_insight|summary_long.update_profile],summary_long.summary_collect,memory_store
""".strip()
}
model_embedding = _default_embedding_client
model_generate = _default_generate_client
model_rerank = _default_rerank_client
db = _default_es_client
class UserConfig(BaseModel):
memory_id: str = ""
C = Wrapper()
def init(config_path):
config = json.loads(config_path)
C.register(Config(**config))
# ## register modules
C.model_embedding = init_instance_by_config(C.model_embedding)
C.model_generate = init_instance_by_config(C.model_generate)
C.model_rerank = init_instance_by_config(C.model_rerank)
C.db = init_instance_by_config(C.db)
## register workers
C.worker = json.loads(C.worker)
## register services
for k,v in C.pipeline.items():
C.pipeline[k] = MemoryService(k)

View file

@ -2,7 +2,7 @@ from elasticsearch import Elasticsearch
from elasticsearch.helpers import bulk
from common.dash_embedding_client import DashEmbeddingClient
from common.logger import Logger
from utils.logger import Logger
from constants.common_constants import ES_ENV_URL_DICT
from enumeration.env_type import EnvType

View file

@ -4,8 +4,8 @@ from http import HTTPStatus
import requests
from common.logger import Logger
from common.timer import Timer
from utils.logger import Logger
from utils.timer import Timer
from enumeration.env_type import EnvType

View file

@ -1,6 +1,6 @@
from pydantic import Field, BaseModel
from model.memory_node import MemoryNode
from node.memory_node import MemoryNode
class MemoryWrapNode(BaseModel):

View file

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

View file

@ -1,12 +1,11 @@
from typing import List, Dict
from pydantic import Field
from pydantic import Field, BaseModel
from model.message import Message
from model.user_attribute import UserAttribute
from request.base_model import RequestBaseModel
from node.message import Message
from node.user_attribute import UserAttribute
class MemoryServiceRequestModel(RequestBaseModel):
class MemoryServiceRequestModel(BaseModel):
user: UserConfig = None
messages: List[Message] = Field(...,

View file

@ -6,17 +6,16 @@ from importlib import import_module
from itertools import zip_longest
from typing import Dict, Any
from worker.memory.base_worker import BaseWorker
from common.context_handler import ContextHandler
from common.logger import Logger
from common.timer import timer, Timer
from worker.base_worker import BaseWorker
from utils.context_handler import ContextHandler
from utils.logger import Logger
from utils.timer import timer, Timer
from common.tool_functions import under_line_to_hump
from constants import common_constants
from constants.common_constants import RESPONSE_EXT_INFO, MAX_WORKERS, PIPELINE
from enumeration.memory_method_enum import MemoryMethodEnum
from request.memory import MemoryServiceRequestModel
from config.env_config import C, Workers
from pipeline.memory import MemoryServiceRequestModel
from cli.cli_config import C
from utils.tool_functions import init_instance_by_config
@ -26,7 +25,7 @@ class MemoryService(object):
self.context_handler = ContextHandler()
# 线程池
self.thread_pool = ThreadPoolExecutor(max_workers=C.THREAD_POOL_MAX_COUNT)
self.thread_pool = ThreadPoolExecutor(max_workers=C.thread_pool_max_count)
# 全部初始化的worker
self.worker_dict: Dict[str, BaseWorker] = {}
@ -40,7 +39,7 @@ class MemoryService(object):
def get_worker(self, worker_name: str, is_multi_thread: bool = False) -> BaseWorker:
return init_instance_by_config(
config = W.get(worker_name),
config = C.worker.get(worker_name),
try_kwargs={
"is_multi_thread": is_multi_thread,
"thread_pool": self.thread_pool

View file

@ -2,7 +2,7 @@ import os
import threading
from typing import Dict, Any
from common.logger import Logger
from utils.logger import Logger
class ContextHandler(object):
def __init__(self):

View file

@ -6,7 +6,7 @@ date: 20221106
import time
from common.logger import Logger
from utils.logger import Logger
class Timer(object):

View file

@ -4,29 +4,8 @@ from datetime import datetime
from typing import Dict, List
from constants.common_constants import WEEKDAYS
from enumeration.env_type import EnvType
from importlib import import_module
global_env_type = None
def get_global_env_type():
global global_env_type
if global_env_type is None:
env = os.environ.get("APP_ENV", "")
if env is None or not env:
raise EnvironmentError("Environment variable APP_ENV must be set")
env = env.split("-")[-1]
if env not in EnvType.__members__.values():
global_env_type = EnvType.DAILY
else:
global_env_type = EnvType(env)
return global_env_type
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:]
@ -165,10 +144,10 @@ def time_to_formatted_str(time: datetime | str | int | float = None,
return return_str
def init_instance_by_config(config, default_module_path=None, try_kwargs={}):
import_module(config.get("module_path", default_module_path))
clazz = getattr(module, config.get("module_name"))
def init_instance_by_config(config: dict, default_module_path=None, try_kwargs={}):
import_module(config.pop("path", default_module_path))
clazz = getattr(module, config.pop("name"))
try:
return clazz(**config.get("kwargs"),**try_kwargs)
return clazz(**config,**try_kwargs)
except:
return clazz(**config.get("kwargs"))
return clazz(**config)

View file

@ -2,8 +2,8 @@ import json
from typing import List, Dict
from enumeration.memory_node_status import MemoryNodeStatus
from model.memory_wrap_node import MemoryWrapNode
from model.user_attribute import UserAttribute
from node.memory_wrap_node import MemoryWrapNode
from node.user_attribute import UserAttribute
class UserProfileHandler(object):

View file

@ -2,9 +2,9 @@ import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from typing import Any, List
from common.context_handler import ContextHandler
from common.logger import Logger
from common.timer import Timer
from utils.context_handler import ContextHandler
from utils.logger import Logger
from utils.timer import Timer
class BaseWorker(object):

View file

@ -3,8 +3,8 @@ from typing import List
from constants.common_constants import INSIGHT_NODES
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from model.memory.memory_wrap_node import MemoryWrapNode
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class EsInsightWorker(MemoryBaseWorker):

View file

@ -4,8 +4,8 @@ from constants.common_constants import KEY_WORD, KEYWORD_OBS_NODES, RECALL_TYPE,
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_recall_type import MemoryRecallType
from enumeration.memory_type_enum import MemoryTypeEnum
from model.memory_wrap_node import MemoryWrapNode
from worker.bailian.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class EsKeywordWorker(MemoryBaseWorker):

View file

@ -3,8 +3,8 @@ from typing import List
from constants.common_constants import NEW, NEW_OBS_NODES
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from model.memory.memory_wrap_node import MemoryWrapNode
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class EsNewObsWorker(MemoryBaseWorker):

View file

@ -3,8 +3,8 @@ from typing import List
from constants.common_constants import REFLECTED, NOT_REFLECTED_OBS_NODES
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from model.memory.memory_wrap_node import MemoryWrapNode
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class EsNotReflectedWorker(MemoryBaseWorker):

View file

@ -2,8 +2,8 @@ from typing import List
from constants.common_constants import ALL_NODES, ALL_MEMORIES
from enumeration.memory_node_status import MemoryNodeStatus
from model.memory.memory_wrap_node import MemoryWrapNode
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class EsRetrieveAllWorker(MemoryBaseWorker):

View file

@ -4,8 +4,8 @@ from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_recall_type import MemoryRecallType
from enumeration.memory_type_enum import MemoryTypeEnum
from model.memory.memory_wrap_node import MemoryWrapNode
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class EsSimilarWorker(MemoryBaseWorker):

View file

@ -4,9 +4,8 @@ from common.tool_functions import time_to_formatted_str
from constants.common_constants import TODAY_OBS_NODES, DT
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from model.memory.memory_wrap_node import MemoryWrapNode
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class EsTodayObsWorker(MemoryBaseWorker):
def __init__(self, es_today_obs_top_k, *args, **kwargs):

View file

@ -1,13 +1,13 @@
from typing import List, Dict
from common.user_profile_handler import UserProfileHandler
from utils.user_profile_handler import UserProfileHandler
from constants import common_constants
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from model.memory.memory_wrap_node import MemoryWrapNode
from model.user_attribute import UserAttribute
from request.memory import MemoryServiceRequestModel
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from node.user_attribute import UserAttribute
from pipeline.memory import MemoryServiceRequestModel
from worker.memory_base_worker import MemoryBaseWorker
class LoadProfileWorker(MemoryBaseWorker):

View file

@ -3,11 +3,11 @@ from typing import List, Dict, Optional
from constants import common_constants
from constants.common_constants import CONFIG, MESSAGES, PROMPT_CONFIG
from enumeration.message_role_enum import MessageRoleEnum
from model.message import Message
from model.user_attribute import UserAttribute
from request.memory import MemoryServiceRequestModel
from worker.memory.base_worker import BaseWorker
from config.env_config import EnvConfig
from node.message import Message
from node.user_attribute import UserAttribute
from pipeline.memory import MemoryServiceRequestModel
from worker.base_worker import BaseWorker
from cli.cli_config import C
class MemoryBaseWorker(BaseWorker):
def __init__(self, **kwargs):
@ -52,19 +52,19 @@ class MemoryBaseWorker(BaseWorker):
@property
def emb_client(self):
return C.emb_client
return C.model_embedding
@property
def gene_client(self):
return C.gene_client
return C.model_generate
@property
def rerank_client(self):
return C.rerank_client
return C.model_rerank
@property
def es_client(self):
return C.es_client
return C.db
@property
def tenant_id(self):

View file

@ -1,12 +1,11 @@
from typing import List
from common.user_profile_handler import UserProfileHandler
from utils.user_profile_handler import UserProfileHandler
from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE
from model.memory_node import MemoryNode
from model.memory.memory_wrap_node import MemoryWrapNode
from model.user_attribute import UserAttribute
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_node import MemoryNode
from node.memory_wrap_node import MemoryWrapNode
from node.user_attribute import UserAttribute
from worker.memory_base_worker import MemoryBaseWorker
class MemoryStoreWorker(MemoryBaseWorker):

View file

@ -2,8 +2,8 @@ import json
from config.bailian_memory_config import BailianMemoryConfig
from constants.common_constants import REQUEST, CONFIG
from request.memory import MemoryServiceRequestModel
from worker.memory.base_worker import BaseWorker
from pipeline.memory import MemoryServiceRequestModel
from worker.base_worker import BaseWorker
class ParseParamsWorker(BaseWorker):

View file

@ -3,7 +3,7 @@ import re
from common.tool_functions import time_to_formatted_str
from constants.common_constants import DATATIME_WORD_LIST, DATATIME_KEY_MAP
from constants.common_constants import EXTRACT_TIME_DICT
from worker.memory.memory_base_worker import MemoryBaseWorker
from worker.memory_base_worker import MemoryBaseWorker
class ExtractTimeWorker(MemoryBaseWorker):

View file

@ -2,8 +2,8 @@ from typing import Dict, List
from constants.common_constants import RELATED_MEMORIES, EXTRACT_TIME_DICT, ALL_ONLINE_NODES, \
TIME_MATCHED
from model.memory.memory_wrap_node import MemoryWrapNode
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class FuseRerankWorker(MemoryBaseWorker):

View file

@ -1,11 +1,11 @@
from typing import List, Dict
from common.user_profile_handler import UserProfileHandler
from utils.user_profile_handler import UserProfileHandler
from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE, KEYWORD_OBS_NODES, ALL_ONLINE_NODES, \
QUERY_KEYWORDS
from enumeration.memory_recall_type import MemoryRecallType
from model.memory.memory_wrap_node import MemoryWrapNode
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class SemanticRankWorker(MemoryBaseWorker):

View file

@ -6,8 +6,8 @@ from constants.common_constants import NEW_INSIGHT_NODES, DT, NOT_REFLECTED_MERG
INSIGHT_VALUE, REFLECTED
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from model.memory.memory_wrap_node import MemoryWrapNode
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class GetInsightWorker(MemoryBaseWorker):

View file

@ -3,8 +3,8 @@ from typing import List
from common.response_text_parser import ResponseTextParser
from constants.common_constants import NEW_OBS_NODES, NOT_REFLECTED_OBS_NODES, REFLECTED, INSIGHT_NODES, INSIGHT_KEY, \
NEW_INSIGHT_KEYS, NOT_REFLECTED_MERGE_NODES
from model.memory.memory_wrap_node import MemoryWrapNode
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class GetReflectionWorker(MemoryBaseWorker):

View file

@ -5,8 +5,8 @@ from constants.common_constants import NEW_OBS_NODES, TODAY_OBS_NODES, MSG_TIME,
MODIFIED_MEMORIES
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from model.memory_wrap_node import MemoryWrapNode
from worker.bailian.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class LongContraRepeatWorker(MemoryBaseWorker):

View file

@ -2,8 +2,8 @@ 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 model.memory.memory_wrap_node import MemoryWrapNode
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class SummaryCollectWorker(MemoryBaseWorker):

View file

@ -2,8 +2,8 @@ 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 model.memory.memory_wrap_node import MemoryWrapNode
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class UpdateInsightWorker(MemoryBaseWorker):

View file

@ -3,9 +3,9 @@ from typing import List
from common.response_text_parser import ResponseTextParser
from constants.common_constants import NEW_OBS_NODES, NEW_USER_PROFILE
from enumeration.memory_type_enum import MemoryTypeEnum
from model.memory.memory_wrap_node import MemoryWrapNode
from model.user_attribute import UserAttribute
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from node.user_attribute import UserAttribute
from worker.memory_base_worker import MemoryBaseWorker
class UpdateProfileWorker(MemoryBaseWorker):

View file

@ -4,8 +4,8 @@ from common.response_text_parser import ResponseTextParser
from constants.common_constants import NEW_OBS_NODES, TODAY_OBS_NODES, MSG_TIME, NEW_OBS_WITH_TIME_NODES, \
MODIFIED_MEMORIES
from enumeration.memory_node_status import MemoryNodeStatus
from model.memory.memory_wrap_node import MemoryWrapNode
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from worker.memory_base_worker import MemoryBaseWorker
class ContraRepeatWorker(MemoryBaseWorker):

View file

@ -7,9 +7,9 @@ from constants.common_constants import REFLECTED, DT, TIME_INFER, NEW, MSG_TIME,
NEW_OBS_WITH_TIME_NODES
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from model.memory.memory_wrap_node import MemoryWrapNode
from model.message import Message
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from node.message import Message
from worker.memory_base_worker import MemoryBaseWorker
class GetObservationWithTimeWorker(MemoryBaseWorker):

View file

@ -7,9 +7,9 @@ from constants.common_constants import REFLECTED, DT, NEW_OBS_NODES, TIME_INFER,
DATATIME_WORD_LIST
from enumeration.memory_node_status import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from model.memory.memory_wrap_node import MemoryWrapNode
from model.message import Message
from worker.memory.memory_base_worker import MemoryBaseWorker
from node.memory_wrap_node import MemoryWrapNode
from node.message import Message
from worker.memory_base_worker import MemoryBaseWorker
class GetObservationWorker(MemoryBaseWorker):

View file

@ -1,6 +1,6 @@
from common.response_text_parser import ResponseTextParser
from enumeration.message_role_enum import MessageRoleEnum
from worker.memory.memory_base_worker import MemoryBaseWorker
from worker.memory_base_worker import MemoryBaseWorker
class InfoFilterWorker(MemoryBaseWorker):

View file

@ -3,12 +3,12 @@ import os
import time
from typing import List, Dict
from common.logger import Logger
from utils.logger import Logger
from constants.common_constants import NEW_USER_PROFILE, MODIFIED_MEMORIES, RELATED_MEMORIES
from enumeration.memory_method_enum import MemoryMethodEnum
from model.memory_node import MemoryNode
from model.user_attribute import UserAttribute
from request.memory import MemoryServiceRequestModel
from node.memory_node import MemoryNode
from node.user_attribute import UserAttribute
from pipeline.memory import MemoryServiceRequestModel
from service.memory_service_bailian import MemoryServiceBailian
"""