diff --git a/v1/app.py b/v1/app.py index d19166c9..e88de5e6 100644 --- a/v1/app.py +++ b/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) diff --git a/v1/em_service.py b/v1/em_service.py deleted file mode 100644 index d57bd81e..00000000 --- a/v1/em_service.py +++ /dev/null @@ -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 diff --git a/v1/pipeline/pipeline_context.py b/v1/pipeline/pipeline_context.py index dffed290..a3dc597e 100644 --- a/v1/pipeline/pipeline_context.py +++ b/v1/pipeline/pipeline_context.py @@ -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"] diff --git a/v1/schema/request.py b/v1/schema/request.py index cba484e8..42206d32 100644 --- a/v1/schema/request.py +++ b/v1/schema/request.py @@ -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) diff --git a/v1/service/experience_maker_service.py b/v1/service/experience_maker_service.py new file mode 100644 index 00000000..22ae644e --- /dev/null +++ b/v1/service/experience_maker_service.py @@ -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() diff --git a/v1/vector_store/__init__.py b/v1/vector_store/__init__.py index 4243e5c6..1e9e6721 100644 --- a/v1/vector_store/__init__.py +++ b/v1/vector_store/__init__.py @@ -1,3 +1,7 @@ from v1.utils.registry import Registry -VECTOR_STORE_REGISTRY = Registry() \ No newline at end of file +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 diff --git a/v1/vector_store/base_vector_store.py b/v1/vector_store/base_vector_store.py index d13586be..3f2fc052 100644 --- a/v1/vector_store/base_vector_store.py +++ b/v1/vector_store/base_vector_store.py @@ -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]: diff --git a/v1/vector_store/chroma_vector_store.py b/v1/vector_store/chroma_vector_store.py new file mode 100644 index 00000000..13a2669c --- /dev/null +++ b/v1/vector_store/chroma_vector_store.py @@ -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 diff --git a/v1/vector_store/es_vector_store.py b/v1/vector_store/es_vector_store.py index f587dc76..02b75b2f 100644 --- a/v1/vector_store/es_vector_store.py +++ b/v1/vector_store/es_vector_store.py @@ -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)