diff --git a/memory_scope/chat/base_memory_chat.py b/memory_scope/chat/base_memory_chat.py index 79a20758..c336df6d 100644 --- a/memory_scope/chat/base_memory_chat.py +++ b/memory_scope/chat/base_memory_chat.py @@ -1,43 +1,16 @@ from abc import ABCMeta, abstractmethod -from memory_scope.constants.common_constants import RELATED_MEMORIES -from memory_scope.enumeration.memory_method_enum import MemoryMethodEnum -from memory_scope.handler.pipeline_handler import PipelineHandler +from memory_scope.chat.memory_service import MemoryService class BaseMemoryChat(metaclass=ABCMeta): - def __init__(self, - user_name: str, - retrieve_pipeline: str, - summary_short_pipeline: str, - summary_long_pipeline: str, - **kwargs): - self.user_name: str = user_name - self.retrieve_pipeline_handler = PipelineHandler(user_name=user_name, - memory_method_type=MemoryMethodEnum.RETRIEVE, - pipeline_str=retrieve_pipeline) - - self.summary_short_pipeline_handler = PipelineHandler(user_name=user_name, - memory_method_type=MemoryMethodEnum.SUMMARY_SHORT, - pipeline_str=summary_short_pipeline) - - self.summary_long_pipeline_handler = PipelineHandler(user_name=user_name, - memory_method_type=MemoryMethodEnum.SUMMARY_LONG, - pipeline_str=summary_long_pipeline) - - def retrieve(self): - self.retrieve_pipeline_handler.run() - return self.retrieve_pipeline_handler.get_context(RELATED_MEMORIES, []) - - def summary_short(self): - self.summary_short_pipeline_handler.run() - - def summary_long(self): - self.summary_long_pipeline_handler.run() + def __init__(self, **kwargs): + self.memory_service = MemoryService(**kwargs) @abstractmethod - def chat(self): + def chat_with_memory(self, query: str): """ + :param query: :return: """ diff --git a/memory_scope/chat/memory_chat.py b/memory_scope/chat/memory_chat.py index e03dc686..4fe6efbb 100644 --- a/memory_scope/chat/memory_chat.py +++ b/memory_scope/chat/memory_chat.py @@ -1,29 +1,51 @@ -from memory_scope.chat.memory_service import MemoryService -from memory_scope.handler.init_handler import InitHandler +import datetime +from typing import List + +from memory_scope.chat.base_memory_chat import BaseMemoryChat +from memory_scope.enumeration.message_role_enum import MessageRoleEnum +from memory_scope.handler.global_context import GLOBAL_CONTEXT +from memory_scope.models.base_model import BaseModel +from memory_scope.node.message import Message +from memory_scope.prompts.prompt_cn import SYSTEM_PROMPT, MEMORY_PROMPT -class MemoryChat(object): +class MemoryChat(BaseMemoryChat): + """ + TODO add agent + """ - def __init__(self, init_handler: InitHandler): - self.init_handler: InitHandler = init_handler + def __init__(self, + generation_model: str, + history_msg_count: int, + **kwargs): + super().__init__(**kwargs) + self.model: BaseModel = GLOBAL_CONTEXT.model_dict[generation_model] + self.history_msg_count: int = history_msg_count - self.memory_service: MemoryService = MemoryService( - retrieve_pipeline=init_handler.retrieve_pipeline, - summary_short_pipeline=init_handler.retrieve_pipeline, - summary_long_pipeline=init_handler.retrieve_pipeline, - ) + self.memory_service.start_summary_short_backend() + self.memory_service.start_summary_long_backend() - def memory_retrieve(self): - pass + self.history_message_list: List[Message] = [] - def memory_summary_short(self): - pass + @staticmethod + def get_system_prompt(related_memories: List[str], time_created: int) -> Message: + system_prompt = SYSTEM_PROMPT + if related_memories: + system_prompt = "\n".join([SYSTEM_PROMPT, MEMORY_PROMPT] + related_memories) + return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt.strip(), time_created=time_created) - def memory_summary_long(self): - pass + def chat_with_memory(self, query: str): + query = query.strip() + if not query: + return + time_created = int(datetime.datetime.now().timestamp()) + new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created) + related_memories: List[str] = self.memory_service.retrieve(message=new_message) + system_message = self.get_system_prompt(related_memories, time_created) + self.history_message_list.append(new_message) + self.history_message_list = self.history_message_list[-self.history_msg_count:] + all_messages = [system_message] + self.history_message_list + return self.model.call(messages=all_messages) - def chat(self): - pass - - def chat_with_memory(self): - pass + def chat_with_memory_stream(self): + raise NotImplementedError diff --git a/memory_scope/chat/memory_service.py b/memory_scope/chat/memory_service.py new file mode 100644 index 00000000..011fe435 --- /dev/null +++ b/memory_scope/chat/memory_service.py @@ -0,0 +1,55 @@ +from memory_scope.constants.common_constants import RELATED_MEMORIES +from memory_scope.enumeration.memory_method_enum import MemoryMethodEnum +from memory_scope.handler.pipeline_handler import PipelineHandler +from memory_scope.node.message import Message + + +class MemoryService(object): + + def __init__(self, + user_name: str, + retrieve_pipeline: str, + retrieve_all_pipeline: str, + summary_short_pipeline: str, + summary_long_pipeline: str, + summary_short_interval_time: int = 60, + summary_short_minimum_count: int = 5, + summary_long_interval_time: int = 60 * 5, + summary_long_minimum_count: int = 5 * 5, + **kwargs): + self.retrieve_pipeline_handler = PipelineHandler(user_name=user_name, + memory_method_type=MemoryMethodEnum.RETRIEVE, + pipeline_str=retrieve_pipeline) + + self.retrieve_all_pipeline_handler = PipelineHandler(user_name=user_name, + memory_method_type=MemoryMethodEnum.RETRIEVE_ALL, + pipeline_str=retrieve_all_pipeline) + + self.summary_short_pipeline_handler = PipelineHandler(user_name=user_name, + memory_method_type=MemoryMethodEnum.SUMMARY_SHORT, + pipeline_str=summary_short_pipeline, + loop_interval_time=summary_short_interval_time, + loop_minimum_count=summary_short_minimum_count) + + self.summary_long_pipeline_handler = PipelineHandler(user_name=user_name, + memory_method_type=MemoryMethodEnum.SUMMARY_LONG, + pipeline_str=summary_long_pipeline, + loop_interval_time=summary_long_interval_time, + loop_minimum_count=summary_long_minimum_count) + + self.kwargs = kwargs + + def retrieve(self, message: Message): + self.retrieve_pipeline_handler.submit_message(message, with_lock=False) + self.summary_short_pipeline_handler.submit_message(message) + self.summary_long_pipeline_handler.submit_message(message) + return self.retrieve_pipeline_handler.run(RELATED_MEMORIES) + + def retrieve_all(self): + return self.retrieve_all_pipeline_handler.run(RELATED_MEMORIES) + + def start_summary_short_backend(self): + self.summary_short_pipeline_handler.start_loop_run() + + def start_summary_long_backend(self): + self.summary_long_pipeline_handler.start_loop_run() diff --git a/memory_scope/cli.py b/memory_scope/cli.py index 2ab398d3..7efac97f 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -1,11 +1,96 @@ +import json +import os +from concurrent.futures import ThreadPoolExecutor +from typing import Dict, Any + import fire -from memory_scope.job import Job +from handler.global_context import GLOBAL_CONTEXT +from memory_scope.enumeration.model_type import ModelType +from memory_scope.utils.logger import Logger +from memory_scope.utils.timer import Timer +from memory_scope.utils.tool_functions import complete_config_name, init_instance_by_config_v2 + + +class CliJob(object): + + def __init__(self, config_path: str): + self.config_path: str = config_path + self.config_base_dir: str = os.path.dirname(config_path) + + self.config: Dict[str, Any] = {} + + self.logger: Logger = Logger.get_logger("memory_chat") + + def init_memory_chat(self): + for chat in self.config["chat_list"]: + memory_chat_config = self.config[chat] + memory_chat = init_instance_by_config_v2(memory_chat_config) + GLOBAL_CONTEXT.memory_chat_dict[chat] = memory_chat + + generation_model = memory_chat_config[ModelType.GENERATION_MODEL.value] + self.init_model(generation_model) + + def init_model(self, model_name: str): + if not model_name or model_name in GLOBAL_CONTEXT.model_dict: + return + + with open(os.path.join(self.config_base_dir, "model", complete_config_name(model_name))) as f: + model_config = json.load(f) + GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config_v2(model_config) + + def init_workers(self): + """ load worker config & init workers + """ + worker_config_name: str = self.config["workers"] + with open(os.path.join(self.config_base_dir, complete_config_name(worker_config_name))) as f: + worker_config_dict = json.load(f) + + for worker_name, worker_config in worker_config_dict.items(): + if worker_name in GLOBAL_CONTEXT.worker_dict: + raise RuntimeError(f"worker_name={worker_name} is repeated!") + + GLOBAL_CONTEXT.worker_dict[worker_name] = init_instance_by_config_v2(worker_config, + suffix_name="worker", + **GLOBAL_CONTEXT.global_configs) + + self.init_model(worker_config.get(ModelType.EMBEDDING_MODEL.value)) + self.init_model(worker_config.get(ModelType.GENERATION_MODEL.value)) + self.init_model(worker_config.get(ModelType.RANK_MODEL.value)) + + def set_global_config(self): + """set global_configs & set apikey into env + """ + + def init_global_content_by_config(self): + with open(complete_config_name(self.config_path)) as f: + self.config = json.load(f) + + GLOBAL_CONTEXT.global_configs = self.config["global_configs"] + self.set_global_config() + GLOBAL_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(GLOBAL_CONTEXT.global_configs["max_workers"])) + + self.init_workers() + GLOBAL_CONTEXT.db_client = init_instance_by_config_v2(self.config["db"]) + GLOBAL_CONTEXT.monitor = init_instance_by_config_v2(self.config["monitor"]) + + self.init_memory_chat() + + def run(self): + with GLOBAL_CONTEXT.thread_pool, Timer("job", log_time=False) as t: + memory_chat = list(GLOBAL_CONTEXT.memory_chat_dict.values())[0] + while True: + query = input("wait for input:") + if query in ["stop", "停止"]: + break + memory_chat.chat_with_memory(query=query) + + self.logger.info(f"chat complete. cost={t.cost_str}") def main(config_path: str): - job = Job(config_path=config_path) - job.init_instance_by_config() + job = CliJob(config_path=config_path) + job.init_global_content_by_config() job.run() diff --git a/memory_scope/config.py b/memory_scope/config.py deleted file mode 100644 index d69c08c5..00000000 --- a/memory_scope/config.py +++ /dev/null @@ -1,34 +0,0 @@ - -import json - -from pipeline.memory_service import MemoryService -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: str): - if self.__dict__.get("_provider", None) is None: - raise AttributeError("Please run __init__ first!") - return getattr(self._provider, key) - -C = Wrapper() - -def init(config_path: str): - config = json.loads(config_path) - C.register(config) - - ## register workers - C.worker = json.loads(C.worker) - - ## register services - for k,v in C.pipeline.items(): - C.pipeline[k] = MemoryService(k) - diff --git a/memory_scope/constants/common_constants.py b/memory_scope/constants/common_constants.py index c276cf02..10768e24 100644 --- a/memory_scope/constants/common_constants.py +++ b/memory_scope/constants/common_constants.py @@ -9,6 +9,8 @@ WORKER = "worker" MEMORY = "memory" +USER_NAME = "user_name" + DEFAULT_SYSTEM_PROMPT = "default_system_prompt" RELATED_MEMORIES = "related_memories" diff --git a/memory_scope/enumeration/memory_method_enum.py b/memory_scope/enumeration/memory_method_enum.py index fa8a7092..b30adf31 100644 --- a/memory_scope/enumeration/memory_method_enum.py +++ b/memory_scope/enumeration/memory_method_enum.py @@ -6,6 +6,8 @@ class MemoryMethodEnum(str, Enum): RETRIEVE = "retrieve" + RETRIEVE_ALL = "retrieve_all" + SUMMARY_SHORT = "summary_short" SUMMARY_LONG = "summary_long" diff --git a/memory_scope/enumeration/model_type.py b/memory_scope/enumeration/model_type.py new file mode 100644 index 00000000..25af770a --- /dev/null +++ b/memory_scope/enumeration/model_type.py @@ -0,0 +1,9 @@ +from enum import Enum + + +class ModelType(str, Enum): + GENERATION_MODEL = "generation_model" + + EMBEDDING_MODEL = "embedding_model" + + RANK_MODEL = "rank_model" diff --git a/memory_scope/handler/pipeline_handler.py b/memory_scope/handler/pipeline_handler.py index 9982cb3d..e6161d50 100644 --- a/memory_scope/handler/pipeline_handler.py +++ b/memory_scope/handler/pipeline_handler.py @@ -1,21 +1,33 @@ import re import threading +import time from concurrent.futures import as_completed from itertools import zip_longest -from typing import Dict, Any +from typing import Dict, Any, List +from memory_scope.constants.common_constants import MESSAGES, USER_NAME from memory_scope.enumeration.memory_method_enum import MemoryMethodEnum from memory_scope.handler.global_context import GLOBAL_CONTEXT +from memory_scope.node.message import Message from memory_scope.utils.logger import Logger from memory_scope.utils.timer import Timer class PipelineHandler(object): - def __init__(self, user_name: str, memory_method_type: MemoryMethodEnum, pipeline_str: str): + def __init__(self, + user_name: str, + memory_method_type: MemoryMethodEnum, + pipeline_str: str, + history_msg_count: int = 3, + loop_interval_time: int = 300, + loop_minimum_count: int = 20): self.user_name: str = user_name self.memory_method_type: MemoryMethodEnum = memory_method_type self.pipeline_str: str = pipeline_str + self.history_msg_count: int = history_msg_count + self.loop_interval_time: int = loop_interval_time + self.loop_minimum_count: int = loop_minimum_count # 日志 self.logger: Logger = Logger.get_logger() @@ -24,11 +36,19 @@ class PipelineHandler(object): self.context: Dict[str, Any] = {} self.context_lock = threading.Lock() + # pipeline run config + self.loop_switch: bool = False + # 解析和打印 pipeline self.pipeline_list: list[list] = [] self._parse_pipeline() self._print_pipeline() + # message list + self.history_message_list: List[Message] = [] + self.current_message_list: List[Message] = [] + self.message_lock = threading.Lock() + def _parse_pipeline(self): # re-match e.g., [a|b],c,[d,e,f|g,h],j pattern = r'(\[[^\]]*\]|[^,]+)' @@ -70,14 +90,8 @@ class PipelineHandler(object): GLOBAL_CONTEXT.worker_dict[w].context = self.context self.logger.info(f"----- {self.user_name} {self.memory_method_type.value} Pipeline End -----") - def get_context(self, key: str, default=None): - return self.context.get(key, default) - - def clear_context(self): - self.context.clear() - @staticmethod - def worker_run(worker_list: list[str]) -> bool: + def _worker_run(worker_list: list[str]) -> bool: for worker_name in worker_list: worker = GLOBAL_CONTEXT.worker_dict[worker_name] worker.run() @@ -85,16 +99,19 @@ class PipelineHandler(object): return False return True - def run(self): + def _run(self, result_key: str = None): with Timer(f"pipeline_{self.user_name}_{self.memory_method_type.value}"): + self.context[MESSAGES] = self.history_message_list + self.current_message_list + self.context[USER_NAME] = self.user_name + for pipeline_part in self.pipeline_list: if len(pipeline_part) == 1: - if not self.worker_run(pipeline_part[0]): + if not self._worker_run(pipeline_part[0]): break else: t_list = [] for worker_list in pipeline_part: - t_list.append(GLOBAL_CONTEXT.thread_pool.submit(self.worker_run, worker_list)) + t_list.append(GLOBAL_CONTEXT.thread_pool.submit(self._worker_run, worker_list)) flag = True for future in as_completed(t_list): @@ -103,3 +120,47 @@ class PipelineHandler(object): break if not flag: break + if result_key: + return self.context.get(result_key) + self.context.clear() + + return None + + def _thread_loop(self): + while self.loop_switch: + time.sleep(self.loop_interval_time) + if len(self.current_message_list) < self.loop_minimum_count: + continue + self._run() + self.context.clear() + self.history_message_list = self.history_message_list.extend(self.current_message_list)[ + -self.history_msg_count:] + with self.message_lock: + self.current_message_list.clear() + + def start_loop_run(self): + if not self.loop_switch: + self.loop_switch = True + return GLOBAL_CONTEXT.thread_pool.submit(self._thread_loop) + + def run(self, result_key: str = None): + self._run() + + # 获取result + result = None + if result_key: + result = self.context.get(result_key) + self.context.clear() + + # 清理 msg + self.history_message_list = self.history_message_list.extend(self.current_message_list)[ + -self.history_msg_count:] + self.current_message_list.clear() + return result + + def submit_message(self, message: Message, with_lock=True): + if with_lock: + with self.message_lock: + self.current_message_list.append(message) + else: + self.current_message_list.append(message) diff --git a/memory_scope/job.py b/memory_scope/job.py deleted file mode 100644 index 253c85ca..00000000 --- a/memory_scope/job.py +++ /dev/null @@ -1,78 +0,0 @@ -import json -import os -from concurrent.futures import ThreadPoolExecutor -from typing import Dict, Any - -from handler.global_context import GLOBAL_CONTEXT -from memory_scope.utils.logger import Logger -from memory_scope.utils.timer import Timer -from memory_scope.utils.tool_functions import complete_config_name, init_instance_by_config_v2 - - -class Job(object): - - def __init__(self, config_path: str): - self.config_path: str = config_path - self.config_base_dir: str = os.path.dirname(config_path) - - self.config: Dict[str, Any] = {} - - self.logger: Logger = Logger.get_logger("memory_chat") - - def init_memory_chat(self): - for chat in self.config["chat_list"]: - memory_chat_config = self.config[chat] - memory_chat = init_instance_by_config_v2(memory_chat_config) - GLOBAL_CONTEXT.memory_chat_dict[chat] = memory_chat - - generation_model = memory_chat_config["generation_model"] - self.init_model(generation_model) - - def init_model(self, model_name: str): - if not model_name or model_name in GLOBAL_CONTEXT.model_dict: - return - - with open(os.path.join(self.config_base_dir, "model", complete_config_name(model_name))) as f: - model_config = json.load(f) - GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config_v2(model_config) - - def init_workers(self): - """ load worker config & init workers - """ - worker_config_name: str = self.config["workers"] - with open(os.path.join(self.config_base_dir, complete_config_name(worker_config_name))) as f: - worker_config_dict = json.load(f) - - for worker_name, worker_config in worker_config_dict.items(): - if worker_name in GLOBAL_CONTEXT.worker_dict: - raise RuntimeError(f"worker_name={worker_name} is repeated!") - - GLOBAL_CONTEXT.worker_dict[worker_name] = init_instance_by_config_v2(worker_config, - suffix_name="worker", - **GLOBAL_CONTEXT.global_configs) - - self.init_model(worker_config.get("embedding_model")) - self.init_model(worker_config.get("generation_model")) - self.init_model(worker_config.get("rank_model")) - - def set_global_config(self): - """set global_configs & set apikey into env - """ - - def init_instance_by_config(self): - with open(complete_config_name(self.config_path)) as f: - self.config = json.load(f) - - GLOBAL_CONTEXT.global_configs = self.config["global_configs"] - self.set_global_config() - GLOBAL_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(GLOBAL_CONTEXT.global_configs["max_workers"])) - - self.init_workers() - GLOBAL_CONTEXT.db_client = init_instance_by_config_v2(self.config["db"]) - GLOBAL_CONTEXT.monitor = init_instance_by_config_v2(self.config["monitor"]) - - self.init_memory_chat() - - def run(self): - with GLOBAL_CONTEXT.thread_pool, Timer("job"): - pass diff --git a/memory_scope/node/message.py b/memory_scope/node/message.py index e458ddab..b6e8f785 100644 --- a/memory_scope/node/message.py +++ b/memory_scope/node/message.py @@ -6,6 +6,4 @@ class Message(BaseModel): content: str = Field(..., description="The body of the message") - time_created: str = Field("", description="Timestamp when the message was created") - - info_score: str = Field("", description="2 > 1 > 0") + time_created: int = Field("", description="Timestamp when the message was created") \ No newline at end of file diff --git a/memory_scope/prompts/prompt_cn.py b/memory_scope/prompts/prompt_cn.py new file mode 100644 index 00000000..f64e5cd8 --- /dev/null +++ b/memory_scope/prompts/prompt_cn.py @@ -0,0 +1,7 @@ +SYSTEM_PROMPT = """ + +""" + +MEMORY_PROMPT = """ + +""" \ No newline at end of file diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index 6cd284be..91d8b6a9 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -5,6 +5,8 @@ from typing import Dict, List from constants.common_constants import WEEKDAYS +from memory_scope.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) @@ -181,3 +183,16 @@ def complete_config_name(config_name: str, suffix: str = ".json"): if not config_name.endswith(suffix): config_name += suffix return config_name + + +def prompt_to_msg(system_prompt: str, few_shot: str, user_query: str): + return [ + { + "role": MessageRoleEnum.SYSTEM.value, + "content": system_prompt.strip(), + }, + { + "role": MessageRoleEnum.USER.value, + "content": "\n".join([x.strip() for x in [few_shot, system_prompt, user_query]]) + }, + ] diff --git a/memory_scope/worker/memory_base_worker.py b/memory_scope/worker/memory_base_worker.py index 0abef156..b17f4e7d 100644 --- a/memory_scope/worker/memory_base_worker.py +++ b/memory_scope/worker/memory_base_worker.py @@ -1,101 +1,53 @@ -from typing import List, Dict, Optional +from typing import List + +from memory_scope.constants.common_constants import MESSAGES, USER_NAME +from memory_scope.handler.global_context import GLOBAL_CONTEXT +from memory_scope.models.base_model import BaseModel +from memory_scope.node.message import Message +from memory_scope.worker.base_worker import BaseWorker -from constants import common_constants -from constants.common_constants import CONFIG, MESSAGES, PROMPT_CONFIG -from enumeration.message_role_enum import MessageRoleEnum -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 -from utils.tool_functions import init_instance_by_config class MemoryBaseWorker(BaseWorker): - def __init__(self, **kwargs): + def __init__(self, + embedding_model: str, + generation_model: str, + rank_model: str, + **kwargs): super(MemoryBaseWorker, self).__init__(**kwargs) + self.embedding_model_name: str = embedding_model + self.generation_model_name: str = generation_model + self.rank_model_name: str = rank_model - @property - def request(self) -> MemoryServiceRequestModel: - return self.get_context(common_constants.REQUEST) + self._embedding_model: BaseModel | None = None + self._generation_model: BaseModel | None = None + self._rank_model: BaseModel | None = None @property def messages(self) -> List[Message]: - messages: List[Message] = self.context_handler.get_context(MESSAGES) - if messages is None: - messages = self.request.messages - self.context_handler.set_context(MESSAGES, messages) - return messages + return self.context[MESSAGES] @messages.setter def messages(self, value): - self.context_handler.set_context(MESSAGES, value) - - def flush(self, context_handler): - super(MemoryBaseWorker, self).flush(context_handler) - self._user_profile_dict: Dict[str, UserAttribute] = {} + self.context[MESSAGES] = value @property - def user_profile_dict(self) -> Dict[str, UserAttribute]: - if not self._user_profile_dict: - self._user_profile_dict = {user_attr.memory_key: user_attr for user_attr in self.request.user.user_profile} - return self._user_profile_dict + def embedding_model(self): + if self._embedding_model is None: + GLOBAL_CONTEXT.model_dict.get(self.embedding_model_name) + return self._embedding_model @property - def request_ext_info(self): - return self.request.user.ext_info - # if not self._request_ext_info: - # self._request_ext_info = self.request.ext_info - # return self._request_ext_info + def generation_model(self): + if self._generation_model is None: + GLOBAL_CONTEXT.model_dict.get(self.generation_model_name) + return self._generation_model @property - def prompt_config(self) -> BailianPromptConfig: - return self.request.user.prompt + def rank_model(self): + if self._rank_model is None: + GLOBAL_CONTEXT.model_dict.get(self.rank_model_name) + return self._rank_model @property - def client(self, model_type: str, model_name:str): - models = C.get(model_type) - models["model_name"] = init_instance_by_config( - config = models.get(model_name), - try_kwargs={ - "is_multi_thread": is_multi_thread, - "thread_pool": self.thread_pool - } - ) - return models["model_name"] - - @property - def emb_client(self, model_name: str): - self.client("model_embedding", model_name) - - @property - def gene_client(self): - self.client("model_generate", model_name) - - @property - def rerank_client(self): - self.client("model_rerank", model_name) - - @property - def es_client(self): - self.client("db", model_name) - - @property - def tenant_id(self): - return self.request.user.tenant_id - - @property - def memory_id(self): - return self.request.user.memory_id - - @staticmethod - 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]]) - }, - ] + def user_name(self): + return self.context[USER_NAME] diff --git a/memory_scope/worker/memory_store_worker.py b/memory_scope/worker/retrieve/memory_store_worker.py similarity index 100% rename from memory_scope/worker/memory_store_worker.py rename to memory_scope/worker/retrieve/memory_store_worker.py diff --git a/memory_scope/worker/parse_params_worker.py b/memory_scope/worker/retrieve/parse_params_worker.py similarity index 100% rename from memory_scope/worker/parse_params_worker.py rename to memory_scope/worker/retrieve/parse_params_worker.py