add unit test

This commit is contained in:
jinli.yl 2025-06-09 20:54:25 +08:00
parent 15da7cc730
commit 79a247e114
30 changed files with 349 additions and 347 deletions

2
.gitignore vendored
View file

@ -17,4 +17,4 @@ log/
.trash/
runs
logs
rag_nodes_index.jsonl

View file

@ -1,5 +0,0 @@
from experiencemaker.config.config_handler import ConfigHandler
agent_wrapper_config = ConfigHandler(module_name="agent_wrapper")
summarizer_config = ConfigHandler(module_name="summarizer")
context_generator_config = ConfigHandler(module_name="context_generator")

View file

@ -1,32 +0,0 @@
from pathlib import Path
from pydantic import BaseModel, Field, model_validator
from experiencemaker.utils.file_handler import FileHandler
class ConfigHandler(BaseModel):
module_name: str = Field(default=...)
config_dict: dict = Field(default_factory=dict)
@model_validator(mode="after")
def register_config(self):
module_config_path: Path = Path(__file__).parent / self.module_name
for config_path in module_config_path.iterdir():
config_name = config_path.stem
config = FileHandler(file_path=config_path).load()
self.config_dict[config_name] = config
return self
def list_config_names(self):
return list(self.config_dict.keys())
def __getattr__(self, item):
if item in self.config_dict:
return self.config_dict[item]
return super().__getattr__(item)
def __getitem__(self, item):
if item in self.config_dict:
return self.config_dict[item]
return super().__getitem__(item)

View file

@ -1 +0,0 @@
{"a": 1}

View file

@ -10,7 +10,7 @@ from experiencemaker.model.base_embedding_model import BaseEmbeddingModel, EMBED
class OpenAICompatibleEmbeddingModel(BaseEmbeddingModel):
api_key: str = Field(default_factory=lambda: os.getenv("OPENAI_API_KEY"), description="api key")
base_url: str = Field(default_factory=lambda: os.getenv("OPENAI_BASE_URL"), description="base url")
model_name: str = Field(default="text-embedding-v3", description="model name")
model_name: str = Field(default="text-embedding-v4", description="model name")
dimensions: int = Field(default=1024, description="dimensions")
encoding_format: Literal["float", "base64"] = Field(default="float", description="encoding_format")
_client: OpenAI = PrivateAttr()
@ -74,3 +74,20 @@ class OpenAICompatibleEmbeddingModel(BaseEmbeddingModel):
EMBEDDING_MODEL_REGISTRY.register(OpenAICompatibleEmbeddingModel, "openai_compatible")
def main():
from experiencemaker.utils.util_function import load_env_keys
load_env_keys()
model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4")
res1 = model.get_embeddings(
"The clothes are of good quality and look good, definitely worth the wait. I love them.")
res2 = model.get_embeddings(["aa", "bb"])
print(res1)
print(res2)
if __name__ == "__main__":
main()
# launch with: python -m experiencemaker.model.openai_compatible_embedding_model

View file

@ -153,3 +153,25 @@ class OpenAICompatibleBaseLLM(BaseLLM):
LLM_REGISTRY.register(OpenAICompatibleBaseLLM, "openai_compatible")
def main():
from experiencemaker.utils.util_function import load_env_keys
from experiencemaker.tool.dashscope_search_tool import DashscopeSearchTool
from experiencemaker.tool.code_tool import CodeTool
from experiencemaker.enumeration.role import Role
load_env_keys()
model_name = "qwen-max-2025-01-25"
# model_name = "qwen3-32b"
llm = OpenAICompatibleBaseLLM(model_name=model_name)
tools: List[BaseTool] = [DashscopeSearchTool(), CodeTool()]
llm.stream_print([Message(role=Role.USER, content="hello")], [])
print("=" * 20)
llm.stream_print([Message(role=Role.USER, content="What's the weather like in Beijing today?")], tools)
if __name__ == "__main__":
main()
# launch with: python -m experiencemaker.model.openai_compatible_llm

View file

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

View file

@ -1,15 +1,15 @@
from abc import ABC
from typing import List
from pydantic import Field
from pydantic import Field, BaseModel
from experiencemaker.module.base_module import BaseModule
from experiencemaker.module.reward_fn.base_reward_fn import BaseRewardFn
from experiencemaker.schema.reward import Reward
from experiencemaker.schema.trajectory import StateMessage, ActionMessage, ToolCall
from experiencemaker.tool.base_tool import BaseTool
class BaseEnvironment(BaseModule):
class BaseEnvironment(BaseModel, ABC):
tools: List[BaseTool] = Field(default_factory=list)
reward_fns: List[BaseRewardFn] = Field(default_factory=list)
current_state: StateMessage = Field(default_factory=StateMessage)
@ -54,15 +54,3 @@ class BaseEnvironment(BaseModule):
def build_info(self, **kwargs):
return {}
def get_tool_info(self,tool_name):
tool_dict = {tool.name: tool for tool in self.tools}
if tool_name in tool_dict:
return f'tool \'{tool_name}\' description is: {tool_dict[tool_name].description}\t' + f'parameters: {str(tool_dict[tool_name].input_schema)}'
else:
return ''
def get_tools_info(self):
tool_dict = {tool.name: tool for tool in self.tools}
return {tool_name:self.get_tool_info(tool_name=tool_name) for tool_name in tool_dict}

