From a76cde5f413e678165bd4b70c1f57759803ca824 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Tue, 10 Jun 2025 12:28:54 +0800 Subject: [PATCH] bug fix --- experiencemaker/model/__init__.py | 13 ++ experiencemaker/model/base_embedding_model.py | 6 +- experiencemaker/model/base_llm.py | 4 - .../openai_compatible_embedding_model.py | 5 +- .../model/openai_compatible_llm.py | 5 +- .../module/agent_wrapper/__init__.py | 7 + .../agent_wrapper/agent_wrapper_mixin.py | 4 +- .../agent_wrapper/simple_agent_wrapper.py | 5 +- .../module/context_generator/__init__.py | 7 + .../base_context_generator.py | 5 - .../simple_context_generator.py | 11 +- experiencemaker/module/prompt/prompt_mixin.py | 5 +- experiencemaker/module/summarizer/__init__.py | 6 + .../module/summarizer/base_summarizer.py | 4 - .../module/summarizer/simple_summarizer.py | 5 +- .../experience_maker_client.py} | 2 +- .../service/experience_maker_service.py | 192 ++++++++++++------ experiencemaker/service/http_service.py | 50 ----- experiencemaker/storage/__init__.py | 8 + experiencemaker/storage/base_vector_store.py | 2 - experiencemaker/storage/es_vector_store.py | 5 +- experiencemaker/storage/file_vector_store.py | 5 +- experiencemaker/tool/__init__.py | 15 +- experiencemaker/tool/base_tool.py | 5 - experiencemaker/tool/code_tool.py | 5 +- experiencemaker/tool/dashscope_search_tool.py | 6 +- experiencemaker/tool/terminate_tool.py | 5 +- experiencemaker/utils/registry.py | 5 +- 28 files changed, 203 insertions(+), 194 deletions(-) rename experiencemaker/{client.py => service/experience_maker_client.py} (96%) delete mode 100644 experiencemaker/service/http_service.py diff --git a/experiencemaker/model/__init__.py b/experiencemaker/model/__init__.py index e69de29b..f3a307df 100644 --- a/experiencemaker/model/__init__.py +++ b/experiencemaker/model/__init__.py @@ -0,0 +1,13 @@ +from experiencemaker.model.base_embedding_model import BaseEmbeddingModel +from experiencemaker.model.base_llm import BaseLLM +from experiencemaker.model.openai_compatible_embedding_model import OpenAICompatibleEmbeddingModel +from experiencemaker.model.openai_compatible_llm import OpenAICompatibleBaseLLM +from experiencemaker.utils.registry import Registry + +EMBEDDING_MODEL_REGISTRY = Registry[BaseEmbeddingModel]("embedding_model") + +EMBEDDING_MODEL_REGISTRY.register(OpenAICompatibleEmbeddingModel, "openai_compatible") + +LLM_REGISTRY = Registry[BaseLLM]("llm") + +LLM_REGISTRY.register(OpenAICompatibleBaseLLM, "openai_compatible") diff --git a/experiencemaker/model/base_embedding_model.py b/experiencemaker/model/base_embedding_model.py index ef09a74e..b015ad9f 100644 --- a/experiencemaker/model/base_embedding_model.py +++ b/experiencemaker/model/base_embedding_model.py @@ -5,7 +5,6 @@ from loguru import logger from pydantic import BaseModel, Field from experiencemaker.schema.vector_store_node import VectorStoreNode -from experiencemaker.utils.registry import Registry class BaseEmbeddingModel(BaseModel, ABC): @@ -65,7 +64,7 @@ class BaseEmbeddingModel(BaseModel, ABC): - nodes (VectorStoreNode | List[VectorStoreNode]): A single node or list of nodes whose embeddings need to be retrieved. Returns: - - (VectorStoreNode | List[VectorStoreNode]): Returns the input nodes with their vector attribute populated with embeddings. + - VectorStoreNode | List[VectorStoreNode]: Returns the input nodes with their vector attribute populated with embeddings. Raises: - RuntimeError: If the input is neither a VectorStoreNode nor a list of VectorStoreNodes, a RuntimeError is raised. @@ -85,6 +84,3 @@ class BaseEmbeddingModel(BaseModel, ABC): else: raise RuntimeError(f"unsupported type={type(nodes)}") - - -EMBEDDING_MODEL_REGISTRY = Registry[BaseEmbeddingModel]("embedding_model") diff --git a/experiencemaker/model/base_llm.py b/experiencemaker/model/base_llm.py index a055e6f9..f78dc6e2 100644 --- a/experiencemaker/model/base_llm.py +++ b/experiencemaker/model/base_llm.py @@ -6,7 +6,6 @@ from pydantic import Field, BaseModel from experiencemaker.schema.trajectory import Message, ActionMessage from experiencemaker.tool.base_tool import BaseTool -from experiencemaker.utils.registry import Registry class BaseLLM(BaseModel, ABC): @@ -109,6 +108,3 @@ class BaseLLM(BaseModel, ABC): raise e return None - - -LLM_REGISTRY = Registry[BaseLLM]("llm") diff --git a/experiencemaker/model/openai_compatible_embedding_model.py b/experiencemaker/model/openai_compatible_embedding_model.py index fe3ff2de..9d59686f 100644 --- a/experiencemaker/model/openai_compatible_embedding_model.py +++ b/experiencemaker/model/openai_compatible_embedding_model.py @@ -4,7 +4,7 @@ from typing import Literal, List from openai import OpenAI from pydantic import Field, PrivateAttr, model_validator -from experiencemaker.model.base_embedding_model import BaseEmbeddingModel, EMBEDDING_MODEL_REGISTRY +from experiencemaker.model.base_embedding_model import BaseEmbeddingModel class OpenAICompatibleEmbeddingModel(BaseEmbeddingModel): @@ -73,9 +73,6 @@ class OpenAICompatibleEmbeddingModel(BaseEmbeddingModel): raise RuntimeError(f"unsupported type={type(input_text)}") -EMBEDDING_MODEL_REGISTRY.register(OpenAICompatibleEmbeddingModel, "openai_compatible") - - def main(): from experiencemaker.utils.util_function import load_env_keys load_env_keys() diff --git a/experiencemaker/model/openai_compatible_llm.py b/experiencemaker/model/openai_compatible_llm.py index 09fedef2..68de4744 100644 --- a/experiencemaker/model/openai_compatible_llm.py +++ b/experiencemaker/model/openai_compatible_llm.py @@ -7,7 +7,7 @@ from openai.types import CompletionUsage from pydantic import Field, PrivateAttr, model_validator from experiencemaker.enumeration.chunk_enum import ChunkEnum -from experiencemaker.model.base_llm import BaseLLM, LLM_REGISTRY +from experiencemaker.model.base_llm import BaseLLM from experiencemaker.schema.trajectory import Message, ActionMessage, ToolCall from experiencemaker.tool.base_tool import BaseTool @@ -152,9 +152,6 @@ class OpenAICompatibleBaseLLM(BaseLLM): print(f"\n{chunk}", end="") -LLM_REGISTRY.register(OpenAICompatibleBaseLLM, "openai_compatible") - - def main(): from experiencemaker.utils.util_function import load_env_keys from experiencemaker.tool.dashscope_search_tool import DashscopeSearchTool diff --git a/experiencemaker/module/agent_wrapper/__init__.py b/experiencemaker/module/agent_wrapper/__init__.py index e69de29b..95e9c073 100644 --- a/experiencemaker/module/agent_wrapper/__init__.py +++ b/experiencemaker/module/agent_wrapper/__init__.py @@ -0,0 +1,7 @@ +from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AgentWrapperMixin +from experiencemaker.module.agent_wrapper.simple_agent_wrapper import SimpleAgentWrapper +from experiencemaker.utils.registry import Registry + +AGENT_WRAPPER_REGISTRY = Registry[AgentWrapperMixin]("agent_wrapper") + +AGENT_WRAPPER_REGISTRY.register(SimpleAgentWrapper, "simple") diff --git a/experiencemaker/module/agent_wrapper/agent_wrapper_mixin.py b/experiencemaker/module/agent_wrapper/agent_wrapper_mixin.py index bcaf4450..58e4b078 100644 --- a/experiencemaker/module/agent_wrapper/agent_wrapper_mixin.py +++ b/experiencemaker/module/agent_wrapper/agent_wrapper_mixin.py @@ -2,17 +2,17 @@ from abc import ABC from pydantic import Field, BaseModel +from experiencemaker.model.base_llm import BaseLLM from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator from experiencemaker.schema.trajectory import Trajectory -from experiencemaker.utils.registry import Registry class AgentWrapperMixin(BaseModel, ABC): context_generator: BaseContextGenerator | None = Field(default=None) + llm: BaseLLM | None = Field(default=None) workspace_id: str = Field(default="") def execute(self, query: str, **kwargs) -> Trajectory: raise NotImplementedError -AGENT_WRAPPER_REGISTRY = Registry[AgentWrapperMixin]("agent_wrapper") diff --git a/experiencemaker/module/agent_wrapper/simple_agent_wrapper.py b/experiencemaker/module/agent_wrapper/simple_agent_wrapper.py index 124ced7c..0a984897 100644 --- a/experiencemaker/module/agent_wrapper/simple_agent_wrapper.py +++ b/experiencemaker/module/agent_wrapper/simple_agent_wrapper.py @@ -1,4 +1,4 @@ -from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AgentWrapperMixin, AGENT_WRAPPER_REGISTRY +from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AgentWrapperMixin from experiencemaker.module.agent_wrapper.simple_agent import SimpleAgent from experiencemaker.schema.trajectory import Trajectory @@ -16,6 +16,3 @@ class SimpleAgentWrapper(SimpleAgent, AgentWrapperMixin): trajectory.answer = messages[-1].content trajectory.done = True return trajectory - - -AGENT_WRAPPER_REGISTRY.register(SimpleAgentWrapper, "simple") diff --git a/experiencemaker/module/context_generator/__init__.py b/experiencemaker/module/context_generator/__init__.py index e69de29b..621fc540 100644 --- a/experiencemaker/module/context_generator/__init__.py +++ b/experiencemaker/module/context_generator/__init__.py @@ -0,0 +1,7 @@ +from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator +from experiencemaker.module.context_generator.simple_context_generator import SimpleContextGenerator +from experiencemaker.utils.registry import Registry + +CONTEXT_GENERATOR_REGISTRY = Registry[BaseContextGenerator]("context_generator") + +CONTEXT_GENERATOR_REGISTRY.register(SimpleContextGenerator, "simple") diff --git a/experiencemaker/module/context_generator/base_context_generator.py b/experiencemaker/module/context_generator/base_context_generator.py index e7839a12..8c49cb0d 100644 --- a/experiencemaker/module/context_generator/base_context_generator.py +++ b/experiencemaker/module/context_generator/base_context_generator.py @@ -3,12 +3,10 @@ from typing import List from pydantic import Field, BaseModel -from experiencemaker.model.base_embedding_model import BaseEmbeddingModel from experiencemaker.model.base_llm import BaseLLM from experiencemaker.schema.trajectory import Trajectory, ContextMessage from experiencemaker.schema.vector_store_node import VectorStoreNode from experiencemaker.storage.base_vector_store import BaseVectorStore -from experiencemaker.utils.registry import Registry class BaseContextGenerator(BaseModel, ABC): @@ -33,6 +31,3 @@ class BaseContextGenerator(BaseModel, ABC): nodes: List[VectorStoreNode] = self._retrieve_by_query(trajectory, query, **kwargs) context_msg: ContextMessage = self._generate_context_message(trajectory, nodes, **kwargs) return context_msg - - -CONTEXT_GENERATOR_REGISTRY = Registry[BaseContextGenerator]("context_generator") diff --git a/experiencemaker/module/context_generator/simple_context_generator.py b/experiencemaker/module/context_generator/simple_context_generator.py index 4beb15a0..9b76a66f 100644 --- a/experiencemaker/module/context_generator/simple_context_generator.py +++ b/experiencemaker/module/context_generator/simple_context_generator.py @@ -1,12 +1,14 @@ from typing import List -from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator, \ - CONTEXT_GENERATOR_REGISTRY +from pydantic import Field + +from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator from experiencemaker.schema.trajectory import Trajectory, ContextMessage from experiencemaker.schema.vector_store_node import VectorStoreNode class SimpleContextGenerator(BaseContextGenerator): + retrieve_top_k: int = Field(default=5) def _build_retrieve_query(self, trajectory: Trajectory, **kwargs) -> str: query = "" @@ -18,7 +20,7 @@ class SimpleContextGenerator(BaseContextGenerator): if not query: return [] - return self.vector_store.retrieve_by_query(query=query, top_k=self.vector_store_top_k) + return self.vector_store.retrieve_by_query(query=query, top_k=self.retrieve_top_k) def _generate_context_message(self, trajectory: Trajectory, @@ -35,6 +37,3 @@ class SimpleContextGenerator(BaseContextGenerator): content += f"- {node.content} {experience}\n" return ContextMessage(content=content.strip()) - - -CONTEXT_GENERATOR_REGISTRY.register(SimpleContextGenerator, "simple") diff --git a/experiencemaker/module/prompt/prompt_mixin.py b/experiencemaker/module/prompt/prompt_mixin.py index 56fef91e..eab0cec4 100644 --- a/experiencemaker/module/prompt/prompt_mixin.py +++ b/experiencemaker/module/prompt/prompt_mixin.py @@ -1,6 +1,7 @@ from pathlib import Path import yaml +from langchain_core.prompts import load_prompt from loguru import logger from pydantic import BaseModel, Field, model_validator @@ -23,13 +24,13 @@ class PromptMixin(BaseModel): else: with self.prompt_file_path.open("r") as f: - for k, v in yaml.load(f, yaml.FullLoader): + load_prompt_dict = yaml.load(f, yaml.FullLoader) + for k, v in load_prompt_dict.items(): if k not in self.prompt_dict: self.prompt_dict[k] = v logger.info(f"add prompt_dict key={k}") else: logger.warning(f"key={k} is already exists in prompt_dict!") - return self def prompt_format(self, prompt_name: str, **kwargs): diff --git a/experiencemaker/module/summarizer/__init__.py b/experiencemaker/module/summarizer/__init__.py index e69de29b..d5974760 100644 --- a/experiencemaker/module/summarizer/__init__.py +++ b/experiencemaker/module/summarizer/__init__.py @@ -0,0 +1,6 @@ +from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer +from experiencemaker.module.summarizer.simple_summarizer import SimpleSummarizer +from experiencemaker.utils.registry import Registry + +SUMMARIZER_REGISTRY = Registry[BaseSummarizer]("summarizer") +SUMMARIZER_REGISTRY.register(SimpleSummarizer, "simple") diff --git a/experiencemaker/module/summarizer/base_summarizer.py b/experiencemaker/module/summarizer/base_summarizer.py index 53ae4a4a..d668f1a9 100644 --- a/experiencemaker/module/summarizer/base_summarizer.py +++ b/experiencemaker/module/summarizer/base_summarizer.py @@ -8,7 +8,6 @@ from experiencemaker.schema.experience import Experience from experiencemaker.schema.trajectory import Trajectory from experiencemaker.schema.vector_store_node import VectorStoreNode from experiencemaker.storage.base_vector_store import BaseVectorStore -from experiencemaker.utils.registry import Registry class BaseSummarizer(BaseModel, ABC): @@ -28,6 +27,3 @@ class BaseSummarizer(BaseModel, ABC): if return_experience: return experiences return [] - - -SUMMARIZER_REGISTRY = Registry[BaseSummarizer]("summarizer") diff --git a/experiencemaker/module/summarizer/simple_summarizer.py b/experiencemaker/module/summarizer/simple_summarizer.py index 3bfb820f..c53b5eab 100644 --- a/experiencemaker/module/summarizer/simple_summarizer.py +++ b/experiencemaker/module/summarizer/simple_summarizer.py @@ -6,7 +6,7 @@ from pydantic import Field from experiencemaker.enumeration.role import Role from experiencemaker.module.prompt.prompt_mixin import PromptMixin -from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer, SUMMARIZER_REGISTRY +from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer from experiencemaker.schema.experience import Experience from experiencemaker.schema.trajectory import Trajectory, Message, ActionMessage from experiencemaker.utils.util_function import get_html_match_content @@ -63,6 +63,3 @@ class SimpleSummarizer(BaseSummarizer, PromptMixin): if experience: experiences.append(experience) return experiences - - -SUMMARIZER_REGISTRY.register(SimpleSummarizer, "simple") diff --git a/experiencemaker/client.py b/experiencemaker/service/experience_maker_client.py similarity index 96% rename from experiencemaker/client.py rename to experiencemaker/service/experience_maker_client.py index f9490d9b..da9c4a41 100644 --- a/experiencemaker/client.py +++ b/experiencemaker/service/experience_maker_client.py @@ -5,7 +5,7 @@ from experiencemaker.schema.response import AgentWrapperResponse, ContextGenerat from experiencemaker.utils.http_client import HttpClient -class ModelServiceClient(HttpClient): +class ExperienceMakerClient(HttpClient): base_url: str = Field(default=...) def call_agent_wrapper(self, request: AgentWrapperRequest): diff --git a/experiencemaker/service/experience_maker_service.py b/experiencemaker/service/experience_maker_service.py index e2a6b98b..bfdc9dc3 100644 --- a/experiencemaker/service/experience_maker_service.py +++ b/experiencemaker/service/experience_maker_service.py @@ -1,34 +1,41 @@ +import argparse +import json from typing import List +import uvicorn +from fastapi import FastAPI from loguru import logger from pydantic import BaseModel, Field, model_validator -from experiencemaker.model.base_embedding_model import BaseEmbeddingModel, EMBEDDING_MODEL_REGISTRY -from experiencemaker.model.base_llm import BaseLLM, LLM_REGISTRY -from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AGENT_WRAPPER_REGISTRY, AgentWrapperMixin -from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator, \ - CONTEXT_GENERATOR_REGISTRY -from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer, SUMMARIZER_REGISTRY +from experiencemaker.utils.util_function import load_env_keys + +load_env_keys() + +from experiencemaker.model import LLM_REGISTRY, EMBEDDING_MODEL_REGISTRY +from experiencemaker.model.base_embedding_model import BaseEmbeddingModel +from experiencemaker.model.base_llm import BaseLLM +from experiencemaker.module.agent_wrapper import AGENT_WRAPPER_REGISTRY +from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AgentWrapperMixin +from experiencemaker.module.context_generator import CONTEXT_GENERATOR_REGISTRY +from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator +from experiencemaker.module.summarizer import SUMMARIZER_REGISTRY +from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer from experiencemaker.schema.experience import Experience from experiencemaker.schema.request import AgentWrapperRequest, ContextGeneratorRequest, SummarizerRequest from experiencemaker.schema.response import AgentWrapperResponse, ContextGeneratorResponse, SummarizerResponse from experiencemaker.schema.trajectory import Trajectory, ContextMessage -from experiencemaker.storage.base_vector_store import BaseVectorStore, VECTOR_STORE_REGISTRY +from experiencemaker.storage import VECTOR_STORE_REGISTRY +from experiencemaker.storage.base_vector_store import BaseVectorStore +from experiencemaker.utils.file_handler import FileHandler class ExperienceMakerService(BaseModel): + workspace_id: str = Field(default="") host: str = Field(default="0.0.0.0") port: int = Field(default=8001) timeout_keep_alive: int = Field(default=600000) limit_concurrency: int = Field(default=32) - llm_config: dict = Field(default_factory=dict) - embedding_model_config: dict = Field(default_factory=dict) - vector_store_config: dict = Field(default_factory=dict) - agent_wrapper_config: dict = Field(default_factory=dict) - context_generator_config: dict = Field(default_factory=dict) - summarizer_config: dict = Field(default_factory=dict) - llm: BaseLLM | None = Field(default=None) embedding_model: BaseEmbeddingModel | None = Field(default=None) vector_store: BaseVectorStore | None = Field(default=None) @@ -41,15 +48,16 @@ class ExperienceMakerService(BaseModel): backend = llm_config.pop("backend", None) assert backend is not None, "llm must have a backend like `openai_compatible`." assert backend in LLM_REGISTRY, f"llm backend={backend} not supported. " \ - f"supported={LLM_REGISTRY.registered_modules}" + f"supported={LLM_REGISTRY.registered_module_names}" llm = LLM_REGISTRY[backend](**llm_config) logger.info(f"llm is inited with backend={backend} params={llm_config}") return llm - def get_llm(self, config: dict, llm: BaseLLM = None) -> BaseLLM: + @classmethod + def get_llm(cls, config: dict, llm: BaseLLM = None) -> BaseLLM: if "llm" in config: llm_config = config.pop("llm") - llm = self.init_llm(llm_config) + llm = cls.init_llm(llm_config) elif llm is None: raise RuntimeError("llm must be provided.") return llm @@ -59,106 +67,166 @@ class ExperienceMakerService(BaseModel): backend = embedding_model_config.pop("backend", None) assert backend is not None, "embedding_model must have a backend like `openai_compatible`." assert backend in EMBEDDING_MODEL_REGISTRY, f"embedding_model backend={backend} not supported. " \ - f"supported={EMBEDDING_MODEL_REGISTRY.registered_modules}" + f"supported={EMBEDDING_MODEL_REGISTRY.registered_module_names}" embedding_model = EMBEDDING_MODEL_REGISTRY[backend](**embedding_model_config) logger.info(f"embedding_model is inited with backend={backend} params={embedding_model_config}") return embedding_model - def get_embedding_model(self, config: dict, embedding_model: BaseEmbeddingModel = None) -> BaseEmbeddingModel: + @classmethod + def get_embedding_model(cls, config: dict, embedding_model: BaseEmbeddingModel = None) -> BaseEmbeddingModel: if "embedding_model" in config: embedding_model_config = config.pop("embedding_model") - embedding_model = self.init_embedding_model(embedding_model_config) + embedding_model = cls.init_embedding_model(embedding_model_config) elif embedding_model is None: raise RuntimeError("embedding_model must be provided.") return embedding_model - def init_vector_store(self, vector_store_config: dict) -> BaseVectorStore: + @classmethod + def init_vector_store(cls, vector_store_config: dict, + embedding_model: BaseEmbeddingModel = None) -> BaseVectorStore: backend = vector_store_config.pop("backend", None) assert backend is not None, "vector_store must have a backend like `elasticsearch`." assert backend in VECTOR_STORE_REGISTRY, f"vector_store backend={backend} not supported. " \ - f"supported={VECTOR_STORE_REGISTRY.registered_modules}" - embedding_model = self.get_embedding_model(vector_store_config, embedding_model=self.embedding_model) + f"supported={VECTOR_STORE_REGISTRY.registered_module_names}" + embedding_model = cls.get_embedding_model(vector_store_config, embedding_model=embedding_model) vector_store = VECTOR_STORE_REGISTRY[backend](**vector_store_config, embedding_model=embedding_model) logger.info(f"vector_store is inited with backend={backend} params={vector_store_config}") return vector_store - def get_vector_store(self, config: dict, vector_store: BaseVectorStore = None) -> BaseVectorStore: + @classmethod + def get_vector_store(cls, config: dict, vector_store: BaseVectorStore = None, + embedding_model: BaseEmbeddingModel = None) -> BaseVectorStore: if "vector_store" in config: vector_store_config = config.pop("vector_store") - vector_store = self.init_vector_store(vector_store_config) + vector_store = cls.init_vector_store(vector_store_config, embedding_model=embedding_model) elif vector_store is None: raise RuntimeError("vector_store must be provided.") return vector_store - def init_context_generator(self, context_generator_config: dict) -> BaseContextGenerator: + @classmethod + def init_context_generator(cls, context_generator_config: dict, data: dict) -> BaseContextGenerator: backend = context_generator_config.pop("backend", None) assert backend is not None, "context_generator must have a backend like `simple`." assert backend in CONTEXT_GENERATOR_REGISTRY, f"context_generator backend={backend} not supported. " \ - f"supported={CONTEXT_GENERATOR_REGISTRY.registered_modules}" - llm = self.get_llm(context_generator_config, llm=self.llm) - vector_store = self.get_vector_store(context_generator_config, vector_store=self.vector_store) + f"supported={CONTEXT_GENERATOR_REGISTRY.registered_module_names}" + + llm = cls.get_llm(context_generator_config, llm=data.get("llm")) + vector_store = cls.get_vector_store(context_generator_config, vector_store=data.get("vector_store"), + embedding_model=data.get("embedding_model")) + context_generator: BaseContextGenerator = CONTEXT_GENERATOR_REGISTRY[backend]( - **context_generator_config, llm=llm, vector_store=vector_store) + **context_generator_config, llm=llm, vector_store=vector_store, workspace_id=data.get("workspace_id", "")) logger.info(f"context_generator is inited with backend={backend} params={context_generator_config}") return context_generator - def init_summarizer(self, summarizer_config: dict) -> BaseSummarizer: + @classmethod + def init_summarizer(cls, summarizer_config: dict, data: dict) -> BaseSummarizer: backend = summarizer_config.pop("backend", None) assert backend is not None, "summarizer must have a backend like `simple`." assert backend in SUMMARIZER_REGISTRY, f"summarizer backend={backend} not supported. " \ - f"supported={SUMMARIZER_REGISTRY.registered_modules}" - llm = self.get_llm(summarizer_config, llm=self.llm) - vector_store = self.get_vector_store(summarizer_config, vector_store=self.vector_store) - summarizer: BaseSummarizer = SUMMARIZER_REGISTRY[backend](**summarizer_config, - llm=llm, vector_store=vector_store) + f"supported={SUMMARIZER_REGISTRY.registered_module_names}" + + llm = cls.get_llm(summarizer_config, llm=data.get("llm")) + vector_store = cls.get_vector_store(summarizer_config, vector_store=data.get("vector_store"), + embedding_model=data.get("embedding_model")) + summarizer: BaseSummarizer = SUMMARIZER_REGISTRY[backend]( + **summarizer_config, llm=llm, vector_store=vector_store, workspace_id=data.get("workspace_id", "")) logger.info(f"summarizer is inited with backend={backend} params={summarizer_config}") return summarizer - def init_agent_wrapper(self, agent_wrapper_config: dict) -> AgentWrapperMixin: + @classmethod + def init_agent_wrapper(cls, agent_wrapper_config: dict, data: dict) -> AgentWrapperMixin: backend = agent_wrapper_config.pop("backend", None) assert backend is not None, "agent_wrapper must have a backend like `simple`." assert backend in AGENT_WRAPPER_REGISTRY, f"agent_wrapper backend={backend} not supported. " \ - f"supported={AGENT_WRAPPER_REGISTRY.registered_modules}" - - llm = self.get_llm(agent_wrapper_config, llm=self.llm) + f"supported={AGENT_WRAPPER_REGISTRY.registered_module_names}" + llm = cls.get_llm(agent_wrapper_config, llm=data.get("llm")) agent_wrapper: AgentWrapperMixin = AGENT_WRAPPER_REGISTRY[backend]( - **agent_wrapper_config, llm=llm, context_generator=self.context_generator) + **agent_wrapper_config, llm=llm, context_generator=data.get("context_generator"), + workspace_id=data.get("workspace_id", "")) logger.info(f"agent_wrapper is inited with backend={backend} params={agent_wrapper_config}") return agent_wrapper - @model_validator(mode="after") - def init_modules(self): - if self.llm_config: - self.llm = self.init_llm(self.llm_config) + @model_validator(mode="before") # noqa + @classmethod + def init_modules(cls, data: dict): + try: + if "llm" in data: + data["llm"] = cls.init_llm(data["llm"]) - if self.embedding_model_config: - self.embedding_model = self.init_embedding_model(self.embedding_model_config) + if "embedding_model" in data: + data["embedding_model"] = cls.init_embedding_model(data["embedding_model"]) - if self.vector_store_config: - self.vector_store = self.init_vector_store(self.vector_store_config) + if "vector_store" in data: + data["vector_store"] = cls.init_vector_store(data["vector_store"], embedding_model=data["embedding_model"]) - if self.context_generator_config: - self.context_generator = self.init_context_generator(self.context_generator_config) + if "context_generator" in data: + data["context_generator"] = cls.init_context_generator(data["context_generator"], data) - if self.summarizer_config: - self.summarizer = self.init_summarizer(self.summarizer_config) + if "summarizer" in data: + data["summarizer"] = cls.init_summarizer(data["summarizer"], data) - if self.agent_wrapper_config: - self.agent_wrapper = self.init_agent_wrapper(self.agent_wrapper_config) + if "agent_wrapper" in data: + data["agent_wrapper"] = cls.init_agent_wrapper(data["agent_wrapper"], data) + except Exception as e: + logger.exception(e.args) + return data def call_agent_wrapper(self, request: AgentWrapperRequest) -> AgentWrapperResponse: - assert self.agent_wrapper is not None, "agent_wrapper must be provided." - trajectory: Trajectory = self.agent_wrapper.execute(request.query, **request.metadata) + assert self.agent_wrapper_ is not None, "agent_wrapper must be provided." + trajectory: Trajectory = self.agent_wrapper_.execute(request.query, **request.metadata) return AgentWrapperResponse(trajectory=trajectory) def call_context_generator(self, request: ContextGeneratorRequest) -> ContextGeneratorResponse: - assert self.context_generator is not None, "context_generator must be provided." - context_msg: ContextMessage = self.context_generator.execute(request.trajectory, **request.metadata) + assert self.context_generator_ is not None, "context_generator must be provided." + context_msg: ContextMessage = self.context_generator_.execute(request.trajectory, **request.metadata) return ContextGeneratorResponse(context_msg=context_msg) def call_summarizer(self, request: SummarizerRequest) -> SummarizerResponse: - assert self.summarizer is not None, "summarizer must be provided." - experiences: List[Experience] = self.summarizer.execute(request.trajectories, request.return_experience, - **request.metadata) + assert self.summarizer_ is not None, "summarizer must be provided." + experiences: List[Experience] = self.summarizer_.execute(request.trajectories, request.return_experience, + **request.metadata) return SummarizerResponse(experiences=experiences) + + +app = FastAPI() +service: ExperienceMakerService | None = None + + +@app.post('/agent_wrapper', response_model=AgentWrapperResponse) +def call_agent_wrapper(request: AgentWrapperRequest): + return service.call_agent_wrapper(request) + + +@app.post('/context_generator', response_model=ContextGeneratorResponse) +def call_context_generator(request: ContextGeneratorRequest): + return service.call_context_generator(request) + + +@app.post('/summarizer', response_model=SummarizerResponse) +def call_summarizer(request: SummarizerRequest): + return service.call_summarizer(request) + + +if __name__ == '__main__': + parser = argparse.ArgumentParser() + parser.add_argument('--config', type=str, help='config dict') + parser.add_argument('--config_path', type=str, help='config load path') + args = parser.parse_args() + + if args.config_path: + parse_config = FileHandler(file_path=args.config_path).load() + elif args.config: + parse_config = json.loads(args.config) + else: + raise RuntimeError("both config and config_path are not specified") + + service = ExperienceMakerService(**parse_config) + uvicorn.run(app, + host=service.host, + port=service.port, + timeout_keep_alive=service.timeout_keep_alive, + limit_concurrency=service.limit_concurrency) + +# launch with: python -m experiencemaker.service.http_service diff --git a/experiencemaker/service/http_service.py b/experiencemaker/service/http_service.py deleted file mode 100644 index 155df9c0..00000000 --- a/experiencemaker/service/http_service.py +++ /dev/null @@ -1,50 +0,0 @@ -import argparse -import json - -import uvicorn -from fastapi import FastAPI - -from experiencemaker.schema.request import AgentWrapperRequest, ContextGeneratorRequest, SummarizerRequest -from experiencemaker.schema.response import AgentWrapperResponse, ContextGeneratorResponse, SummarizerResponse -from experiencemaker.service.experience_maker_service import ExperienceMakerService -from experiencemaker.utils.file_handler import FileHandler - -app = FastAPI() -service: ExperienceMakerService | None = None - -@app.post('/agent_wrapper', response_model=AgentWrapperResponse) -def call_agent_wrapper(request: AgentWrapperRequest): - return service.call_agent_wrapper(request) - - -@app.post('/context_generator', response_model=ContextGeneratorResponse) -def call_context_generator(request: ContextGeneratorRequest): - return service.call_context_generator(request) - - -@app.post('/summarizer', response_model=SummarizerResponse) -def call_summarizer(request: SummarizerRequest): - return service.call_summarizer(request) - - -# launch with: python -m experiencemaker.service.http_service - -if __name__ == '__main__': - parser = argparse.ArgumentParser() - parser.add_argument('--config', type=str, help='config dict') - parser.add_argument('--config_path', type=str, help='config load path') - args = parser.parse_args() - - if args.config_path: - config = FileHandler(file_path=args.config_path).load() - elif args.config: - config = json.loads(args.config) - else: - raise RuntimeError("both config and config_path are not specified") - - service = ExperienceMakerService(**config) - uvicorn.run(app, - host=service.host, - port=service.port, - timeout_keep_alive=service.timeout_keep_alive, - limit_concurrency=service.limit_concurrency) diff --git a/experiencemaker/storage/__init__.py b/experiencemaker/storage/__init__.py index e69de29b..eb979b19 100644 --- a/experiencemaker/storage/__init__.py +++ b/experiencemaker/storage/__init__.py @@ -0,0 +1,8 @@ +from experiencemaker.storage.base_vector_store import BaseVectorStore +from experiencemaker.storage.es_vector_store import EsVectorStore +from experiencemaker.storage.file_vector_store import FileVectorStore +from experiencemaker.utils.registry import Registry + +VECTOR_STORE_REGISTRY = Registry[BaseVectorStore]("vector_store") +VECTOR_STORE_REGISTRY.register(EsVectorStore, "elasticsearch") +VECTOR_STORE_REGISTRY.register(FileVectorStore, "local_file") diff --git a/experiencemaker/storage/base_vector_store.py b/experiencemaker/storage/base_vector_store.py index e4644aff..d858be73 100644 --- a/experiencemaker/storage/base_vector_store.py +++ b/experiencemaker/storage/base_vector_store.py @@ -5,7 +5,6 @@ from pydantic import BaseModel, Field from experiencemaker.model.base_embedding_model import BaseEmbeddingModel from experiencemaker.schema.vector_store_node import VectorStoreNode -from experiencemaker.utils.registry import Registry class BaseVectorStore(BaseModel, ABC): @@ -27,4 +26,3 @@ class BaseVectorStore(BaseModel, ABC): raise NotImplementedError -VECTOR_STORE_REGISTRY = Registry[BaseVectorStore]("vector_store") diff --git a/experiencemaker/storage/es_vector_store.py b/experiencemaker/storage/es_vector_store.py index 1e366057..3dcba383 100644 --- a/experiencemaker/storage/es_vector_store.py +++ b/experiencemaker/storage/es_vector_store.py @@ -8,7 +8,7 @@ from pydantic import Field, PrivateAttr, model_validator from experiencemaker.model.openai_compatible_embedding_model import OpenAICompatibleEmbeddingModel from experiencemaker.schema.vector_store_node import VectorStoreNode -from experiencemaker.storage.base_vector_store import BaseVectorStore, VECTOR_STORE_REGISTRY +from experiencemaker.storage.base_vector_store import BaseVectorStore class EsVectorStore(BaseVectorStore): @@ -170,9 +170,6 @@ class EsVectorStore(BaseVectorStore): return nodes -VECTOR_STORE_REGISTRY.register(EsVectorStore, "elasticsearch") - - def main(): from experiencemaker.utils.util_function import load_env_keys load_env_keys() diff --git a/experiencemaker/storage/file_vector_store.py b/experiencemaker/storage/file_vector_store.py index 41d0ca31..d06c91a6 100644 --- a/experiencemaker/storage/file_vector_store.py +++ b/experiencemaker/storage/file_vector_store.py @@ -9,7 +9,7 @@ from pydantic import Field, model_validator, PrivateAttr from experiencemaker.model.openai_compatible_embedding_model import OpenAICompatibleEmbeddingModel from experiencemaker.schema.vector_store_node import VectorStoreNode -from experiencemaker.storage.base_vector_store import BaseVectorStore, VECTOR_STORE_REGISTRY +from experiencemaker.storage.base_vector_store import BaseVectorStore class FileVectorStore(BaseVectorStore): @@ -133,9 +133,6 @@ class FileVectorStore(BaseVectorStore): return nodes[:top_k] -VECTOR_STORE_REGISTRY.register(FileVectorStore, "local_file") - - def main(): from experiencemaker.utils.util_function import load_env_keys load_env_keys() diff --git a/experiencemaker/tool/__init__.py b/experiencemaker/tool/__init__.py index e9084df5..fa9158a0 100644 --- a/experiencemaker/tool/__init__.py +++ b/experiencemaker/tool/__init__.py @@ -1,13 +1,12 @@ +from experiencemaker.tool.base_tool import BaseTool from experiencemaker.tool.code_tool import CodeTool from experiencemaker.tool.dashscope_search_tool import DashscopeSearchTool from experiencemaker.tool.terminate_tool import TerminateTool -# from experiencemaker.tool.python_tools.code_tool import CodeTool -# from experiencemaker.tool.python_tools.dashscope_search_tool import DashscopeSearchTool -# from experiencemaker.tool.python_tools.terminate_tool import TerminateTool -# from experiencemaker.utils.registry import Registry +from experiencemaker.utils.registry import Registry -# TOOL_REGISTRY = Registry("tools") -# TOOL_REGISTRY.register(CodeTool) -# TOOL_REGISTRY.register(DashscopeSearchTool) -# TOOL_REGISTRY.register(TerminateTool) +TOOL_REGISTRY = Registry[BaseTool]("tool") + +TOOL_REGISTRY.register(CodeTool, "code") +TOOL_REGISTRY.register(DashscopeSearchTool, "web_search") +TOOL_REGISTRY.register(TerminateTool, "terminate") diff --git a/experiencemaker/tool/base_tool.py b/experiencemaker/tool/base_tool.py index 6c9a1b1b..4b3ffae5 100644 --- a/experiencemaker/tool/base_tool.py +++ b/experiencemaker/tool/base_tool.py @@ -3,8 +3,6 @@ from abc import ABC from loguru import logger from pydantic import BaseModel, Field -from experiencemaker.utils.registry import Registry - class BaseTool(BaseModel, ABC): tool_id: str = Field(default="") @@ -79,6 +77,3 @@ class BaseTool(BaseModel, ABC): def get_cache_id(self, **kwargs) -> str: raise NotImplementedError - - -TOOL_REGISTRY = Registry[BaseTool]("tool") diff --git a/experiencemaker/tool/code_tool.py b/experiencemaker/tool/code_tool.py index 99d5bedb..8b61b85f 100644 --- a/experiencemaker/tool/code_tool.py +++ b/experiencemaker/tool/code_tool.py @@ -1,7 +1,7 @@ import sys from io import StringIO -from experiencemaker.tool.base_tool import BaseTool, TOOL_REGISTRY +from experiencemaker.tool.base_tool import BaseTool class CodeTool(BaseTool): @@ -35,9 +35,6 @@ class CodeTool(BaseTool): return result -TOOL_REGISTRY.register(CodeTool, "code") - - if __name__ == '__main__': tool = CodeTool() print(tool.execute(code="print('Hello World')")) diff --git a/experiencemaker/tool/dashscope_search_tool.py b/experiencemaker/tool/dashscope_search_tool.py index fbfdf0ba..8048411e 100644 --- a/experiencemaker/tool/dashscope_search_tool.py +++ b/experiencemaker/tool/dashscope_search_tool.py @@ -6,7 +6,7 @@ from dashscope.api_entities.dashscope_response import Message from loguru import logger from pydantic import Field -from experiencemaker.tool.base_tool import BaseTool, TOOL_REGISTRY +from experiencemaker.tool.base_tool import BaseTool class DashscopeSearchTool(BaseTool): @@ -140,10 +140,6 @@ Extract the original content related to the user's question directly from the co else: return result - -TOOL_REGISTRY.register(DashscopeSearchTool, "web_search") - - def main(): from experiencemaker.utils.util_function import load_env_keys load_env_keys() diff --git a/experiencemaker/tool/terminate_tool.py b/experiencemaker/tool/terminate_tool.py index 07bc6c13..fb2412e7 100644 --- a/experiencemaker/tool/terminate_tool.py +++ b/experiencemaker/tool/terminate_tool.py @@ -1,4 +1,4 @@ -from experiencemaker.tool.base_tool import BaseTool, TOOL_REGISTRY +from experiencemaker.tool.base_tool import BaseTool class TerminateTool(BaseTool): @@ -19,6 +19,3 @@ class TerminateTool(BaseTool): def execute(self, status: str): self.success = status in ["success", "failure"] return f"The interaction has been completed with status: {status}" - - -TOOL_REGISTRY.register(TerminateTool, "terminate") diff --git a/experiencemaker/utils/registry.py b/experiencemaker/utils/registry.py index da9e8634..909102fa 100644 --- a/experiencemaker/utils/registry.py +++ b/experiencemaker/utils/registry.py @@ -12,7 +12,7 @@ class Registry(Generic[T]): self.module_dict: Dict[str, T] = {} @property - def registered_modules(self) -> List[str]: + def registered_module_names(self) -> List[str]: return sorted(self.module_dict.keys()) def register(self, module: T, module_name: str = None): @@ -38,3 +38,6 @@ class Registry(Generic[T]): def __getitem__(self, module_name: str) -> T: assert module_name in self.module_dict, f"{module_name} not found in {self.name}" return self.module_dict[module_name] + + def __contains__(self, module_name: str): + return module_name in self.module_dict \ No newline at end of file