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}>(.*?){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