diff --git a/.gitignore b/.gitignore
index 69361c03..061dfd98 100644
--- a/.gitignore
+++ b/.gitignore
@@ -17,5 +17,4 @@ log/
.trash/
runs
logs
-alfworld_data
-beyondagent/dataset/appworld/data
+rag_nodes_index.jsonl
diff --git a/doc/quick_start.md b/doc/quick_start.md
new file mode 100644
index 00000000..ce963070
--- /dev/null
+++ b/doc/quick_start.md
@@ -0,0 +1,148 @@
+# Quick Start
+
+Here is a simple user guide.
+
+### Step0: Configuration
+- start es
+- set env APIKEY / host
+- python -m model_service -port 8000 -config simple/path
+
+### Step1: Own an agent
+
+Assume you have a runnable [agent code](./mxc_agent.py).
+Here, we use a basic LLM combined with a simple react framework including three tools(code, web_search, terminate) as an example.
+
+```python
+class MxcAgent(BaseModel):
+ llm: BaseLLM | None = Field(default=None)
+ max_steps: int = Field(default=10)
+ tools: List[BaseTool] = Field(default_factory=list)
+
+ def think(self, query: str, **kwargs) -> bool:
+ ...
+
+ def act(self, **kwargs):
+ ...
+
+ def run(self, query: str, **kwargs) -> List[Message]:
+ messages: List[Message] = []
+
+ for i in range(self.max_steps):
+ should_act: bool = self.think(query, messages=messages, **kwargs)
+ if should_act:
+ self.act(messages=messages, **kwargs)
+ else:
+ break
+ return messages
+
+
+query = "Analyze Xiaomi Corporation
+agent = MxcAgent(llm=OpenAICompatibleBaseLLM(model_name="qwen3-32b", temperature=0.0001),
+ max_steps=10,
+ tools=[CodeTool(), DashscopeSearchTool(), TerminateTool()])
+messages = agent.run(query=query)
+answer = messages[-1].content
+print(answer)
+```
+
+### Step2: Implement AgentWrapper
+
+In order to utilize the **context generator** and **summarizer** capabilities of beyondagent, please inherit from **MxcAgent** and **BaseAgentWrapperMixin** to implement the AgentWrapper.
+
+Here, you need to customize two parts:
+- how to integrate the content message(insight) generated by the `self.context_generator` into the context.
+- implement the execute function to output the trajectory.
+
+Below is a simple example of integrating **trajectory-level insight** into the context.
+
+```python
+from beyondagent.core.module.agent_wrapper.base_agent_wrapper import BaseAgentWrapperMixin
+
+class MxcAgentWrapper(MxcAgent, BaseAgentWrapperMixin):
+ def execute(self, query: str, **kwargs) -> Trajectory:
+ trajectory = Trajectory(steps=messages, query=query)
+ context_msg = self.context_generator.execute(trajectory=trajectory)
+ new_query = f"""
+previous insight:
+{context_msg.content}
+Please consider the helpful parts from these in answering the question, to make the response more comprehensive and substantial.
+
+user query:
+{query}
+ """.strip()
+
+ messages = self.run(new_query, **kwargs)
+ return Trajectory(query=query, steps=messages, answer=messages[-1].content, done=True)
+
+```
+
+### Step3: Run AgentRunner with insight
+
+Once you have completed the implementation of the AgentWrapper class, you will be able to utilize the capabilities of
+beyondagent.
+
+Here is an example using **SimpleAgentRunner**.
+We first executed two historical tasks, then summarized the experience and made it persistent.
+Finally, we utilized the historical experience in a new task.
+
+[insights demo](./insight.json)
+
+
+```python
+from beyondagent.core.module.runner.simple_agent_runner import SimpleAgentRunner
+
+
+
+mxc_agent_wrapper = MxcAgentWrapper(llm=OpenAICompatibleBaseLLM(model_name="qwen3-32b", temperature=0.0001),
+ max_steps=10,
+ tools=[CodeTool(), DashscopeSearchTool(), TerminateTool()])
+agent_runner = SimpleAgentRunner(agent_wrapper=mxc_agent_wrapper, summarizer="default", context_generator="default")
+
+# historical tasks
+agent_runner.rollout_trajectory(query="Analyze the company Tesla.")
+agent_runner.rollout_trajectory(query="Analyze the company Apple.")
+
+# summary insights and store them
+agent_runner.summary_and_store()
+
+# run agent with historical insights
+trajectory = agent_runner.rollout_trajectory(query="Analyze the company Xiaomi Corporation.")
+```
+
+
+### Step4: Evaluation(Optional)
+
+If we have a reward function that allows us to compare the performance before and after adding context, we can try this
+part.
+
+Use `run_agent` to obtain the answer from the original agent (answer1), and use `run_agent_wrapper` to get the answer
+with added insights and experience (answer2).
+
+Here, the reward function is used to compare and score the two answers. The `reward.reward_value` indicates the win rate
+of answer2.
+
+```python
+# task
+query = "Analyze Xiaomi Corporation."
+
+# run agent
+agent = MxcAgent(llm=OpenAICompatibleBaseLLM(model_name="qwen3-32b", temperature=0.0001),
+ max_steps=10,
+ tools=[CodeTool(), DashscopeSearchTool(), TerminateTool()])
+messages = agent.run(query=query)
+answer1 = messages[-1].content
+
+# agent runner: Assume we already have some historical experience.
+mxc_agent_wrapper = MxcAgentWrapper(llm=OpenAICompatibleBaseLLM(model_name="qwen3-32b", temperature=0.0001),
+ max_steps=10,
+ tools=[CodeTool(), DashscopeSearchTool(), TerminateTool()])
+agent_runner = SimpleAgentRunner(agent_wrapper=mxc_agent_wrapper, summarizer="default", context_generator="default")
+trajectory = agent_runner.rollout_trajectory(query=query)
+answer2 = trajectory.answer
+
+# pair-wise LLM evaluation
+from beyondagent.core.module.reward_fn.simple_reward_fn import SimpleRewardFn
+reward_fn = SimpleRewardFn(llm=OpenAICompatibleBaseLLM(model_name="qwen3-32b", temperature=0.0001))
+reward = reward_fn.execute(query=query, answer1=answer1, answer2=answer2, eval_times=5)
+print(f"final reward={reward.reward_value}")
+```
\ No newline at end of file
diff --git a/experiencemaker/service/model_service_client.py b/experiencemaker/client.py
similarity index 100%
rename from experiencemaker/service/model_service_client.py
rename to experiencemaker/client.py
diff --git a/experiencemaker/model/__init__.py b/experiencemaker/model/__init__.py
index 5afcdb48..e69de29b 100644
--- a/experiencemaker/model/__init__.py
+++ b/experiencemaker/model/__init__.py
@@ -1,10 +0,0 @@
-from experiencemaker.model.openai_compatible_embedding_model import OpenAICompatibleEmbeddingModel
-from experiencemaker.model.openai_compatible_llm import OpenAICompatibleBaseLLM
-
-from experiencemaker.utils.registry import Registry
-
-LLM_REGISTRY = Registry("llm")
-LLM_REGISTRY.register(OpenAICompatibleBaseLLM, "openai_compatible")
-
-EMBEDDING_MODEL_REGISTRY = Registry("embedding_model")
-EMBEDDING_MODEL_REGISTRY.register(OpenAICompatibleEmbeddingModel, "openai_compatible")
diff --git a/experiencemaker/model/base_embedding_model.py b/experiencemaker/model/base_embedding_model.py
index 3cde45c3..ef09a74e 100644
--- a/experiencemaker/model/base_embedding_model.py
+++ b/experiencemaker/model/base_embedding_model.py
@@ -5,6 +5,7 @@ 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):
@@ -84,3 +85,6 @@ class BaseEmbeddingModel(BaseModel, ABC):
else:
raise RuntimeError(f"unsupported type={type(nodes)}")
+
+
+EMBEDDING_MODEL_REGISTRY = Registry[BaseEmbeddingModel]("embedding_model")
diff --git a/experiencemaker/model/base_llm.py b/experiencemaker/model/base_llm.py
index f78dc6e2..a055e6f9 100644
--- a/experiencemaker/model/base_llm.py
+++ b/experiencemaker/model/base_llm.py
@@ -6,6 +6,7 @@ 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):
@@ -108,3 +109,6 @@ class BaseLLM(BaseModel, ABC):
raise e
return None
+
+
+LLM_REGISTRY = Registry[BaseLLM]("llm")
diff --git a/experiencemaker/model/openai_compatible_embedding_model.py b/experiencemaker/model/openai_compatible_embedding_model.py
index 79eedee6..fe3ff2de 100644
--- a/experiencemaker/model/openai_compatible_embedding_model.py
+++ b/experiencemaker/model/openai_compatible_embedding_model.py
@@ -4,13 +4,13 @@ from typing import Literal, List
from openai import OpenAI
from pydantic import Field, PrivateAttr, model_validator
-from experiencemaker.model.base_embedding_model import BaseEmbeddingModel
+from experiencemaker.model.base_embedding_model import BaseEmbeddingModel, EMBEDDING_MODEL_REGISTRY
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()
@@ -71,3 +71,23 @@ class OpenAICompatibleEmbeddingModel(BaseEmbeddingModel):
else:
# If the input type is neither a string nor a list of strings, throw an exception
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()
+
+ 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 03814862..09fedef2 100644
--- a/experiencemaker/model/openai_compatible_llm.py
+++ b/experiencemaker/model/openai_compatible_llm.py
@@ -7,7 +7,7 @@ from openai.types import CompletionUsage
from pydantic import Field, PrivateAttr, model_validator
from experiencemaker.enumeration.chunk_enum import ChunkEnum
-from experiencemaker.model.base_llm import BaseLLM
+from experiencemaker.model.base_llm import BaseLLM, LLM_REGISTRY
from experiencemaker.schema.trajectory import Message, ActionMessage, ToolCall
from experiencemaker.tool.base_tool import BaseTool
@@ -150,3 +150,28 @@ class OpenAICompatibleBaseLLM(BaseLLM):
elif chunk_enum is ChunkEnum.ERROR:
print(f"\n{chunk}", end="")
+
+
+LLM_REGISTRY.register(OpenAICompatibleBaseLLM, "openai_compatible")
+
+
+def main():
+ from experiencemaker.utils.util_function import load_env_keys
+ from experiencemaker.tool.dashscope_search_tool import DashscopeSearchTool
+ 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/agent_wrapper/agent_wrapper_mixin.py b/experiencemaker/module/agent_wrapper/agent_wrapper_mixin.py
new file mode 100644
index 00000000..bcaf4450
--- /dev/null
+++ b/experiencemaker/module/agent_wrapper/agent_wrapper_mixin.py
@@ -0,0 +1,18 @@
+from abc import ABC
+
+from pydantic import Field, BaseModel
+
+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)
+ workspace_id: str = Field(default="")
+
+ def execute(self, query: str, **kwargs) -> Trajectory:
+ raise NotImplementedError
+
+
+AGENT_WRAPPER_REGISTRY = Registry[AgentWrapperMixin]("agent_wrapper")
diff --git a/experiencemaker/module/agent_wrapper/base_agent_wrapper.py b/experiencemaker/module/agent_wrapper/base_agent_wrapper.py
index 47844d04..93f6303f 100644
--- a/experiencemaker/module/agent_wrapper/base_agent_wrapper.py
+++ b/experiencemaker/module/agent_wrapper/base_agent_wrapper.py
@@ -1,26 +1,21 @@
-from abc import ABC
from typing import List
-from loguru import logger
from pydantic import Field
-from experiencemaker.module.agent_wrapper.base_agent_wrapper_mixin import BaseAgentWrapperMixin
+from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AgentWrapperMixin
from experiencemaker.module.environment.base_environment import BaseEnvironment
from experiencemaker.schema.trajectory import Trajectory, Message, StateMessage, ActionMessage, ContextMessage
-class BaseAgentWrapper(BaseAgentWrapperMixin, ABC):
+class BaseAgentWrapper(AgentWrapperMixin):
max_steps: int = Field(default=10)
enable_exploration: bool = Field(default=False)
trajectory: Trajectory | None = Field(default_factory=Trajectory)
def reset(self):
- self.trajectory = Trajectory()
+ self.trajectory.reset()
- def before_execute_hook(self, query: str, **kwargs):
- self.trajectory.query = query
-
- def after_step_hook(self, action_msg: ActionMessage, next_state: StateMessage, **kwargs):
+ def after_step_hook(self, **kwargs):
raise NotImplementedError
def build_messages(self,
@@ -32,7 +27,7 @@ class BaseAgentWrapper(BaseAgentWrapperMixin, ABC):
def explore_messages(self, messages: List[Message], **kwargs) -> List[Message]:
raise NotImplementedError
- def action_parser(self, action_msg: ActionMessage, **kwargs) -> ActionMessage:
+ def action_parser(self, action_msg: ActionMessage) -> ActionMessage: # noqa
return action_msg
def generate_action(self,
@@ -48,13 +43,10 @@ class BaseAgentWrapper(BaseAgentWrapperMixin, ABC):
action_msg: ActionMessage = self.llm.chat(messages, tools=env.tools)
return self.action_parser(action_msg)
- def after_execute_hook(self, **kwargs):
- return self.trajectory
-
def execute(self, query: str, env: BaseEnvironment = None, **kwargs) -> Trajectory:
- self.before_execute_hook(query=query, **kwargs)
-
+ self.trajectory.query = query
current_state = env.current_state
+
for i in range(self.max_steps):
self.trajectory.current_step = i
@@ -62,17 +54,12 @@ class BaseAgentWrapper(BaseAgentWrapperMixin, ABC):
context_msg: ContextMessage | None = None
if self.context_generator:
context_msg = self.context_generator.execute(trajectory=self.trajectory, **kwargs)
- if context_msg.content:
- logger.info(f"step{i}.context_msg={context_msg.content}")
# generate action
action_msg = self.generate_action(state=current_state, context_msg=context_msg, env=env, **kwargs)
- logger.info(f"step{i} ====== reasoning_content ======\n{action_msg.reasoning_content}\n\n"
- f"====== content ======\n{action_msg.content}\ntool_calls={action_msg.tool_calls}")
# generate next state
next_state, reward, done, info = env.step(action_msg, trajectory=self.trajectory, **kwargs)
- logger.info(f"step{i}.next_state={next_state.content} reward={reward.reward_value} done={done} info={info}")
self.after_step_hook(step_index=i,
current_state=current_state,
@@ -88,4 +75,4 @@ class BaseAgentWrapper(BaseAgentWrapperMixin, ABC):
current_state = next_state
- return self.after_execute_hook(**kwargs)
+ return self.trajectory
diff --git a/experiencemaker/module/agent_wrapper/base_agent_wrapper_mixin.py b/experiencemaker/module/agent_wrapper/base_agent_wrapper_mixin.py
deleted file mode 100644
index 9dc7b82a..00000000
--- a/experiencemaker/module/agent_wrapper/base_agent_wrapper_mixin.py
+++ /dev/null
@@ -1,27 +0,0 @@
-from abc import ABC
-
-from pydantic import Field
-
-from experiencemaker.module.base_module import BaseModule
-from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator
-from experiencemaker.schema.trajectory import Trajectory, ActionMessage, Message
-
-
-class BaseAgentWrapperMixin(BaseModule, ABC):
- context_generator: BaseContextGenerator | None = Field(default=None)
-
- def execute(self, query: str, **kwargs) -> Trajectory:
- raise NotImplementedError
-
-
-class MockAgentWrapper(BaseAgentWrapperMixin):
-
- def execute(self, query: str, **kwargs) -> Trajectory:
- user_message = Message(content=query)
- answer_message = ActionMessage(content="hello world")
-
- traj = Trajectory(steps=[user_message, answer_message],
- done=True,
- query=query,
- answer=answer_message.content)
- return traj
diff --git a/experiencemaker/module/agent_wrapper/naive_agent_wrapper.py b/experiencemaker/module/agent_wrapper/naive_agent_wrapper.py
deleted file mode 100644
index f5e769bf..00000000
--- a/experiencemaker/module/agent_wrapper/naive_agent_wrapper.py
+++ /dev/null
@@ -1,62 +0,0 @@
-import datetime
-
-from experiencemaker.module.agent_wrapper.base_agent_wrapper_v2 import BaseAgentWrapperV2
-from loguru import logger
-
-from experiencemaker.module.environment.base_environment import BaseEnvironment
-from experiencemaker.schema.trajectory import Message, StateMessage, ContextMessage, ActionMessage
-
-
-class NaiveAgentWrapper(BaseAgentWrapperV2):
-
- def generate_action(self,
- state: StateMessage,
- context_msg: ContextMessage | None,
- env: BaseEnvironment,
- **kwargs) -> ActionMessage:
-
- tool_names = [x.name for x in env.tools]
- if self.trajectory.current_step == 0:
- now_time = datetime.datetime.now().strftime('%Y-%m-%d %H:%M:%S')
- insight_tag: bool = True if context_msg is not None and context_msg.content else False
- user_prompt = self.prompt_handler.prompt_format(
- prompt_name="role_prompt",
- insight_tag=insight_tag,
- time=now_time,
- tools=", ".join(tool_names),
- previous_insight=context_msg.content if insight_tag else "",
- query=self.trajectory.query)
- # When using the reasoning models of Qwen3 or DeepSeek R1, it is not recommended to use system prompt.
- self.trajectory.steps.append(Message(content=user_prompt))
-
- elif self.trajectory.metadata.get("has_terminate_tool") is True:
- user_prompt = self.prompt_handler.final_prompt.format(query=self.trajectory.query)
- self.trajectory.steps.append(Message(content=user_prompt))
-
- else:
- user_prompt = self.prompt_handler.next_prompt.format(query=self.trajectory.query)
- self.trajectory.steps.append(Message(content=user_prompt))
-
- if self.trajectory.metadata.get("has_terminate_tool") is True:
- action_msg: ActionMessage = self.llm.chat(self.trajectory.steps)
- logger.info(f"step{self.trajectory.current_step} size={len(self.trajectory.steps)} user_prompt={user_prompt}")
-
- else:
- action_msg: ActionMessage = self.llm.chat(messages, tools=env.tools)
- logger.info(f"step{self.trajectory.current_step} size={len(messages)} user_prompt={user_prompt} "
- f"tool_names={tool_names}")
-
- for tool in action_msg.tool_calls:
- if tool.name == "terminate":
- self.trajectory.metadata["has_terminate_tool"] = True
- break
-
- self.trajectory.add_step(action_msg)
- return action_msg
-
-
- def after_step_hook(self, action_msg: ActionMessage, next_state: StateMessage, done: bool = False, **kwargs):
- if done:
- self.trajectory.answer = action_msg.content
- else:
- self.trajectory.add_step(next_state)
diff --git a/experiencemaker/module/agent_wrapper/simple_agent.py b/experiencemaker/module/agent_wrapper/simple_agent.py
new file mode 100644
index 00000000..b2db9b24
--- /dev/null
+++ b/experiencemaker/module/agent_wrapper/simple_agent.py
@@ -0,0 +1,98 @@
+import datetime
+from pathlib import Path
+from typing import List
+
+from loguru import logger
+from pydantic import Field, BaseModel
+
+from experiencemaker.model.base_llm import BaseLLM
+from experiencemaker.module.prompt.prompt_mixin import PromptMixin
+from experiencemaker.schema.trajectory import Message, ActionMessage, ToolCall, StateMessage
+from experiencemaker.tool import CodeTool, DashscopeSearchTool, TerminateTool
+from experiencemaker.tool.base_tool import BaseTool
+
+
+class SimpleAgentContext(BaseModel):
+ current_step: int = Field(default=-1)
+ query: str = Field(default="")
+ previous_experience: str = Field(default="")
+ messages: List[Message] = Field(default_factory=list)
+ metadata: dict = Field(default_factory=dict)
+ has_terminate_tool: bool = Field(default=False)
+
+
+class SimpleAgent(PromptMixin):
+ llm: BaseLLM | None = Field(default=None)
+ max_steps: int = Field(default=10)
+ tools: List[BaseTool] = [CodeTool(), DashscopeSearchTool(), TerminateTool()]
+ prompt_file_path: Path = Path(__file__).parent / "simple_agent_prompt.yaml"
+
+ def think(self, context: SimpleAgentContext):
+ now_time = datetime.datetime.now().strftime('%Y-%m-%d %H:%M:%S')
+ tool_names = [x.name for x in self.tools]
+
+ if context.current_step == 0:
+ user_prompt = self.prompt_format(prompt_name="role_prompt",
+ experience_tag=False if context.previous_experience else True,
+ time=now_time,
+ tools=", ".join(tool_names),
+ previous_insight=context.previous_experience,
+ query=context.query)
+
+ elif context.has_terminate_tool:
+ user_prompt = self.prompt_format(prompt_name="final_prompt", query=context.query)
+
+ else:
+ user_prompt = self.prompt_format(prompt_name="next_prompt", query=context.query)
+
+ context.messages.append(Message(content=user_prompt))
+ logger.info(f"step.{context.current_step} user_prompt={user_prompt}")
+
+ if context.has_terminate_tool:
+ action_msg: ActionMessage = self.llm.chat(context.messages)
+
+ else:
+ action_msg: ActionMessage = self.llm.chat(context.messages, tools=self.tools)
+ for tool in action_msg.tool_calls:
+ if tool.name == "terminate":
+ context.has_terminate_tool = True
+ break
+
+ context.messages.append(action_msg)
+ action_msg_context: str = action_msg.content + "\n\n" + action_msg.reasoning_content
+ logger.info(f"step.{context.current_step} action_msg_context={action_msg_context} "
+ f"tool_calls={action_msg.tool_calls}")
+ return True if action_msg.tool_calls else False
+
+ def act(self, context: SimpleAgentContext):
+ action_msg = context.messages[-1]
+ assert isinstance(action_msg, ActionMessage)
+
+ tool_dict = {tool.name: tool for tool in self.tools}
+
+ new_tool_calls: List[ToolCall] = []
+ for tool_call in action_msg.tool_calls:
+ if tool_call.name not in tool_dict:
+ continue
+
+ new_tool_call = tool_call.model_copy(deep=True)
+ tool = tool_dict[tool_call.name]
+ new_tool_call.result = tool.execute(**tool_call.argument_dict)
+ new_tool_calls.append(new_tool_call)
+
+ state_msg = StateMessage(tool_calls=new_tool_calls)
+ context.messages.append(state_msg)
+ logger.info(f"step.{context.current_step} state_msg_context={state_msg.content}")
+
+ def run(self, query: str, previous_experience: str) -> List[Message]:
+ context: SimpleAgentContext = SimpleAgentContext(query=query, previous_experience=previous_experience)
+
+ for i in range(self.max_steps):
+ context.current_step = i
+
+ should_act: bool = self.think(context)
+ if should_act:
+ self.act(context)
+ else:
+ break
+ return context.messages
diff --git a/experiencemaker/module/agent_wrapper/simple_agent_prompt.yaml b/experiencemaker/module/agent_wrapper/simple_agent_prompt.yaml
new file mode 100644
index 00000000..514c2ebe
--- /dev/null
+++ b/experiencemaker/module/agent_wrapper/simple_agent_prompt.yaml
@@ -0,0 +1,28 @@
+role_prompt: |
+ You are a helpful assistant named BeyondAgent.
+ The current time is {time}.
+
+ Please proactively choose the most suitable tool or combination of tools based on the user's question, including {tools} etc.
+ For complex tasks, you can break down the problem step by step and use different tools to solve it incrementally.
+ Please determine the response language based on the language of the user's question.
+ [experience_tag]
+ [experience_tag]Previous Experience
+ [experience_tag]{previous_experience}
+ [experience_tag]Please consider the helpful parts from these in answering the question, to make the response more comprehensive and substantial.
+
+ User's question
+ {query}
+
+next_prompt: |
+ User's question
+ {query}
+
+ Plan the most suitable tool or combination of tools based on the context and the user's question.
+ For complex tasks, break down the problem and use different tools step by step to solve it, don't give up easily.
+ Try calling the same tool multiple times with different parameters to obtain information from various perspectives.
+ If the task is completed and the user's question can now be answered, use the **terminate** tool.
+
+final_prompt: |
+ Please integrate the context and provide a complete answer to the user's question:
+ {query}
+
diff --git a/experiencemaker/module/agent_wrapper/simple_agent_wrapper.py b/experiencemaker/module/agent_wrapper/simple_agent_wrapper.py
new file mode 100644
index 00000000..124ced7c
--- /dev/null
+++ b/experiencemaker/module/agent_wrapper/simple_agent_wrapper.py
@@ -0,0 +1,21 @@
+from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AgentWrapperMixin, AGENT_WRAPPER_REGISTRY
+from experiencemaker.module.agent_wrapper.simple_agent import SimpleAgent
+from experiencemaker.schema.trajectory import Trajectory
+
+
+class SimpleAgentWrapper(SimpleAgent, AgentWrapperMixin):
+
+ def execute(self, query: str, **kwargs) -> Trajectory:
+ trajectory = Trajectory(query=query)
+ context_msg = self.context_generator.execute(trajectory=trajectory)
+ previous_experience = context_msg.content
+
+ messages = self.run(query, previous_experience)
+
+ trajectory.steps = messages
+ trajectory.answer = messages[-1].content
+ trajectory.done = True
+ return trajectory
+
+
+AGENT_WRAPPER_REGISTRY.register(SimpleAgentWrapper, "simple")
diff --git a/experiencemaker/module/base_module.py b/experiencemaker/module/base_module.py
deleted file mode 100644
index 4cad23c9..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/context_generator/base_context_generator.py b/experiencemaker/module/context_generator/base_context_generator.py
index 200b46de..e7839a12 100644
--- a/experiencemaker/module/context_generator/base_context_generator.py
+++ b/experiencemaker/module/context_generator/base_context_generator.py
@@ -1,16 +1,20 @@
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.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(BaseModule, ABC):
+class BaseContextGenerator(BaseModel, ABC):
vector_store: BaseVectorStore | None = Field(default=None)
+ llm: BaseLLM | None = Field(default=None)
+ workspace_id: str = Field(default="")
def _build_retrieve_query(self, trajectory: Trajectory, **kwargs) -> str:
raise NotImplementedError
@@ -31,7 +35,4 @@ class BaseContextGenerator(BaseModule, ABC):
return context_msg
-class MockContextGenerator(BaseContextGenerator):
-
- def execute(self, trajectory: Trajectory, **kwargs) -> ContextMessage:
- return ContextMessage(content="mock context")
+CONTEXT_GENERATOR_REGISTRY = Registry[BaseContextGenerator]("context_generator")
diff --git a/experiencemaker/module/context_generator/simple_context_generator.py b/experiencemaker/module/context_generator/simple_context_generator.py
new file mode 100644
index 00000000..4beb15a0
--- /dev/null
+++ b/experiencemaker/module/context_generator/simple_context_generator.py
@@ -0,0 +1,40 @@
+from typing import List
+
+from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator, \
+ CONTEXT_GENERATOR_REGISTRY
+from experiencemaker.schema.trajectory import Trajectory, ContextMessage
+from experiencemaker.schema.vector_store_node import VectorStoreNode
+
+
+class SimpleContextGenerator(BaseContextGenerator):
+
+ def _build_retrieve_query(self, trajectory: Trajectory, **kwargs) -> str:
+ query = ""
+ if trajectory.current_step == 0:
+ query = trajectory.query
+ return query
+
+ def _retrieve_by_query(self, trajectory: Trajectory, query: str, **kwargs) -> List[VectorStoreNode]:
+ if not query:
+ return []
+
+ return self.vector_store.retrieve_by_query(query=query, top_k=self.vector_store_top_k)
+
+ def _generate_context_message(self,
+ trajectory: Trajectory,
+ nodes: List[VectorStoreNode],
+ **kwargs) -> ContextMessage:
+ if not nodes:
+ return ContextMessage(content="")
+
+ content = ""
+ for node in nodes:
+ experience = node.metadata.get("experience", "")
+ if not experience:
+ continue
+
+ content += f"- {node.content} {experience}\n"
+ return ContextMessage(content=content.strip())
+
+
+CONTEXT_GENERATOR_REGISTRY.register(SimpleContextGenerator, "simple")
diff --git a/experiencemaker/module/context_generator/step_context_generator.py b/experiencemaker/module/context_generator/step_context_generator.py
new file mode 100644
index 00000000..2313dbbb
--- /dev/null
+++ b/experiencemaker/module/context_generator/step_context_generator.py
@@ -0,0 +1,393 @@
+import json
+import re
+from typing import List, Dict, Any, Optional
+from loguru import logger
+from pydantic import Field, model_validator
+
+from experiencemaker.enumeration.role import Role
+from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator
+from experiencemaker.schema.trajectory import Trajectory, ContextMessage, Message
+from experiencemaker.schema.vector_store_node import VectorStoreNode
+from experiencemaker.storage.es_vector_store import EsVectorStore
+from experiencemaker.storage.file_vector_store import FileVectorStore
+
+
+class StepContextGenerator(BaseContextGenerator):
+ """
+ Step-level context generator that retrieves and utilizes step-level experiences
+ from the experience store to provide relevant context for agent execution
+ """
+
+ # Vector Store Configuration
+ vector_store_type: str = Field(default="file_vector_store")
+ vector_store_hosts: str | List[str] = Field(default="http://localhost:9200")
+ vector_store_index_name: str = Field(default="step_experience_store")
+ store_dir: str = Field(default="./step_experiences/")
+
+ # Retrieval Configuration
+ vector_retrieve_top_k: int = Field(default=15)
+ final_top_k: int = Field(default=5)
+ min_score_threshold: float = Field(default=0.3)
+
+ # Feature Switches
+ enable_llm_rerank: bool = Field(default=True)
+ enable_context_rewrite: bool = Field(default=True)
+ enable_score_filter: bool = Field(default=True)
+
+ @model_validator(mode="after")
+ def init_vector_store(self):
+ """Initialize vector store based on configuration"""
+ if self.vector_store_type == "file_vector_store":
+ self.vector_store = FileVectorStore(
+ embedding_model=self.embedding_model,
+ index_name=self.vector_store_index_name,
+ store_dir=self.store_dir
+ )
+ elif self.vector_store_type == "es_vector_store":
+ self.vector_store = EsVectorStore(
+ embedding_model=self.embedding_model,
+ index_name=self.vector_store_index_name,
+ hosts=self.vector_store_hosts
+ )
+ else:
+ raise ValueError(f"Unknown vector store type: {self.vector_store_type}")
+
+ return self
+
+ def _build_retrieve_query(self, trajectory: Trajectory, **kwargs) -> str:
+ """Build retrieval query from trajectory"""
+ # Use the original query as base
+ base_query = trajectory.query
+
+ # Optionally enhance with current step context if available
+ current_context = kwargs.get("current_context", "")
+ if current_context:
+ base_query = f"{base_query} {current_context}"
+
+ return base_query
+
+ def vector_retrieve(self, query: str, top_k: int = 10) -> List[VectorStoreNode]:
+ """Vector similarity retrieval from experience store"""
+ if not query:
+ logger.warning("Empty query provided for vector retrieval")
+ return []
+
+ try:
+ retrieved_nodes = self.vector_store.retrieve_by_query(
+ query=query,
+ top_k=top_k
+ )
+ logger.info(f"Vector retrieval found {len(retrieved_nodes)} candidates")
+ return retrieved_nodes
+
+ except Exception as e:
+ logger.error(f"Error in vector retrieval: {e}")
+ return []
+
+ def llm_rerank(self, query: str, candidates: List[VectorStoreNode]) -> List[VectorStoreNode]:
+ """LLM-based reranking of candidate experiences"""
+ if not self.enable_llm_rerank or not candidates:
+ return candidates
+
+ try:
+ # Format candidates for LLM evaluation
+ candidates_text = self._format_candidates_for_rerank(candidates)
+
+ prompt = self.prompt_handler.experience_rerank_prompt.format(
+ query=query,
+ candidates=candidates_text,
+ num_candidates=len(candidates)
+ )
+
+ response = self.llm.chat([Message(role=Role.USER, content=prompt)])
+
+ # Parse reranking results
+ reranked_indices = self._parse_rerank_response(response.content)
+
+ # Reorder candidates based on LLM ranking
+ if reranked_indices:
+ reranked_candidates = []
+ for idx in reranked_indices:
+ if 0 <= idx < len(candidates):
+ reranked_candidates.append(candidates[idx])
+ return reranked_candidates
+
+ return candidates
+
+ except Exception as e:
+ logger.error(f"Error in LLM reranking: {e}")
+ return candidates
+
+ def llm_rewrite_context(self, query: str, context_content: str, trajectory: Trajectory) -> str:
+ """LLM-based context rewriting to make experiences more relevant and actionable for current task"""
+ if not self.enable_query_rewrite or not context_content:
+ return context_content
+
+ try:
+ # Extract current trajectory context
+ current_context = self._extract_trajectory_context(trajectory)
+
+ prompt = self.prompt_handler.context_rewrite_prompt.format(
+ current_query=query,
+ current_context=current_context,
+ original_context=context_content
+ )
+
+ response = self.llm.chat([Message(role=Role.USER, content=prompt)])
+
+ # Extract rewritten context from JSON
+ rewritten_context = self._parse_json_response(response.content, "rewritten_context")
+
+ if rewritten_context and rewritten_context.strip():
+ logger.info("Context successfully rewritten for current task")
+ return rewritten_context.strip()
+
+ return context_content
+
+ except Exception as e:
+ logger.error(f"Error in context rewriting: {e}")
+ return context_content
+
+ def score_based_filter(self, experiences: List[VectorStoreNode],
+ min_score: float) -> List[VectorStoreNode]:
+ """Filter experiences based on quality scores"""
+ if not self.enable_score_filter:
+ return experiences
+
+ filtered_experiences = []
+
+ for exp in experiences:
+ # Get confidence score from metadata
+ confidence = exp.metadata.get("confidence", 0.5)
+ validation_score = exp.metadata.get("validation_score", 0.5)
+
+ # Calculate combined score
+ combined_score = (confidence + validation_score) / 2
+
+ if combined_score >= min_score:
+ filtered_experiences.append(exp)
+ else:
+ logger.debug(f"Filtered out experience with score {combined_score:.2f}")
+
+ logger.info(f"Score filtering: {len(filtered_experiences)}/{len(experiences)} experiences retained")
+ return filtered_experiences
+
+ def hybrid_retrieve(self, query: str, trajectory: Trajectory, top_k: int = 5) -> List[VectorStoreNode]:
+ """Hybrid retrieval strategy combining multiple approaches"""
+ logger.info(f"Starting hybrid retrieval for query: '{query}'")
+
+ # Step 1: Vector retrieval to get candidates
+ candidates = self.vector_retrieve(query, self.vector_retrieve_top_k)
+
+ if not candidates:
+ logger.warning("No candidates found in vector retrieval")
+ return []
+
+ # Step 2: LLM reranking (optional)
+ reranked = self.llm_rerank(query, candidates)
+
+ # Step 3: Score-based filtering (optional)
+ filtered = self.score_based_filter(reranked, self.min_score_threshold)
+
+ # Step 4: Return top-k results
+ final_results = filtered[:top_k]
+ logger.info(f"Hybrid retrieval completed: {len(final_results)} experiences selected")
+
+ return final_results
+
+ def retrieve_by_query(self, trajectory: Trajectory, query: str, **kwargs) -> List[VectorStoreNode]:
+ """Retrieve experiences by query (implements base class method)"""
+ return self.hybrid_retrieve(query, trajectory, self.final_top_k)
+
+ def generate_context_message(self,
+ trajectory: Trajectory,
+ nodes: List[VectorStoreNode],
+ **kwargs) -> ContextMessage:
+ """Generate context message from retrieved experiences"""
+ if not nodes:
+ return ContextMessage(content="")
+
+ try:
+ # Format retrieved experiences
+ formatted_experiences = self._format_experiences_for_context(nodes)
+
+ prompt = self.prompt_handler.context_generation_prompt.format(
+ query=trajectory.query,
+ current_step=kwargs.get("current_step", ""),
+ retrieved_experiences=formatted_experiences,
+ num_experiences=len(nodes)
+ )
+
+ response = self.llm.chat([Message(role=Role.USER, content=prompt)])
+
+ # Extract generated context from JSON
+ context_content = self._parse_json_response(response.content, "context")
+
+ if not context_content:
+ # Fallback to simple formatting
+ context_content = self._create_context(nodes)
+
+ return ContextMessage(content=context_content)
+
+ except Exception as e:
+ logger.error(f"Error generating context message: {e}")
+ return ContextMessage(content=self._create_context(nodes))
+
+ def build_context_messages(self, task: str, experiences: List[VectorStoreNode], trajectory: Trajectory) -> List[
+ Message]:
+ """Build context messages from experiences for agent consumption"""
+ if not experiences:
+ return []
+
+ messages = []
+
+ # Create initial context content with experiences
+ system_content = "You have access to the following relevant experiences from previous executions:\n\n"
+
+ for i, exp in enumerate(experiences, 1):
+ condition = exp.content
+ experience_content = exp.metadata.get("experience", "")
+ tags = exp.metadata.get("tags", [])
+
+ system_content += f"**Experience {i}:**\n"
+ system_content += f"When to use: {condition}\n"
+ system_content += f"Experience: {experience_content}\n"
+ system_content += f"Tags: {', '.join(tags)}\n\n"
+
+ system_content += "Consider these experiences when planning and executing your approach."
+
+ # Rewrite the complete context to make it more relevant to current task
+ if self.enable_context_rewrite:
+ system_content = self.llm_rewrite_context(task, system_content, trajectory)
+
+ messages.append(Message(role=Role.SYSTEM, content=system_content))
+
+ return messages
+
+ def get_best_experiences(self, task: str, trajectory: Trajectory, max_count: int = 3) -> List[Message]:
+ """Get the best relevant experiences for a task as formatted messages"""
+ experiences = self.hybrid_retrieve(task, trajectory, max_count)
+ return self.build_context_messages(task, experiences, trajectory)
+
+ def _extract_trajectory_context(self, trajectory: Trajectory) -> str:
+ """Extract relevant context from trajectory for query enhancement"""
+ context_parts = []
+
+ # Add recent steps if available
+ if trajectory.steps:
+ recent_steps = trajectory.steps[-3:] # Last 3 steps
+ step_summaries = []
+ for step in recent_steps:
+ step_summary = step.content[:100] + "..." if len(step.content) > 100 else step.content
+ step_summaries.append(f"- {step.role.value}: {step_summary}")
+
+ if step_summaries:
+ context_parts.append("Recent steps:\n" + "\n".join(step_summaries))
+
+ # Add metadata if available
+ if trajectory.metadata:
+ relevant_metadata = {k: v for k, v in trajectory.metadata.items()
+ if k in ["domain", "task_type", "difficulty"]}
+ if relevant_metadata:
+ context_parts.append(f"Task metadata: {relevant_metadata}")
+
+ return "\n\n".join(context_parts)
+
+ def _format_candidates_for_rerank(self, candidates: List[VectorStoreNode]) -> str:
+ """Format candidates for LLM reranking"""
+ formatted_candidates = []
+
+ for i, candidate in enumerate(candidates):
+ condition = candidate.content
+ experience = candidate.metadata.get("experience", "")
+ tags = candidate.metadata.get("tags", [])
+ confidence = candidate.metadata.get("confidence", 0.5)
+
+ candidate_text = f"Candidate {i}:\n"
+ candidate_text += f"Condition: {condition}\n"
+ candidate_text += f"Experience: {experience}\n"
+ candidate_text += f"Tags: {', '.join(tags)}\n"
+ candidate_text += f"Confidence: {confidence}\n"
+
+ formatted_candidates.append(candidate_text)
+
+ return "\n---\n".join(formatted_candidates)
+
+ def _parse_rerank_response(self, response: str) -> List[int]:
+ """Parse LLM reranking response to extract ranked indices"""
+ try:
+ # Try to extract JSON format
+ json_pattern = r'```json\s*([\s\S]*?)\s*```'
+ json_blocks = re.findall(json_pattern, response)
+
+ if json_blocks:
+ parsed = json.loads(json_blocks[0])
+ if isinstance(parsed, dict) and "ranked_indices" in parsed:
+ return parsed["ranked_indices"]
+ elif isinstance(parsed, list):
+ return parsed
+
+ # Try to extract numbers from text
+ numbers = re.findall(r'\b\d+\b', response)
+ return [int(num) for num in numbers]
+
+ except Exception as e:
+ logger.error(f"Error parsing rerank response: {e}")
+ return []
+
+ def _format_experiences_for_context(self, experiences: List[VectorStoreNode]) -> str:
+ """Format experiences for context generation"""
+ formatted_experiences = []
+
+ for i, exp in enumerate(experiences, 1):
+ condition = exp.content
+ experience_content = exp.metadata.get("experience", "")
+ experience_type = exp.metadata.get("experience_type", "general")
+ tags = exp.metadata.get("tags", [])
+
+ exp_text = f"Experience {i} ({experience_type}):\n"
+ exp_text += f"When to use: {condition}\n"
+ exp_text += f"Experience: {experience_content}\n"
+ exp_text += f"Tags: {', '.join(tags)}"
+
+ formatted_experiences.append(exp_text)
+
+ return "\n\n---\n\n".join(formatted_experiences)
+
+ def _create_context(self, experiences: List[VectorStoreNode]) -> str:
+ """Create simple context when LLM generation fails"""
+ if not experiences:
+ return ""
+
+ context = "Here are some relevant experiences that might help:\n\n"
+
+ for i, exp in enumerate(experiences, 1):
+ condition = exp.content
+ experience_content = exp.metadata.get("experience", "")
+
+ context += f"{i}. **When**: {condition}\n"
+ context += f" **Experience**: {experience_content}\n\n"
+
+ return context
+
+ def _parse_json_response(self, response: str, key: str) -> str:
+ """Parse JSON response to extract specific key"""
+ try:
+ # Try to extract JSON blocks
+ json_pattern = r'```json\s*([\s\S]*?)\s*```'
+ json_blocks = re.findall(json_pattern, response)
+
+ if json_blocks:
+ parsed = json.loads(json_blocks[0])
+ if isinstance(parsed, dict) and key in parsed:
+ return parsed[key]
+
+ # Fallback: try to parse the entire response as JSON
+ parsed = json.loads(response)
+ if isinstance(parsed, dict) and key in parsed:
+ return parsed[key]
+
+ except json.JSONDecodeError:
+ logger.warning(f"Failed to parse JSON response for key '{key}'")
+
+ return ""
\ No newline at end of file
diff --git a/experiencemaker/module/environment/base_environment.py b/experiencemaker/module/environment/base_environment.py
index 9e52d3c1..cca3f4b2 100644
--- a/experiencemaker/module/environment/base_environment.py
+++ b/experiencemaker/module/environment/base_environment.py
@@ -1,18 +1,18 @@
+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 | None = Field(default=None)
+ current_state: StateMessage = Field(default_factory=StateMessage)
metadata: dict = Field(default_factory=dict, description="add query / answer and etc for reward calculating!")
def reset(self):
@@ -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/evaluator/base_evaluator.py b/experiencemaker/module/evaluator/base_evaluator.py
index 13e67575..e2de7d44 100644
--- a/experiencemaker/module/evaluator/base_evaluator.py
+++ b/experiencemaker/module/evaluator/base_evaluator.py
@@ -2,7 +2,7 @@ from abc import ABC
from pydantic import BaseModel, Field
-from experiencemaker.module.agent_wrapper.base_agent_wrapper import BaseAgentWrapper
+from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AgentWrapperMixin
from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator
from experiencemaker.module.environment.base_environment import BaseEnvironment
from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer
@@ -10,7 +10,7 @@ from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer
class BaseEvaluator(BaseModel, ABC):
data_path: str = Field(default="")
- agent_wrapper: BaseAgentWrapper | None = Field(default=None)
+ agent_wrapper: AgentWrapperMixin | None = Field(default=None)
context_generator: BaseContextGenerator | None = Field(default=None)
summarizer: BaseSummarizer | None = Field(default=None)
env: BaseEnvironment | None = Field(default=None)
diff --git a/experiencemaker/module/prompt/__init__.py b/experiencemaker/module/prompt/__init__.py
new file mode 100644
index 00000000..e69de29b
diff --git a/experiencemaker/utils/prompt_handler.py b/experiencemaker/module/prompt/prompt_mixin.py
similarity index 51%
rename from experiencemaker/utils/prompt_handler.py
rename to experiencemaker/module/prompt/prompt_mixin.py
index fe68956e..56fef91e 100644
--- a/experiencemaker/utils/prompt_handler.py
+++ b/experiencemaker/module/prompt/prompt_mixin.py
@@ -1,40 +1,36 @@
-import os
+from pathlib import Path
import yaml
from loguru import logger
-from pydantic import BaseModel, Field
+from pydantic import BaseModel, Field, model_validator
-class PromptHandler(BaseModel):
- dir_path: str = Field(default="")
+class PromptMixin(BaseModel):
+ prompt_file_path: Path | str = Field(default=None)
prompt_dict: dict = Field(default_factory=dict)
- def add_prompt_file(self, file_name: str):
- prompt_path = os.path.join(self.dir_path, file_name + ".yaml")
- self._add_prompt_file(prompt_path)
+ @model_validator(mode="after")
+ def init_prompt(self):
+ if self.prompt_dict:
+ logger.info(f"load prompt from dict, keys={self.prompt_dict.keys()}")
- def _add_prompt_file(self, prompt_path: str):
- if os.path.exists(prompt_path):
- with open(prompt_path) as f:
- prompt_dict: dict = yaml.load(f, yaml.FullLoader)
- self.update_prompt_dict(prompt_dict)
- else:
- logger.warning(f"prompt_path={prompt_path} not exists!")
+ if self.prompt_file_path is not None:
+ if isinstance(self.prompt_file_path, str):
+ self.prompt_file_path = Path(self.prompt_file_path)
- def update_prompt_dict(self, prompt_dict: dict):
- self.prompt_dict.update(prompt_dict)
+ if not self.prompt_file_path.exists():
+ logger.warning(f"prompt_file_path={self.prompt_file_path} not exists!")
- def __getitem__(self, key: str):
- return self.prompt_dict[key]
+ else:
+ with self.prompt_file_path.open("r") as f:
+ for k, v in yaml.load(f, yaml.FullLoader):
+ 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!")
- 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)
+ return self
def prompt_format(self, prompt_name: str, **kwargs):
prompt = self.prompt_dict[prompt_name]
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..e0db5abc 100644
--- a/experiencemaker/module/runner/base_runner.py
+++ b/experiencemaker/module/runner/base_runner.py
@@ -2,7 +2,7 @@ from typing import List
from pydantic import BaseModel, Field
-from experiencemaker.module.agent_wrapper.base_agent_wrapper import BaseAgentWrapper
+from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AgentWrapperMixin
from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator
from experiencemaker.module.environment.base_environment import BaseEnvironment
from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer
@@ -10,7 +10,7 @@ from experiencemaker.schema.trajectory import Trajectory
class BaseRunner(BaseModel):
- agent_wrapper: BaseAgentWrapper | None = Field(default=None)
+ agent_wrapper: AgentWrapperMixin | None = Field(default=None)
context_generator: BaseContextGenerator | None = Field(default=None)
summarizer: BaseSummarizer | None = Field(default=None)
env: BaseEnvironment | None = Field(default=None)
@@ -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/summarizer/base_summarizer.py b/experiencemaker/module/summarizer/base_summarizer.py
index 2ccb2c6c..53ae4a4a 100644
--- a/experiencemaker/module/summarizer/base_summarizer.py
+++ b/experiencemaker/module/summarizer/base_summarizer.py
@@ -1,37 +1,33 @@
+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.schema.trajectory import Trajectory, Sample, SummaryMessage
+from experiencemaker.model.base_llm import BaseLLM
+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(BaseModule):
+class BaseSummarizer(BaseModel, ABC):
vector_store: BaseVectorStore | None = Field(default=None)
+ llm: BaseLLM | None = Field(default=None)
+ workspace_id: str = Field(default="")
- def extract_samples(self, trajectories: List[Trajectory], **kwargs) -> List[Sample]:
+ def _extract_experiences(self, trajectories: List[Trajectory], **kwargs) -> List[Experience]:
raise NotImplementedError
- def insert_into_vector_store(self, samples: List[Sample], **kwargs):
- raise NotImplementedError
+ def execute(self, trajectories: List[Trajectory], return_experience: bool = True, **kwargs) -> List[Experience]:
+ experiences: List[Experience] = self._extract_experiences(trajectories, **kwargs)
- def execute(self, trajectories: List[Trajectory], return_samples: bool = True, **kwargs) -> List[Sample]:
- samples: List[Sample] = self.extract_samples(trajectories, **kwargs)
- self.insert_into_vector_store(samples, **kwargs)
-
- if return_samples:
- return samples
+ nodes: List[VectorStoreNode] = [x.to_vector_store_node() for x in experiences]
+ self.vector_store.insert(nodes, **kwargs)
+ if return_experience:
+ return experiences
return []
-class MockSummarizer(BaseSummarizer):
-
- def execute(self, trajectories: List[Trajectory], return_samples: bool = True, **kwargs) -> List[Sample]:
- tip_message = SummaryMessage(content="I am a mock summarizer.")
-
- if return_samples:
- return [Sample(steps=[tip_message])]
-
- return []
+SUMMARIZER_REGISTRY = Registry[BaseSummarizer]("summarizer")
diff --git a/experiencemaker/module/summarizer/simple_summarizer.py b/experiencemaker/module/summarizer/simple_summarizer.py
new file mode 100644
index 00000000..3bfb820f
--- /dev/null
+++ b/experiencemaker/module/summarizer/simple_summarizer.py
@@ -0,0 +1,68 @@
+from pathlib import Path
+from typing import List
+
+from loguru import logger
+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.schema.experience import Experience
+from experiencemaker.schema.trajectory import Trajectory, Message, ActionMessage
+from experiencemaker.utils.util_function import get_html_match_content
+
+
+class SimpleSummarizer(BaseSummarizer, PromptMixin):
+ max_retries: int = Field(default=5, description="max retries")
+ prompt_file_path: Path = Path(__file__).parent / "simple_summarizer_prompt.yaml"
+
+ def _extract_trajectory_experience(self, trajectory: Trajectory) -> Experience | None:
+ step_content_collector: List[str] = []
+
+ for step in trajectory.steps:
+ step_index = len(step_content_collector)
+
+ if step.role is Role.ASSISTANT:
+ line = f"### step.{step_index} role={step.role.value} content=\n{step.content}\n{step.reasoning_content}\n"
+ if step.tool_calls:
+ for tool_call in step.tool_calls:
+ line += f" - tool call={tool_call.name}\n params={tool_call.arguments}\n"
+ step_content_collector.append(line)
+
+ elif step.role is Role.USER:
+ line = f"### step.{step_index} role={step.role.value} content=\n{step.content}\n"
+ step_content_collector.append(line)
+
+ elif step.role is Role.TOOL:
+ line = f"### step.{step_index} role={step.role.value} tool call result=\n{step.content}\n"
+ step_content_collector.append(line)
+
+ prompt = self.prompt_format(prompt_name="summary_prompt",
+ query=trajectory.query,
+ execution_process="\n".join(step_content_collector).strip(),
+ answer=trajectory.answer)
+
+ for i in range(self.max_retries):
+ action_message: ActionMessage = self.llm.chat(messages=[Message(content=prompt)])
+ experience_str = get_html_match_content(action_message.content, key="experience")
+ condition_str = get_html_match_content(action_message.content, key="condition")
+ if experience_str and condition_str:
+ return Experience(experience_workspace_id=self.workspace_id,
+ experience_role=self.llm.model_name,
+ experience_desc=condition_str,
+ experience_content=experience_str)
+ else:
+ logger.warning(f"action_message.content={action_message.content} re.search failed.")
+
+ return None
+
+ def _extract_experiences(self, trajectories: List[Trajectory], **kwargs) -> List[Experience]:
+ experiences: List[Experience] = []
+ for trajectory in trajectories:
+ experience: Experience = self._extract_trajectory_experience(trajectory)
+ if experience:
+ experiences.append(experience)
+ return experiences
+
+
+SUMMARIZER_REGISTRY.register(SimpleSummarizer, "simple")
diff --git a/experiencemaker/module/summarizer/simple_summarizer_prompt.yaml b/experiencemaker/module/summarizer/simple_summarizer_prompt.yaml
new file mode 100644
index 00000000..3a0f7283
--- /dev/null
+++ b/experiencemaker/module/summarizer/simple_summarizer_prompt.yaml
@@ -0,0 +1,22 @@
+summary_prompt: |
+ # Role
+ You are a helpful assistant named BeyondAgent.
+
+ # User Question
+ {query}
+
+ # Execution Process
+ {execution_process}
+
+ # Answer
+ {answer}
+
+ # Task
+ Reflect on the strengths and weaknesses of the **Execution Process** and the **Answer** based on the **User Question**.
+ Finally, summarize generalized experience from handling such problems to accumulate experience for future similar tasks.
+ The experience should be broadly applicable, such as how to use tools effectively or approaches to solving certain types of problems.
+ Also, specify the conditions or scenarios in which these experience are applicable.
+
+ # Output Format
+ Output the scenarios or conditions in which applying this experience would be particularly effective...
+ Output generalized experience, concise content is required...
\ No newline at end of file
diff --git a/experiencemaker/module/summarizer/step_summarizer.py b/experiencemaker/module/summarizer/step_summarizer.py
new file mode 100644
index 00000000..991ca68a
--- /dev/null
+++ b/experiencemaker/module/summarizer/step_summarizer.py
@@ -0,0 +1,580 @@
+import re
+import uuid
+import json
+from typing import List, Dict, Any, Optional, Tuple
+from datetime import datetime
+from loguru import logger
+from pydantic import Field, model_validator
+
+from experiencemaker.enumeration.role import Role
+from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer
+from experiencemaker.schema.trajectory import Trajectory, Sample, SummaryMessage, Message
+from experiencemaker.schema.vector_store_node import VectorStoreNode
+from experiencemaker.storage.es_vector_store import EsVectorStore
+from experiencemaker.storage.file_vector_store import FileVectorStore
+
+
+class StepSummarizer(BaseSummarizer):
+ """
+ Step-level experience extractor that focuses on extracting reusable experiences
+ from individual steps or step sequences in trajectories
+ """
+
+ # Vector Store 配置
+ vector_store_type: str = Field(default="file_vector_store")
+ vector_store_hosts: str | List[str] = Field(default="http://localhost:9200")
+ vector_store_index_name: str = Field(default="step_experience_store")
+ store_dir: str = Field(default="./step_experiences/")
+
+ # 功能开关
+ enable_step_segmentation: bool = Field(default=False)
+ enable_similarity_search: bool = Field(default=False)
+ enable_experience_validation: bool = Field(default=True)
+
+ # llm retries
+ max_retries: int = Field(default=3)
+
+ @model_validator(mode="after")
+ def init_vector_store(self):
+ """initialize"""
+ if self.vector_store_type == "file_vector_store":
+ self.vector_store = FileVectorStore(
+ embedding_model=self.embedding_model,
+ index_name=self.vector_store_index_name,
+ store_dir=self.store_dir
+ )
+ elif self.vector_store_type == "es_vector_store":
+ self.vector_store = EsVectorStore(
+ embedding_model=self.embedding_model,
+ index_name=self.vector_store_index_name,
+ hosts=self.vector_store_hosts
+ )
+ else:
+ raise ValueError(f"Unknown vector store type: {self.vector_store_type}")
+
+ return self
+
+ def extract_step_experiences_from_success(self, trajectories: List[Trajectory], **kwargs) -> List[SummaryMessage]:
+ """Extract step-level experiences from successful samples"""
+ logger.info(f"Extracting step experiences from {len(trajectories)} successful trajectories")
+
+ all_experiences = []
+ for trajectory in trajectories:
+ step_sequences = self._segment_trajectory_into_steps(trajectory)
+
+ for step_seq in step_sequences:
+ try:
+ prompt = self.prompt_handler.success_step_experience_prompt.format(
+ query=trajectory.query,
+ step_sequence=self._format_step_sequence(step_seq),
+ context=self._get_trajectory_context(trajectory, step_seq),
+ outcome="successful"
+ )
+
+ experience = self._extract_with_llm(prompt, "success")
+ if experience:
+ all_experiences.extend(experience)
+
+ except Exception as e:
+ logger.error(f"Error extracting success experience: {e}")
+ continue
+
+ return all_experiences
+
+ def extract_step_experiences_from_failure(self, trajectories: List[Trajectory], **kwargs) -> List[SummaryMessage]:
+ """Extract step-level experiences from failed samples"""
+ logger.info(f"Extracting step experiences from {len(trajectories)} failed trajectories")
+
+ all_experiences = []
+ for trajectory in trajectories:
+ step_sequences = self._segment_trajectory_into_steps(trajectory)
+
+ for step_seq in step_sequences:
+ try:
+ prompt = self.prompt_handler.failure_step_experience_prompt.format(
+ query=trajectory.query,
+ step_sequence=self._format_step_sequence(step_seq),
+ context=self._get_trajectory_context(trajectory, step_seq),
+ outcome="failed"
+ )
+
+ experience = self._extract_with_llm(prompt, "failure")
+ if experience:
+ all_experiences.extend(experience)
+
+ except Exception as e:
+ logger.error(f"Error extracting failure experience: {e}")
+ continue
+
+ return all_experiences
+
+ def extract_step_experiences_from_comparison(self,
+ success_trajectories: List[Trajectory],
+ failure_trajectories: List[Trajectory],
+ **kwargs) -> List[SummaryMessage]:
+ """Extract step-level experiences from comparative samples"""
+ logger.info(f"Extracting comparative step experiences from {len(success_trajectories)} success "
+ f"and {len(failure_trajectories)} failure trajectories")
+
+ all_experiences = []
+
+ # Find similar step sequences for comparison
+ similar_step_pairs = self._find_similar_step_sequences(success_trajectories, failure_trajectories)
+
+ for success_steps, failure_steps, similarity_score in similar_step_pairs:
+ try:
+ prompt = self.prompt_handler.comparative_step_experience_prompt.format(
+ success_steps=self._format_step_sequence(success_steps),
+ failure_steps=self._format_step_sequence(failure_steps),
+ similarity_score=similarity_score
+ )
+
+ experience = self._extract_with_llm(prompt, "comparative")
+ if experience:
+ all_experiences.extend(experience)
+
+ except Exception as e:
+ logger.error(f"Error extracting comparative experience: {e}")
+ continue
+
+ return all_experiences
+
+ def extract_step_experiences_general(self, trajectories: List[Trajectory], **kwargs) -> List[SummaryMessage]:
+ """Extract general step experiences when no labels are provided"""
+ logger.info(f"Extracting general step experiences from {len(trajectories)} trajectories")
+
+ all_experiences = []
+
+ for trajectory in trajectories:
+ step_sequences = self._segment_trajectory_into_steps(trajectory)
+
+ for step_seq in step_sequences:
+ try:
+ prompt = self.prompt_handler.general_step_experience_prompt.format(
+ query=trajectory.query,
+ step_sequence=self._format_step_sequence(step_seq),
+ context=self._get_trajectory_context(trajectory, step_seq)
+ )
+
+ experience = self._extract_with_llm(prompt, "general")
+ if experience:
+ all_experiences.extend(experience)
+
+ except Exception as e:
+ logger.error(f"Error extracting general experience: {e}")
+ continue
+
+ return all_experiences
+
+ def validate_experiences(self, experiences: List[SummaryMessage], **kwargs) -> List[SummaryMessage]:
+ """Validate the quality and validity of extracted experiences"""
+ if not self.enable_experience_validation:
+ return experiences
+
+ logger.info(f"Validating {len(experiences)} extracted experiences")
+
+ validated_experiences = []
+
+ for experience in experiences:
+ try:
+ validation_result = self._validate_single_experience(experience)
+
+ if validation_result["is_valid"]:
+ # Add validation info to metadata
+ experience.metadata.update({
+ "validation_score": validation_result["score"],
+ "validation_feedback": validation_result["feedback"],
+ "validated_at": datetime.now().isoformat()
+ })
+ validated_experiences.append(experience)
+ else:
+ logger.warning(f"Experience validation failed: {validation_result['reason']}")
+
+ except Exception as e:
+ logger.error(f"Error validating experience: {e}")
+ continue
+
+ logger.info(f"Validated {len(validated_experiences)} out of {len(experiences)} experiences")
+ return validated_experiences
+
+ def store_experiences(self, experiences: List[SummaryMessage], **kwargs):
+ """Store experiences into vector storage"""
+ if not experiences:
+ logger.warning("No experiences to store")
+ return
+
+ # Deduplication
+ unique_experiences = self._deduplicate_experiences(experiences)
+ logger.info(f"Storing {len(unique_experiences)} unique experiences (deduplicated from {len(experiences)})")
+
+ # Convert to storage nodes
+ nodes = []
+ for exp in unique_experiences:
+ node = VectorStoreNode(
+ content=exp.content,
+ metadata={
+ **exp.metadata,
+ "stored_at": datetime.now().isoformat(),
+ "experience_type": "step_level"
+ }
+ )
+ nodes.append(node)
+
+ # Store to vector database
+ refresh_index = kwargs.get("refresh_index", True)
+ self.vector_store.insert(nodes, refresh_index=refresh_index)
+ logger.info(f"Successfully stored {len(nodes)} step experiences")
+
+ def execute(self, trajectories: List[Trajectory], **kwargs) -> List[Sample]:
+ """Execute complete step-level experience extraction pipeline"""
+ logger.info(f"Starting step-level experience extraction pipeline for {len(trajectories)} trajectories")
+
+ all_experiences = []
+
+ # Classify trajectories based on trajectory.done
+ success_trajectories = [traj for traj in trajectories if traj.done]
+ failure_trajectories = [traj for traj in trajectories if not traj.done]
+
+ # Process success and failure samples separately
+ if success_trajectories:
+ success_experiences = self.extract_step_experiences_from_success(success_trajectories, **kwargs)
+ all_experiences.extend(success_experiences)
+
+ if failure_trajectories:
+ failure_experiences = self.extract_step_experiences_from_failure(failure_trajectories, **kwargs)
+ all_experiences.extend(failure_experiences)
+
+ # Comparative analysis (if similarity search is enabled)
+ if success_trajectories and failure_trajectories and self.enable_similarity_search:
+ comparative_experiences = self.extract_step_experiences_from_comparison(
+ success_trajectories, failure_trajectories, **kwargs
+ )
+ all_experiences.extend(comparative_experiences)
+
+ # Validate experiences
+ if self.enable_experience_validation:
+ validated_experiences = self.validate_experiences(all_experiences, **kwargs)
+ else:
+ validated_experiences = all_experiences
+
+ # Store experiences
+ if validated_experiences:
+ self.store_experiences(validated_experiences, **kwargs)
+
+ # Construct return result
+ return [Sample(steps=validated_experiences)]
+
+ # ========== Helper Methods ==========
+
+ def _segment_trajectory_into_steps(self, trajectory: Trajectory) -> List[List[Message]]:
+ """Segment trajectory into meaningful step sequences"""
+ if not self.enable_step_segmentation:
+ # If segmentation is not enabled, return the entire trajectory as one step sequence
+ return [trajectory.steps]
+
+ try:
+ # Use LLM for segmentation
+ trajectory_content = self._format_trajectory_content(trajectory)
+
+ prompt = self.prompt_handler.step_segmentation_prompt.format(
+ query=trajectory.query,
+ trajectory_content=trajectory_content,
+ total_steps=len(trajectory.steps)
+ )
+
+ response = self.llm.chat([Message(role=Role.USER, content=prompt)])
+
+ # Parse segmentation points
+ segment_points = self._parse_segmentation_response(response.content)
+
+ # Segment trajectory based on split points
+ step_sequences = []
+ start_idx = 0
+
+ for end_idx in segment_points:
+ if start_idx < end_idx <= len(trajectory.steps):
+ step_sequences.append(trajectory.steps[start_idx:end_idx])
+ start_idx = end_idx
+
+ # Add remaining steps
+ if start_idx < len(trajectory.steps):
+ step_sequences.append(trajectory.steps[start_idx:])
+
+ return step_sequences if step_sequences else [trajectory.steps]
+
+ except Exception as e:
+ logger.error(f"Error in step segmentation: {e}, falling back to whole trajectory")
+ return [trajectory.steps]
+
+ def _parse_segmentation_response(self, response: str) -> List[int]:
+ """Parse segmentation response to extract split point positions"""
+ segment_points = []
+
+ # Try to extract JSON format split points
+ json_pattern = r'```json\s*([\s\S]*?)\s*```'
+ json_blocks = re.findall(json_pattern, response)
+
+ if json_blocks:
+ try:
+ parsed = json.loads(json_blocks[0])
+ if isinstance(parsed, dict) and "segment_points" in parsed:
+ segment_points = parsed["segment_points"]
+ elif isinstance(parsed, list):
+ segment_points = parsed
+ except json.JSONDecodeError:
+ pass
+
+ # If JSON parsing fails, try to extract numbers
+ if not segment_points:
+ numbers = re.findall(r'\b\d+\b', response)
+ segment_points = [int(num) for num in numbers if int(num) > 0]
+
+ return sorted(list(set(segment_points))) # Remove duplicates and sort
+
+ def _format_step_sequence(self, step_sequence: List[Message]) -> str:
+ """Format step sequence to string"""
+ formatted_steps = []
+ for i, step in enumerate(step_sequence):
+ step_info = f"Step {i + 1} [{step.role.value}]:"
+
+ if hasattr(step, 'reasoning_content') and step.reasoning_content:
+ step_info += f"\nReasoning: {step.reasoning_content}"
+
+ step_info += f"\nContent: {step.content}"
+
+ if hasattr(step, 'tool_calls') and step.tool_calls:
+ for tool_call in step.tool_calls:
+ step_info += f"\nTool: {tool_call.name}({tool_call.arguments})"
+
+ formatted_steps.append(step_info)
+
+ return "\n\n".join(formatted_steps)
+
+ def _get_trajectory_context(self, trajectory: Trajectory, step_sequence: List[Message]) -> str:
+ """Get context of step sequence within trajectory"""
+ # Find position of step sequence in trajectory
+ start_idx = 0
+ for i, step in enumerate(trajectory.steps):
+ if step == step_sequence[0]:
+ start_idx = i
+ break
+
+ # Extract before and after context
+ context_before = trajectory.steps[max(0, start_idx - 2):start_idx]
+ context_after = trajectory.steps[start_idx + len(step_sequence):start_idx + len(step_sequence) + 2]
+
+ context = f"Query: {trajectory.query}\n"
+
+ if context_before:
+ context += "Previous steps:\n" + "\n".join([f"- {step.content[:100]}..." for step in context_before]) + "\n"
+
+ if context_after:
+ context += "Following steps:\n" + "\n".join([f"- {step.content[:100]}..." for step in context_after])
+
+ return context
+
+ def _format_trajectory_content(self, trajectory: Trajectory) -> str:
+ """Format trajectory content to string"""
+ content = ""
+ for i, step in enumerate(trajectory.steps):
+ content += f"Step {i + 1} ({step.role.value}):\n{step.content}\n\n"
+ return content
+
+ def _find_similar_step_sequences(self, success_trajectories: List[Trajectory],
+ failure_trajectories: List[Trajectory]) -> List[Tuple]:
+ """Use embedding model to find similar step sequences for comparison"""
+ if not self.enable_similarity_search:
+ return []
+
+ try:
+ similar_pairs = []
+
+ # Get step sequences from success and failure trajectories
+ success_step_sequences = []
+ for traj in success_trajectories:
+ sequences = self._segment_trajectory_into_steps(traj)
+ success_step_sequences.extend(sequences)
+
+ failure_step_sequences = []
+ for traj in failure_trajectories:
+ sequences = self._segment_trajectory_into_steps(traj)
+ failure_step_sequences.extend(sequences)
+
+ # Limit comparison count to avoid computation overload
+ max_sequences = 5
+ success_step_sequences = success_step_sequences[:max_sequences]
+ failure_step_sequences = failure_step_sequences[:max_sequences]
+
+ if not success_step_sequences or not failure_step_sequences:
+ return []
+
+ # Generate text representations of step sequences for embedding
+ success_texts = [self._format_step_sequence(seq) for seq in success_step_sequences]
+ failure_texts = [self._format_step_sequence(seq) for seq in failure_step_sequences]
+
+ # Get embeddings
+ success_embeddings = self.embedding_model.get_embeddings(success_texts)
+ failure_embeddings = self.embedding_model.get_embeddings(failure_texts)
+
+ # Calculate similarity and find most similar pairs
+ for i, s_emb in enumerate(success_embeddings):
+ for j, f_emb in enumerate(failure_embeddings):
+ similarity = self._calculate_cosine_similarity(s_emb, f_emb)
+
+ if similarity > 0.3: # Similarity threshold
+ similar_pairs.append((
+ success_step_sequences[i],
+ failure_step_sequences[j],
+ similarity
+ ))
+
+ # Return top 3 most similar pairs
+ return sorted(similar_pairs, key=lambda x: x[2], reverse=True)[:3]
+
+ except Exception as e:
+ logger.error(f"Error finding similar step sequences: {e}")
+ return []
+
+ def _calculate_cosine_similarity(self, embedding1: List[float], embedding2: List[float]) -> float:
+ """Calculate cosine similarity between two embedding vectors"""
+ try:
+ import numpy as np
+
+ vec1 = np.array(embedding1)
+ vec2 = np.array(embedding2)
+
+ # Calculate cosine similarity
+ dot_product = np.dot(vec1, vec2)
+ norm1 = np.linalg.norm(vec1)
+ norm2 = np.linalg.norm(vec2)
+
+ if norm1 == 0 or norm2 == 0:
+ return 0.0
+
+ return dot_product / (norm1 * norm2)
+
+ except Exception as e:
+ logger.error(f"Error calculating cosine similarity: {e}")
+ return 0.0
+ import json
+
+ def _extract_with_llm(self, prompt: str, experience_type: str) -> List[SummaryMessage]:
+ for attempt in range(self.max_retries):
+ try:
+ response = self.llm.chat([Message(role=Role.USER, content=prompt)])
+ experiences = self._parse_experience_response(response.content, experience_type)
+
+ if experiences:
+ return experiences
+
+ except Exception as e:
+ logger.warning(f"Attempt {attempt + 1} failed for experience extraction: {e}")
+
+ logger.error(f"Failed to extract experience after {self.max_retries} attempts")
+ return []
+
+ def _parse_experience_response(self, response: str, experience_type: str) -> List[SummaryMessage]:
+ """解析经验抽取响应"""
+ experiences = []
+
+ try:
+ # 尝试提取JSON格式的经验
+ json_pattern = r'```json\s*([\s\S]*?)\s*```'
+ json_blocks = re.findall(json_pattern, response)
+
+ for block in json_blocks:
+ try:
+ parsed = json.loads(block)
+ if isinstance(parsed, list):
+ for exp_data in parsed:
+ experience = self._create_experience_message(exp_data, experience_type)
+ if experience:
+ experiences.append(experience)
+ else:
+ experience = self._create_experience_message(parsed, experience_type)
+ if experience:
+ experiences.append(experience)
+ except json.JSONDecodeError:
+ continue
+
+ except Exception as e:
+ logger.error(f"Error parsing experience response: {e}")
+
+ return experiences
+
+ def _create_experience_message(self, exp_data: Dict[str, Any], experience_type: str) -> Optional[SummaryMessage]:
+ """创建经验消息对象"""
+ try:
+ condition = exp_data.get("when_to_use", exp_data.get("condition", ""))
+ experience_content = exp_data.get("experience", exp_data.get("tip_content", exp_data.get("tips", "")))
+
+ if not condition or not experience_content:
+ return None
+
+ metadata = {
+ "experience": experience_content,
+ "experience_type": experience_type,
+ "tags": exp_data.get("tags", []),
+ "confidence": exp_data.get("confidence", 0.5),
+ "extracted_at": datetime.now().isoformat(),
+ "experience_id": str(uuid.uuid4())
+ }
+
+ return SummaryMessage(content=condition, metadata=metadata)
+
+ except Exception as e:
+ logger.error(f"Error creating experience message: {e}")
+ return None
+
+ def _validate_single_experience(self, experience: SummaryMessage) -> Dict[str, Any]:
+ """验证单个经验的有效性"""
+ try:
+ prompt = self.prompt_handler.experience_validation_prompt.format(
+ condition=experience.content,
+ experience_content=experience.metadata.get("experience", ""),
+ experience_type=experience.metadata.get("experience_type", ""),
+ tags=experience.metadata.get("tags", [])
+ )
+
+ response = self.llm.chat([Message(role=Role.USER, content=prompt)])
+
+ # 解析验证结果
+ is_valid = "valid" in response.content.lower() and "invalid" not in response.content.lower()
+ score_match = re.search(r'score[:\s]*([0-9.]+)', response.content.lower())
+ score = float(score_match.group(1)) if score_match else 0.5
+
+ return {
+ "is_valid": is_valid and score > 0.3,
+ "score": score,
+ "feedback": response.content,
+ "reason": "" if is_valid else "Low validation score or marked as invalid"
+ }
+
+ except Exception as e:
+ logger.error(f"Error validating experience: {e}")
+ return {"is_valid": False, "score": 0.0, "feedback": "", "reason": str(e)}
+
+ def _deduplicate_experiences(self, experiences: List[SummaryMessage]) -> List[SummaryMessage]:
+ unique_experiences = []
+ seen_contents = set()
+
+ for exp in experiences:
+ content_hash = hash(exp.content)
+
+ if content_hash not in seen_contents:
+ seen_contents.add(content_hash)
+ unique_experiences.append(exp)
+
+ return unique_experiences
+
+ def extract_samples(self, trajectories: List[Trajectory], **kwargs) -> List[Sample]:
+ experiences = self.execute(trajectories, **kwargs)
+ return [Sample(steps=experiences)] if experiences else []
+
+ def insert_into_vector_store(self, samples: List[Sample], **kwargs):
+ all_experiences = []
+ for sample in samples:
+ all_experiences.extend(sample.steps)
+
+ if all_experiences:
+ self.store_experiences(all_experiences, **kwargs)
\ 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/experience.py b/experiencemaker/schema/experience.py
new file mode 100644
index 00000000..017a6fb1
--- /dev/null
+++ b/experiencemaker/schema/experience.py
@@ -0,0 +1,83 @@
+import datetime
+from typing import List
+from uuid import uuid4
+
+from pydantic import BaseModel, Field
+
+from experiencemaker.schema.vector_store_node import VectorStoreNode
+
+
+class ExperienceFunctionArg(BaseModel):
+ arg_name: str = Field(default=..., description="argument name")
+ arg_type: str = Field(default=..., description="argument type, like: 'str', 'int', 'bool'")
+ required: bool = Field(default=True, description="whether the argument is required")
+
+
+class ExperienceFunction(BaseModel):
+ func_code: str = Field(default=..., description="function code")
+ func_name: str = Field(default=..., description="function name")
+ func_args: List[ExperienceFunctionArg] = Field(default_factory=list, description="function arguments")
+
+
+class Experience(BaseModel):
+ experience_id: str = Field(default_factory=lambda: uuid4().hex, description="experience unique id")
+ experience_workspace_id: str = Field(default="", description="unique workspace id")
+ experience_role: str = Field(default="", description="experience role")
+ experience_desc: str = Field(default="", description="use condition/purpose. It will be used in vector matching")
+ experience_content: str | bytes = Field(default="", description="content of the experience")
+ experience_function: ExperienceFunction | None = Field(default=None, description="experience function(optional)")
+ experience_score: float = Field(default=0.0, description="score of the experience")
+ experience_created_time: str = Field(default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S"))
+ experience_modified_time: str = Field(default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S"))
+ metadata: dict = Field(default_factory=dict, description="additional metadata")
+
+ def to_vector_store_node(self) -> VectorStoreNode:
+ metadata: dict = {
+ "experience_role": self.experience_role,
+ "experience_content": self.experience_content,
+ "experience_function": self.experience_function.model_dump(),
+ "experience_score": self.experience_score,
+ "experience_created_time": self.experience_created_time,
+ "experience_modified_time": self.experience_modified_time,
+ "metadata": self.metadata,
+ }
+ return VectorStoreNode(
+ unique_id=self.experience_id,
+ workspace_id=self.experience_workspace_id,
+ content=self.experience_desc,
+ metadata=metadata)
+
+ @classmethod
+ def from_vector_store_node(cls, node: VectorStoreNode) -> "Experience":
+ return cls(
+ experience_id=node.unique_id,
+ experience_workspace_id=node.workspace_id,
+ experience_role=node.metadata.get("experience_role", ""),
+ experience_desc=node.content,
+ experience_content=node.metadata.get("experience_content", ""),
+ experience_function=node.metadata.get("experience_function", None),
+ experience_score=node.metadata.get("experience_score", 0.0),
+ experience_created_time=node.metadata.get("experience_created_time", ""),
+ experience_modified_time=node.metadata.get("experience_modified_time", ""),
+ metadata=node.metadata.get("metadata", {}))
+
+
+if __name__ == "__main__":
+ e1 = Experience(
+ experience_workspace_id="w_1024",
+ experience_role="qwen3",
+ experience_desc="test desc",
+ experience_content="test content",
+ experience_function=ExperienceFunction(
+ func_code="def a():\n return",
+ func_name="a",
+ func_args=[ExperienceFunctionArg(arg_name="x", arg_type="str", required=True)]
+ ),
+ experience_score=0.99,
+ metadata={"haha": 1}
+ )
+ print(e1.model_dump_json(indent=2))
+ v1 = e1.to_vector_store_node()
+ print(v1.model_dump_json(indent=2))
+ e2 = Experience.from_vector_store_node(v1)
+ print(e2.model_dump_json(indent=2))
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/schema/request.py b/experiencemaker/schema/request.py
index 71f6445e..57e23fcc 100644
--- a/experiencemaker/schema/request.py
+++ b/experiencemaker/schema/request.py
@@ -1,13 +1,12 @@
from abc import ABC
from typing import List
-from pydantic import Field
+from pydantic import BaseModel, Field
-from experiencemaker.schema.module_loader import ModuleLoader
from experiencemaker.schema.trajectory import Trajectory
-class BaseRequest(ModuleLoader, ABC):
+class BaseRequest(BaseModel, ABC):
metadata: dict = Field(default_factory=dict)
@@ -21,4 +20,4 @@ class ContextGeneratorRequest(BaseRequest):
class SummarizerRequest(BaseRequest):
trajectories: List[Trajectory] = Field(default_factory=dict)
- return_samples: bool = Field(default=False)
+ return_experience: bool = Field(default=True)
diff --git a/experiencemaker/schema/response.py b/experiencemaker/schema/response.py
index 8a0e3caa..feed1b44 100644
--- a/experiencemaker/schema/response.py
+++ b/experiencemaker/schema/response.py
@@ -3,7 +3,8 @@ from typing import List
from pydantic import BaseModel, Field
-from experiencemaker.schema.trajectory import Trajectory, ContextMessage, Sample
+from experiencemaker.schema.experience import Experience
+from experiencemaker.schema.trajectory import Trajectory, ContextMessage
class BaseResponse(BaseModel, ABC):
@@ -20,4 +21,4 @@ class ContextGeneratorResponse(BaseResponse):
class SummarizerResponse(BaseResponse):
- extract_samples: List[Sample] = Field(default_factory=list)
+ experiences: List[Experience] = Field(default_factory=list)
diff --git a/experiencemaker/schema/trajectory.py b/experiencemaker/schema/trajectory.py
index 8842eb08..14018b56 100644
--- a/experiencemaker/schema/trajectory.py
+++ b/experiencemaker/schema/trajectory.py
@@ -49,19 +49,12 @@ class Message(BaseModel):
content: str | bytes = Field(default="")
reasoning_content: str = Field(default="")
tool_calls: List[ToolCall] = Field(default_factory=list)
- timestamp: str = Field(
- default_factory=lambda: datetime.datetime.now().strftime(
- "%Y-%m-%d %H:%M:%S.%f",
- ),
- )
+ timestamp: str = Field(default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S.%f"))
metadata: dict = Field(default_factory=dict)
@property
def simple_dict(self) -> dict:
- result = {
- "role": self.role.value,
- "content": self.content,
- }
+ result = {"role": self.role.value, "content": self.content}
if self.tool_calls:
result["tool_calls"] = [x.simple_dict for x in self.tool_calls]
return result
@@ -77,10 +70,7 @@ class StateMessage(Message):
@property
def simple_dict(self) -> dict:
- result = {
- "role": self.role.value,
- "content": self.content,
- }
+ result = super().simple_dict
if self.tool_call_id:
result["tool_call_id"] = self.tool_call_id
return result
@@ -110,6 +100,7 @@ class Sample(BaseModel):
class Trajectory(BaseModel):
id: str = Field(default_factory=lambda: uuid4().hex)
steps: List[Message] = Field(default_factory=list)
+ current_step: int = Field(default=0)
done: bool = Field(default=False)
query: str = Field(default="")
@@ -120,7 +111,9 @@ class Trajectory(BaseModel):
self.steps.append(step)
def reset(self):
+ self.id = uuid4().hex
self.steps.clear()
+ self.current_step = 0
self.done = False
self.query = ""
self.answer = ""
diff --git a/experiencemaker/service/experience_maker_service.py b/experiencemaker/service/experience_maker_service.py
new file mode 100644
index 00000000..e2a6b98b
--- /dev/null
+++ b/experiencemaker/service/experience_maker_service.py
@@ -0,0 +1,164 @@
+from typing import List
+
+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.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 ExperienceMakerService(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_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)
+ agent_wrapper: AgentWrapperMixin | None = Field(default=None)
+ context_generator: BaseContextGenerator | None = Field(default=None)
+ summarizer: BaseSummarizer | None = Field(default=None)
+
+ @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_modules}"
+ 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:
+ if "llm" in config:
+ llm_config = config.pop("llm")
+ llm = self.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_modules}"
+ 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:
+ if "embedding_model" in config:
+ embedding_model_config = config.pop("embedding_model")
+ embedding_model = self.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:
+ 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)
+ 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:
+ if "vector_store" in config:
+ vector_store_config = config.pop("vector_store")
+ vector_store = self.init_vector_store(vector_store_config)
+ 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:
+ 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)
+ 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
+
+ def init_summarizer(self, summarizer_config: 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)
+ logger.info(f"summarizer is inited with backend={backend} params={summarizer_config}")
+ return summarizer
+
+ def init_agent_wrapper(self, agent_wrapper_config: 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)
+
+ agent_wrapper: AgentWrapperMixin = AGENT_WRAPPER_REGISTRY[backend](
+ **agent_wrapper_config, llm=llm, context_generator=self.context_generator)
+ 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)
+
+ if self.embedding_model_config:
+ self.embedding_model = self.init_embedding_model(self.embedding_model_config)
+
+ if self.vector_store_config:
+ self.vector_store = self.init_vector_store(self.vector_store_config)
+
+ if self.context_generator_config:
+ self.context_generator = self.init_context_generator(self.context_generator_config)
+
+ if self.summarizer_config:
+ self.summarizer = self.init_summarizer(self.summarizer_config)
+
+ if self.agent_wrapper_config:
+ self.agent_wrapper = self.init_agent_wrapper(self.agent_wrapper_config)
+
+ 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)
+ 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)
+ 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)
+ return SummarizerResponse(experiences=experiences)
diff --git a/experiencemaker/service/http_service.py b/experiencemaker/service/http_service.py
new file mode 100644
index 00000000..155df9c0
--- /dev/null
+++ b/experiencemaker/service/http_service.py
@@ -0,0 +1,50 @@
+import argparse
+import json
+
+import uvicorn
+from fastapi import FastAPI
+
+from experiencemaker.schema.request import AgentWrapperRequest, ContextGeneratorRequest, SummarizerRequest
+from experiencemaker.schema.response import AgentWrapperResponse, ContextGeneratorResponse, SummarizerResponse
+from experiencemaker.service.experience_maker_service import ExperienceMakerService
+from experiencemaker.utils.file_handler import FileHandler
+
+app = FastAPI()
+service: ExperienceMakerService | None = None
+
+@app.post('/agent_wrapper', response_model=AgentWrapperResponse)
+def call_agent_wrapper(request: AgentWrapperRequest):
+ return service.call_agent_wrapper(request)
+
+
+@app.post('/context_generator', response_model=ContextGeneratorResponse)
+def call_context_generator(request: ContextGeneratorRequest):
+ return service.call_context_generator(request)
+
+
+@app.post('/summarizer', response_model=SummarizerResponse)
+def call_summarizer(request: SummarizerRequest):
+ return service.call_summarizer(request)
+
+
+# launch with: python -m experiencemaker.service.http_service
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser()
+ parser.add_argument('--config', type=str, help='config dict')
+ parser.add_argument('--config_path', type=str, help='config load path')
+ args = parser.parse_args()
+
+ if args.config_path:
+ config = FileHandler(file_path=args.config_path).load()
+ elif args.config:
+ config = json.loads(args.config)
+ else:
+ raise RuntimeError("both config and config_path are not specified")
+
+ service = ExperienceMakerService(**config)
+ uvicorn.run(app,
+ host=service.host,
+ port=service.port,
+ timeout_keep_alive=service.timeout_keep_alive,
+ limit_concurrency=service.limit_concurrency)
diff --git a/experiencemaker/service/model_service.py b/experiencemaker/service/model_service.py
deleted file mode 100644
index 3f00afe0..00000000
--- a/experiencemaker/service/model_service.py
+++ /dev/null
@@ -1,42 +0,0 @@
-from experiencemaker.utils.logger import init_logger
-init_logger()
-
-from typing import List
-from fastapi import FastAPI
-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
-import uvicorn
-
-app = FastAPI()
-
-
-@app.post('/agent_wrapper', response_model=AgentWrapperResponse)
-def call_agent_wrapper(request: AgentWrapperRequest):
- module: BaseAgentWrapper = request.load_from_path()
- trajectory: Trajectory = module.execute(request.query, **request.metadata)
- return AgentWrapperResponse(trajectory=trajectory)
-
-
-@app.post('/context_generator', response_model=ContextGeneratorResponse)
-def call_context_generator(request: ContextGeneratorRequest):
- module: BaseContextGenerator = request.load_from_path()
- context_msg: ContextMessage = module.execute(request.trajectory, **request.metadata)
- return ContextGeneratorResponse(context_msg=context_msg)
-
-
-@app.post('/summarizer', response_model=SummarizerResponse)
-def call_summarizer(request: SummarizerRequest):
- module: BaseSummarizer = request.load_from_path()
- samples: List[Sample] = module.execute(request.trajectories, request.return_samples, **request.metadata)
- return SummarizerResponse(extract_samples=samples)
-
-
-if __name__ == '__main__':
- uvicorn.run(app, host="0.0.0.0", port=8000, timeout_keep_alive=600000, limit_concurrency=32)
-
-# launch with:
-# python -m experiencemaker.service.model_service
diff --git a/experiencemaker/storage/base_vector_store.py b/experiencemaker/storage/base_vector_store.py
index 9674c6cb..e4644aff 100644
--- a/experiencemaker/storage/base_vector_store.py
+++ b/experiencemaker/storage/base_vector_store.py
@@ -5,6 +5,7 @@ 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):
@@ -24,3 +25,6 @@ class BaseVectorStore(BaseModel, ABC):
def retrieve_by_query(self, query: str, top_k: int = 3, **kwargs) -> List[VectorStoreNode]:
raise NotImplementedError
+
+
+VECTOR_STORE_REGISTRY = Registry[BaseVectorStore]("vector_store")
diff --git a/experiencemaker/storage/es_vector_store.py b/experiencemaker/storage/es_vector_store.py
index 92526f9b..1e366057 100644
--- a/experiencemaker/storage/es_vector_store.py
+++ b/experiencemaker/storage/es_vector_store.py
@@ -6,8 +6,9 @@ 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
+from experiencemaker.storage.base_vector_store import BaseVectorStore, VECTOR_STORE_REGISTRY
class EsVectorStore(BaseVectorStore):
@@ -60,7 +61,7 @@ class EsVectorStore(BaseVectorStore):
node = VectorStoreNode(**doc["_source"])
node.unique_id = doc["_id"]
if "_score" in doc:
- node.metadata["score"] = doc["_score"] - 1
+ node.metadata["_score"] = doc["_score"] - 1
return node
def exist_id(self, doc_id: str):
@@ -103,7 +104,7 @@ class EsVectorStore(BaseVectorStore):
self.retrieve_filters.clear()
return self
- def insert(self, nodes: VectorStoreNode | List[VectorStoreNode], refresh_index: bool = False, **kwargs):
+ def insert(self, nodes: VectorStoreNode | List[VectorStoreNode], refresh_index: bool = True, **kwargs):
if isinstance(nodes, VectorStoreNode):
nodes = [nodes]
@@ -118,7 +119,7 @@ class EsVectorStore(BaseVectorStore):
if refresh_index:
self.refresh_index()
- def update(self, nodes: VectorStoreNode | List[VectorStoreNode], refresh_index: bool = False, **kwargs):
+ def update(self, nodes: VectorStoreNode | List[VectorStoreNode], refresh_index: bool = True, **kwargs):
if isinstance(nodes, VectorStoreNode):
nodes = [nodes]
@@ -167,3 +168,77 @@ class EsVectorStore(BaseVectorStore):
self.retrieve_filters.clear()
return nodes
+
+
+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 4544a7c9..41d0ca31 100644
--- a/experiencemaker/storage/file_vector_store.py
+++ b/experiencemaker/storage/file_vector_store.py
@@ -7,8 +7,9 @@ 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
+from experiencemaker.storage.base_vector_store import BaseVectorStore, VECTOR_STORE_REGISTRY
class FileVectorStore(BaseVectorStore):
@@ -39,22 +40,22 @@ class FileVectorStore(BaseVectorStore):
self.index_path.touch(exist_ok=True)
def load(self) -> List[VectorStoreNode]:
+ nodes = []
with self._thread_lock:
- nodes = []
with open(self.index_path) as f:
for line in f:
if line.strip():
nodes.append(VectorStoreNode(**json.loads(line)))
- return nodes
+ return nodes
def _load(self) -> List[VectorStoreNode]:
+ nodes = []
with self._thread_lock:
- nodes = []
with open(self.index_path) as f:
for line in f:
if line.strip():
nodes.append(VectorStoreNode(**json.loads(line)))
- return nodes
+ return nodes
def _dump(self, nodes: List[VectorStoreNode]):
with self._thread_lock:
@@ -130,3 +131,61 @@ class FileVectorStore(BaseVectorStore):
nodes = sorted(nodes, key=lambda x: x.metadata["score"], reverse=True)
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()
+
+ 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/__init__.py b/experiencemaker/tool/__init__.py
index f342c5ec..4df01cc7 100644
--- a/experiencemaker/tool/__init__.py
+++ b/experiencemaker/tool/__init__.py
@@ -1,6 +1,6 @@
-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.tool.code_tool import CodeTool
+from experiencemaker.tool.dashscope_search_tool import DashscopeSearchTool
+from experiencemaker.tool.terminate_tool import TerminateTool
from experiencemaker.utils.registry import Registry
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/file_handler.py b/experiencemaker/utils/file_handler.py
index 66850ac9..8ca8ac5c 100644
--- a/experiencemaker/utils/file_handler.py
+++ b/experiencemaker/utils/file_handler.py
@@ -7,7 +7,7 @@ from pydantic import BaseModel, Field, PrivateAttr
class FileHandler(BaseModel):
- file_path: str = Field(default=...)
+ file_path: str | Path = Field(default=...)
_obj: Any = PrivateAttr()
def __init__(self, **kwargs):
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/registry.py b/experiencemaker/utils/registry.py
index 65199cda..da9e8634 100644
--- a/experiencemaker/utils/registry.py
+++ b/experiencemaker/utils/registry.py
@@ -1,13 +1,21 @@
-from typing import Dict, Any, List
+from typing import Dict, List
+
+from typing import TypeVar, Generic
+
+T = TypeVar('T')
-class Registry(object):
+class Registry(Generic[T]):
def __init__(self, name: str):
self.name: str = name
- self.module_dict: Dict[str, Any] = {}
+ self.module_dict: Dict[str, T] = {}
- def register(self, module, module_name: str = None):
+ @property
+ def registered_modules(self) -> List[str]:
+ return sorted(self.module_dict.keys())
+
+ def register(self, module: T, module_name: str = None):
if module_name is None:
module_name = module.__name__
@@ -16,7 +24,7 @@ class Registry(object):
self.module_dict[module_name] = module
- def batch_register(self, modules: List[Any] | Dict[str, Any]):
+ def batch_register(self, modules: List[T] | Dict[str, T]):
if isinstance(modules, list):
module_name_dict = {m.__name__: m for m in modules}
@@ -27,6 +35,6 @@ class Registry(object):
raise NotImplementedError("Input must be a list or a dictionary.")
self.module_dict.update(module_name_dict)
- def __getitem__(self, module_name: str):
+ 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]
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