View file

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

View file

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

View file

@ -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
<rule>
List the rules that could be used for comparison...
</rule>
<rule_based_comparison>
Conduct a step-by-step comparison according to the rules...
</rule_based_comparison>
<result>
Output only the name of the better plan, either **Plan1** or **Plan2**.
</result>

View file

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

View file

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

View file

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

View file

@ -4,17 +4,15 @@ import uvicorn
from fastapi import FastAPI
from loguru import logger
from experiencemaker.model import LLM_REGISTRY, EMBEDDING_MODEL_REGISTRY
from experiencemaker.model.base_embedding_model import BaseEmbeddingModel
from experiencemaker.model.base_llm import BaseLLM
from experiencemaker.model.base_embedding_model import BaseEmbeddingModel, EMBEDDING_MODEL_REGISTRY
from experiencemaker.model.base_llm import BaseLLM, LLM_REGISTRY
from experiencemaker.module.agent_wrapper.base_agent_wrapper import BaseAgentWrapper
from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator
from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer
from experiencemaker.schema.request import AgentWrapperRequest, ContextGeneratorRequest, SummarizerRequest
from experiencemaker.schema.response import AgentWrapperResponse, ContextGeneratorResponse, SummarizerResponse
from experiencemaker.schema.trajectory import ContextMessage, Trajectory, Sample
from experiencemaker.storage import VECTOR_STORE_REGISTRY
from experiencemaker.storage.base_vector_store import BaseVectorStore
from experiencemaker.storage.base_vector_store import BaseVectorStore, VECTOR_STORE_REGISTRY
app = FastAPI()
from pydantic import BaseModel, Field, model_validator
@ -29,7 +27,6 @@ class ExperienceMakerHttpService(BaseModel):
llm_config: dict = Field(default_factory=dict)
embedding_model_config: dict = Field(default_factory=dict)
vector_store_config: dict = Field(default_factory=dict)
agent_wrapper_config: dict = Field(default_factory=dict)
context_generator_config: dict = Field(default_factory=dict)
summarizer_config: dict = Field(default_factory=dict)
@ -37,7 +34,6 @@ class ExperienceMakerHttpService(BaseModel):
llm: BaseLLM | None = Field(default=None)
embedding_model: BaseEmbeddingModel | None = Field(default=None)
vector_store: BaseVectorStore | None = Field(default=None)
agent_wrapper: BaseAgentWrapper | None = Field(default=None)
context_generator: BaseContextGenerator | None = Field(default=None)
summarizer: BaseSummarizer | None = Field(default=None)

View file

@ -6,6 +6,7 @@ from elasticsearch.helpers import bulk
from loguru import logger
from pydantic import Field, PrivateAttr, model_validator
from experiencemaker.model.openai_compatible_embedding_model import OpenAICompatibleEmbeddingModel
from experiencemaker.schema.vector_store_node import VectorStoreNode
from experiencemaker.storage.base_vector_store import BaseVectorStore, VECTOR_STORE_REGISTRY
@ -170,3 +171,74 @@ class EsVectorStore(BaseVectorStore):
VECTOR_STORE_REGISTRY.register(EsVectorStore, "elasticsearch")
def main():
from experiencemaker.utils.util_function import load_env_keys
load_env_keys()
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=1024)
index_name = "rag_nodes_index"
hosts = "http://11.160.132.46:8200"
es = EsVectorStore(hosts=hosts, embedding_model=embedding_model, index_name=index_name)
es.delete_index()
es.create_index()
sample_nodes = [
VectorStoreNode(
workspace_id="w1",
content="Artificial intelligence is a technology that simulates human intelligence.",
metadata={
"node_type": "n1",
}
),
VectorStoreNode(
workspace_id="w1",
content="AI is the future of mankind.",
metadata={
"node_type": "n1",
}
),
VectorStoreNode(
workspace_id="w1",
content="I want to eat fish!",
metadata={
"node_type": "n2",
}
),
VectorStoreNode(
workspace_id="w2",
content="The bigger the storm, the more expensive the fish.",
metadata={
"node_type": "n1",
}
),
]
es.insert(sample_nodes, refresh_index=True)
logger.info("=" * 20)
results = es.add_term_filter(key="workspace_id", value="w1") \
.add_term_filter(key="metadata.node_type", value="n1") \
.retrieve_by_query("What is AI?", top_k=5)
for r in results:
logger.info(r.model_dump(exclude={"vector"}))
logger.info("=" * 20)
logger.info("=" * 20)
results = es.add_term_filter(key="workspace_id", value="w1") \
.retrieve_by_query("What is AI?", top_k=5)
for r in results:
logger.info(r.model_dump(exclude={"vector"}))
logger.info("=" * 20)
logger.info("=" * 20)
results = es.retrieve_by_query("What is AI?", top_k=5)
for r in results:
logger.info(r.model_dump(exclude={"vector"}))
logger.info("=" * 20)
if __name__ == "__main__":
main()
# launch with: python -m experiencemaker.storage.es_vector_store

