mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-11 22:51:10 +00:00
[dev] change dummy_generation to dashscope_generation
This commit is contained in:
parent
0e07add988
commit
8282350ab4
15 changed files with 80 additions and 74 deletions
|
|
@ -6,7 +6,7 @@ memory_chat:
|
|||
cli_memory_chat:
|
||||
class: chat.cli_memory_chat
|
||||
memory_service: memory_chat_service
|
||||
generation_model: dummy_generation
|
||||
generation_model: dashscope_generation
|
||||
|
||||
memory_service:
|
||||
memory_chat_service:
|
||||
|
|
@ -19,8 +19,7 @@ memory_service:
|
|||
description: "read session messages of the user"
|
||||
read_memory:
|
||||
class: memory.operation.read_memory
|
||||
# workflow: set_query,retrieve_memory1,[extract_time|semantic_rank],fuse_rerank
|
||||
workflow: dummy
|
||||
workflow: set_query,[extract_time|retrieve_memory1,semantic_rank],fuse_rerank
|
||||
description: "read related memories of the user"
|
||||
list_memory:
|
||||
class: memory.operation.read_memory
|
||||
|
|
@ -56,6 +55,7 @@ worker:
|
|||
generation_model_top_k: 1
|
||||
semantic_rank:
|
||||
class: memory.worker.read.semantic_rank_worker
|
||||
rank_model: dashscope_rank
|
||||
fuse_rerank:
|
||||
class: memory.worker.read.fuse_rerank_worker
|
||||
fuse_score_threshold: 0.1
|
||||
|
|
@ -142,7 +142,7 @@ models:
|
|||
model_name: dummy_generation_model
|
||||
|
||||
memory_store:
|
||||
class: storage.llama_index_es_memory_store
|
||||
class: storage.llama_index_es_memory_store_sync
|
||||
embedding_model: dashscope_embedding
|
||||
index_name: memory_index
|
||||
es_url: http://localhost:9200
|
||||
|
|
|
|||
|
|
@ -154,11 +154,12 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
if refresh_time and refresh_time.isdigit():
|
||||
refresh_time = int(refresh_time)
|
||||
while True:
|
||||
time.sleep(refresh_time)
|
||||
result = self.memory_service.do_operation(op_name=command, **kwargs)
|
||||
os.system("clear")
|
||||
self.print_logo()
|
||||
questionary.print(result)
|
||||
time.sleep(refresh_time)
|
||||
|
||||
else:
|
||||
result = self.memory_service.do_operation(op_name=command, **kwargs)
|
||||
questionary.print(result)
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import datetime
|
||||
import sys
|
||||
|
||||
sys.path.append(".") # noqa: E402
|
||||
|
|
@ -22,7 +23,8 @@ class MemoryScope(object):
|
|||
|
||||
def __init__(self):
|
||||
self.config: Dict[str, Any] = {}
|
||||
self.logger: Logger = Logger.get_logger("cli_job", to_stream=False)
|
||||
datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S')
|
||||
self.logger: Logger = Logger.get_logger(f"cli_job_{datetime_suffix}", to_stream=False)
|
||||
|
||||
def load_config(self, path: str):
|
||||
with open(path) as f:
|
||||
|
|
@ -36,7 +38,8 @@ class MemoryScope(object):
|
|||
atexit.register(self.shutdown) # register clean up function
|
||||
return self
|
||||
|
||||
def shutdown(self):
|
||||
@staticmethod
|
||||
def shutdown():
|
||||
print('Gracefully executing the shutdown function...')
|
||||
G_CONTEXT.memory_store.close()
|
||||
G_CONTEXT.monitor.close()
|
||||
|
|
@ -68,9 +71,8 @@ class MemoryScope(object):
|
|||
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]]
|
||||
embedding_model_conf = self.config["models"][memory_store_config[ModelEnum.EMBEDDING_MODEL.value]]
|
||||
G_CONTEXT.memory_store = init_instance_by_config(memory_store_config, embedding_model_conf=embedding_model_conf)
|
||||
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"])
|
||||
|
|
@ -78,10 +80,12 @@ class MemoryScope(object):
|
|||
# set worker config
|
||||
G_CONTEXT.worker_config = self.config["worker"]
|
||||
|
||||
def get_default_service(self):
|
||||
@property
|
||||
def default_service(self):
|
||||
return list(G_CONTEXT.memory_service_dict.values())[0]
|
||||
|
||||
def get_default_chat_handle(self):
|
||||
@property
|
||||
def default_chat_handle(self):
|
||||
return list(G_CONTEXT.memory_chat_dict.values())[0]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -25,10 +25,10 @@ class ReadMemory(BaseWorkflow, BaseOperation):
|
|||
self.init_workers(**kwargs)
|
||||
|
||||
def run_operation(self, **kwargs):
|
||||
self.context.clear()
|
||||
max_count = 1 + self.his_msg_count
|
||||
self.context[CHAT_MESSAGES] = [x.copy(deep=True) for x in self.chat_messages[-max_count:]]
|
||||
self.context[CHAT_KWARGS] = kwargs
|
||||
self.run_workflow()
|
||||
result = self.context.get(RESULT)
|
||||
self.context.clear()
|
||||
return result
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ class ReadMessage(BaseOperation):
|
|||
name: str,
|
||||
description: str,
|
||||
chat_messages: List[Message],
|
||||
contextual_msg_count: int = 6, # for the current context dialogue
|
||||
contextual_msg_count: int, # for the current context dialogue
|
||||
**kwargs):
|
||||
super().__init__(name=name, description=description, **kwargs)
|
||||
self.chat_messages: List[Message] = chat_messages
|
||||
|
|
|
|||
|
|
@ -32,14 +32,18 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
self.logger: Logger = Logger.get_logger()
|
||||
|
||||
def submit_async_task(self, fn, *args, **kwargs):
|
||||
# if self.is_multi_thread:
|
||||
# raise RuntimeError(f"async_task is not allowed in multi_thread condition")
|
||||
if self.is_multi_thread:
|
||||
raise RuntimeError(f"async_task is not allowed in multi_thread condition")
|
||||
|
||||
self.async_task_list.append((fn, args, kwargs))
|
||||
|
||||
async def _async_gather(self):
|
||||
return await asyncio.gather(*[fn(*args, **kwargs) for fn, args, kwargs in self.async_task_list])
|
||||
|
||||
def gather_async_result(self):
|
||||
if self.is_multi_thread:
|
||||
raise RuntimeError(f"async_task is not allowed in multi_thread condition")
|
||||
|
||||
results = asyncio.run(self._async_gather())
|
||||
self.async_task_list.clear()
|
||||
return results
|
||||
|
|
@ -50,6 +54,7 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
def gather_thread_result(self):
|
||||
for future in as_completed(self.thread_task_list):
|
||||
yield future.result()
|
||||
self.thread_task_list.clear()
|
||||
|
||||
@abstractmethod
|
||||
def _run(self):
|
||||
|
|
|
|||
|
|
@ -95,6 +95,8 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
if node.memory_id in self.contex_memory_dict:
|
||||
continue
|
||||
self.contex_memory_dict[node.memory_id] = node
|
||||
self.logger.info(f"add to memory context memory id={node.memory_id} content={node.content} "
|
||||
f"status={node.status}")
|
||||
self.set_context(key, [n.memory_id for n in nodes])
|
||||
|
||||
def save_memories(self, keys: str | List[str] = None):
|
||||
|
|
|
|||
|
|
@ -40,7 +40,7 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
def _run(self):
|
||||
# parse input
|
||||
extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT)
|
||||
memory_node_list: List[MemoryNode] = self.get_context(RANKED_MEMORY_NODES)
|
||||
memory_node_list: List[MemoryNode] = self.get_memories(RANKED_MEMORY_NODES)
|
||||
if not memory_node_list:
|
||||
self.logger.warning(f"ranked memory nodes is empty!")
|
||||
return
|
||||
|
|
|
|||
|
|
@ -26,7 +26,7 @@ class PrintMemoryWorker(MemoryBaseWorker):
|
|||
line = f" {i} {node.content}"
|
||||
expired_content_list.append(line)
|
||||
|
||||
if MemoryTypeEnum(node.memory_type) in [MemoryTypeEnum.OBSERVATION, MemoryTypeEnum.OBS_CUSTOMIZED]:
|
||||
elif MemoryTypeEnum(node.memory_type) in [MemoryTypeEnum.OBSERVATION, MemoryTypeEnum.OBS_CUSTOMIZED]:
|
||||
j += 1
|
||||
dt_handler = DatetimeHandler(node.timestamp)
|
||||
dt = dt_handler.datetime_format("%Y%m%d %H:%M:%S")
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ from memory_scope.utils.timer import timer
|
|||
class RetrieveMemoryWorker(MemoryBaseWorker):
|
||||
|
||||
@timer
|
||||
async def retrieve_from_observation(self, query: str) -> List[MemoryNode]:
|
||||
def retrieve_from_observation(self, query: str) -> List[MemoryNode]:
|
||||
if not self.retrieve_obs_top_k:
|
||||
return []
|
||||
|
||||
|
|
@ -22,11 +22,11 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
|
|||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
}
|
||||
return self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_obs_top_k,
|
||||
filter_dict=filter_dict)
|
||||
top_k=self.retrieve_obs_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
@timer
|
||||
async def retrieve_from_insight_and_profile(self, query: str) -> List[MemoryNode]:
|
||||
def retrieve_from_insight_and_profile(self, query: str) -> List[MemoryNode]:
|
||||
if not self.retrieve_ins_pf_top_k:
|
||||
return []
|
||||
|
||||
|
|
@ -37,11 +37,11 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
|
|||
"memory_type": MemoryTypeEnum.INSIGHT.value,
|
||||
}
|
||||
return self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_ins_pf_top_k,
|
||||
filter_dict=filter_dict)
|
||||
top_k=self.retrieve_ins_pf_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
@timer
|
||||
async def retrieve_expired_memory(self, query: str) -> List[MemoryNode]:
|
||||
def retrieve_expired_memory(self, query: str) -> List[MemoryNode]:
|
||||
if not self.retrieve_expired_top_k:
|
||||
return []
|
||||
|
||||
|
|
@ -52,23 +52,23 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
|
|||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
}
|
||||
return self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_expired_top_k,
|
||||
filter_dict=filter_dict)
|
||||
top_k=self.retrieve_expired_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
def _run(self):
|
||||
query, _ = self.get_context(QUERY_WITH_TS)
|
||||
self.submit_async_task(self.retrieve_from_observation, query=query)
|
||||
self.submit_async_task(self.retrieve_from_insight_and_profile, query=query)
|
||||
self.submit_async_task(self.retrieve_expired_memory, query=query)
|
||||
self.submit_thread_task(self.retrieve_from_observation, query=query)
|
||||
self.submit_thread_task(self.retrieve_from_insight_and_profile, query=query)
|
||||
self.submit_thread_task(self.retrieve_expired_memory, query=query)
|
||||
|
||||
memory_node_list: List[MemoryNode] = []
|
||||
for result in self.gather_async_result():
|
||||
for result in self.gather_thread_result():
|
||||
if result:
|
||||
memory_node_list.extend(result)
|
||||
self.logger.info(f"memory_node_list.size={len(memory_node_list)}")
|
||||
|
||||
memory_node_list = sorted(memory_node_list, key=lambda x: x.score_similar, reverse=True)
|
||||
for node in memory_node_list:
|
||||
self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} type={node.memory_type}"
|
||||
f"status={node.status}")
|
||||
self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} "
|
||||
f"type={node.memory_type} status={node.status}")
|
||||
self.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list)
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ from memory_scope.utils.timer import timer
|
|||
class LoadMemoryWorker(MemoryBaseWorker):
|
||||
|
||||
@timer
|
||||
async def retrieve_not_reflected_memory(self, query: str):
|
||||
def retrieve_not_reflected_memory(self, query: str):
|
||||
if not self.retrieve_not_reflected_top_k:
|
||||
return
|
||||
|
||||
|
|
@ -24,13 +24,13 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"obs_reflected": False,
|
||||
}
|
||||
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_not_reflected_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_memories(NOT_REFLECTED_NODES, nodes)
|
||||
|
||||
@timer
|
||||
async def retrieve_not_updated_memory(self, query: str):
|
||||
def retrieve_not_updated_memory(self, query: str):
|
||||
if not self.retrieve_not_updated_top_k:
|
||||
return
|
||||
|
||||
|
|
@ -41,13 +41,13 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"obs_updated": False,
|
||||
}
|
||||
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_not_updated_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_memories(NOT_UPDATED_NODES, nodes)
|
||||
|
||||
@timer
|
||||
async def retrieve_insight_memory(self, query: str):
|
||||
def retrieve_insight_memory(self, query: str):
|
||||
if not self.retrieve_insight_top_k:
|
||||
return
|
||||
|
||||
|
|
@ -57,13 +57,13 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": MemoryTypeEnum.INSIGHT.value,
|
||||
}
|
||||
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_insight_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_memories(INSIGHT_NODES, nodes)
|
||||
|
||||
@timer
|
||||
async def retrieve_today_memory(self):
|
||||
def retrieve_today_memory(self):
|
||||
if not self.today_obs_top_k:
|
||||
return
|
||||
|
||||
|
|
@ -80,7 +80,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"dt": dt_handler.datetime_format(),
|
||||
}
|
||||
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=message.content,
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=message.content,
|
||||
top_k=self.today_obs_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
|
|
@ -88,14 +88,8 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
|
||||
def _run(self):
|
||||
mock_query = "-"
|
||||
# self.submit_thread_task(self.retrieve_not_reflected_memory, query=mock_query)
|
||||
# self.submit_thread_task(self.retrieve_not_updated_memory, query=mock_query)
|
||||
# self.submit_thread_task(self.retrieve_insight_memory, query=mock_query)
|
||||
# self.submit_thread_task(self.retrieve_today_memory)
|
||||
# self.gather_thread_result()
|
||||
|
||||
self.submit_async_task(self.retrieve_not_reflected_memory, query=mock_query)
|
||||
self.submit_async_task(self.retrieve_not_updated_memory, query=mock_query)
|
||||
self.submit_async_task(self.retrieve_insight_memory, query=mock_query)
|
||||
self.submit_async_task(self.retrieve_today_memory)
|
||||
self.gather_async_result()
|
||||
self.submit_thread_task(self.retrieve_not_reflected_memory, query=mock_query)
|
||||
self.submit_thread_task(self.retrieve_not_updated_memory, query=mock_query)
|
||||
self.submit_thread_task(self.retrieve_insight_memory, query=mock_query)
|
||||
self.submit_thread_task(self.retrieve_today_memory)
|
||||
self.gather_thread_result()
|
||||
|
|
|
|||
|
|
@ -2,15 +2,15 @@ from typing import Dict, List, Any, Optional, cast
|
|||
|
||||
from llama_index.core import VectorStoreIndex
|
||||
from llama_index.core.schema import TextNode, NodeWithScore
|
||||
from llama_index.vector_stores.elasticsearch import ElasticsearchStore, AsyncDenseVectorStrategy
|
||||
from llama_index.vector_stores.elasticsearch import AsyncDenseVectorStrategy
|
||||
|
||||
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.storage.base_memory_store import BaseMemoryStore
|
||||
from memory_scope.storage.llama_index_sync_elasticsearch import SyncElasticsearchStore
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
from memory_scope.storage.llama_index_sync_elasticsearch import SyncElasticsearchStore
|
||||
|
||||
class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy):
|
||||
def _hybrid(self, query: str, knn: Dict[str, Any], filter: List[Dict[str, Any]], top_k: int) -> Dict[str, Any]:
|
||||
|
|
@ -124,7 +124,7 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str
|
|||
return result
|
||||
|
||||
|
||||
class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
||||
class LlamaIndexEsMemoryStoreSync(BaseMemoryStore):
|
||||
def __init__(self,
|
||||
embedding_model: BaseModel,
|
||||
index_name: str,
|
||||
|
|
@ -134,11 +134,13 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
|
||||
self.embedding_model: BaseModel = embedding_model
|
||||
self.es_store = SyncElasticsearchStore(index_name=index_name,
|
||||
es_url=es_url,
|
||||
retrieval_strategy=_AsyncDenseVectorStrategy(hybrid=use_hybrid),
|
||||
**kwargs)
|
||||
self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store,
|
||||
embed_model=self.embedding_model.model)
|
||||
es_url=es_url,
|
||||
retrieval_strategy=_AsyncDenseVectorStrategy(hybrid=use_hybrid),
|
||||
**kwargs)
|
||||
# use /dev/null
|
||||
with open(os.devnull, 'w') as devnull:
|
||||
self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store,
|
||||
embed_model=self.embedding_model.model)
|
||||
self.index.build_index_from_nodes([TextNode(text="text")])
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
|
|
@ -192,13 +194,14 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
if not nodes:
|
||||
self.logger.warning("empty nodes!")
|
||||
return
|
||||
self.logger.info(f"update_memories nodes={nodes}")
|
||||
|
||||
if isinstance(nodes, MemoryNode):
|
||||
nodes = [nodes]
|
||||
|
||||
# emb & insert new memories
|
||||
# TODO batch insert
|
||||
new_memories = [n for n in nodes if n.status == MemoryNodeStatus.NEW]
|
||||
new_memories = [n for n in nodes if n.status == MemoryNodeStatus.NEW.value]
|
||||
if new_memories:
|
||||
for n in new_memories:
|
||||
n.status = MemoryNodeStatus.ACTIVE.value
|
||||
|
|
@ -206,7 +209,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
|
||||
# emb & update new memories
|
||||
# TODO insert overwrite
|
||||
c_modified_memories = [n for n in nodes if n.status == MemoryNodeStatus.CONTENT_MODIFIED]
|
||||
c_modified_memories = [n for n in nodes if n.status == MemoryNodeStatus.CONTENT_MODIFIED.value]
|
||||
if c_modified_memories:
|
||||
for n in c_modified_memories:
|
||||
n.status = MemoryNodeStatus.ACTIVE.value
|
||||
|
|
@ -215,7 +218,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
|
||||
# update new memories
|
||||
# TODO no emb
|
||||
modified_memories = [n for n in nodes if n.status == MemoryNodeStatus.MODIFIED]
|
||||
modified_memories = [n for n in nodes if n.status == MemoryNodeStatus.MODIFIED.value]
|
||||
if modified_memories:
|
||||
for n in modified_memories:
|
||||
n.status = MemoryNodeStatus.ACTIVE.value
|
||||
|
|
@ -223,10 +226,9 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
self.insert(n)
|
||||
|
||||
# set memories expired
|
||||
expired_memories = [n for n in nodes if n.status == MemoryNodeStatus.EXPIRED]
|
||||
expired_memories = [n for n in nodes if n.status == MemoryNodeStatus.EXPIRED.value]
|
||||
if expired_memories:
|
||||
for n in expired_memories:
|
||||
n.status = MemoryNodeStatus.ACTIVE.value
|
||||
self.delete(n)
|
||||
self.insert(n)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,11 +1,9 @@
|
|||
import init_test
|
||||
from memory_scope.cli import MemoryScope
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memory_scope.scheme.message import Message
|
||||
|
||||
ms = MemoryScope().load_config("config/demo_config_no_stream.yaml")
|
||||
memory_service = ms.get_default_service()
|
||||
memory_chat = ms.get_default_chat_handle()
|
||||
memory_service = ms.default_service
|
||||
memory_chat = ms.default_chat_handle
|
||||
|
||||
# new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name="我", content="我的爱好是弹琴并且喜欢看电影。")
|
||||
# memory_service.add_messages(new_message)
|
||||
|
|
|
|||
|
|
@ -2,7 +2,8 @@ import unittest
|
|||
|
||||
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.storage.llama_index_es_memory_store_sync import LlamaIndexEsMemoryStore
|
||||
from memory_scope.storage.llama_index_es_memory_store_sync import LlamaIndexEsMemoryStoreSync
|
||||
|
||||
|
||||
class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
||||
"""Tests for LLIEmbedding"""
|
||||
|
|
@ -22,7 +23,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
|||
"use_hybrid": True
|
||||
|
||||
}
|
||||
self.es_store = LlamaIndexEsMemoryStore(**config)
|
||||
self.es_store = LlamaIndexEsMemoryStoreSync(**config)
|
||||
self.data = [
|
||||
MemoryNode(
|
||||
content="The lives of two mob hitmen, a boxer, a gangster and his wife, "
|
||||
|
|
|
|||
|
|
@ -2,13 +2,12 @@ import sys
|
|||
|
||||
sys.path.append(".")
|
||||
|
||||
import asyncio
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
|
||||
|
||||
from memory_scope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore
|
||||
from memory_scope.storage.llama_index_es_memory_store_sync import LlamaIndexEsMemoryStore as SyncLlamaIndexEsMemoryStore
|
||||
from memory_scope.storage.llama_index_es_memory_store_sync import \
|
||||
LlamaIndexEsMemoryStoreSync as SyncLlamaIndexEsMemoryStore
|
||||
from memory_scope.utils.logger import Logger
|
||||
logger = Logger.get_logger("default")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue