From 79a247e11411c0ccc7bc7adc7c9b1dbf31959d46 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Mon, 9 Jun 2025 20:54:25 +0800 Subject: [PATCH] add unit test --- .gitignore | 2 +- experiencemaker/config/__init__.py | 5 -- .../config/agent_wrapper/simple.json | 1 - experiencemaker/config/config_handler.py | 32 ------- .../config/context_generator/simple.json | 1 - experiencemaker/config/summarizer/simple.json | 1 - .../openai_compatible_embedding_model.py | 19 +++- .../model/openai_compatible_llm.py | 22 +++++ experiencemaker/module/base_module.py | 48 ---------- .../module/environment/base_environment.py | 18 +--- .../module/reward_fn/base_reward_fn.py | 7 +- .../reward_fn/simple_compare_reward_fn.py | 61 +++++++++++++ .../simple_compare_reward_fn_prompt.yaml | 30 +++++++ experiencemaker/module/runner/base_runner.py | 10 +-- .../module/trainner/base_trainner.py | 17 +--- experiencemaker/schema/module_loader.py | 19 ---- experiencemaker/service.py | 10 +-- experiencemaker/storage/es_vector_store.py | 72 +++++++++++++++ experiencemaker/storage/file_vector_store.py | 56 ++++++++++++ experiencemaker/tool/base_tool.py | 5 ++ experiencemaker/tool/code_tool.py | 4 +- experiencemaker/tool/dashscope_search_tool.py | 10 ++- experiencemaker/tool/mcp_tool.py | 89 ++++++++----------- experiencemaker/tool/terminate_tool.py | 3 +- experiencemaker/utils/http_client.py | 6 +- experiencemaker/utils/logger.py | 11 --- experiencemaker/utils/prompt_handler.py | 72 --------------- experiencemaker/utils/test_key.py | 33 ------- experiencemaker/utils/trajectory_utils.py | 18 ---- experiencemaker/utils/util_function.py | 14 ++- 30 files changed, 349 insertions(+), 347 deletions(-) delete mode 100644 experiencemaker/config/__init__.py delete mode 100644 experiencemaker/config/agent_wrapper/simple.json delete mode 100644 experiencemaker/config/config_handler.py delete mode 100644 experiencemaker/config/context_generator/simple.json delete mode 100644 experiencemaker/config/summarizer/simple.json delete mode 100644 experiencemaker/module/base_module.py create mode 100644 experiencemaker/module/reward_fn/simple_compare_reward_fn.py create mode 100644 experiencemaker/module/reward_fn/simple_compare_reward_fn_prompt.yaml delete mode 100644 experiencemaker/schema/module_loader.py delete mode 100644 experiencemaker/utils/logger.py delete mode 100644 experiencemaker/utils/prompt_handler.py delete mode 100644 experiencemaker/utils/test_key.py delete mode 100644 experiencemaker/utils/trajectory_utils.py diff --git a/.gitignore b/.gitignore index 8424f1ef..061dfd98 100644 --- a/.gitignore +++ b/.gitignore @@ -17,4 +17,4 @@ log/ .trash/ runs logs - +rag_nodes_index.jsonl diff --git a/experiencemaker/config/__init__.py b/experiencemaker/config/__init__.py deleted file mode 100644 index 8eea6cd5..00000000 --- a/experiencemaker/config/__init__.py +++ /dev/null @@ -1,5 +0,0 @@ -from experiencemaker.config.config_handler import ConfigHandler - -agent_wrapper_config = ConfigHandler(module_name="agent_wrapper") -summarizer_config = ConfigHandler(module_name="summarizer") -context_generator_config = ConfigHandler(module_name="context_generator") diff --git a/experiencemaker/config/agent_wrapper/simple.json b/experiencemaker/config/agent_wrapper/simple.json deleted file mode 100644 index 9e26dfee..00000000 --- a/experiencemaker/config/agent_wrapper/simple.json +++ /dev/null @@ -1 +0,0 @@ -{} \ No newline at end of file diff --git a/experiencemaker/config/config_handler.py b/experiencemaker/config/config_handler.py deleted file mode 100644 index 1e471de3..00000000 --- a/experiencemaker/config/config_handler.py +++ /dev/null @@ -1,32 +0,0 @@ -from pathlib import Path - -from pydantic import BaseModel, Field, model_validator - -from experiencemaker.utils.file_handler import FileHandler - - -class ConfigHandler(BaseModel): - module_name: str = Field(default=...) - config_dict: dict = Field(default_factory=dict) - - @model_validator(mode="after") - def register_config(self): - module_config_path: Path = Path(__file__).parent / self.module_name - for config_path in module_config_path.iterdir(): - config_name = config_path.stem - config = FileHandler(file_path=config_path).load() - self.config_dict[config_name] = config - return self - - def list_config_names(self): - return list(self.config_dict.keys()) - - def __getattr__(self, item): - if item in self.config_dict: - return self.config_dict[item] - return super().__getattr__(item) - - def __getitem__(self, item): - if item in self.config_dict: - return self.config_dict[item] - return super().__getitem__(item) diff --git a/experiencemaker/config/context_generator/simple.json b/experiencemaker/config/context_generator/simple.json deleted file mode 100644 index 9e26dfee..00000000 --- a/experiencemaker/config/context_generator/simple.json +++ /dev/null @@ -1 +0,0 @@ -{} \ No newline at end of file diff --git a/experiencemaker/config/summarizer/simple.json b/experiencemaker/config/summarizer/simple.json deleted file mode 100644 index 4fec8f75..00000000 --- a/experiencemaker/config/summarizer/simple.json +++ /dev/null @@ -1 +0,0 @@ -{"a": 1} \ No newline at end of file diff --git a/experiencemaker/model/openai_compatible_embedding_model.py b/experiencemaker/model/openai_compatible_embedding_model.py index bec37d61..fe3ff2de 100644 --- a/experiencemaker/model/openai_compatible_embedding_model.py +++ b/experiencemaker/model/openai_compatible_embedding_model.py @@ -10,7 +10,7 @@ from experiencemaker.model.base_embedding_model import BaseEmbeddingModel, EMBED class OpenAICompatibleEmbeddingModel(BaseEmbeddingModel): api_key: str = Field(default_factory=lambda: os.getenv("OPENAI_API_KEY"), description="api key") base_url: str = Field(default_factory=lambda: os.getenv("OPENAI_BASE_URL"), description="base url") - model_name: str = Field(default="text-embedding-v3", description="model name") + model_name: str = Field(default="text-embedding-v4", description="model name") dimensions: int = Field(default=1024, description="dimensions") encoding_format: Literal["float", "base64"] = Field(default="float", description="encoding_format") _client: OpenAI = PrivateAttr() @@ -74,3 +74,20 @@ class OpenAICompatibleEmbeddingModel(BaseEmbeddingModel): EMBEDDING_MODEL_REGISTRY.register(OpenAICompatibleEmbeddingModel, "openai_compatible") + + +def main(): + from experiencemaker.utils.util_function import load_env_keys + load_env_keys() + + model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4") + res1 = model.get_embeddings( + "The clothes are of good quality and look good, definitely worth the wait. I love them.") + res2 = model.get_embeddings(["aa", "bb"]) + print(res1) + print(res2) + + +if __name__ == "__main__": + main() + # launch with: python -m experiencemaker.model.openai_compatible_embedding_model diff --git a/experiencemaker/model/openai_compatible_llm.py b/experiencemaker/model/openai_compatible_llm.py index d4f5d0a3..09fedef2 100644 --- a/experiencemaker/model/openai_compatible_llm.py +++ b/experiencemaker/model/openai_compatible_llm.py @@ -153,3 +153,25 @@ class OpenAICompatibleBaseLLM(BaseLLM): 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 + from experiencemaker.tool.code_tool import CodeTool + from experiencemaker.enumeration.role import Role + + load_env_keys() + model_name = "qwen-max-2025-01-25" + # model_name = "qwen3-32b" + llm = OpenAICompatibleBaseLLM(model_name=model_name) + tools: List[BaseTool] = [DashscopeSearchTool(), CodeTool()] + + llm.stream_print([Message(role=Role.USER, content="hello")], []) + print("=" * 20) + llm.stream_print([Message(role=Role.USER, content="What's the weather like in Beijing today?")], tools) + + +if __name__ == "__main__": + main() + # launch with: python -m experiencemaker.model.openai_compatible_llm diff --git a/experiencemaker/module/base_module.py b/experiencemaker/module/base_module.py deleted file mode 100644 index e45efcba..00000000 --- a/experiencemaker/module/base_module.py +++ /dev/null @@ -1,48 +0,0 @@ -from abc import ABC - -from loguru import logger -from pydantic import BaseModel, Field, model_validator - -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.utils.prompt_handler import PromptHandler - - -class BaseModule(BaseModel, ABC): - prompt_dir: str | None = Field(default=None) - prompt_file: str | None = Field(default=None) - - prompt_handler: PromptHandler | None = Field(default=None) - llm: BaseLLM | None = Field(default=None) - embedding_model: BaseEmbeddingModel | None = Field(default=None) - - @model_validator(mode="before") # noqa - @classmethod - def init_model(cls, data: dict): - if "llm" in data and isinstance(data["llm"], dict): - backend = data["llm"].pop("backend", None) - assert backend is not None, "llm must have a backend" - module = LLM_REGISTRY[backend] - params = data["llm"] - data["llm"] = module(**params) - logger.info(f"{cls.__name__} load llm.backend={backend} params={params}") - - if "embedding_model" in data and isinstance(data["embedding_model"], dict): - backend = data["embedding_model"].pop("backend", None) - assert backend is not None, "embedding_model must have a backend" - module = EMBEDDING_MODEL_REGISTRY[backend] - params = data["embedding_model"] - data["embedding_model"] = module(**params) - logger.info(f"{cls.__name__} load embedding_model.backend={backend} params={params}") - - if "prompt_dir" in data: - handler = PromptHandler(dir_path=data.get("prompt_dir")) - data["prompt_handler"] = handler - if data.get("prompt_file"): - handler.add_prompt_file(data.get("prompt_file")) - - return data - - def execute(self, **kwargs): - raise NotImplementedError diff --git a/experiencemaker/module/environment/base_environment.py b/experiencemaker/module/environment/base_environment.py index a4efd655..cca3f4b2 100644 --- a/experiencemaker/module/environment/base_environment.py +++ b/experiencemaker/module/environment/base_environment.py @@ -1,15 +1,15 @@ +from abc import ABC from typing import List -from pydantic import Field +from pydantic import Field, BaseModel -from experiencemaker.module.base_module import BaseModule from experiencemaker.module.reward_fn.base_reward_fn import BaseRewardFn from experiencemaker.schema.reward import Reward from experiencemaker.schema.trajectory import StateMessage, ActionMessage, ToolCall from experiencemaker.tool.base_tool import BaseTool -class BaseEnvironment(BaseModule): +class BaseEnvironment(BaseModel, ABC): tools: List[BaseTool] = Field(default_factory=list) reward_fns: List[BaseRewardFn] = Field(default_factory=list) current_state: StateMessage = Field(default_factory=StateMessage) @@ -54,15 +54,3 @@ class BaseEnvironment(BaseModule): def build_info(self, **kwargs): return {} - - def get_tool_info(self,tool_name): - tool_dict = {tool.name: tool for tool in self.tools} - if tool_name in tool_dict: - - return f'tool \'{tool_name}\' description is: {tool_dict[tool_name].description}\t' + f'parameters: {str(tool_dict[tool_name].input_schema)}' - else: - return '' - - def get_tools_info(self): - tool_dict = {tool.name: tool for tool in self.tools} - return {tool_name:self.get_tool_info(tool_name=tool_name) for tool_name in tool_dict} \ No newline at end of file diff --git a/experiencemaker/module/reward_fn/base_reward_fn.py b/experiencemaker/module/reward_fn/base_reward_fn.py index 84cb0d5f..000f380e 100644 --- a/experiencemaker/module/reward_fn/base_reward_fn.py +++ b/experiencemaker/module/reward_fn/base_reward_fn.py @@ -1,11 +1,12 @@ from abc import ABC -from experiencemaker.module.base_module import BaseModule +from pydantic import BaseModel + from experiencemaker.schema.reward import Reward from experiencemaker.schema.trajectory import Trajectory -class BaseRewardFn(BaseModule, ABC): +class BaseRewardFn(BaseModel, ABC): - def execute(self, trajectory: Trajectory, ground_truth=None, **kwargs) -> Reward: + def execute(self, trajectory: Trajectory = None, ground_truth=None, **kwargs) -> Reward: raise NotImplementedError diff --git a/experiencemaker/module/reward_fn/simple_compare_reward_fn.py b/experiencemaker/module/reward_fn/simple_compare_reward_fn.py new file mode 100644 index 00000000..7c77af94 --- /dev/null +++ b/experiencemaker/module/reward_fn/simple_compare_reward_fn.py @@ -0,0 +1,61 @@ +import datetime +from pathlib import Path + +from loguru import logger +from pydantic import Field + +from experiencemaker.model.base_llm import BaseLLM +from experiencemaker.module.prompt.prompt_mixin import PromptMixin +from experiencemaker.module.reward_fn.base_reward_fn import BaseRewardFn +from experiencemaker.schema.reward import Reward +from experiencemaker.schema.trajectory import Trajectory, Message +from experiencemaker.utils.util_function import get_html_match_content + + +class SimpleCompareRewardFn(BaseRewardFn, PromptMixin): + llm: BaseLLM | None = Field(default=None) + eval_times: int = Field(default=5) + prompt_file_path: Path = Path(__file__).parent / "simple_compare_reward_fn_prompt.yaml" + + def execute(self, trajectory: Trajectory = None, comp_traj: Trajectory = None, **kwargs) -> Reward: + query = trajectory.query + answer1 = trajectory.answer + answer2 = comp_traj.answer + logger.info("=" * 10 + f"answer1\n{answer1}\n" + "=" * 10 + f"answer2\n{answer2}\n") + + valid_cnt = 0 + better_cnt = 0 + for i in range(self.eval_times): + now_time = datetime.datetime.now().strftime('%Y-%m-%d %H:%M:%S') + user_prompt = self.prompt_handler.compare_prompt.format( + now_time=now_time, + query=query, + answer1=answer1, + answer2=answer2) + + messages = [Message(content=user_prompt)] + action_msg = self.llm.chat(messages=messages) + + rule: str = get_html_match_content(action_msg.content, "rule") + rule_based_comparison: str = get_html_match_content(action_msg.content, "rule_based_comparison") + result: str | None = get_html_match_content(action_msg.content, "result") + logger.info(f"round.{i} rule={rule} rule_based_comparison={rule_based_comparison} result={result}") + if result: + result = result.lower() + if "plan1" in result and "plan2" in result: + logger.warning(f"both plan exists in result={result}") + + elif "plan1" in result: + valid_cnt += 1 + better_cnt += 1 + + elif "plan2" in result: + valid_cnt += 1 + + else: + logger.warning(f"no plan exists in result={result}") + + reward_value = 0 + if valid_cnt > 0: + reward_value = better_cnt / valid_cnt + return Reward(reward_value=reward_value) diff --git a/experiencemaker/module/reward_fn/simple_compare_reward_fn_prompt.yaml b/experiencemaker/module/reward_fn/simple_compare_reward_fn_prompt.yaml new file mode 100644 index 00000000..24e57a3f --- /dev/null +++ b/experiencemaker/module/reward_fn/simple_compare_reward_fn_prompt.yaml @@ -0,0 +1,30 @@ +compare_prompt: | + # Role + You are a helpful assistant named BeyondAgent. + current time: {now_time} + + # User Question + {query} + + # Plan1 Answer + {answer1} + + # Plan2 Answer + {answer2} + + # Task + Based on the **User Question**, determine which plan provides a better answer. Response steps: + 1. Consider which comparison rules apply, the rules for comparison can be macro-level dimensions or micro-level details. + 2. Conduct a step-by-step comparison according to the rules. + 3. Combine all results to arrive at a final answer. + + # Output Example + + List the rules that could be used for comparison... + + + Conduct a step-by-step comparison according to the rules... + + + Output only the name of the better plan, either **Plan1** or **Plan2**. + \ No newline at end of file diff --git a/experiencemaker/module/runner/base_runner.py b/experiencemaker/module/runner/base_runner.py index 03f2c552..f5f80a76 100644 --- a/experiencemaker/module/runner/base_runner.py +++ b/experiencemaker/module/runner/base_runner.py @@ -18,12 +18,10 @@ class BaseRunner(BaseModel): def reset(self): self.traj_buffer.clear() + self.env.reset() - def rollout_trajectory(self, user_query: str, **kwargs): + def rollout_trajectory(self, query: str, **kwargs): raise NotImplementedError - def summary(self): - raise NotImplementedError - - def start_backend_summary(self): - raise NotImplementedError + def summary(self, **kwargs): + raise NotImplementedError \ No newline at end of file diff --git a/experiencemaker/module/trainner/base_trainner.py b/experiencemaker/module/trainner/base_trainner.py index 99810723..dc597b60 100644 --- a/experiencemaker/module/trainner/base_trainner.py +++ b/experiencemaker/module/trainner/base_trainner.py @@ -7,19 +7,8 @@ from experiencemaker.schema.trajectory import Trajectory class BaseTrainner(BaseModel, ABC): - """ - load model/prompt - data - env - -> 新cpt - - off policy/onpolicy - """ traj_buffer: List[Trajectory] = Field() - def __init__(self, **kwargs): - super().__init__(**kwargs) - def fit(self): return @@ -29,17 +18,17 @@ class BaseTrainner(BaseModel, ABC): class BaseContextTrainner(BaseTrainner): """ - load model/prompt/db/buffer -> 新cpt + load model/prompt/db/buffer -> new cpt """ class BaseSummaryTrainner(BaseTrainner): """ - load model/prompt/db/buffer -> 新cpt + load model/prompt/db/buffer -> new cpt """ class BasePolicyTrainner(BaseTrainner): """ - load model/prompt/db/buffer -> 新cpt + load model/prompt/db/buffer -> new cpt """ diff --git a/experiencemaker/schema/module_loader.py b/experiencemaker/schema/module_loader.py deleted file mode 100644 index 58be25d0..00000000 --- a/experiencemaker/schema/module_loader.py +++ /dev/null @@ -1,19 +0,0 @@ -from importlib import import_module - -from pydantic import BaseModel, Field - -from experiencemaker.utils.file_handler import FileHandler - - -class ModuleLoader(BaseModel): - class_path: str = Field(default=...) - class_name: str = Field(default=...) - config_path: str = Field(default="") - config: dict = Field(default_factory=dict) - - def load_from_config(self): - return getattr(import_module(self.class_path), self.class_name)(**self.config) - - def load_from_path(self): - self.config = FileHandler(file_path=self.config_path).load() - return self.load_from_config() diff --git a/experiencemaker/service.py b/experiencemaker/service.py index 8c4a93e2..24fe1113 100644 --- a/experiencemaker/service.py +++ b/experiencemaker/service.py @@ -4,17 +4,15 @@ import uvicorn from fastapi import FastAPI from loguru import logger -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.model.base_embedding_model import BaseEmbeddingModel, EMBEDDING_MODEL_REGISTRY +from experiencemaker.model.base_llm import BaseLLM, LLM_REGISTRY from experiencemaker.module.agent_wrapper.base_agent_wrapper import BaseAgentWrapper from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer from experiencemaker.schema.request import AgentWrapperRequest, ContextGeneratorRequest, SummarizerRequest from experiencemaker.schema.response import AgentWrapperResponse, ContextGeneratorResponse, SummarizerResponse from experiencemaker.schema.trajectory import ContextMessage, Trajectory, Sample -from experiencemaker.storage import VECTOR_STORE_REGISTRY -from experiencemaker.storage.base_vector_store import BaseVectorStore +from experiencemaker.storage.base_vector_store import BaseVectorStore, VECTOR_STORE_REGISTRY app = FastAPI() from pydantic import BaseModel, Field, model_validator @@ -29,7 +27,6 @@ class ExperienceMakerHttpService(BaseModel): 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) @@ -37,7 +34,6 @@ class ExperienceMakerHttpService(BaseModel): llm: BaseLLM | None = Field(default=None) embedding_model: BaseEmbeddingModel | None = Field(default=None) vector_store: BaseVectorStore | None = Field(default=None) - agent_wrapper: BaseAgentWrapper | None = Field(default=None) context_generator: BaseContextGenerator | None = Field(default=None) summarizer: BaseSummarizer | None = Field(default=None) diff --git a/experiencemaker/storage/es_vector_store.py b/experiencemaker/storage/es_vector_store.py index d3865a86..1e366057 100644 --- a/experiencemaker/storage/es_vector_store.py +++ b/experiencemaker/storage/es_vector_store.py @@ -6,6 +6,7 @@ from elasticsearch.helpers import bulk from loguru import logger 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 @@ -170,3 +171,74 @@ class EsVectorStore(BaseVectorStore): VECTOR_STORE_REGISTRY.register(EsVectorStore, "elasticsearch") + + +def main(): + from experiencemaker.utils.util_function import load_env_keys + load_env_keys() + + embedding_model = OpenAICompatibleEmbeddingModel(dimensions=1024) + index_name = "rag_nodes_index" + hosts = "http://11.160.132.46:8200" + es = EsVectorStore(hosts=hosts, embedding_model=embedding_model, index_name=index_name) + es.delete_index() + es.create_index() + + sample_nodes = [ + VectorStoreNode( + workspace_id="w1", + content="Artificial intelligence is a technology that simulates human intelligence.", + metadata={ + "node_type": "n1", + } + ), + VectorStoreNode( + workspace_id="w1", + content="AI is the future of mankind.", + metadata={ + "node_type": "n1", + } + ), + VectorStoreNode( + workspace_id="w1", + content="I want to eat fish!", + metadata={ + "node_type": "n2", + } + ), + VectorStoreNode( + workspace_id="w2", + content="The bigger the storm, the more expensive the fish.", + metadata={ + "node_type": "n1", + } + ), + ] + + es.insert(sample_nodes, refresh_index=True) + + logger.info("=" * 20) + results = es.add_term_filter(key="workspace_id", value="w1") \ + .add_term_filter(key="metadata.node_type", value="n1") \ + .retrieve_by_query("What is AI?", top_k=5) + for r in results: + logger.info(r.model_dump(exclude={"vector"})) + logger.info("=" * 20) + + logger.info("=" * 20) + results = es.add_term_filter(key="workspace_id", value="w1") \ + .retrieve_by_query("What is AI?", top_k=5) + for r in results: + logger.info(r.model_dump(exclude={"vector"})) + logger.info("=" * 20) + + logger.info("=" * 20) + results = es.retrieve_by_query("What is AI?", top_k=5) + for r in results: + logger.info(r.model_dump(exclude={"vector"})) + logger.info("=" * 20) + + +if __name__ == "__main__": + main() + # launch with: python -m experiencemaker.storage.es_vector_store diff --git a/experiencemaker/storage/file_vector_store.py b/experiencemaker/storage/file_vector_store.py index 3fb3c8e9..41d0ca31 100644 --- a/experiencemaker/storage/file_vector_store.py +++ b/experiencemaker/storage/file_vector_store.py @@ -7,6 +7,7 @@ from typing import List, Any from loguru import logger 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 @@ -133,3 +134,58 @@ class FileVectorStore(BaseVectorStore): VECTOR_STORE_REGISTRY.register(FileVectorStore, "local_file") + + +def main(): + from experiencemaker.utils.util_function import load_env_keys + load_env_keys() + + embedding_model = OpenAICompatibleEmbeddingModel(dimensions=1024) + index_name = "rag_nodes_index" + client = FileVectorStore(embedding_model=embedding_model, index_name=index_name) + client.delete_index() + client.create_index() + + sample_nodes = [ + VectorStoreNode( + workspace_id="w1", + content="Artificial intelligence is a technology that simulates human intelligence.", + metadata={ + "node_type": "n1", + } + ), + VectorStoreNode( + workspace_id="w1", + content="AI is the future of mankind.", + metadata={ + "node_type": "n1", + } + ), + VectorStoreNode( + workspace_id="w1", + content="I want to eat fish!", + metadata={ + "node_type": "n2", + } + ), + VectorStoreNode( + workspace_id="w2", + content="The bigger the storm, the more expensive the fish.", + metadata={ + "node_type": "n1", + } + ), + ] + + client.insert(sample_nodes) + + logger.info("=" * 20) + results = client.retrieve_by_query("What is AI?", top_k=5) + for r in results: + logger.info(r.model_dump(exclude={"vector"})) + logger.info("=" * 20) + + +if __name__ == "__main__": + main() + # launch with: python -m experiencemaker.storage.file_vector_store diff --git a/experiencemaker/tool/base_tool.py b/experiencemaker/tool/base_tool.py index 4b3ffae5..6c9a1b1b 100644 --- a/experiencemaker/tool/base_tool.py +++ b/experiencemaker/tool/base_tool.py @@ -3,6 +3,8 @@ 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="") @@ -77,3 +79,6 @@ 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 c63c4250..99d5bedb 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 +from experiencemaker.tool.base_tool import BaseTool, TOOL_REGISTRY class CodeTool(BaseTool): @@ -35,6 +35,8 @@ class CodeTool(BaseTool): return result +TOOL_REGISTRY.register(CodeTool, "code") + if __name__ == '__main__': tool = CodeTool() diff --git a/experiencemaker/tool/dashscope_search_tool.py b/experiencemaker/tool/dashscope_search_tool.py index 1492a2aa..fbfdf0ba 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 +from experiencemaker.tool.base_tool import BaseTool, TOOL_REGISTRY class DashscopeSearchTool(BaseTool): @@ -140,9 +140,13 @@ 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.test_key import set_key - set_key() + from experiencemaker.utils.util_function import load_env_keys + load_env_keys() query = "What is artificial intelligence?" tool = DashscopeSearchTool(stream_print=True) diff --git a/experiencemaker/tool/mcp_tool.py b/experiencemaker/tool/mcp_tool.py index dd175c0a..1ffdbd43 100644 --- a/experiencemaker/tool/mcp_tool.py +++ b/experiencemaker/tool/mcp_tool.py @@ -1,11 +1,10 @@ import asyncio -from typing import List, Optional +from typing import List -from loguru import logger from mcp import ClientSession from mcp.client.sse import sse_client -from pydantic import Field +from pydantic import Field, model_validator from experiencemaker.tool.base_tool import BaseTool @@ -14,35 +13,11 @@ class MCPTool(BaseTool): server_url: str = Field(..., description="MCP server URL") tool_name_list: List[str] = Field(default_factory=list) cache_tools: dict = Field(default_factory=dict, alias="cache_tools") - cache_tools_info: Optional[dict] = Field(default=None, alias="cache_tools_info") - class Config: - underscore_attrs_are_private = True - - def __init__(self, **data): - super().__init__(**data) + @model_validator(mode="after") + def refresh_tools(self): self.refresh() - - def get_tool_name_list(self) -> List[str]: - return self.tool_name_list - - def get_server_info(self): - return self.cache_tools_info - - def refresh(self): - self.cache_tools.clear() - self.tool_name_list.clear() - - if "sse" in self.server_url: - original_tool_list = asyncio.run(self._get_tools()) - self.cache_tools_info = original_tool_list.tools - - for tool in self.cache_tools_info: - self.cache_tools[tool.name] = tool - self.tool_name_list.append(tool.name) - else: - # TODO: Implement non-SSE refresh logic - logger.warning("Non-SSE refresh not implemented yet") + return self async def _get_tools(self): async with sse_client(url=self.server_url) as streams: @@ -51,41 +26,51 @@ class MCPTool(BaseTool): tools = await session.list_tools() return tools - def input_schema(self, tool_name: str) -> dict: - return self.cache_tools.get(tool_name, {}).inputSchema + def refresh(self): + self.tool_name_list.clear() + self.cache_tools.clear() - def output_schema(self, tool_name: str) -> dict: - # TODO: Implement output schema logic - return {} + if "sse" in self.server_url: + original_tool_list = asyncio.run(self._get_tools()) + for tool in original_tool_list.tools: + self.cache_tools[tool.name] = tool + self.tool_name_list.append(tool.name) + else: + raise NotImplementedError("Non-SSE refresh not implemented yet") + + @property + def input_schema(self) -> dict: + return {x: self.cache_tools[x].inputSchema for x in self.cache_tools} + + @property + def output_schema(self) -> dict: + raise NotImplementedError("Output schema not implemented yet") def get_tool_description(self, tool_name: str, schema: bool = False) -> str: + if tool_name not in self.cache_tools: + raise RuntimeError(f"Tool {tool_name} not found") + tool = self.cache_tools.get(tool_name) - if not tool: - return "" - - description = f'tool \'{tool_name}\' description is:'+ tool.description + description = f"tool={tool_name} description={tool.description}\n" if schema: - description += f"\nInput Schema: {self.input_schema(tool_name)}" - description += f"\nOutput Schema: {self.output_schema(tool_name)}" - return description - - async def _execute(self, **kwargs): - tool_name = kwargs.get('tool_name') - args = kwargs.get('args', {}) + description += f"input_schema={self.input_schema[tool_name]}\n" \ + f"output_schema={self.output_schema[tool_name]}\n" + return description.strip() + async def async_execute(self, tool_name: str, **kwargs): if "sse" in self.server_url: async with sse_client(url=self.server_url) as streams: async with ClientSession(streams[0], streams[1]) as session: await session.initialize() - results = await session.call_tool(tool_name, args) + results = await session.call_tool(tool_name, kwargs) return results.content[0].text, results.isError - else: - return "Failed to connect to the tool", False - def execute(self, **kwargs): - return asyncio.run(self._execute(**kwargs)) + else: + raise NotImplementedError("Non-SSE execute not implemented yet") + + def _execute(self, **kwargs): + return asyncio.run(self.async_execute(**kwargs)) def get_cache_id(self, **kwargs) -> str: # Implement a method to generate a unique cache ID based on the input return f"{kwargs.get('tool_name')}_{hash(frozenset(kwargs.get('args', {}).items()))}" - diff --git a/experiencemaker/tool/terminate_tool.py b/experiencemaker/tool/terminate_tool.py index 765835c9..07bc6c13 100644 --- a/experiencemaker/tool/terminate_tool.py +++ b/experiencemaker/tool/terminate_tool.py @@ -1,4 +1,4 @@ -from experiencemaker.tool.base_tool import BaseTool +from experiencemaker.tool.base_tool import BaseTool, TOOL_REGISTRY class TerminateTool(BaseTool): @@ -21,3 +21,4 @@ class TerminateTool(BaseTool): return f"The interaction has been completed with status: {status}" +TOOL_REGISTRY.register(TerminateTool, "terminate") diff --git a/experiencemaker/utils/http_client.py b/experiencemaker/utils/http_client.py index b5eb2d22..feadb053 100644 --- a/experiencemaker/utils/http_client.py +++ b/experiencemaker/utils/http_client.py @@ -111,6 +111,8 @@ class HttpClient(BaseModel): retry_sleep_time *= self.retry_time_multiplier time.sleep(retry_sleep_time) + return None + def request_stream(self, data: str = None, json_data: dict = None, @@ -137,7 +139,7 @@ class HttpClient(BaseModel): http_enum=http_enum, **kwargs) - return + return None except Exception as e: logger.exception(f"{self.__class__.__name__} {i}th request failed with args={e.args}") @@ -150,3 +152,5 @@ class HttpClient(BaseModel): retry_sleep_time *= self.retry_time_multiplier time.sleep(retry_sleep_time) + + return None diff --git a/experiencemaker/utils/logger.py b/experiencemaker/utils/logger.py deleted file mode 100644 index 36c4f69f..00000000 --- a/experiencemaker/utils/logger.py +++ /dev/null @@ -1,11 +0,0 @@ -from best_logger import register_logger - -def init_logger(): - register_logger( - mods=["agent", "context", "summary"], - non_console_mods=[], - auto_clean_mods=[], - base_log_path=f"logs/default" - ) - - diff --git a/experiencemaker/utils/prompt_handler.py b/experiencemaker/utils/prompt_handler.py deleted file mode 100644 index 30c48b87..00000000 --- a/experiencemaker/utils/prompt_handler.py +++ /dev/null @@ -1,72 +0,0 @@ -import os - -import yaml -from loguru import logger -from pydantic import BaseModel, Field, model_validator - - -class PromptHandler(BaseModel): - file_path: str = Field(default="") - prompt_dict: dict = Field(default_factory=dict) - - @model_validator(mode="after") - def init_prompt_dict(self): - if self.prompt_dict: - logger.info(f"load prompt from dict, keys={self.prompt_dict.keys()}") - - if self.file_path: - self.load_file_prompt() - logger.info(f"load prompt from file_path, keys={self.prompt_dict.keys()}") - - def load_file_prompt(self): - if not os.path.exists(self.file_path): - raise RuntimeError(f"file_path={self.file_path} not exists!") - - with open(self.file_path) as f: - prompt_dict: dict = yaml.load(f, yaml.FullLoader) - self.prompt_dict.update(prompt_dict) - - def __getitem__(self, key: str): - return self.prompt_dict[key] - - def __setitem__(self, key: str, value: str): - self.prompt_dict[key] = value - - def __getattr__(self, key: str): - if key in self.prompt_dict: - return self.prompt_dict[key] - - return super().__getattr__(key) - - def prompt_format(self, prompt_name: str, **kwargs): - prompt = self.prompt_dict[prompt_name] - - flag_kwargs = {k: v for k, v in kwargs.items() if isinstance(v, bool)} - other_kwargs = {k: v for k, v in kwargs.items() if not isinstance(v, bool)} - - if flag_kwargs: - split_prompt = [] - for line in prompt.strip().split("\n"): - hit = False - hit_flag = True - for key, flag in kwargs.items(): - if not line.startswith(f"[{key}]"): - continue - - else: - hit = True - hit_flag = flag - line = line.strip(f"[{key}]") - break - - if not hit: - split_prompt.append(line) - elif hit_flag: - split_prompt.append(line) - - prompt = "\n".join(split_prompt) - - if other_kwargs: - prompt = prompt.format(**other_kwargs) - - return prompt diff --git a/experiencemaker/utils/test_key.py b/experiencemaker/utils/test_key.py deleted file mode 100644 index 86abed92..00000000 --- a/experiencemaker/utils/test_key.py +++ /dev/null @@ -1,33 +0,0 @@ -import json -import os - -from experiencemaker.schema.module_loader import ModuleLoader - - -def load_env_keys(): - if os.path.exists(".env"): - with open(".env") as f: - config = json.load(f) - for k, v in config.items(): - os.environ[k] = v - - -agent_wrapper_loader = ModuleLoader( - class_path="experiencemaker.module.agent_wrapper.naive_agent_wrapper", - class_name="NaiveAgentWrapper", - config_path="beyondagent/config/agent_wrapper/naive_agent_wrapper.json") - -context_generator_loader = ModuleLoader( - class_path="experiencemaker.module.context_generator.simple_context_generator", - class_name="SimpleContextGenerator", - config_path="beyondagent/config/context_generator/simple_context_generator.json") - -summarizer_loader = ModuleLoader( - class_path="experiencemaker.module.summarizer.simple_summarizer", - class_name="SimpleSummarizer", - config_path="beyondagent/config/summarizer/simple_summarizer.json") - -env_loader = ModuleLoader( - class_path="experiencemaker.module.environment.simple_environment", - class_name="SimpleEnvironment", - config_path="beyondagent/config/environment/simple_environment.json") \ No newline at end of file diff --git a/experiencemaker/utils/trajectory_utils.py b/experiencemaker/utils/trajectory_utils.py deleted file mode 100644 index 66ba4545..00000000 --- a/experiencemaker/utils/trajectory_utils.py +++ /dev/null @@ -1,18 +0,0 @@ -from typing import List - -from experiencemaker.schema.trajectory import Message, StateMessage, ActionMessage - - -def format_trajectory_steps(steps: List[Message]) -> str: - format_steps = [] - step_idx = 0 - single_step = [] - for idx, step in enumerate(steps): - if isinstance(step, ActionMessage): - step_idx += 1 - single_step.append(f"** STEP {step_idx} **\n{step.content}") - elif isinstance(step, StateMessage): - single_step.append(f"{step.content}") - format_steps.append("\n".join(single_step)) - single_step = [] - return "\n\n".join(format_steps) \ No newline at end of file diff --git a/experiencemaker/utils/util_function.py b/experiencemaker/utils/util_function.py index 3427cf26..e699d3ef 100644 --- a/experiencemaker/utils/util_function.py +++ b/experiencemaker/utils/util_function.py @@ -1,5 +1,7 @@ +import json +import os import re - +from loguru import logger def get_html_match_content(content: str, key: str): pattern = rf"<{key}>(.*?)" @@ -7,3 +9,13 @@ def get_html_match_content(content: str, key: str): if match: return match.group(1) return None + + +def load_env_keys(): + if os.path.exists(".env"): + with open(".env") as f: + config = json.load(f) + for k, v in config.items(): + os.environ[k] = v + else: + logger.warning(".env file not found~") \ No newline at end of file