mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
update service
This commit is contained in:
parent
3f742eb3b2
commit
0bb3068ad2
9 changed files with 303 additions and 125 deletions
35
v1/app.py
35
v1/app.py
|
|
@ -1,52 +1,39 @@
|
|||
import sys
|
||||
from concurrent.futures.thread import ThreadPoolExecutor
|
||||
|
||||
import uvicorn
|
||||
from fastapi import FastAPI
|
||||
|
||||
from v1.config.config_parser import ConfigParser
|
||||
from v1.em_service import EMService
|
||||
from v1.schema.request import RetrieverRequest, SummarizerRequest, VectorStoreRequest, AgentRequest
|
||||
from v1.schema.response import RetrieverResponse, SummarizerResponse, VectorStoreResponse, AgentResponse
|
||||
from v1.service.experience_maker_service import ExperienceMakerService
|
||||
|
||||
app = FastAPI()
|
||||
config_parser = ConfigParser(sys.argv[1:])
|
||||
global_app_config = config_parser.get_app_config()
|
||||
thread_pool: ThreadPoolExecutor = ThreadPoolExecutor(max_workers=global_app_config.thread_pool.max_workers)
|
||||
|
||||
service = ExperienceMakerService(sys.argv[1:])
|
||||
|
||||
@app.post('/retriever', response_model=RetrieverResponse)
|
||||
def call_retriever(request: RetrieverRequest):
|
||||
app_config = config_parser.get_app_config(**request.config)
|
||||
ems = EMService(app_config=app_config, thread_pool=thread_pool)
|
||||
return ems.call_retriever(request)
|
||||
return service(api="retriever", request=request)
|
||||
|
||||
|
||||
@app.post('/summarizer', response_model=SummarizerResponse)
|
||||
def call_summarizer(request: SummarizerRequest):
|
||||
app_config = config_parser.get_app_config(**request.config)
|
||||
ems = EMService(app_config=app_config, thread_pool=thread_pool)
|
||||
return ems.call_summarizer(request)
|
||||
return service(api="summarizer", request=request)
|
||||
|
||||
|
||||
@app.post('/vector_store', response_model=VectorStoreResponse)
|
||||
def call_vector_store(request: VectorStoreRequest):
|
||||
app_config = config_parser.get_app_config(**request.config)
|
||||
ems = EMService(app_config=app_config, thread_pool=thread_pool)
|
||||
return ems.call_vector_store(request)
|
||||
return service(api="vector_store", request=request)
|
||||
|
||||
|
||||
@app.post('/agent', response_model=AgentResponse)
|
||||
def call_agent(request: AgentRequest):
|
||||
app_config = config_parser.get_app_config(**request.config)
|
||||
ems = EMService(app_config=app_config, thread_pool=thread_pool)
|
||||
return ems.call_agent(request)
|
||||
return service(api="agent", request=request)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
uvicorn.run(app=app,
|
||||
host=global_app_config.http_service.host,
|
||||
port=global_app_config.http_service.port,
|
||||
timeout_keep_alive=global_app_config.http_service.timeout_keep_alive,
|
||||
limit_concurrency=global_app_config.http_service.limit_concurrency,
|
||||
workers=global_app_config.http_service.workers)
|
||||
host=service.http_service_config.host,
|
||||
port=service.http_service_config.port,
|
||||
timeout_keep_alive=service.http_service_config.timeout_keep_alive,
|
||||
limit_concurrency=service.http_service_config.limit_concurrency,
|
||||
workers=service.http_service_config.workers)
|
||||
|
|
|
|||
|
|
@ -1,80 +0,0 @@
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from v1.pipeline.pipeline import Pipeline
|
||||
from v1.pipeline.pipeline_context import PipelineContext
|
||||
from v1.schema.app_config import AppConfig
|
||||
from v1.schema.request import SummarizerRequest, RetrieverRequest, VectorStoreRequest, AgentRequest
|
||||
from v1.schema.response import SummarizerResponse, RetrieverResponse, VectorStoreResponse, AgentResponse
|
||||
|
||||
|
||||
class EMService(object):
|
||||
def __init__(self, app_config: AppConfig, thread_pool: ThreadPoolExecutor):
|
||||
self.context: PipelineContext = PipelineContext(app_config=app_config, thread_pool=thread_pool)
|
||||
|
||||
def __call__(self, service: str, **kwargs) -> dict:
|
||||
if service == "retriever":
|
||||
response = self.call_retriever(RetrieverRequest(**kwargs))
|
||||
|
||||
elif service == "summarizer":
|
||||
response = self.call_summarizer(SummarizerRequest(**kwargs))
|
||||
|
||||
elif service == "vector_store":
|
||||
response = self.call_vector_store(VectorStoreRequest(**kwargs))
|
||||
|
||||
elif service == "agent":
|
||||
response = self.call_agent(AgentRequest(**kwargs))
|
||||
|
||||
else:
|
||||
raise Exception(f"Invalid service={service}")
|
||||
|
||||
return response.model_dump()
|
||||
|
||||
def call_retriever(self, request: RetrieverRequest) -> RetrieverResponse:
|
||||
self.context.set_context("request", request)
|
||||
response = RetrieverResponse()
|
||||
self.context.set_context("response", response)
|
||||
try:
|
||||
Pipeline(pipeline=self.context.app_config.api.retriever, context=self.context)()
|
||||
except Exception as e:
|
||||
logger.exception(f"call_retriever encounter error={e.args}")
|
||||
response.success = False
|
||||
response.metadata["error"] = str(e)
|
||||
return response
|
||||
|
||||
def call_summarizer(self, request: SummarizerRequest) -> SummarizerResponse:
|
||||
self.context.set_context("request", request)
|
||||
response = SummarizerResponse()
|
||||
self.context.set_context("request", response)
|
||||
try:
|
||||
Pipeline(pipeline=self.context.app_config.api.summarizer, context=self.context)()
|
||||
except Exception as e:
|
||||
logger.exception(f"call_summarizer encounter error={e.args}")
|
||||
response.success = False
|
||||
response.metadata["error"] = str(e)
|
||||
return response
|
||||
|
||||
def call_vector_store(self, request: VectorStoreRequest) -> VectorStoreResponse:
|
||||
self.context.set_context("request", request)
|
||||
response = VectorStoreResponse()
|
||||
self.context.set_context("request", response)
|
||||
try:
|
||||
Pipeline(pipeline=self.context.app_config.api.vector_store, context=self.context)()
|
||||
except Exception as e:
|
||||
logger.exception(f"call_vector_store encounter error={e.args}")
|
||||
response.success = False
|
||||
response.metadata["error"] = str(e)
|
||||
return response
|
||||
|
||||
def call_agent(self, request: AgentRequest) -> AgentResponse:
|
||||
self.context.set_context("request", request)
|
||||
response = AgentResponse()
|
||||
self.context.set_context("request", response)
|
||||
try:
|
||||
Pipeline(pipeline=self.context.app_config.api.agent, context=self.context)()
|
||||
except Exception as e:
|
||||
logger.exception(f"call_agent encounter error={e.args}")
|
||||
response.success = False
|
||||
response.metadata["error"] = str(e)
|
||||
return response
|
||||
|
|
@ -1,26 +1,37 @@
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Any
|
||||
from typing import Dict
|
||||
|
||||
from v1.schema.app_config import AppConfig
|
||||
from v1.vector_store.base_vector_store import BaseVectorStore
|
||||
|
||||
|
||||
class PipelineContext(object):
|
||||
|
||||
def __init__(self, app_config: AppConfig, thread_pool: ThreadPoolExecutor):
|
||||
self.app_config: AppConfig = app_config
|
||||
self.thread_pool: ThreadPoolExecutor = thread_pool
|
||||
self.context: dict = {}
|
||||
def __init__(self, **kwargs):
|
||||
self._context: dict = {**kwargs}
|
||||
|
||||
def set_context(self, key: str, value: Any):
|
||||
self.context[key] = value
|
||||
def __getattr__(self, key: str, default=None):
|
||||
return self._context.get(key, default)
|
||||
|
||||
def get_context(self, key: str):
|
||||
return self.context.get(key)
|
||||
def __setattr__(self, key: str, value):
|
||||
self._context[key] = value
|
||||
|
||||
@property
|
||||
def request(self):
|
||||
return self.get_context("request")
|
||||
return self._context["request"]
|
||||
|
||||
@property
|
||||
def response(self):
|
||||
return self.get_context("response")
|
||||
return self._context["response"]
|
||||
|
||||
@property
|
||||
def app_config(self) -> AppConfig:
|
||||
return self._context["app_config"]
|
||||
|
||||
@property
|
||||
def thread_pool(self) -> ThreadPoolExecutor:
|
||||
return self._context["thread_pool"]
|
||||
|
||||
@property
|
||||
def vector_store_dict(self) -> Dict[str, BaseVectorStore]:
|
||||
return self._context["vector_store_dict"]
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ from pydantic import BaseModel, Field
|
|||
from v1.schema.message import Message, Trajectory
|
||||
|
||||
|
||||
class BaseRequest(BaseModel, ABC):
|
||||
class BaseRequest(BaseModel):
|
||||
workspace_id: str = Field(default=...)
|
||||
config: dict = Field(default_factory=dict)
|
||||
metadata: dict | None = Field(default=None)
|
||||
|
|
|
|||
79
v1/service/experience_maker_service.py
Normal file
79
v1/service/experience_maker_service.py
Normal file
|
|
@ -0,0 +1,79 @@
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import List
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from v1.config.config_parser import ConfigParser
|
||||
from v1.pipeline.pipeline import Pipeline
|
||||
from v1.pipeline.pipeline_context import PipelineContext
|
||||
from v1.schema.app_config import AppConfig, HttpServiceConfig
|
||||
from v1.schema.request import SummarizerRequest, RetrieverRequest, VectorStoreRequest, AgentRequest, BaseRequest
|
||||
from v1.schema.response import SummarizerResponse, RetrieverResponse, VectorStoreResponse, AgentResponse
|
||||
from v1.vector_store import VECTOR_STORE_REGISTRY
|
||||
|
||||
|
||||
class ExperienceMakerService:
|
||||
|
||||
def __init__(self, args: List[str]):
|
||||
self.config_parser = ConfigParser(args)
|
||||
self.init_app_config: AppConfig = self.config_parser.get_app_config()
|
||||
self.thread_pool = ThreadPoolExecutor(max_workers=self.init_app_config.thread_pool.max_workers)
|
||||
|
||||
# The vectorstore is initialized at the very beginning and then used directly afterwards.
|
||||
self.vector_store_dict: dict = {}
|
||||
for name, config in self.init_app_config.vector_store.items():
|
||||
assert config.backend in VECTOR_STORE_REGISTRY, f"backend={config.backend} is not existed"
|
||||
vector_store_cls = VECTOR_STORE_REGISTRY[config.backend]
|
||||
self.vector_store_dict[name] = vector_store_cls(**config.params)
|
||||
|
||||
@property
|
||||
def http_service_config(self) -> HttpServiceConfig:
|
||||
return self.init_app_config.http_service
|
||||
|
||||
def __call__(self, api: str, request: dict | BaseRequest) -> dict:
|
||||
if isinstance(request, dict):
|
||||
request = BaseRequest(**request)
|
||||
app_config: AppConfig = self.config_parser.get_app_config(**request.config)
|
||||
|
||||
if api == "retriever":
|
||||
if isinstance(request, dict):
|
||||
request = RetrieverRequest(**request)
|
||||
response = RetrieverResponse()
|
||||
pipeline = app_config.api.retriever
|
||||
|
||||
elif api == "summarizer":
|
||||
if isinstance(request, dict):
|
||||
request = SummarizerRequest(**request)
|
||||
response = SummarizerResponse()
|
||||
pipeline = app_config.api.summarizer
|
||||
|
||||
elif api == "vector_store":
|
||||
if isinstance(request, dict):
|
||||
request = VectorStoreRequest(**request)
|
||||
response = VectorStoreResponse()
|
||||
pipeline = app_config.api.vector_store
|
||||
|
||||
elif api == "agent":
|
||||
if isinstance(request, dict):
|
||||
request = AgentRequest(**request)
|
||||
response = AgentResponse()
|
||||
pipeline = app_config.api.agent
|
||||
|
||||
else:
|
||||
raise RuntimeError(f"Invalid service.api={api}")
|
||||
|
||||
try:
|
||||
context = PipelineContext(app_config=app_config,
|
||||
thread_pool=self.thread_pool,
|
||||
request=request,
|
||||
response=response,
|
||||
vector_store_dict=self.vector_store_dict)
|
||||
pipeline = Pipeline(pipeline=pipeline, context=context)
|
||||
pipeline()
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f"api={api} encounter error={e.args}")
|
||||
response.success = False
|
||||
response.metadata["error"] = str(e)
|
||||
|
||||
return response.model_dump()
|
||||
|
|
@ -1,3 +1,7 @@
|
|||
from v1.utils.registry import Registry
|
||||
|
||||
VECTOR_STORE_REGISTRY = Registry()
|
||||
VECTOR_STORE_REGISTRY = Registry()
|
||||
|
||||
from v1.vector_store.es_vector_store import EsVectorStore
|
||||
from v1.vector_store.chroma_vector_store import ChromaVectorStore
|
||||
from v1.vector_store.file_vector_store import FileVectorStore
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ from v1.schema.vector_node import VectorNode
|
|||
|
||||
|
||||
class BaseVectorStore(BaseModel, ABC):
|
||||
embedding_model: BaseEmbeddingModel = Field(default=...)
|
||||
embedding_model: BaseEmbeddingModel | None = Field(default=None)
|
||||
|
||||
@staticmethod
|
||||
def _load_from_path(path: str | Path, workspace_id: str, **kwargs) -> Iterable[VectorNode]:
|
||||
|
|
|
|||
180
v1/vector_store/chroma_vector_store.py
Normal file
180
v1/vector_store/chroma_vector_store.py
Normal file
|
|
@ -0,0 +1,180 @@
|
|||
from typing import List, Iterable
|
||||
|
||||
import chromadb
|
||||
from chromadb import Collection
|
||||
from chromadb.config import Settings
|
||||
from loguru import logger
|
||||
from pydantic import Field, PrivateAttr, model_validator
|
||||
|
||||
from v1.embedding_model.openai_compatible_embedding_model import OpenAICompatibleEmbeddingModel
|
||||
from v1.schema.vector_node import VectorNode
|
||||
from v1.vector_store import VECTOR_STORE_REGISTRY
|
||||
from v1.vector_store.base_vector_store import BaseVectorStore
|
||||
|
||||
|
||||
@VECTOR_STORE_REGISTRY.register("chroma")
|
||||
class ChromaVectorStore(BaseVectorStore):
|
||||
store_dir: str = Field(default="./chroma_vector_store")
|
||||
collections: dict = Field(default_factory=dict)
|
||||
_client: chromadb.Client = PrivateAttr()
|
||||
|
||||
@model_validator(mode="after")
|
||||
def init_client(self):
|
||||
self._client = chromadb.Client(Settings(persist_directory=self.store_dir))
|
||||
return self
|
||||
|
||||
def _get_collection(self, workspace_id: str) -> Collection:
|
||||
if workspace_id not in self.collections:
|
||||
self.collections[workspace_id] = self._client.get_or_create_collection(workspace_id)
|
||||
return self.collections[workspace_id]
|
||||
|
||||
def exist_workspace(self, workspace_id: str, **kwargs) -> bool:
|
||||
return workspace_id in [c.name for c in self._client.list_collections()]
|
||||
|
||||
def _delete_workspace(self, workspace_id: str, **kwargs):
|
||||
self._client.delete_collection(workspace_id)
|
||||
if workspace_id in self.collections:
|
||||
del self.collections[workspace_id]
|
||||
|
||||
def _create_workspace(self, workspace_id: str, **kwargs):
|
||||
self.collections[workspace_id] = self._client.get_or_create_collection(workspace_id)
|
||||
|
||||
def _iter_workspace_nodes(self, workspace_id: str, max_size: int = 10000, **kwargs) -> Iterable[VectorNode]:
|
||||
collection: Collection = self._get_collection(workspace_id)
|
||||
results = collection.peek(limit=max_size)
|
||||
for i in range(len(results["ids"])):
|
||||
node = VectorNode(workspace_id=workspace_id,
|
||||
unique_id=results["ids"][i],
|
||||
content=results["documents"][i],
|
||||
metadata=results["metadatas"][i])
|
||||
yield node
|
||||
|
||||
def retrieve_by_query(self, query: str, workspace_id: str, top_k: int = 1, **kwargs) -> List[VectorNode]:
|
||||
if not self.exist_workspace(workspace_id):
|
||||
return []
|
||||
|
||||
collection: Collection = self._get_collection(workspace_id)
|
||||
query_vector = self.embedding_model.get_embeddings(query)
|
||||
results = collection.query(query_embeddings=[query_vector], n_results=top_k)
|
||||
nodes = []
|
||||
for i in range(len(results["ids"][0])):
|
||||
node = VectorNode(workspace_id=workspace_id,
|
||||
unique_id=results["ids"][0][i],
|
||||
content=results["documents"][0][i],
|
||||
metadata=results["metadatas"][0][i])
|
||||
nodes.append(node)
|
||||
return nodes
|
||||
|
||||
def insert(self, nodes: VectorNode | List[VectorNode], workspace_id: str, refresh: bool = True, **kwargs):
|
||||
if isinstance(nodes, VectorNode):
|
||||
nodes = [nodes]
|
||||
|
||||
embedded_nodes = [node for node in nodes if node.vector]
|
||||
not_embedded_nodes = [node for node in nodes if not node.vector]
|
||||
now_embedded_nodes = self.embedding_model.get_node_embeddings(not_embedded_nodes)
|
||||
all_nodes = embedded_nodes + now_embedded_nodes
|
||||
|
||||
collection: Collection = self._get_collection(workspace_id)
|
||||
collection.add(ids=[n.unique_id for n in all_nodes],
|
||||
embeddings=[n.vector for n in all_nodes],
|
||||
documents=[n.content for n in all_nodes],
|
||||
metadatas=[n.metadata for n in all_nodes])
|
||||
|
||||
def update(self, nodes: VectorNode | List[VectorNode], workspace_id: str, refresh: bool = True, **kwargs):
|
||||
if isinstance(nodes, VectorNode):
|
||||
nodes = [nodes]
|
||||
collection: Collection = self._get_collection(workspace_id)
|
||||
collection.delete(ids=[node.unique_id for node in nodes])
|
||||
self.insert(nodes, workspace_id, refresh=refresh)
|
||||
|
||||
|
||||
def main():
|
||||
from experiencemaker.utils.util_function import load_env_keys
|
||||
load_env_keys()
|
||||
load_env_keys("../../.env")
|
||||
|
||||
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=1024) # OpenAI text-embedding-ada-002的维度
|
||||
workspace_id = "chroma_test_index"
|
||||
|
||||
chroma_store = ChromaVectorStore(
|
||||
embedding_model=embedding_model,
|
||||
store_dir="./chroma_test_db"
|
||||
)
|
||||
|
||||
if chroma_store.exist_workspace(workspace_id):
|
||||
chroma_store.delete_workspace(workspace_id)
|
||||
chroma_store.create_workspace(workspace_id)
|
||||
|
||||
sample_nodes = [
|
||||
VectorNode(
|
||||
unique_id="node1",
|
||||
workspace_id=workspace_id,
|
||||
content="Artificial intelligence is a technology that simulates human intelligence.",
|
||||
metadata={
|
||||
"node_type": "n1",
|
||||
"category": "tech"
|
||||
}
|
||||
),
|
||||
VectorNode(
|
||||
unique_id="node2",
|
||||
workspace_id=workspace_id,
|
||||
content="AI is the future of mankind.",
|
||||
metadata={
|
||||
"node_type": "n1",
|
||||
"category": "tech"
|
||||
}
|
||||
),
|
||||
VectorNode(
|
||||
unique_id="node3",
|
||||
workspace_id=workspace_id,
|
||||
content="I want to eat fish!",
|
||||
metadata={
|
||||
"node_type": "n2",
|
||||
"category": "food"
|
||||
}
|
||||
),
|
||||
VectorNode(
|
||||
unique_id="node4",
|
||||
workspace_id=workspace_id,
|
||||
content="The bigger the storm, the more expensive the fish.",
|
||||
metadata={
|
||||
"node_type": "n1",
|
||||
"category": "food"
|
||||
}
|
||||
),
|
||||
]
|
||||
|
||||
chroma_store.insert(sample_nodes, workspace_id=workspace_id)
|
||||
|
||||
logger.info("=" * 20)
|
||||
results = chroma_store.retrieve_by_query("What is AI?", top_k=5, workspace_id=workspace_id)
|
||||
for r in results:
|
||||
logger.info(r.model_dump(exclude={"vector"}))
|
||||
logger.info("=" * 20)
|
||||
|
||||
node2_update = VectorNode(
|
||||
unique_id="node2",
|
||||
workspace_id=workspace_id,
|
||||
content="AI is the future of humanity and technology.",
|
||||
metadata={
|
||||
"node_type": "n1",
|
||||
"category": "tech",
|
||||
"updated": True
|
||||
}
|
||||
)
|
||||
chroma_store.update(node2_update, workspace_id=workspace_id)
|
||||
|
||||
logger.info("Updated Result:")
|
||||
results = chroma_store.retrieve_by_query("fish?", top_k=10, workspace_id=workspace_id)
|
||||
for r in results:
|
||||
logger.info(r.model_dump(exclude={"vector"}))
|
||||
logger.info("=" * 20)
|
||||
|
||||
chroma_store.dump_workspace(workspace_id=workspace_id)
|
||||
|
||||
chroma_store.delete_workspace(workspace_id=workspace_id)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
# launch with: python -m experiencemaker.storage.chroma_vector_store
|
||||
|
|
@ -53,12 +53,9 @@ class EsVectorStore(BaseVectorStore):
|
|||
return self._client.indices.create(index=workspace_id, body=body)
|
||||
|
||||
def _iter_workspace_nodes(self, workspace_id: str, max_size: int = 10000, **kwargs) -> Iterable[VectorNode]:
|
||||
response = self._client.search(
|
||||
index=workspace_id,
|
||||
body={"query": {"match_all": {}}},
|
||||
scroll='5m',
|
||||
size=max_size
|
||||
)
|
||||
response = self._client.search(index=workspace_id,
|
||||
body={"query": {"match_all": {}}, "size": max_size},
|
||||
scroll='5m')
|
||||
|
||||
for doc in response['hits']['hits']:
|
||||
yield self.doc2node(doc)
|
||||
|
|
@ -228,7 +225,7 @@ def main():
|
|||
for r in results:
|
||||
logger.info(r.model_dump(exclude={"vector"}))
|
||||
logger.info("=" * 20)
|
||||
|
||||
es.dump_workspace(workspace_id=workspace_id)
|
||||
es.delete_workspace(workspace_id=workspace_id)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue