This commit is contained in:
jinli.yl 2025-06-10 12:28:54 +08:00
parent 7192619070
commit a76cde5f41
28 changed files with 203 additions and 194 deletions

View file

@ -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")

View file

@ -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")

View file

@ -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")

View file

@ -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()

View file

@ -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<error>{chunk}</error>", 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

View file

@ -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")

View file

@ -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")

View file

@ -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")

View file

@ -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")

View file

@ -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")

View file

@ -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")

View file

@ -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):

View file

@ -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")

View file

@ -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")

View file

@ -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")

View file

@ -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):

View file

@ -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

View file

@ -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)

View file

@ -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")

View file

@ -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")

View file

@ -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()

View file

@ -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()

View file

@ -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")

View file

@ -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")

View file

@ -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')"))

View file

@ -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()

View file

@ -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")

View file

@ -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