mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-10 03:30:56 +00:00
289 lines
14 KiB
Python
289 lines
14 KiB
Python
import argparse
|
|
import copy
|
|
import json
|
|
import types
|
|
from typing import List
|
|
|
|
import uvicorn
|
|
from fastapi import FastAPI
|
|
from loguru import logger
|
|
from pydantic import BaseModel, Field, model_validator
|
|
|
|
from experiencemaker.utils.util_function import load_env_keys
|
|
|
|
load_env_keys()
|
|
|
|
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 AgentWrapperMixin, AGENT_WRAPPER_REGISTRY
|
|
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.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
|
|
|
|
|
|
class EMService(BaseModel):
|
|
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: BaseLLM | None = Field(default=None)
|
|
embedding_model: BaseEmbeddingModel | None = Field(default=None)
|
|
vector_store: BaseVectorStore | None = Field(default=None)
|
|
agent_wrapper: AgentWrapperMixin | None = Field(default=None)
|
|
context_generator: BaseContextGenerator | None = Field(default=None)
|
|
summarizer: BaseSummarizer | None = Field(default=None)
|
|
|
|
origin_config: dict = Field(default_factory=dict)
|
|
|
|
@staticmethod
|
|
def init_llm(llm_config: dict) -> BaseLLM:
|
|
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_module_names}"
|
|
llm = LLM_REGISTRY[backend](**llm_config)
|
|
logger.info(f"llm is inited with backend={backend} params={llm_config}")
|
|
return llm
|
|
|
|
@classmethod
|
|
def get_llm(cls, config: dict, llm: BaseLLM = None) -> BaseLLM:
|
|
if "llm" in config:
|
|
llm_config = config.pop("llm")
|
|
llm = cls.init_llm(llm_config)
|
|
elif llm is None:
|
|
raise RuntimeError("llm must be provided.")
|
|
return llm
|
|
|
|
@staticmethod
|
|
def init_embedding_model(embedding_model_config: dict) -> BaseEmbeddingModel:
|
|
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_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
|
|
|
|
@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 = cls.init_embedding_model(embedding_model_config)
|
|
elif embedding_model is None:
|
|
raise RuntimeError("embedding_model must be provided.")
|
|
return embedding_model
|
|
|
|
@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_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
|
|
|
|
@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 = 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
|
|
|
|
@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_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)
|
|
logger.info(f"context_generator is inited with backend={backend} params={context_generator_config}")
|
|
return context_generator
|
|
|
|
@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_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)
|
|
logger.info(f"summarizer is inited with backend={backend} params={summarizer_config}")
|
|
return summarizer
|
|
|
|
@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_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=data.get("context_generator"))
|
|
logger.info(f"agent_wrapper is inited with backend={backend} params={agent_wrapper_config}")
|
|
return agent_wrapper
|
|
|
|
@classmethod
|
|
def init_class_by_config(cls, data: dict):
|
|
try:
|
|
if "llm" in data:
|
|
data["llm"] = cls.init_llm(data["llm"])
|
|
|
|
if "embedding_model" in data:
|
|
data["embedding_model"] = cls.init_embedding_model(data["embedding_model"])
|
|
|
|
if "vector_store" in data:
|
|
data["vector_store"] = cls.init_vector_store(data["vector_store"],
|
|
embedding_model=data["embedding_model"])
|
|
|
|
if "context_generator" in data:
|
|
data["context_generator"] = cls.init_context_generator(data["context_generator"], data)
|
|
|
|
if "summarizer" in data:
|
|
data["summarizer"] = cls.init_summarizer(data["summarizer"], data)
|
|
|
|
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
|
|
|
|
@model_validator(mode="before") # noqa
|
|
@classmethod
|
|
def init_modules(cls, data: dict):
|
|
origin_config = copy.deepcopy(data)
|
|
data = cls.init_class_by_config(data)
|
|
data["origin_config"] = origin_config
|
|
return data
|
|
|
|
def call_agent_wrapper(self, request: AgentWrapperRequest) -> AgentWrapperResponse:
|
|
if "em_config" in request.metadata:
|
|
new_config = copy.deepcopy(self.origin_config)
|
|
new_config.update(request.metadata["em_config"])
|
|
data = EMService.init_class_by_config(new_config)
|
|
agent_wrapper = data["agent_wrapper"]
|
|
else:
|
|
assert self.agent_wrapper is not None, "agent_wrapper must be provided."
|
|
agent_wrapper = self.agent_wrapper
|
|
|
|
trajectory: Trajectory = agent_wrapper.execute(query=request.query,
|
|
workspace_id=request.workspace_id,
|
|
**request.metadata)
|
|
return AgentWrapperResponse(trajectory=trajectory)
|
|
|
|
def call_context_generator(self, request: ContextGeneratorRequest) -> ContextGeneratorResponse:
|
|
if "em_config" in request.metadata:
|
|
new_config = copy.deepcopy(self.origin_config)
|
|
new_config.update(request.metadata["em_config"])
|
|
data = EMService.init_class_by_config(new_config)
|
|
context_generator = data["context_generator"]
|
|
else:
|
|
assert self.context_generator is not None, "context_generator must be provided."
|
|
context_generator = self.context_generator
|
|
|
|
context_msg: ContextMessage = context_generator.execute(trajectory=request.trajectory,
|
|
workspace_id=request.workspace_id,
|
|
**request.metadata)
|
|
return ContextGeneratorResponse(context_msg=context_msg)
|
|
|
|
def call_summarizer(self, request: SummarizerRequest) -> SummarizerResponse:
|
|
if "em_config" in request.metadata:
|
|
new_config = copy.deepcopy(self.origin_config)
|
|
new_config.update(request.metadata["em_config"])
|
|
data = EMService.init_class_by_config(new_config)
|
|
summarizer = data["summarizer"]
|
|
else:
|
|
assert self.summarizer is not None, "summarizer must be provided."
|
|
summarizer = self.summarizer
|
|
|
|
experiences: List[Experience] = summarizer.execute(trajectories=request.trajectories,
|
|
workspace_id=request.workspace_id,
|
|
**request.metadata)
|
|
return SummarizerResponse(experiences=experiences)
|
|
|
|
|
|
app = FastAPI()
|
|
service: EMService | 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()
|
|
field_dict = EMService.model_fields
|
|
assert isinstance(field_dict, dict)
|
|
|
|
json_keys = []
|
|
for key, info in field_dict.items():
|
|
if info.annotation in [int, str, bool]:
|
|
parser.add_argument(f"--{key}", type=info.annotation, default=info.default)
|
|
|
|
elif isinstance(info.annotation, types.UnionType) and issubclass(info.annotation.__args__[0], BaseModel):
|
|
parser.add_argument(f"--{key}", type=str, default=None)
|
|
json_keys.append(key)
|
|
|
|
elif info.annotation in [dict, list]:
|
|
logger.warning(f"skip key={key} info.annotation={info.annotation}")
|
|
continue
|
|
|
|
else:
|
|
raise NotImplementedError(f"key={key} annotation={info.annotation} is not supported.")
|
|
|
|
args: argparse.Namespace = parser.parse_args()
|
|
service_kwargs = {k: json.loads(v) if k in json_keys else v for k, v in args.__dict__.items() if v is not None}
|
|
logger.info(f"service.kwargs={json.dumps(service_kwargs, indent=2, ensure_ascii=False)}")
|
|
|
|
service = EMService(**service_kwargs)
|
|
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.em_service \
|
|
--port=8001 \
|
|
--llm='{"backend": "openai_compatible", "model_name": "qwen3-32b", "temperature": 0.6}' \
|
|
--embedding_model='{"backend": "openai_compatible", "model_name": "text-embedding-v4", "dimensions": 1024}' \
|
|
--vector_store='{"backend": "elasticsearch"}' \
|
|
--agent_wrapper='{"backend": "simple"}' \
|
|
--context_generator='{"backend": "simple"}' \
|
|
--summarizer='{"backend": "simple"}'
|
|
"""
|