mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-09 03:20:54 +00:00
[dev] add new config format
This commit is contained in:
parent
84166dfc1e
commit
3c423118a5
21 changed files with 978 additions and 309 deletions
0
memoryscope/argument/__init__.py
Normal file
0
memoryscope/argument/__init__.py
Normal file
|
|
@ -1,19 +1,17 @@
|
|||
global_config:
|
||||
language: cn
|
||||
max_workers: 5
|
||||
|
||||
logger_config:
|
||||
language: en
|
||||
thread_pool_max_workers: 5
|
||||
logger_name: memoryscope
|
||||
logger_suffix: time
|
||||
logger_name_time_suffix: %Y%m%d_%H%M%S
|
||||
|
||||
memory_chat:
|
||||
cli_memory_chat:
|
||||
class: chat.cli_memory_chat
|
||||
memory_service: memory_scope_service
|
||||
generation_model: dashscope_generation
|
||||
memory_service: memoryscope_service
|
||||
generation_model: generation_model
|
||||
|
||||
memory_service:
|
||||
memory_scope_service:
|
||||
memoryscope_service:
|
||||
class: memory.service.memory_scope_service
|
||||
memory_operations:
|
||||
read_message:
|
||||
|
|
@ -46,13 +44,13 @@ memory_service:
|
|||
workflow: add_memory
|
||||
description: "add a single observation"
|
||||
|
||||
summary_observation_memory:
|
||||
class: memory.operation.summary_observation_op
|
||||
consolidate_memory:
|
||||
class: memory.operation.consolidate_memory_op
|
||||
workflow: info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat,store_memory
|
||||
description: "summary user's observation memory"
|
||||
interval_time: 1
|
||||
|
||||
summary_insight_memory:
|
||||
reflect_and_reconsolidate:
|
||||
class: memory.operation.backend_operation
|
||||
workflow: load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat,store_memory
|
||||
description: "summary user's insight memory"
|
||||
|
|
@ -61,9 +59,9 @@ memory_service:
|
|||
worker:
|
||||
dummy:
|
||||
class: memory.worker.dummy_worker
|
||||
generation_model: dashscope_generation
|
||||
embedding_model: dashscope_embedding
|
||||
rank_model: dashscope_rank
|
||||
generation_model: generation_model
|
||||
embedding_model: embedding_model
|
||||
rank_model: rank_model
|
||||
read_message:
|
||||
class: memory.worker.frontend.read_message_worker
|
||||
set_query:
|
||||
|
|
@ -74,12 +72,10 @@ worker:
|
|||
retrieve_ins_top_k: 100
|
||||
extract_time:
|
||||
class: memory.worker.frontend.extract_time_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
generation_model: generation_model
|
||||
semantic_rank:
|
||||
class: memory.worker.frontend.semantic_rank_worker
|
||||
rank_model: dashscope_rank
|
||||
rank_model: rank_model
|
||||
fuse_rerank:
|
||||
class: memory.worker.frontend.fuse_rerank_worker
|
||||
fuse_score_threshold: 0.01
|
||||
|
|
@ -99,9 +95,9 @@ worker:
|
|||
class: memory.worker.frontend.print_memory_worker
|
||||
retrieve_all_memory:
|
||||
class: memory.worker.frontend.retrieve_memory_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_top_k: 100
|
||||
retrieve_expired_top_k: 100
|
||||
retrieve_obs_top_k: 1000
|
||||
retrieve_ins_top_k: 1000
|
||||
retrieve_expired_top_k: 1000
|
||||
delete_memory:
|
||||
class: memory.worker.backend.update_memory_worker
|
||||
method: delete_memory
|
||||
|
|
@ -113,27 +109,19 @@ worker:
|
|||
method: from_query
|
||||
info_filter:
|
||||
class: memory.worker.backend.info_filter_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
generation_model: generation_model
|
||||
load_today_memory:
|
||||
class: memory.worker.backend.load_memory_worker
|
||||
retrieve_today_top_k: 100
|
||||
get_observation:
|
||||
class: memory.worker.backend.get_observation_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
generation_model: generation_model
|
||||
get_observation_with_time:
|
||||
class: memory.worker.backend.get_observation_with_time_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
generation_model: generation_model
|
||||
contra_repeat:
|
||||
class: memory.worker.backend.contra_repeat_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
generation_model: generation_model
|
||||
store_memory:
|
||||
class: memory.worker.backend.update_memory_worker
|
||||
method: from_memory_key
|
||||
|
|
@ -145,37 +133,31 @@ worker:
|
|||
retrieve_insight_top_k: 100
|
||||
get_reflection_subject:
|
||||
class: memory.worker.backend.get_reflection_subject_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model: generation_model
|
||||
reflect_obs_cnt_threshold: 10
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
update_insight:
|
||||
class: memory.worker.backend.update_insight_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
rank_model: dashscope_rank
|
||||
generation_model: generation_model
|
||||
rank_model: rank_model
|
||||
long_contra_repeat:
|
||||
class: memory.worker.backend.long_contra_repeat_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
generation_model: generation_model
|
||||
|
||||
models:
|
||||
dashscope_generation:
|
||||
model:
|
||||
generation_model:
|
||||
class: models.llama_index_generation_model
|
||||
module_name: dashscope_generation
|
||||
model_name: qwen-max
|
||||
max_tokens: 2000
|
||||
dashscope_embedding:
|
||||
embedding_model:
|
||||
class: models.llama_index_embedding_model
|
||||
module_name: dashscope_embedding
|
||||
model_name: text-embedding-v2
|
||||
dashscope_rank:
|
||||
rank_model:
|
||||
class: models.llama_index_rank_model
|
||||
module_name: dashscope_rank
|
||||
model_name: gte-rerank
|
||||
top_n: 10
|
||||
top_n: 500
|
||||
dummy_generation:
|
||||
class: models.dummy_generation_model
|
||||
module_name: dummy_generation
|
||||
|
|
@ -183,10 +165,11 @@ models:
|
|||
|
||||
memory_store:
|
||||
class: storage.llama_index_es_memory_store
|
||||
embedding_model: dashscope_embedding
|
||||
embedding_model: embedding_model
|
||||
index_name: memory_index
|
||||
es_url: http://localhost:9200
|
||||
use_hybrid: true
|
||||
retrieve_type: dense
|
||||
hybrid_alpha: 1.0
|
||||
|
||||
monitor:
|
||||
class: storage.dummy_monitor
|
||||
182
memoryscope/argument/default_arguments.py
Normal file
182
memoryscope/argument/default_arguments.py
Normal file
|
|
@ -0,0 +1,182 @@
|
|||
DEFAULT_GLOBAL_ARGUMENTS = {
|
||||
"language": "en",
|
||||
"thread_pool_max_workers": 5,
|
||||
"logger_name": "memoryscope",
|
||||
"logger_name_time_suffix": "%Y%m%d_%H%M%S"
|
||||
}
|
||||
|
||||
DEFAULT_MEMORY_CHAT_ARGUMENTS = {
|
||||
"cli_memory_chat": {
|
||||
"class": "chat.cli_memory_chat",
|
||||
"memory_service": "memoryscope_service",
|
||||
"generation_model": "generation_model"
|
||||
}
|
||||
}
|
||||
|
||||
DEFAULT_MEMORY_SERVICE_ARGUMENTS = {
|
||||
"memoryscope_service": {
|
||||
"class": "memory.service.memory_scope_service",
|
||||
"memory_operations": {
|
||||
"read_message": {
|
||||
"class": "memory.operation.frontend_operation",
|
||||
"workflow": "read_message",
|
||||
"description": "read short memory"
|
||||
},
|
||||
"retrieve_memory": {
|
||||
"class": "memory.operation.frontend_operation",
|
||||
"workflow": "set_query,[extract_time|retrieve_obs_ins,semantic_rank],fuse_rerank",
|
||||
"description": "retrieve long-term memory"
|
||||
},
|
||||
"list_memory": {
|
||||
"class": "memory.operation.frontend_operation",
|
||||
"workflow": "set_query,retrieve_top_memory,print_memory",
|
||||
"description": "read all long-term memory of the user"
|
||||
},
|
||||
"delete_memory": {
|
||||
"class": "memory.operation.frontend_operation",
|
||||
"workflow": "set_query,retrieve_all_memory,delete_memory",
|
||||
"description": "delete a single long-term memory"
|
||||
},
|
||||
"delete_all": {
|
||||
"class": "memory.operation.frontend_operation",
|
||||
"workflow": "set_query,retrieve_all_memory,delete_all",
|
||||
"description": "delete all long-term memory"
|
||||
},
|
||||
"add_memory": {
|
||||
"class": "memory.operation.frontend_operation",
|
||||
"workflow": "add_memory",
|
||||
"description": "add a single observation"
|
||||
},
|
||||
"consolidate_memory": {
|
||||
"class": "memory.operation.consolidate_memory_op",
|
||||
"workflow": "info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat,"
|
||||
"store_memory",
|
||||
"description": "summary user's observation memory",
|
||||
"interval_time": 1
|
||||
},
|
||||
"reflect_and_reconsolidate": {
|
||||
"class": "memory.operation.backend_operation",
|
||||
"workflow": "load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat,"
|
||||
"store_memory",
|
||||
"description": "summary user's insight memory",
|
||||
"interval_time": 15
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
DEFAULT_WORKER_ARGUMENTS = {
|
||||
"dummy": {
|
||||
"class": "memory.worker.dummy_worker",
|
||||
"generation_model": "generation_model",
|
||||
"embedding_model": "embedding_model",
|
||||
"rank_model": "rank_model"
|
||||
},
|
||||
"read_message": {
|
||||
"class": "memory.worker.frontend.read_message_worker"
|
||||
},
|
||||
"set_query": {
|
||||
"class": "memory.worker.frontend.set_query_worker"
|
||||
},
|
||||
"retrieve_obs_ins": {
|
||||
"class": "memory.worker.frontend.retrieve_memory_worker",
|
||||
"retrieve_obs_top_k": 100,
|
||||
"retrieve_ins_top_k": 100
|
||||
},
|
||||
"extract_time": {
|
||||
"class": "memory.worker.frontend.extract_time_worker",
|
||||
"generation_model": "generation_model"
|
||||
},
|
||||
"semantic_rank": {
|
||||
"class": "memory.worker.frontend.semantic_rank_worker",
|
||||
"rank_model": "rank_model"
|
||||
},
|
||||
"fuse_rerank": {
|
||||
"class": "memory.worker.frontend.fuse_rerank_worker",
|
||||
"fuse_score_threshold": 0.01,
|
||||
"fuse_ratio_dict": {
|
||||
"conversation": 0.5,
|
||||
"observation": 1,
|
||||
"obs_customized": 1.2,
|
||||
"insight": 2
|
||||
},
|
||||
"fuse_time_ratio": 2,
|
||||
"fuse_rerank_top_k": 10
|
||||
},
|
||||
"retrieve_top_memory": {
|
||||
"class": "memory.worker.frontend.retrieve_memory_worker",
|
||||
"retrieve_obs_top_k": 100,
|
||||
"retrieve_ins_top_k": 100,
|
||||
"retrieve_expired_top_k": 100
|
||||
},
|
||||
"print_memory": {
|
||||
"class": "memory.worker.frontend.print_memory_worker"
|
||||
},
|
||||
"retrieve_all_memory": {
|
||||
"class": "memory.worker.frontend.retrieve_memory_worker",
|
||||
"retrieve_obs_top_k": 1000,
|
||||
"retrieve_ins_top_k": 1000,
|
||||
"retrieve_expired_top_k": 1000
|
||||
},
|
||||
"delete_memory": {
|
||||
"class": "memory.worker.backend.update_memory_worker",
|
||||
"method": "delete_memory"
|
||||
},
|
||||
"delete_all": {
|
||||
"class": "memory.worker.backend.update_memory_worker",
|
||||
"method": "delete_all"
|
||||
},
|
||||
"add_memory": {
|
||||
"class": "memory.worker.backend.update_memory_worker",
|
||||
"method": "from_query"
|
||||
},
|
||||
"info_filter": {
|
||||
"class": "memory.worker.backend.info_filter_worker",
|
||||
"generation_model": "generation_model"
|
||||
},
|
||||
"load_today_memory": {
|
||||
"class": "memory.worker.backend.load_memory_worker",
|
||||
"retrieve_today_top_k": 100
|
||||
},
|
||||
"get_observation": {
|
||||
"class": "memory.worker.backend.get_observation_worker",
|
||||
"generation_model": "generation_model"
|
||||
},
|
||||
"get_observation_with_time": {
|
||||
"class": "memory.worker.backend.get_observation_with_time_worker",
|
||||
"generation_model": "generation_model"
|
||||
},
|
||||
"contra_repeat": {
|
||||
"class": "memory.worker.backend.contra_repeat_worker",
|
||||
"generation_model": "generation_model"
|
||||
},
|
||||
"store_memory": {
|
||||
"class": "memory.worker.backend.update_memory_worker",
|
||||
"method": "from_memory_key",
|
||||
"memory_key": "all"
|
||||
},
|
||||
"load_obs_and_insight": {
|
||||
"class": "memory.worker.backend.load_memory_worker",
|
||||
"retrieve_not_reflected_top_k": 100,
|
||||
"retrieve_not_updated_top_k": 100,
|
||||
"retrieve_insight_top_k": 100
|
||||
},
|
||||
"get_reflection_subject": {
|
||||
"class": "memory.worker.backend.get_reflection_subject_worker",
|
||||
"generation_model": "generation_model",
|
||||
"reflect_obs_cnt_threshold": 10
|
||||
},
|
||||
"update_insight": {
|
||||
"class": "memory.worker.backend.update_insight_worker",
|
||||
"generation_model": "generation_model",
|
||||
"rank_model": "rank_model"
|
||||
},
|
||||
"long_contra_repeat": {
|
||||
"class": "memory.worker.backend.long_contra_repeat_worker",
|
||||
"generation_model": "generation_model"
|
||||
}
|
||||
}
|
||||
|
||||
DEFAULT_MONITOR_ARGUMENTS = {
|
||||
"class": "storage.dummy_monitor"
|
||||
}
|
||||
28
memoryscope/argument/init_handler.py
Normal file
28
memoryscope/argument/init_handler.py
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
class InitializationHandler(object):
|
||||
|
||||
def __init__(self):
|
||||
self.file_path: str = __file__
|
||||
|
||||
self.global_config_dict: dict = {}
|
||||
|
||||
self.memory_chat_dict: dict = {}
|
||||
|
||||
self.memory_service_dict: dict = {}
|
||||
|
||||
self.worker_dict: dict = {}
|
||||
|
||||
self.model_dict: dict = {}
|
||||
|
||||
self.memory_store: dict = {}
|
||||
|
||||
self.monitor: dict = {}
|
||||
|
||||
def update_by_arguments(self):
|
||||
pass
|
||||
|
||||
def load_from_config(self):
|
||||
pass
|
||||
|
||||
|
||||
def load_from_file(self):
|
||||
pass
|
||||
61
memoryscope/argument/memoryscope_arguments.py
Normal file
61
memoryscope/argument/memoryscope_arguments.py
Normal file
|
|
@ -0,0 +1,61 @@
|
|||
from dataclasses import dataclass, field
|
||||
from typing import Literal, Dict
|
||||
|
||||
|
||||
@dataclass
|
||||
class MemoryscopeArguments(object):
|
||||
language: Literal["cn", "en"] = field(default="en", metadata={"help": "support en & cn now"})
|
||||
|
||||
thread_pool_max_workers: int = field(default=5, metadata={"help": "thread pool max workers"})
|
||||
|
||||
logger_name: str = field(default="memoryscope")
|
||||
|
||||
logger_name_time_suffix: str = field(default="%Y%m%d_%H%M%S")
|
||||
|
||||
memory_chat_class: str = field(default="chat.api_memory_chat", metadata={
|
||||
"help": "The memory chat class for dynamic import: chat.cli_memory_chat, chat.api_memory_chat"})
|
||||
|
||||
human_name: str = field(default="user", metadata={"help": "en: user, cn: 用户"})
|
||||
|
||||
assistant_name: str = field(default="AI")
|
||||
|
||||
consolidate_memory_interval_time: int = field(default=1, metadata={
|
||||
"help": "If you feel that the token consumption is relatively high, please increase the time interval."})
|
||||
|
||||
reflect_and_reconsolidate_interval_time: int = field(default=15, metadata={
|
||||
"help": "If you feel that the token consumption is relatively high, please increase the time interval."})
|
||||
|
||||
worker_params: Dict[str, dict] = field(default_factory=lambda: {}, metadata={
|
||||
"help": "dict format: worker_name -> param_key -> param_value"})
|
||||
|
||||
generation_backend: str = field(default="openai_generation", metadata={
|
||||
"help": "global generation backend: openai_generation, dashscope_generation, etc."})
|
||||
|
||||
generation_model: str = field(default="gpt-4o", metadata={
|
||||
"help": "global generation model: gpt-4o, gpt-4, qwen-max, etc."})
|
||||
|
||||
generation_params: dict = field(default_factory=lambda: {}, metadata={
|
||||
"help": "global generation params: max_tokens, top_p, temperature, etc."})
|
||||
|
||||
embedding_backend: str = field(default="openai_embedding", metadata={
|
||||
"help": "global embedding backend: openai_embedding, dashscope_embedding, etc."})
|
||||
|
||||
embedding_model: str = field(default="gpt-4o", metadata={
|
||||
"help": "global embedding model: text-embedding-ada-002, text-embedding-v2, etc."})
|
||||
|
||||
embedding_params: dict = field(default_factory=lambda: {})
|
||||
|
||||
rank_backend: str = field(default="dashscope_rank", metadata={"help": "global rank backend: dashscope_rank, etc."})
|
||||
|
||||
rank_model: str = field(default="gte-rerank", metadata={"help": "global rank model: gte-rerank, etc."})
|
||||
|
||||
rank_params: dict = field(default_factory=lambda: {})
|
||||
|
||||
es_index_name: str = field(default="memory_index")
|
||||
|
||||
es_url: str = field(default="http://localhost:9200")
|
||||
|
||||
# TODO at xianzhe
|
||||
retrieve_type: str = field(default="dense", metadata={"help": "es_retrieve_type: dense, sparse, hybrid"})
|
||||
|
||||
hybrid_alpha: float | None = field(default=1.0, metadata={"help": ""})
|
||||
322
memoryscope/chat/api_memory_chat.py
Normal file
322
memoryscope/chat/api_memory_chat.py
Normal file
|
|
@ -0,0 +1,322 @@
|
|||
import os
|
||||
import time
|
||||
from typing import List
|
||||
|
||||
import questionary
|
||||
|
||||
from memoryscope.chat.base_memory_chat import BaseMemoryChat
|
||||
from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME
|
||||
from memoryscope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memoryscope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.models.base_model import BaseModel
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.utils.prompt_handler import PromptHandler
|
||||
from memoryscope.utils.tool_functions import char_logo
|
||||
|
||||
|
||||
class ApiMemoryChat(BaseMemoryChat):
|
||||
|
||||
def __init__(self,
|
||||
memory_service: str,
|
||||
generation_model: str,
|
||||
stream: bool = True,
|
||||
human_name: str = DEFAULT_HUMAN_NAME[G_CONTEXT.language],
|
||||
assistant_name: str = "AI",
|
||||
**kwargs):
|
||||
|
||||
self._memory_service: BaseMemoryService | str = memory_service
|
||||
self._generation_model: BaseModel | str = generation_model
|
||||
self.generation_model_kwargs: dict = kwargs.pop("generation_model_kwargs", {})
|
||||
|
||||
self.stream: bool = stream
|
||||
self.human_name: str = human_name
|
||||
self.assistant_name: str = assistant_name
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
self._logo = char_logo("MemoryScope")
|
||||
self._prompt_handler: PromptHandler | None = None
|
||||
G_CONTEXT.meta_data.update({
|
||||
"human_name": human_name,
|
||||
"assistant_name": assistant_name,
|
||||
})
|
||||
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
@property
|
||||
def prompt_handler(self) -> PromptHandler:
|
||||
"""
|
||||
Lazy initialization property for the prompt handler.
|
||||
|
||||
This property ensures that the `_prompt_handler` attribute is only instantiated when it is first accessed.
|
||||
It uses the current file's path and additional keyword arguments for configuration.
|
||||
|
||||
Returns:
|
||||
PromptHandler: An instance of the PromptHandler configured for this CLI session.
|
||||
"""
|
||||
if self._prompt_handler is None:
|
||||
self._prompt_handler = PromptHandler(__file__, **self.kwargs)
|
||||
return self._prompt_handler
|
||||
|
||||
def print_logo(self):
|
||||
"""
|
||||
Prints the logo of the CLI application to the console.
|
||||
|
||||
The logo is composed of multiple lines, which are iterated through
|
||||
and printed one by one to provide a visual identity for the chat interface.
|
||||
"""
|
||||
for line in self._logo:
|
||||
print(line)
|
||||
|
||||
@property
|
||||
def memory_service(self) -> BaseMemoryService:
|
||||
"""
|
||||
Property to access the memory service. If the service is initially set as a string,
|
||||
it will be looked up in the memory service dictionary of global context, initialized,
|
||||
and then returned as an instance of `BaseMemoryService`. Ensures the memory service
|
||||
is properly started before use.
|
||||
|
||||
Returns:
|
||||
BaseMemoryService: An active memory service instance.
|
||||
|
||||
Raises:
|
||||
ValueError: If the declaration of memory service is not found in the memory service dictionary of global context.
|
||||
"""
|
||||
if isinstance(self._memory_service, str):
|
||||
if self._memory_service not in G_CONTEXT.memory_service_conf_dict:
|
||||
raise ValueError("Missing declaration of memory_service in yaml configuration: " + self._memory_service)
|
||||
self._memory_service = G_CONTEXT.memory_service_conf_dict[self._memory_service]
|
||||
self._memory_service.init_service()
|
||||
self._memory_service.start_backend_service()
|
||||
return self._memory_service
|
||||
|
||||
@property
|
||||
def generation_model(self) -> BaseModel:
|
||||
"""
|
||||
Property to get the generation model. If the model is set as a string, it will be resolved from the global
|
||||
context's model dictionary.
|
||||
|
||||
Raises:
|
||||
ValueError: If the declaration of generation model is not found in the model dictionary of global context .
|
||||
|
||||
Returns:
|
||||
BaseModel: An actual generation model instance.
|
||||
"""
|
||||
if isinstance(self._generation_model, str):
|
||||
if self._generation_model not in G_CONTEXT.model_conf_dict:
|
||||
raise ValueError(f"Missing declaration of generation model in yaml config: {self._generation_model}")
|
||||
self._generation_model = G_CONTEXT.model_conf_dict[self._generation_model]
|
||||
return self._generation_model
|
||||
|
||||
def chat_with_memory(self, query: str, remember_response: bool = False) -> ModelResponse | ModelResponseGen:
|
||||
"""
|
||||
Engages in a conversation with the AI model, utilizing conversation memory.
|
||||
The function sends the user's query, incorporates conversation history and memory,
|
||||
and optionally remembers the AI's response based on the user's preference.
|
||||
|
||||
Args:
|
||||
query (str): The user's input or query for the AI.
|
||||
remember_response (bool, optional): Flag indicating whether to save the AI's response to memory.
|
||||
Defaults to False.
|
||||
|
||||
Returns:
|
||||
- ModelResponse: In non-streaming mode, returns a complete AI response.
|
||||
- ModelResponseGen: In streaming mode, returns a generator yielding AI response parts.
|
||||
|
||||
Side Effects:
|
||||
- Updates the conversation memory with the query of user and (optionally) the response of AI.
|
||||
- Retrieves and includes historical messages and memory content in the context of conversation.
|
||||
"""
|
||||
new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=self.human_name, content=query)
|
||||
self.memory_service.add_messages(new_message)
|
||||
|
||||
messages: List[Message] = []
|
||||
|
||||
# Incorporate memory into the system prompt if available
|
||||
system_prompt = self.prompt_handler.system_prompt
|
||||
memories: str = self.memory_service.retrieve_memory()
|
||||
if memories:
|
||||
memory_prompt = self.prompt_handler.memory_prompt
|
||||
system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt, memories]])
|
||||
messages.append(Message(role=MessageRoleEnum.SYSTEM, content=system_prompt))
|
||||
|
||||
# Include past conversation history in the message list
|
||||
history_messages = self.memory_service.read_message()
|
||||
if history_messages:
|
||||
messages.extend(history_messages)
|
||||
|
||||
# Append the current user's message to the conversation context
|
||||
messages.append(new_message)
|
||||
self.logger.info(f"messages={messages}")
|
||||
|
||||
# Invoke the Language Model with the constructed message context, respecting streaming setting
|
||||
generated = self.generation_model.call(messages=messages, stream=self.stream, **self.generation_model_kwargs)
|
||||
|
||||
# In non-streaming interactions, explicitly save the AI's reply to memory if instructed
|
||||
if remember_response:
|
||||
assert not self.stream # Ensure we're not in streaming mode when remembering responses
|
||||
generated.message.role_name = self.assistant_name
|
||||
self.memory_service.add_messages(generated.message)
|
||||
|
||||
# Return the AI's response directly or as a generator based on the streaming mode
|
||||
return generated
|
||||
|
||||
@staticmethod
|
||||
def parse_query_command(query: str):
|
||||
"""
|
||||
Parses the user's input query command, separating it into the command and its associated keyword arguments.
|
||||
|
||||
Args:
|
||||
query (str): The raw input string from the user which includes the command and its arguments.
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing the command (str) as the first element and a dictionary (kwargs) of keyword
|
||||
arguments as the second element.
|
||||
"""
|
||||
query_split = query.lstrip("/").lower().split(" ") # Split and preprocess the input command
|
||||
command = query_split[0] # Extract the command
|
||||
args = query_split[1:] # Extract the arguments following the command
|
||||
kwargs = {} # Initialize dictionary to hold keyword arguments
|
||||
|
||||
for arg in args:
|
||||
# Skip if no arguments exist (unnecessary check due to prior assignment, but retained as per original)
|
||||
if not args:
|
||||
continue
|
||||
arg_split = arg.split("=") # Split argument into key-value pair
|
||||
if len(arg_split) >= 2: # Ensure there's both a key and value
|
||||
k = arg_split[0] # Extract key
|
||||
v = arg_split[1] # Extract value
|
||||
if k and v: # Only add to kwargs if both key and value are non-empty
|
||||
kwargs[k] = v
|
||||
|
||||
return command, kwargs # Return the parsed command and keyword arguments
|
||||
|
||||
def process_commands(self, query: str) -> bool:
|
||||
"""
|
||||
Parses and executes commands from user input in the CLI chat interface.
|
||||
Supports operations like exiting, clearing screen, showing help, toggling stream mode,
|
||||
executing predefined memory operations, and handling unknown commands.
|
||||
|
||||
Args:
|
||||
query (str): The user's input command string.
|
||||
|
||||
Returns:
|
||||
bool: Indicates whether to continue running the CLI after processing the command.
|
||||
"""
|
||||
continue_run = True
|
||||
command, kwargs = self.parse_query_command(query)
|
||||
|
||||
# Print prompt for AI's response
|
||||
questionary.print("> ", end="", style="fg:yellow")
|
||||
questionary.print(f"{self.assistant_name}: ", end="", style="bold")
|
||||
|
||||
if command == "exit":
|
||||
self.memory_service.stop_backend_service()
|
||||
continue_run = False
|
||||
|
||||
elif command == "clear":
|
||||
os.system("clear")
|
||||
|
||||
elif command == "help":
|
||||
questionary.print("CLI commands", "bold")
|
||||
for cmd, desc in self.USER_COMMANDS.items():
|
||||
questionary.print(text=f" /{cmd}:", style="bold")
|
||||
questionary.print(text=f" {desc}")
|
||||
|
||||
elif command == "stream":
|
||||
self.stream = not self.stream
|
||||
questionary.print(f"set stream: {self.stream}")
|
||||
|
||||
elif command in self.memory_service.op_description_dict:
|
||||
refresh_time = kwargs.pop("refresh_time", "")
|
||||
if refresh_time and refresh_time.isdigit():
|
||||
refresh_time = int(refresh_time)
|
||||
self.memory_service.stop_backend_service()
|
||||
while True:
|
||||
result = self.memory_service.do_operation(op_name=command, **kwargs)
|
||||
os.system("clear")
|
||||
self.print_logo()
|
||||
if result:
|
||||
if isinstance(result, list):
|
||||
result = "\n".join([str(x) for x in result])
|
||||
questionary.print(result)
|
||||
else:
|
||||
questionary.print(f"command={command} result is empty! kwargs={kwargs}")
|
||||
time.sleep(refresh_time)
|
||||
|
||||
else:
|
||||
result = self.memory_service.do_operation(op_name=command, **kwargs)
|
||||
if result:
|
||||
if isinstance(result, list):
|
||||
result = "\n".join([str(x) for x in result])
|
||||
questionary.print(result)
|
||||
else:
|
||||
questionary.print(f"command={command} result is empty! kwargs={kwargs}")
|
||||
|
||||
else:
|
||||
questionary.print(f"Unknown command={command} received.")
|
||||
|
||||
return continue_run
|
||||
|
||||
def run(self):
|
||||
"""
|
||||
Runs the CLI chat loop, which handles user input, processes commands,
|
||||
communicates with the AI model, manages conversation memory, and controls
|
||||
the chat session including streaming responses, command execution, and error handling.
|
||||
|
||||
The loop continues until the user explicitly chooses to exit.
|
||||
"""
|
||||
self.print_logo()
|
||||
self.USER_COMMANDS.update(self.memory_service.op_description_dict)
|
||||
|
||||
while True:
|
||||
try:
|
||||
query = questionary.text(message=f"{self.human_name}:", multiline=False, qmark=">").unsafe_ask()
|
||||
if not query:
|
||||
continue
|
||||
|
||||
query: str = query.strip()
|
||||
|
||||
# Handle special commands prefixed with '/'
|
||||
if query.startswith("/"):
|
||||
if self.process_commands(query=query):
|
||||
continue
|
||||
else:
|
||||
break
|
||||
|
||||
# Print prompt for AI's response
|
||||
questionary.print("> ", end="", style="fg:yellow")
|
||||
questionary.print(f"{self.assistant_name}: ", end="", style="bold")
|
||||
|
||||
# Fetch and display AI's response, with support for streaming
|
||||
self.memory_service.start_backend_service()
|
||||
if self.stream:
|
||||
model_response = None
|
||||
for model_response in self.chat_with_memory(query=query):
|
||||
questionary.print(model_response.delta, end="")
|
||||
questionary.print("")
|
||||
|
||||
else:
|
||||
model_response = self.chat_with_memory(query=query)
|
||||
questionary.print(model_response.message.content)
|
||||
|
||||
# Append AI's response to the conversation memory
|
||||
model_response.message.role_name = self.assistant_name
|
||||
self.memory_service.add_messages(model_response.message)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
# Handle user interruption and confirm exit
|
||||
questionary.print("User interrupt occurred.")
|
||||
is_exit = questionary.confirm("Continue exit?").unsafe_ask()
|
||||
if is_exit:
|
||||
self.memory_service.stop_backend_service()
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
# Log and handle any unanticipated exceptions
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
self.logger.exception(f"An exception occurred when running cli memory chat. args={e.args}.")
|
||||
continue
|
||||
|
|
@ -1,6 +1,9 @@
|
|||
from abc import ABCMeta, abstractmethod
|
||||
from typing import List
|
||||
|
||||
from memoryscope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.utils.logger import Logger
|
||||
|
||||
|
||||
class BaseMemoryChat(metaclass=ABCMeta):
|
||||
|
|
@ -9,13 +12,19 @@ class BaseMemoryChat(metaclass=ABCMeta):
|
|||
It outlines the method to initiate a chat session leveraging memory data, which concrete subclasses must implement.
|
||||
"""
|
||||
|
||||
def __init__(self, generation_stream: bool = True, **kwargs):
|
||||
self.generation_stream: bool = generation_stream
|
||||
self.kwargs: dict = kwargs
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
@abstractmethod
|
||||
def chat_with_memory(self, query: str):
|
||||
def chat_with_memory(self, query: str, role_name: str = ""):
|
||||
"""
|
||||
Initiates a chat interaction using the memory service, with the provided query as input.
|
||||
|
||||
Args:
|
||||
query (str): The user's query or message to start the chat.
|
||||
role_name (str): The role's name.
|
||||
|
||||
Returns:
|
||||
This method should return the chat response generated after processing the query
|
||||
|
|
@ -23,6 +32,9 @@ class BaseMemoryChat(metaclass=ABCMeta):
|
|||
subclass.
|
||||
"""
|
||||
|
||||
def add_message(self, messages: List[Message] | Message):
|
||||
self.memory_service.add_messages(messages)
|
||||
|
||||
@property
|
||||
def memory_service(self) -> BaseMemoryService:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -8,11 +8,10 @@ from memoryscope.chat.base_memory_chat import BaseMemoryChat
|
|||
from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME
|
||||
from memoryscope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memoryscope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.models.base_model import BaseModel
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.utils.prompt_handler import PromptHandler
|
||||
from memoryscope.utils.tool_functions import char_logo
|
||||
|
||||
|
|
@ -26,48 +25,33 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
"exit": "Exit the CLI.",
|
||||
"clear": "Clear the command history.",
|
||||
"help": "Display available CLI commands and their descriptions.",
|
||||
"stream": "Toggle between getting streamed responses from the model."
|
||||
}
|
||||
|
||||
def __init__(self,
|
||||
memory_service: str,
|
||||
generation_model: str,
|
||||
stream: bool = True,
|
||||
human_name: str = DEFAULT_HUMAN_NAME[G_CONTEXT.language],
|
||||
assistant_name: str = "AI",
|
||||
context: MemoryscopeContext,
|
||||
human_name: str = None,
|
||||
assistant_name: str = None,
|
||||
**kwargs):
|
||||
"""
|
||||
Initializes the CLI chat instance with specified services, models, and personalized settings.
|
||||
|
||||
Args:
|
||||
memory_service (str | BaseMemoryService): The memory service to be used for storing conversation history.
|
||||
generation_model (str | BaseModel): The model responsible for generating AI responses.
|
||||
stream (bool, optional): Flag indicating whether responses should be streamed. Defaults to True.
|
||||
human_name (str, optional): The name assigned to the human user. Defaults to a language-specific user.
|
||||
assistant_name (str, optional): The name of the AI assistant. Defaults to "AI".
|
||||
**kwargs: Additional keyword arguments for flexibility or future extensions.
|
||||
super().__init__(**kwargs)
|
||||
|
||||
Side Effects:
|
||||
- Updates global context with human and AI names.
|
||||
- Initializes logging for the instance.
|
||||
"""
|
||||
self._memory_service: BaseMemoryService | str = memory_service
|
||||
self._generation_model: BaseModel | str = generation_model
|
||||
self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {})
|
||||
self.context: MemoryscopeContext = context
|
||||
self.generation_model_kwargs: dict = kwargs.pop("generation_model_kwargs", {})
|
||||
|
||||
self.stream: bool = stream
|
||||
self.human_name: str = human_name
|
||||
if not self.human_name:
|
||||
self.human_name = DEFAULT_HUMAN_NAME[self.context.language]
|
||||
|
||||
self.assistant_name: str = assistant_name
|
||||
self.kwargs: dict = kwargs
|
||||
if not self.assistant_name:
|
||||
self.assistant_name = "AI"
|
||||
|
||||
self._logo = char_logo("MemoryScope")
|
||||
self._prompt_handler: PromptHandler | None = None
|
||||
G_CONTEXT.meta_data.update({
|
||||
"human_name": human_name,
|
||||
"assistant_name": assistant_name,
|
||||
})
|
||||
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
@property
|
||||
def prompt_handler(self) -> PromptHandler:
|
||||
|
|
@ -81,7 +65,7 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
PromptHandler: An instance of the PromptHandler configured for this CLI session.
|
||||
"""
|
||||
if self._prompt_handler is None:
|
||||
self._prompt_handler = PromptHandler(__file__, **self.kwargs)
|
||||
self._prompt_handler = PromptHandler(__file__, prompt_file="memory_chat_prompt", **self.kwargs)
|
||||
return self._prompt_handler
|
||||
|
||||
def print_logo(self):
|
||||
|
|
@ -98,7 +82,7 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
def memory_service(self) -> BaseMemoryService:
|
||||
"""
|
||||
Property to access the memory service. If the service is initially set as a string,
|
||||
it will be looked up in the memory service dictionary of global context, initialized,
|
||||
it will be looked up in the memory service dictionary of context, initialized,
|
||||
and then returned as an instance of `BaseMemoryService`. Ensures the memory service
|
||||
is properly started before use.
|
||||
|
||||
|
|
@ -106,13 +90,15 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
BaseMemoryService: An active memory service instance.
|
||||
|
||||
Raises:
|
||||
ValueError: If the declaration of memory service is not found in the memory service dictionary of global context.
|
||||
ValueError: If the declaration of memory service is not found in the memory service dictionary of context.
|
||||
"""
|
||||
if isinstance(self._memory_service, str):
|
||||
if self._memory_service not in G_CONTEXT.memory_service_dict:
|
||||
raise ValueError("Missing declaration of memory_service in yaml configuration: " + self._memory_service)
|
||||
self._memory_service = G_CONTEXT.memory_service_dict[self._memory_service]
|
||||
self._memory_service.init_service()
|
||||
if self._memory_service not in self.context.memory_service_dict:
|
||||
raise ValueError(f"Missing declaration of memory_service in context: {self._memory_service}")
|
||||
|
||||
self._memory_service: BaseMemoryService = self.context.memory_service_dict[self._memory_service]
|
||||
# init service & update kwargs
|
||||
self._memory_service.init_service(human_name=self.human_name, assistant_name=self.assistant_name)
|
||||
self._memory_service.start_backend_service()
|
||||
return self._memory_service
|
||||
|
||||
|
|
@ -123,37 +109,21 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
context's model dictionary.
|
||||
|
||||
Raises:
|
||||
ValueError: If the declaration of generation model is not found in the model dictionary of global context .
|
||||
ValueError: If the declaration of generation model is not found in the model dictionary of context .
|
||||
|
||||
Returns:
|
||||
BaseModel: An actual generation model instance.
|
||||
"""
|
||||
if isinstance(self._generation_model, str):
|
||||
if self._generation_model not in G_CONTEXT.model_dict:
|
||||
if self._generation_model not in self.context.model_dict:
|
||||
raise ValueError(f"Missing declaration of generation model in yaml config: {self._generation_model}")
|
||||
self._generation_model = G_CONTEXT.model_dict[self._generation_model]
|
||||
self._generation_model = self.context.model_dict[self._generation_model]
|
||||
return self._generation_model
|
||||
|
||||
def chat_with_memory(self, query: str, remember_response: bool = False) -> ModelResponse | ModelResponseGen:
|
||||
"""
|
||||
Engages in a conversation with the AI model, utilizing conversation memory.
|
||||
The function sends the user's query, incorporates conversation history and memory,
|
||||
and optionally remembers the AI's response based on the user's preference.
|
||||
|
||||
Args:
|
||||
query (str): The user's input or query for the AI.
|
||||
remember_response (bool, optional): Flag indicating whether to save the AI's response to memory.
|
||||
Defaults to False.
|
||||
|
||||
Returns:
|
||||
- ModelResponse: In non-streaming mode, returns a complete AI response.
|
||||
- ModelResponseGen: In streaming mode, returns a generator yielding AI response parts.
|
||||
|
||||
Side Effects:
|
||||
- Updates the conversation memory with the query of user and (optionally) the response of AI.
|
||||
- Retrieves and includes historical messages and memory content in the context of conversation.
|
||||
"""
|
||||
new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=self.human_name, content=query)
|
||||
def chat_with_memory(self, query: str, role_name: str = "") -> ModelResponse | ModelResponseGen:
|
||||
if not role_name:
|
||||
role_name = self.human_name
|
||||
new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query)
|
||||
self.memory_service.add_messages(new_message)
|
||||
|
||||
messages: List[Message] = []
|
||||
|
|
@ -176,16 +146,9 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
self.logger.info(f"messages={messages}")
|
||||
|
||||
# Invoke the Language Model with the constructed message context, respecting streaming setting
|
||||
generated = self.generation_model.call(messages=messages, stream=self.stream, **self.generation_model_kwargs)
|
||||
|
||||
# In non-streaming interactions, explicitly save the AI's reply to memory if instructed
|
||||
if remember_response:
|
||||
assert not self.stream # Ensure we're not in streaming mode when remembering responses
|
||||
generated.message.role_name = self.assistant_name
|
||||
self.memory_service.add_messages(generated.message)
|
||||
|
||||
# Return the AI's response directly or as a generator based on the streaming mode
|
||||
return generated
|
||||
return self.generation_model.call(messages=messages,
|
||||
stream=self.generation_stream,
|
||||
**self.generation_model_kwargs)
|
||||
|
||||
@staticmethod
|
||||
def parse_query_command(query: str):
|
||||
|
|
@ -249,10 +212,6 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
questionary.print(text=f" /{cmd}:", style="bold")
|
||||
questionary.print(text=f" {desc}")
|
||||
|
||||
elif command == "stream":
|
||||
self.stream = not self.stream
|
||||
questionary.print(f"set stream: {self.stream}")
|
||||
|
||||
elif command in self.memory_service.op_description_dict:
|
||||
refresh_time = kwargs.pop("refresh_time", "")
|
||||
if refresh_time and refresh_time.isdigit():
|
||||
|
|
@ -314,14 +273,13 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
questionary.print("> ", end="", style="fg:yellow")
|
||||
questionary.print(f"{self.assistant_name}: ", end="", style="bold")
|
||||
|
||||
# Fetch and display AI's response, with support for streaming
|
||||
# Fetch and display AI's response
|
||||
self.memory_service.start_backend_service()
|
||||
if self.stream:
|
||||
if self.generation_stream:
|
||||
model_response = None
|
||||
for model_response in self.chat_with_memory(query=query):
|
||||
questionary.print(model_response.delta, end="")
|
||||
questionary.print("")
|
||||
|
||||
else:
|
||||
model_response = self.chat_with_memory(query=query)
|
||||
questionary.print(model_response.message.content)
|
||||
|
|
|
|||
|
|
@ -1,105 +1,18 @@
|
|||
import datetime
|
||||
import sys
|
||||
|
||||
import questionary
|
||||
from memoryscope.chat.base_memory_chat import BaseMemoryChat
|
||||
from memoryscope.memoryscope import MemoryScope
|
||||
|
||||
sys.path.append(".") # noqa: E402
|
||||
|
||||
import json
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, Any
|
||||
|
||||
import fire
|
||||
import yaml
|
||||
import atexit
|
||||
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
from memoryscope.enumeration.model_enum import ModelEnum
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.utils.timer import timer
|
||||
from memoryscope.utils.tool_functions import init_instance_by_config, camelcase_to_underscore
|
||||
|
||||
|
||||
class MemoryScope(object):
|
||||
|
||||
def __init__(self):
|
||||
self.config: Dict[str, Any] = {}
|
||||
datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S')
|
||||
class_name = camelcase_to_underscore(self.__class__.__name__)
|
||||
self.logger: Logger = Logger.get_logger(f"{class_name}_{datetime_suffix}", to_stream=False)
|
||||
|
||||
def load_config(self, path: str):
|
||||
with open(path) as f:
|
||||
if path.endswith("yaml"):
|
||||
self.config = yaml.load(f, yaml.FullLoader)
|
||||
elif path.endswith("json"):
|
||||
self.config = json.load(f)
|
||||
else:
|
||||
raise RuntimeError("not supported config file type!")
|
||||
self.init_global_content_by_config()
|
||||
atexit.register(self.shutdown) # register clean up function
|
||||
return self
|
||||
|
||||
@staticmethod
|
||||
def shutdown():
|
||||
questionary.print('Gracefully executing the shutdown function...')
|
||||
G_CONTEXT.memory_store.close()
|
||||
G_CONTEXT.monitor.close()
|
||||
G_CONTEXT.thread_pool.shutdown()
|
||||
|
||||
def set_global_config(self):
|
||||
G_CONTEXT.global_config = global_config = self.config["global_config"]
|
||||
G_CONTEXT.language = LanguageEnum(global_config["language"])
|
||||
G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(global_config["max_workers"]))
|
||||
|
||||
@timer
|
||||
def init_global_content_by_config(self):
|
||||
# set global config
|
||||
self.set_global_config()
|
||||
|
||||
# init memory_chat
|
||||
for name, conf in self.config["memory_chat"].items():
|
||||
G_CONTEXT.memory_chat_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# set memory_service
|
||||
for name, conf in self.config["memory_service"].items():
|
||||
G_CONTEXT.memory_service_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# init models
|
||||
for name, conf in self.config["models"].items():
|
||||
G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# init vector_store
|
||||
if "memory_store" not in self.config:
|
||||
raise RuntimeError("memory_store config is required!")
|
||||
memory_store_config = self.config["memory_store"]
|
||||
embedding_model = G_CONTEXT.model_dict[memory_store_config[ModelEnum.EMBEDDING_MODEL.value]]
|
||||
G_CONTEXT.memory_store = init_instance_by_config(memory_store_config, embedding_model=embedding_model)
|
||||
|
||||
# init monitor
|
||||
G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
|
||||
|
||||
# set worker config
|
||||
G_CONTEXT.worker_config = self.config["worker"]
|
||||
|
||||
@property
|
||||
def default_chat_handle(self):
|
||||
return list(G_CONTEXT.memory_chat_dict.values())[0]
|
||||
|
||||
@property
|
||||
def default_service(self):
|
||||
return self.default_chat_handle.memory_service
|
||||
|
||||
|
||||
class CliJob(MemoryScope):
|
||||
|
||||
def run(self, config: str):
|
||||
self.load_config(config)
|
||||
self.init_global_content_by_config()
|
||||
self.default_chat_handle.run()
|
||||
def cli_job(config_path: str):
|
||||
ms = MemoryScope(config_path=config_path)
|
||||
memory_chat: BaseMemoryChat = ms.default_memory_chat
|
||||
memory_chat.run()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
cli_job = CliJob()
|
||||
fire.Fire(cli_job.run)
|
||||
fire.Fire(cli_job)
|
||||
|
|
|
|||
|
|
@ -1,46 +1,40 @@
|
|||
import threading
|
||||
from abc import ABCMeta, abstractmethod
|
||||
from typing import List, Dict
|
||||
|
||||
from memoryscope.memory.operation.base_operation import BaseOperation
|
||||
from memoryscope.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.utils.logger import Logger
|
||||
|
||||
|
||||
class BaseMemoryService(metaclass=ABCMeta):
|
||||
"""
|
||||
An abstract base class for managing memory operations within a multi-threaded context.
|
||||
An abstract base class for managing memory operations within a multithreaded context.
|
||||
It sets up the infrastructure for operation handling, message storage, and synchronization,
|
||||
along with logging capabilities and customizable configurations.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
memory_operations: Dict[str, dict],
|
||||
retrieve_memory_key: str = "retrieve_memory",
|
||||
read_message_key: str = "read_message",
|
||||
**kwargs):
|
||||
def __init__(self, memory_operations: Dict[str, dict], context: MemoryscopeContext, **kwargs):
|
||||
"""
|
||||
Initializes the BaseMemoryService with operation definitions, keys for memory access,
|
||||
and additional keyword arguments for flexibility.
|
||||
|
||||
Args:
|
||||
memory_operations (Dict[str, dict]): A dictionary defining available memory operations.
|
||||
retrieve_memory_key (str): The key indicating a retrieve memory operation. Defaults to "retrieve_memory".
|
||||
read_message_key (str): The key for reading messages. Defaults to "read_message".
|
||||
**kwargs: Additional parameters to customize service behavior.
|
||||
"""
|
||||
self.memory_operations: Dict[str, dict] = memory_operations
|
||||
self.retrieve_memory_key: str = retrieve_memory_key
|
||||
self.read_message_key: str = read_message_key
|
||||
self.memory_operations_conf: Dict[str, dict] = memory_operations
|
||||
self.context: MemoryscopeContext = context
|
||||
|
||||
self._operation_dict: Dict[str, BaseOperation] = {}
|
||||
self._op_description_dict: Dict[str, str] = {}
|
||||
self.chat_messages: List[Message] = []
|
||||
self.message_lock = threading.Lock()
|
||||
|
||||
self.logger = Logger.get_logger()
|
||||
self.kwargs = kwargs
|
||||
|
||||
def update_kwargs(self, **kwargs):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def add_messages(self, messages: List[Message] | Message):
|
||||
raise NotImplementedError
|
||||
|
|
|
|||
|
|
@ -28,6 +28,9 @@ class MemoryScopeService(BaseMemoryService):
|
|||
self.contextual_msg_min_count: int = contextual_msg_min_count
|
||||
assert history_msg_count >= contextual_msg_max_count >= contextual_msg_min_count
|
||||
|
||||
self.chat_messages: List[Message] = []
|
||||
self.message_lock = threading.Lock()
|
||||
|
||||
def add_messages(self, messages: List[Message] | Message):
|
||||
"""
|
||||
Adds a single message or a list of messages to the chat history, ensuring the message list
|
||||
|
|
|
|||
|
|
@ -87,7 +87,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
BaseModel: The embedding model used for converting text into vector representations.
|
||||
"""
|
||||
if isinstance(self._embedding_model, str):
|
||||
self._embedding_model = G_CONTEXT.model_dict[self._embedding_model]
|
||||
self._embedding_model = G_CONTEXT.model_conf_dict[self._embedding_model]
|
||||
# ⭐ Retrieve the actual model instance when the attribute is a string reference
|
||||
return self._embedding_model
|
||||
|
||||
|
|
@ -101,7 +101,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
BaseModel: The model used for text generation.
|
||||
"""
|
||||
if isinstance(self._generation_model, str):
|
||||
self._generation_model = G_CONTEXT.model_dict[self._generation_model]
|
||||
self._generation_model = G_CONTEXT.model_conf_dict[self._generation_model]
|
||||
# ⭐ Retrieve the model instance if currently a string reference
|
||||
return self._generation_model
|
||||
|
||||
|
|
@ -115,7 +115,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
BaseModel: The rank model instance used for ranking tasks.
|
||||
"""
|
||||
if isinstance(self._rank_model, str):
|
||||
self._rank_model = G_CONTEXT.model_dict[self._rank_model] # Fetch model instance if string reference
|
||||
self._rank_model = G_CONTEXT.model_conf_dict[self._rank_model] # Fetch model instance if string reference
|
||||
return self._rank_model
|
||||
|
||||
@property
|
||||
|
|
@ -128,7 +128,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
BaseMemoryStore: The memory store instance used for inserting, updating, retrieving and deleting operations.
|
||||
"""
|
||||
if self._memory_store is None:
|
||||
self._memory_store = G_CONTEXT.memory_store
|
||||
self._memory_store = G_CONTEXT.memory_store_conf
|
||||
return self._memory_store
|
||||
|
||||
@property
|
||||
|
|
@ -141,7 +141,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
BaseMonitor: The monitoring component instance.
|
||||
"""
|
||||
if self._monitor is None:
|
||||
self._monitor = G_CONTEXT.monitor
|
||||
self._monitor = G_CONTEXT.monitor_conf
|
||||
return self._monitor
|
||||
|
||||
@property
|
||||
|
|
|
|||
198
memoryscope/memoryscope.py
Normal file
198
memoryscope/memoryscope.py
Normal file
|
|
@ -0,0 +1,198 @@
|
|||
import datetime
|
||||
import json
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import yaml
|
||||
|
||||
from memoryscope.argument import default_arguments
|
||||
from memoryscope.argument.memoryscope_arguments import MemoryscopeArguments
|
||||
from memoryscope.chat.base_memory_chat import BaseMemoryChat
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
from memoryscope.enumeration.model_enum import ModelEnum
|
||||
from memoryscope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
||||
class MemoryScope(object):
|
||||
|
||||
def __init__(self,
|
||||
arguments: MemoryscopeArguments | None = None,
|
||||
config: dict | None = None,
|
||||
config_path: str = ""):
|
||||
|
||||
self.global_conf: dict = {}
|
||||
self.memory_chat_conf_dict: dict = {}
|
||||
self.memory_service_conf_dict: dict = {}
|
||||
self.worker_conf_dict: dict = {}
|
||||
self.model_conf_dict: dict = {}
|
||||
self.memory_store_conf: dict = {}
|
||||
self.monitor_conf: dict = {}
|
||||
|
||||
self.context: MemoryscopeContext = MemoryscopeContext()
|
||||
|
||||
if arguments:
|
||||
self._init_by_arguments(arguments=arguments)
|
||||
elif config:
|
||||
self._init_by_config(config=config)
|
||||
elif config_path:
|
||||
self._init_by_config_path(config_path=config_path)
|
||||
else:
|
||||
raise RuntimeError("At least one of arguments, config, or file_path must not be empty!")
|
||||
|
||||
self.logger = self._init_logger()
|
||||
|
||||
self._init_context_by_config()
|
||||
|
||||
def _init_by_arguments(self, arguments: MemoryscopeArguments):
|
||||
# prepare global
|
||||
self.global_conf = {
|
||||
"language": arguments.language,
|
||||
"thread_pool_max_workers": arguments.thread_pool_max_workers,
|
||||
"logger_name": arguments.logger_name,
|
||||
"logger_name_time_suffix": arguments.logger_name_time_suffix,
|
||||
}
|
||||
|
||||
# prepare memory chat
|
||||
self.memory_chat_conf_dict = default_arguments.DEFAULT_MEMORY_CHAT_ARGUMENTS.copy()
|
||||
memory_chat_config = list(self.memory_chat_conf_dict.values())[0]
|
||||
memory_chat_config.update({
|
||||
"class": arguments.memory_chat_class,
|
||||
"human_name": arguments.human_name,
|
||||
"assistant_name": arguments.assistant_name,
|
||||
})
|
||||
|
||||
# prepare memory service
|
||||
self.memory_service_conf_dict = default_arguments.DEFAULT_MEMORY_SERVICE_ARGUMENTS.copy()
|
||||
memory_service_config = list(self.memory_service_conf_dict.values())[0]
|
||||
memory_service_config.update({
|
||||
"human_name": arguments.human_name,
|
||||
"assistant_name": arguments.assistant_name,
|
||||
})
|
||||
memory_service_config["memory_operations"]["consolidate_memory"]["interval_time"] = \
|
||||
arguments.consolidate_memory_interval_time
|
||||
memory_service_config["memory_operations"]["reflect_and_reconsolidate"]["interval_time"] = \
|
||||
arguments.reflect_and_reconsolidate_interval_time
|
||||
|
||||
# prepare memory service
|
||||
self.worker_conf_dict = default_arguments.DEFAULT_WORKER_ARGUMENTS.copy()
|
||||
if arguments.worker_params:
|
||||
for worker_name, kv_dict in arguments.worker_params.items():
|
||||
if worker_name not in self.worker_conf_dict:
|
||||
continue
|
||||
self.worker_conf_dict[worker_name].update(kv_dict)
|
||||
|
||||
# prepare models
|
||||
self.model_conf_dict = {
|
||||
"generation_model": {
|
||||
"class": "models.llama_index_generation_model",
|
||||
"module_name": arguments.generation_backend,
|
||||
"model_name": arguments.generation_model,
|
||||
**arguments.generation_params,
|
||||
},
|
||||
"embedding_model": {
|
||||
"class": "models.llama_index_embedding_model",
|
||||
"module_name": arguments.embedding_backend,
|
||||
"model_name": arguments.embedding_model,
|
||||
**arguments.embedding_params,
|
||||
},
|
||||
"rank_model": {
|
||||
"class": "models.llama_index_rank_model",
|
||||
"module_name": arguments.rank_backend,
|
||||
"model_name": arguments.rank_model,
|
||||
**arguments.rank_params,
|
||||
},
|
||||
}
|
||||
|
||||
# prepare memory store
|
||||
self.memory_store_conf = {
|
||||
"class": "storage.llama_index_es_memory_store",
|
||||
"embedding_model": "embedding_model",
|
||||
"index_name": arguments.es_index_name,
|
||||
"es_url": arguments.es_url,
|
||||
"retrieve_type": arguments.retrieve_type,
|
||||
"hybrid_alpha": arguments.hybrid_alpha,
|
||||
}
|
||||
|
||||
self.monitor_conf = default_arguments.DEFAULT_MONITOR_ARGUMENTS.copy()
|
||||
|
||||
def _init_by_config(self, config: dict):
|
||||
self.global_conf = config["global_config"]
|
||||
self.memory_service_conf_dict = config["memory_service"]
|
||||
self.worker_conf_dict = config["worker"]
|
||||
self.model_conf_dict = config["model"]
|
||||
self.memory_store_conf = config["memory_store"]
|
||||
|
||||
# not necessary
|
||||
self.memory_chat_conf_dict = config.get("memory_chat")
|
||||
self.monitor_conf = config.get("monitor")
|
||||
|
||||
def _init_by_config_path(self, config_path: str):
|
||||
with open(config_path) as f:
|
||||
if config_path.endswith("yaml"):
|
||||
config = yaml.load(f, yaml.FullLoader)
|
||||
elif config_path.endswith("json"):
|
||||
config = json.load(f)
|
||||
else:
|
||||
raise RuntimeError("not supported config file type!")
|
||||
return self._init_by_config(config)
|
||||
|
||||
def _init_logger(self) -> Logger:
|
||||
logger_name = self.global_conf.get("logger_name")
|
||||
assert logger_name, "logger_name is empty!"
|
||||
logger_name_time_suffix = self.global_conf.get("logger_name_time_suffix")
|
||||
if logger_name_time_suffix:
|
||||
suffix = datetime.datetime.now().strftime(logger_name_time_suffix)
|
||||
logger_name = f"{logger_name}_{suffix}"
|
||||
return Logger.get_logger(logger_name, to_stream=False)
|
||||
|
||||
def _init_context_by_config(self):
|
||||
# set global config
|
||||
self.context.language = LanguageEnum(self.global_conf["language"])
|
||||
self.context.thread_pool = ThreadPoolExecutor(max_workers=self.global_conf["max_workers"])
|
||||
|
||||
# init memory_chat
|
||||
if self.memory_chat_conf_dict:
|
||||
for name, conf in self.memory_chat_conf_dict.items():
|
||||
self.context.memory_chat_dict[name] = init_instance_by_config(conf, name=name, context=self.context)
|
||||
|
||||
# set memory_service
|
||||
assert self.memory_service_conf_dict
|
||||
for name, conf in self.memory_service_conf_dict.items():
|
||||
self.context.memory_service_dict[name] = init_instance_by_config(conf, name=name, context=self.context)
|
||||
|
||||
# init models
|
||||
assert self.model_conf_dict
|
||||
for name, conf in self.model_conf_dict.items():
|
||||
self.context.model_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# init vector_store
|
||||
assert self.memory_store_conf
|
||||
emb_model_name: str = self.memory_store_conf[ModelEnum.EMBEDDING_MODEL.value]
|
||||
embedding_model = self.context.model_dict[emb_model_name]
|
||||
self.context.memory_store = init_instance_by_config(self.memory_store_conf, embedding_model=embedding_model)
|
||||
|
||||
# init monitor
|
||||
if self.monitor_conf:
|
||||
self.context.monitor = init_instance_by_config(self.monitor_conf)
|
||||
|
||||
# set worker config
|
||||
self.context.worker_config = self.worker_conf_dict
|
||||
|
||||
def close(self):
|
||||
for _, service in self.context.memory_service_dict.items():
|
||||
service.stop_backend_service()
|
||||
self.context.memory_store.close()
|
||||
self.context.thread_pool.shutdown()
|
||||
|
||||
if self.context.monitor:
|
||||
self.context.monitor.close()
|
||||
|
||||
@property
|
||||
def default_memory_chat(self) -> BaseMemoryChat:
|
||||
return list(self.context.memory_chat_dict.values())[0]
|
||||
|
||||
@property
|
||||
def default_service(self) -> BaseMemoryService:
|
||||
return list(self.context.memory_service_dict.values())[0]
|
||||
27
memoryscope/memoryscope_context.py
Normal file
27
memoryscope/memoryscope_context.py
Normal file
|
|
@ -0,0 +1,27 @@
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
|
||||
|
||||
@dataclass
|
||||
class MemoryscopeContext(object):
|
||||
"""
|
||||
The context class archives all configs utilized by store, monitor, services and workers.
|
||||
"""
|
||||
|
||||
language: LanguageEnum = LanguageEnum.EN
|
||||
|
||||
thread_pool: ThreadPoolExecutor | None = None
|
||||
|
||||
memory_store = None
|
||||
|
||||
monitor = None
|
||||
|
||||
memory_chat_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> memory_chat"})
|
||||
|
||||
memory_service_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> memory_service"})
|
||||
|
||||
model_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> model"})
|
||||
|
||||
worker_conf_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> worker_conf"})
|
||||
|
|
@ -1,33 +0,0 @@
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, Any
|
||||
|
||||
from memoryscope.chat.base_memory_chat import BaseMemoryChat
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
from memoryscope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.models.base_model import BaseModel
|
||||
from memoryscope.storage.base_memory_store import BaseMemoryStore
|
||||
from memoryscope.storage.base_monitor import BaseMonitor
|
||||
|
||||
|
||||
class GlobalContext(object):
|
||||
"""
|
||||
The GlobalContext class archives all configs utilized by store, monitor, services and workers.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.global_config: Dict[str, Any] = {}
|
||||
self.worker_config: Dict[str, Dict[str, Any]] = {}
|
||||
|
||||
self.memory_service_dict: Dict[str, BaseMemoryService] = {}
|
||||
self.model_dict: Dict[str, BaseModel] = {}
|
||||
self.memory_chat_dict: Dict[str, BaseMemoryChat] = {}
|
||||
|
||||
self.memory_store: BaseMemoryStore | None = None
|
||||
self.monitor: BaseMonitor | None = None
|
||||
self.thread_pool: ThreadPoolExecutor | None = None
|
||||
self.language: LanguageEnum = LanguageEnum.EN
|
||||
|
||||
self.meta_data: Dict[str, Any] = {}
|
||||
|
||||
|
||||
G_CONTEXT = GlobalContext()
|
||||
|
|
@ -36,7 +36,7 @@ class MemoryHandler(object):
|
|||
BaseMemoryStore: The memory store instance associated with this worker.
|
||||
"""
|
||||
if self._memory_store is None:
|
||||
self._memory_store = G_CONTEXT.memory_store
|
||||
self._memory_store = G_CONTEXT.memory_store_conf
|
||||
return self._memory_store
|
||||
|
||||
def clear(self):
|
||||
|
|
|
|||
|
|
@ -1,19 +1,25 @@
|
|||
import json
|
||||
import os.path
|
||||
from pathlib import Path
|
||||
from typing import Dict
|
||||
|
||||
import yaml
|
||||
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
|
||||
|
||||
class PromptHandler(object):
|
||||
"""
|
||||
The `PromptHandler` class manages prompt messages by loading them from YAML or JSON files and dictionaries,
|
||||
supporting language selection based on a global context, and providing dictionary-like access to the prompt messages.
|
||||
supporting language selection based on a context, and providing dictionary-like access to the prompt messages.
|
||||
"""
|
||||
|
||||
def __init__(self, class_path: str, prompt_file: str = "", prompt_dict: dict = None, **kwargs):
|
||||
def __init__(self,
|
||||
class_path: str,
|
||||
prompt_file: str = "",
|
||||
prompt_dict: dict = None,
|
||||
language_enum: LanguageEnum = LanguageEnum.EN,
|
||||
**kwargs):
|
||||
"""
|
||||
Initializes the PromptHandler with paths to prompt sources and additional keyword arguments.
|
||||
|
||||
|
|
@ -21,30 +27,32 @@ class PromptHandler(object):
|
|||
class_path (str): The path to the class where prompts are utilized.
|
||||
prompt_file (str, optional): The path to an external file containing prompts. Defaults to "".
|
||||
prompt_dict (dict, optional): A dictionary directly containing prompt definitions. Defaults to None.
|
||||
language_enum (LanguageEnum): context language.
|
||||
**kwargs: Additional keyword arguments that might be used in prompt handling.
|
||||
"""
|
||||
self._class_path: str = class_path
|
||||
self._prompt_dict: Dict[str, str] = {}
|
||||
class_path: Path = Path(class_path)
|
||||
self._class_dir: Path = class_path.parent
|
||||
self._class_name: str = class_path.stem
|
||||
self._language_enum: LanguageEnum = language_enum
|
||||
self.kwargs = kwargs
|
||||
|
||||
file_path = self._class_path.strip(".py")
|
||||
|
||||
self.add_prompt_file(file_path)
|
||||
self._prompt_dict: Dict[str, str] = {}
|
||||
|
||||
self.add_prompt_file((self._class_dir / self._class_name).__str__(), raise_exception=False)
|
||||
if prompt_file:
|
||||
self.add_prompt_file(prompt_file)
|
||||
|
||||
self.add_prompt_file((self._class_dir / prompt_file).__str__())
|
||||
if prompt_dict:
|
||||
self.add_prompt_dict(prompt_dict)
|
||||
|
||||
@staticmethod
|
||||
def file_path_completion(file_path: str) -> str:
|
||||
def file_path_completion(file_path: str, raise_exception: bool = True) -> str:
|
||||
"""
|
||||
Attempts to complete the given file path by appending either a `.yaml` or `.json` extension
|
||||
based on the existence of the respective file. If neither exists, an exception is raised.
|
||||
|
||||
Args:
|
||||
file_path (str): The base path of the file to be completed.
|
||||
raise_exception (bool): If the file cannot be found, report an error.
|
||||
|
||||
Returns:
|
||||
str: The completed file path with the appropriate extension.
|
||||
|
|
@ -61,9 +69,10 @@ class PromptHandler(object):
|
|||
if os.path.exists(f"{file_path}.json"):
|
||||
return f"{file_path}.json"
|
||||
|
||||
raise RuntimeError(f"{file_path}/yaml/json is not exists!")
|
||||
if raise_exception:
|
||||
raise RuntimeError(f"{file_path}/yaml/json is not exists!")
|
||||
|
||||
def add_prompt_file(self, file_path: str):
|
||||
def add_prompt_file(self, file_path: str, raise_exception: bool = True):
|
||||
"""
|
||||
Adds prompt messages from a YAML or JSON file to the internal dictionary.
|
||||
|
||||
|
|
@ -72,8 +81,11 @@ class PromptHandler(object):
|
|||
|
||||
Args:
|
||||
file_path (str): The path to the YAML or JSON file containing the prompts.
|
||||
raise_exception (bool): If the file cannot be found, report an error.
|
||||
"""
|
||||
file_path = self.file_path_completion(file_path)
|
||||
file_path = self.file_path_completion(file_path, raise_exception=raise_exception)
|
||||
if not file_path:
|
||||
return
|
||||
|
||||
prompt_dict = {}
|
||||
|
||||
|
|
@ -102,9 +114,9 @@ class PromptHandler(object):
|
|||
RuntimeError: If a prompt message for the current language is not found.
|
||||
"""
|
||||
for key, language_dict in prompt_dict.items():
|
||||
prompts = language_dict.get(G_CONTEXT.language)
|
||||
prompts = language_dict.get(self._language_enum.value)
|
||||
if not prompts:
|
||||
raise RuntimeError(f"{key}.prompt.{G_CONTEXT.language} is empty!")
|
||||
raise RuntimeError(f"{key}.prompt.{self._language_enum.value} is empty!")
|
||||
self._prompt_dict[key] = prompts.strip()
|
||||
|
||||
@property
|
||||
|
|
|
|||
|
|
@ -47,10 +47,7 @@ def camelcase_to_underscore(name: str) -> str:
|
|||
return re.sub(r'(?<!^)(?=[A-Z])', '_', name).lower()
|
||||
|
||||
|
||||
def init_instance_by_config(config: dict,
|
||||
default_class_path: str = "memoryscope",
|
||||
suffix_name: str = "",
|
||||
**kwargs):
|
||||
def init_instance_by_config(config: dict, default_class_dir: str = "memoryscope", **kwargs):
|
||||
"""
|
||||
Initialize an instance of a class specified in the configuration dictionary.
|
||||
|
||||
|
|
@ -62,12 +59,9 @@ def init_instance_by_config(config: dict,
|
|||
Args:
|
||||
config (dict): A dictionary containing the configuration, including
|
||||
the 'class' key that specifies the class's module path.
|
||||
default_class_path (str, optional): The default module path prefix
|
||||
default_class_dir (str, optional): The default module path prefix
|
||||
to use if not explicitly defined in
|
||||
'config'. Defaults to "memory_scope".
|
||||
suffix_name (str, optional): A string to append to the class name,
|
||||
ensuring the final class name ends with it.
|
||||
Defaults to "".
|
||||
**kwargs: Additional keyword arguments to pass to the class constructor.
|
||||
|
||||
Returns:
|
||||
|
|
@ -80,21 +74,22 @@ def init_instance_by_config(config: dict,
|
|||
raise RuntimeError("empty class path!")
|
||||
user_defined: bool = config_copy.pop("user_defined", False)
|
||||
|
||||
class_name_split = origin_class_path.split(".")
|
||||
class_name: str = class_name_split[-1]
|
||||
if suffix_name and not class_name.lower().endswith(suffix_name.lower()):
|
||||
class_name = f"{class_name}_{suffix_name}"
|
||||
class_name_split[-1] = class_name
|
||||
class_path_list = []
|
||||
if not user_defined and default_class_dir and not origin_class_path.startswith(default_class_dir):
|
||||
class_path_list.append(default_class_dir)
|
||||
|
||||
class_paths = []
|
||||
if not user_defined and default_class_path and not origin_class_path.startswith(default_class_path):
|
||||
class_paths.append(default_class_path)
|
||||
class_paths.extend(class_name_split)
|
||||
module = import_module(".".join(class_paths))
|
||||
class_path_split = origin_class_path.split(".")
|
||||
class_file_name: str = class_path_split[-1]
|
||||
|
||||
cls_name = underscore_to_camelcase(class_name)
|
||||
class_name = underscore_to_camelcase(class_file_name)
|
||||
if class_name == class_file_name:
|
||||
class_path_list.extend(class_path_split[-1:])
|
||||
else:
|
||||
class_path_list.extend(class_path_split)
|
||||
|
||||
module = import_module(".".join(class_path_list))
|
||||
config_copy.update(kwargs)
|
||||
return getattr(module, cls_name)(**config_copy)
|
||||
return getattr(module, class_name)(**config_copy)
|
||||
|
||||
|
||||
def prompt_to_msg(system_prompt: str,
|
||||
|
|
|
|||
3
tests/other/read_prompt.yaml
Normal file
3
tests/other/read_prompt.yaml
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
a:
|
||||
cn: c
|
||||
en: e
|
||||
11
tests/other/read_yaml.py
Normal file
11
tests/other/read_yaml.py
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
import sys
|
||||
|
||||
sys.path.append(".") # noqa: E402
|
||||
|
||||
from memoryscope.utils.prompt_handler import PromptHandler
|
||||
|
||||
if __name__ == "__main__":
|
||||
file_path: str = __file__
|
||||
print(file_path)
|
||||
handler = PromptHandler(__file__, "read_prompt")
|
||||
print(handler.prompt_dict)
|
||||
Loading…
Add table
Reference in a new issue