View file

@ -7,6 +7,7 @@ from typing import List, Any
from loguru import logger
from pydantic import Field, model_validator, PrivateAttr
from experiencemaker.model.openai_compatible_embedding_model import OpenAICompatibleEmbeddingModel
from experiencemaker.schema.vector_store_node import VectorStoreNode
from experiencemaker.storage.base_vector_store import BaseVectorStore, VECTOR_STORE_REGISTRY
@ -133,3 +134,58 @@ class FileVectorStore(BaseVectorStore):
VECTOR_STORE_REGISTRY.register(FileVectorStore, "local_file")
def main():
from experiencemaker.utils.util_function import load_env_keys
load_env_keys()
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=1024)
index_name = "rag_nodes_index"
client = FileVectorStore(embedding_model=embedding_model, index_name=index_name)
client.delete_index()
client.create_index()
sample_nodes = [
VectorStoreNode(
workspace_id="w1",
content="Artificial intelligence is a technology that simulates human intelligence.",
metadata={
"node_type": "n1",
}
),
VectorStoreNode(
workspace_id="w1",
content="AI is the future of mankind.",
metadata={
"node_type": "n1",
}
),
VectorStoreNode(
workspace_id="w1",
content="I want to eat fish!",
metadata={
"node_type": "n2",
}
),
VectorStoreNode(
workspace_id="w2",
content="The bigger the storm, the more expensive the fish.",
metadata={
"node_type": "n1",
}
),
]
client.insert(sample_nodes)
logger.info("=" * 20)
results = client.retrieve_by_query("What is AI?", top_k=5)
for r in results:
logger.info(r.model_dump(exclude={"vector"}))
logger.info("=" * 20)
if __name__ == "__main__":
main()
# launch with: python -m experiencemaker.storage.file_vector_store

View file

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

View file

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

View file

@ -6,7 +6,7 @@ from dashscope.api_entities.dashscope_response import Message
from loguru import logger
from pydantic import Field
from experiencemaker.tool.base_tool import BaseTool
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)

View file

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

View file

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

View file

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

View file

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

View file

@ -1,72 +0,0 @@
import os
import yaml
from loguru import logger
from pydantic import BaseModel, Field, model_validator
class PromptHandler(BaseModel):
file_path: str = Field(default="")
prompt_dict: dict = Field(default_factory=dict)
@model_validator(mode="after")
def init_prompt_dict(self):
if self.prompt_dict:
logger.info(f"load prompt from dict, keys={self.prompt_dict.keys()}")
if self.file_path:
self.load_file_prompt()
logger.info(f"load prompt from file_path, keys={self.prompt_dict.keys()}")
def load_file_prompt(self):
if not os.path.exists(self.file_path):
raise RuntimeError(f"file_path={self.file_path} not exists!")
with open(self.file_path) as f:
prompt_dict: dict = yaml.load(f, yaml.FullLoader)
self.prompt_dict.update(prompt_dict)
def __getitem__(self, key: str):
return self.prompt_dict[key]
def __setitem__(self, key: str, value: str):
self.prompt_dict[key] = value
def __getattr__(self, key: str):
if key in self.prompt_dict:
return self.prompt_dict[key]
return super().__getattr__(key)
def prompt_format(self, prompt_name: str, **kwargs):
prompt = self.prompt_dict[prompt_name]
flag_kwargs = {k: v for k, v in kwargs.items() if isinstance(v, bool)}
other_kwargs = {k: v for k, v in kwargs.items() if not isinstance(v, bool)}
if flag_kwargs:
split_prompt = []
for line in prompt.strip().split("\n"):
hit = False
hit_flag = True
for key, flag in kwargs.items():
if not line.startswith(f"[{key}]"):
continue
else:
hit = True
hit_flag = flag
line = line.strip(f"[{key}]")
break
if not hit:
split_prompt.append(line)
elif hit_flag:
split_prompt.append(line)
prompt = "\n".join(split_prompt)
if other_kwargs:
prompt = prompt.format(**other_kwargs)
return prompt

View file

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

View file

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

View file

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