diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml new file mode 100644 index 00000000..1795bc17 --- /dev/null +++ b/.pre-commit-config.yaml @@ -0,0 +1,85 @@ +repos: + - repo: https://github.com/pre-commit/pre-commit-hooks + rev: v6.0.0 + hooks: + - id: check-ast + exclude: ^(test/|cookbook/) + - id: check-yaml + - id: check-xml + - id: check-toml + - id: check-json + - id: detect-private-key + - id: trailing-whitespace + - repo: https://github.com/asottile/add-trailing-comma + rev: v4.0.0 + hooks: + - id: add-trailing-comma + exclude: ^(test/|cookbook/) + - repo: https://github.com/psf/black + rev: 25.9.0 + hooks: + - id: black + exclude: ^(test/|cookbook/) + args: [--line-length=120] + - repo: https://github.com/PyCQA/flake8 + rev: 7.3.0 + hooks: + - id: flake8 + exclude: ^(test/|cookbook/) + args: [ + "--extend-ignore=E203", + "--max-line-length=120" + ] + - repo: https://github.com/pylint-dev/pylint + rev: v4.0.2 + hooks: + - id: pylint + exclude: + (?x)( + ^docs + | ^test/ + | ^cookbook/ + | pb2\.py$ + | grpc\.py$ + | \.demo$ + | \.md$ + | \.html$ + ) + args: [ + --disable=W0511, + --disable=W0718, + --disable=W0122, + --disable=C0103, + --disable=R0913, + --disable=R0917, + --disable=E0401, + --disable=E1101, + --disable=E1111, + --disable=C0415, + --disable=W0603, + --disable=R1705, + --disable=R0914, + --disable=E0601, + --disable=W0602, + --disable=W0604, + --disable=R0801, + --disable=R0902, + --disable=R0903, + --disable=R0904, + --disable=C0123, + --disable=W0231, + --disable=W1113, + --disable=W0221, + --disable=R0401, + --disable=W0632, + --disable=W0123, + --disable=C3001, + --disable=R1702, + --disable=R0912, + --max-line-length=120, + ] + - repo: https://github.com/regebro/pyroma + rev: "5.0" + hooks: + - id: pyroma + args: [--min=10, .] diff --git a/README.md b/README.md index 3f60b397..58d108af 100644 --- a/README.md +++ b/README.md @@ -3,8 +3,8 @@

- Python Version - PyPI Version + Python Version + PyPI Version License GitHub Stars

@@ -72,22 +72,14 @@ Learn more about how to use tool memory from [tool memory](docs/tool_memory/tool ## 📰 Latest Updates -- **[2025-10]** 🚀 ReMe v0.1.10 released! Core enhancement: direct Python import support. You can now use ReMe without starting an HTTP or MCP service - simply `from reme_ai import ReMeApp` and call methods directly in your Python code. -- **[2025-10]** 🔧 Tool Memory support is now available! Enables data-driven tool selection and parameter optimization through historical performance tracking. Check out the [Tool Memory Guide](docs/tool_memory/tool_memory.md) and [benchmark results](docs/tool_memory/tool_bench.md). -- **[2025-09]** 🎉 ReMe v0.1.9 has been officially released, adding support for asynchronous operations. It has also been - integrated into the memory service of agentscope-runtime. -- **[2025-09]** 🎉 ReMe v0.1 officially released, integrating task memory and personal memory. If you want to use the - original memoryscope project, you can find it - in [MemoryScope](https://github.com/agentscope-ai/ReMe/tree/memoryscope_branch). -- **[2025-09]** 🧪 We validated the effectiveness of task memory extraction and reuse in agents in appworld, bfcl(v3), - and frozenlake environments. For more information, - check [appworld exp](docs/cookbook/appworld/quickstart.md), [bfcl exp](docs/cookbook/bfcl/quickstart.md), - and [frozenlake exp](docs/cookbook/frozenlake/quickstart.md). -- **[2025-08]** 🚀 MCP protocol support is now available -> [MCP Quick Start](docs/mcp_quick_start.md). -- **[2025-06]** 🚀 Multiple backend vector storage support (Elasticsearch & - ChromaDB) -> [Vector DB quick start](docs/vector_store_api_guide.md). -- **[2024-09]** 🧠 [MemoryScope](https://github.com/agentscope-ai/ReMe/tree/memoryscope_branch) v0.1 released, - personalized and time-aware memory storage and usage. +- **[2025-10]** 🚀 Direct Python import support: use `from reme_ai import ReMeApp` without HTTP/MCP service +- **[2025-10]** 🔧 Tool Memory: data-driven tool selection and parameter optimization ([Guide](docs/tool_memory/tool_memory.md)) +- **[2025-09]** 🎉 Async operations support, integrated into agentscope-runtime +- **[2025-09]** 🎉 Task memory and personal memory integration +- **[2025-09]** 🧪 Validated effectiveness in appworld, bfcl(v3), and frozenlake ([Experiments](docs/cookbook)) +- **[2025-08]** 🚀 MCP protocol support ([Quick Start](docs/mcp_quick_start.md)) +- **[2025-06]** 🚀 Multiple backend vector storage (Elasticsearch & ChromaDB) ([Guide](docs/vector_store_api_guide.md)) +- **[2024-09]** 🧠 Personalized and time-aware memory storage --- @@ -194,7 +186,7 @@ async def main(): ] ) print(result) - + # Retriever: Get relevant memories result = await app.async_execute( name="retrieve_task_memory", @@ -327,7 +319,7 @@ async def main(): ] ) print(result) - + # Memory Retrieval: Get personal memory fragments result = await app.async_execute( name="retrieve_personal_memory", @@ -477,7 +469,7 @@ async def main(): ] ) print(result) - + # Generate usage guidelines from history result = await app.async_execute( name="summary_tool_memory", @@ -485,7 +477,7 @@ async def main(): tool_names="web_search" ) print(result) - + # Retrieve tool guidelines before use result = await app.async_execute( name="retrieve_tool_memory", @@ -651,7 +643,7 @@ async def main(): path="./docs/library/" ) print(result) - + # Query relevant memories result = await app.async_execute( name="retrieve_task_memory", @@ -679,7 +671,7 @@ We tested ReMe on Appworld using qwen3-8b: | with ReMe | 0.109 **(+2.6%)** | 0.175 **(+3.5%)** | 0.281 **(+5.3%)** | Pass@K measures the probability that at least one of the K generated samples successfully completes the task ( -score=1). +score=1). The current experiment uses an internal AppWorld environment, which may have slight differences. You can find more details on reproducing the experiment in [quickstart.md](docs/cookbook/appworld/quickstart.md). diff --git a/cookbook/appworld/appworld_react_agent.py b/cookbook/appworld/appworld_react_agent.py index c487f1a5..25b93985 100644 --- a/cookbook/appworld/appworld_react_agent.py +++ b/cookbook/appworld/appworld_react_agent.py @@ -1,5 +1,7 @@ +# flake8: noqa: E402, E501 import os from typing import List + from tqdm import tqdm os.environ["APPWORLD_ROOT"] = "." @@ -7,7 +9,6 @@ from dotenv import load_dotenv load_dotenv("../../../.env") -import re import time import json import ray @@ -18,26 +19,28 @@ from jinja2 import Template from loguru import logger from openai import OpenAI -from prompt import PROMPT_TEMPLATE, PROMPT_TEMPLATE_WITH_EXPERIENCE +from prompt import PROMPT_TEMPLATE_WITH_EXPERIENCE @ray.remote class AppworldReactAgent: """A minimal ReAct Agent for AppWorld tasks.""" - def __init__(self, - index: int, - task_ids: List[str], - experiment_name: str, - model_name: str = "qwen3-8b", - temperature: float = 0.9, - max_interactions: int = 30, - max_response_size: int = 2048, - num_runs: int = 1, - use_task_memory: bool = False, - make_task_memory: bool = False, - api_url: str = "http://0.0.0.0:8002/", - workspace_id: str="appworld_v1"): + def __init__( + self, + index: int, + task_ids: List[str], + experiment_name: str, + model_name: str = "qwen3-8b", + temperature: float = 0.9, + max_interactions: int = 30, + max_response_size: int = 2048, + num_runs: int = 1, + use_task_memory: bool = False, + make_task_memory: bool = False, + api_url: str = "http://0.0.0.0:8002/", + workspace_id: str = "appworld_v1", + ): self.index: int = index self.task_ids: List[str] = task_ids @@ -62,7 +65,8 @@ class AppworldReactAgent: messages=messages, temperature=self.temperature, extra_body={"enable_thinking": False}, - seed=0) + seed=0, + ) return response.choices[0].message.content @@ -72,13 +76,17 @@ class AppworldReactAgent: return "call llm error" - def prompt_messages(self,world: AppWorld) -> list[dict]: + def prompt_messages(self, world: AppWorld) -> list[dict]: if self.use_task_memory: task_memory = self.get_task_memory(world.task.instruction) logger.info(f"loaded task_memory: {task_memory}") - dictionary = {"supervisor": world.task.supervisor, "instruction": world.task.instruction, "experience": task_memory} + dictionary = { + "supervisor": world.task.supervisor, + "instruction": world.task.instruction, + "experience": task_memory, + } else: - dictionary = {"supervisor": world.task.supervisor, "instruction": world.task.instruction ,"experience": ""} + dictionary = {"supervisor": world.task.supervisor, "instruction": world.task.instruction, "experience": ""} print(dictionary) prompt = Template(PROMPT_TEMPLATE_WITH_EXPERIENCE.lstrip()).render(dictionary) messages: list[dict] = [] @@ -96,7 +104,7 @@ class AppworldReactAgent: # messages.append({"role": role_type, "content": None}) # last_start = match.span()[1] # messages[-1]["content"] = prompt[last_start:] - messages.append({"role":"user", "content":prompt}) + messages.append({"role": "user", "content": prompt}) return messages @staticmethod @@ -122,7 +130,7 @@ class AppworldReactAgent: output = world.execute(code) if len(output) > self.max_response_size: # logger.warning(f"output exceed max size={len(output)}") - output = output[:self.max_response_size] + output = output[: self.max_response_size] history.append({"role": "user", "content": output}) if world.task_completed(): @@ -164,7 +172,7 @@ class AppworldReactAgent: json={ "workspace_id": self.workspace_id, "query": query, - } + }, ) result = self.handle_api_response(response) @@ -181,26 +189,28 @@ class AppworldReactAgent: if not result: print("No results to summarize") return - + # Prepare trajectories from results trajectories = [] for r in result: if "task_history" in r: - trajectories.append({ - "messages": r["task_history"], - "score": float(r.get("uplift_score", 0.0)) - }) - + trajectories.append( + { + "messages": r["task_history"], + "score": float(r.get("uplift_score", 0.0)), + }, + ) + if not trajectories: print("No trajectories to summarize") return - + response = requests.post( url=f"{self.api_url}summary_task_memory", json={ "workspace_id": self.workspace_id, - "trajectories": trajectories - } + "trajectories": trajectories, + }, ) result = self.handle_api_response(response) @@ -222,4 +232,4 @@ def main(): if __name__ == "__main__": - main() \ No newline at end of file + main() diff --git a/cookbook/appworld/prompt.py b/cookbook/appworld/prompt.py index ac03b8bd..9db1bc41 100644 --- a/cookbook/appworld/prompt.py +++ b/cookbook/appworld/prompt.py @@ -1,3 +1,4 @@ +# flake8: noqa: E402, E501 # This is a basic prompt template containing all the necessary onboarding information to solve AppWorld tasks. It explains the role of the agent and the supervisor, how to explore the API documentation, how to operate the interactive coding environment and call APIs via a simple task, and provides key instructions and disclaimers. # You can adapt it as needed by your agent. You can also choose to bypass API docs app and build your own API retrieval, e.g., for FullCodeRefl, IPFunCall, etc, we asked an LLM to predict relevant APIs separately and put its documentation directly in the prompt. diff --git a/cookbook/appworld/run_appworld.py b/cookbook/appworld/run_appworld.py index 30dbc2e4..430b5b4e 100644 --- a/cookbook/appworld/run_appworld.py +++ b/cookbook/appworld/run_appworld.py @@ -1,8 +1,9 @@ +# flake8: noqa: E402 import os import time -import requests import ray +import requests from ray import logger os.environ["APPWORLD_ROOT"] = "." @@ -35,7 +36,7 @@ def delete_workspace(workspace_id: str, api_url: str = "http://0.0.0.0:8002/"): json={ "workspace_id": workspace_id, "action": "delete", - } + }, ) result = handle_api_response(response) @@ -51,7 +52,7 @@ def dump_memory(workspace_id: str, path: str = "./", api_url: str = "http://0.0. "workspace_id": workspace_id, "action": "dump", "path": path, - } + }, ) result = handle_api_response(response) @@ -67,7 +68,7 @@ def load_memory(workspace_id: str, path: str = "docs/library", api_url: str = "h "workspace_id": workspace_id, "action": "load", "path": path, - } + }, ) result = handle_api_response(response) @@ -75,7 +76,16 @@ def load_memory(workspace_id: str, path: str = "docs/library", api_url: str = "h print(f"Memory loaded from {path}") -def run_agent(dataset_name: str, experiment_suffix: str, max_workers: int, num_runs: int = 1, use_task_memory: bool = False, make_task_memory: bool = False, workspace_id: str="appworld_v1", api_url: str = "http://0.0.0.0:8002/") : +def run_agent( + dataset_name: str, + experiment_suffix: str, + max_workers: int, + num_runs: int = 1, + use_task_memory: bool = False, + make_task_memory: bool = False, + workspace_id: str = "appworld_v1", + api_url: str = "http://0.0.0.0:8002/", +): experiment_name = dataset_name + "_" + experiment_suffix path: Path = Path(f"./exp_result") path.mkdir(parents=True, exist_ok=True) @@ -93,14 +103,16 @@ def run_agent(dataset_name: str, experiment_suffix: str, max_workers: int, num_r for i in range(max_workers): # Assign tasks to each worker, ensuring each task runs num_runs times worker_task_ids = task_ids[i::max_workers] - actor = AppworldReactAgent.remote(index=i, - task_ids=worker_task_ids, - experiment_name=experiment_name, - num_runs=num_runs, - use_task_memory=use_task_memory, - make_task_memory=make_task_memory, - workspace_id=workspace_id, - api_url=api_url) + actor = AppworldReactAgent.remote( + index=i, + task_ids=worker_task_ids, + experiment_name=experiment_name, + num_runs=num_runs, + use_task_memory=use_task_memory, + make_task_memory=make_task_memory, + workspace_id=workspace_id, + api_url=api_url, + ) future = actor.execute.remote() future_list.append(future) time.sleep(1) @@ -119,14 +131,16 @@ def run_agent(dataset_name: str, experiment_suffix: str, max_workers: int, num_r else: for index, task_id in enumerate(task_ids): - agent = AppworldReactAgent(index=index, - task_ids=[task_id], - experiment_name=experiment_name, - num_runs=num_runs, - use_task_memory=use_task_memory, - make_task_memory=make_task_memory, - workspace_id=workspace_id, - api_url=api_url) + agent = AppworldReactAgent( + index=index, + task_ids=[task_id], + experiment_name=experiment_name, + num_runs=num_runs, + use_task_memory=use_task_memory, + make_task_memory=make_task_memory, + workspace_id=workspace_id, + api_url=api_url, + ) task_results = agent.execute() if isinstance(task_results, list): result.extend(task_results) @@ -140,14 +154,14 @@ def main(): num_runs = 1 # Run each task once workspace_id = "appworld" api_url = "http://0.0.0.0:8002/" - + if max_workers > 1: ray.init(num_cpus=8) - + # Clean up workspace before starting logger.info("Deleting workspace...") delete_workspace(workspace_id=workspace_id, api_url=api_url) - + # First run to build task memories logger.info("Start load experiments to build task memories") load_memory(workspace_id=workspace_id, api_url=api_url) @@ -160,19 +174,30 @@ def main(): # Run experiments with task memory logger.info("Start running experiments with task memory") - run_agent(dataset_name="dev", experiment_suffix=f"with-memory", - max_workers=max_workers, num_runs=1, - use_task_memory=True, make_task_memory=False, - workspace_id=workspace_id, api_url=api_url) + run_agent( + dataset_name="dev", + experiment_suffix=f"with-memory", + max_workers=max_workers, + num_runs=1, + use_task_memory=True, + make_task_memory=False, + workspace_id=workspace_id, + api_url=api_url, + ) # Run experiments without task memory logger.info("Start running experiments without task memory") - run_agent(dataset_name="dev", experiment_suffix=f"no-memory", - max_workers=max_workers, num_runs=1, - use_task_memory=False, make_task_memory=False, - workspace_id=workspace_id, api_url=api_url) - + run_agent( + dataset_name="dev", + experiment_suffix=f"no-memory", + max_workers=max_workers, + num_runs=1, + use_task_memory=False, + make_task_memory=False, + workspace_id=workspace_id, + api_url=api_url, + ) if __name__ == "__main__": - main() \ No newline at end of file + main() diff --git a/cookbook/appworld/run_exp_statistic.py b/cookbook/appworld/run_exp_statistic.py index 32535043..c40241e6 100644 --- a/cookbook/appworld/run_exp_statistic.py +++ b/cookbook/appworld/run_exp_statistic.py @@ -1,8 +1,8 @@ import json -from pathlib import Path from collections import defaultdict -import pandas as pd +from pathlib import Path +import pandas as pd from loguru import logger @@ -24,7 +24,7 @@ def calculate_best_at_k(scores: list, k: int) -> float: group_maxs = [] for i in range(0, len(scores), k): - group = scores[i:i + k] + group = scores[i : i + k] group_maxs.append(max(group)) return sum(group_maxs) / len(group_maxs) @@ -36,8 +36,8 @@ def calculate_pass_at_k(scores: list, k: int) -> float: group_maxs = [] for i in range(0, len(scores), k): - group = scores[i:i + k] - is_pass = 1.0 if max(group) >=1.0 else 0.0 + group = scores[i : i + k] + is_pass = 1.0 if max(group) >= 1.0 else 0.0 group_maxs.append(is_pass) return sum(group_maxs) / len(group_maxs) @@ -61,7 +61,7 @@ def get_possible_k_values(total_runs: int) -> list: def run_exp_statistic(): - path: Path = Path(f"./exp_result") + path: Path = Path("./exp_result") # Store results for all experiments all_results = {} @@ -134,10 +134,10 @@ def run_exp_statistic(): # Create and display table if all_results: df = pd.DataFrame(list(all_results.values())) - df = df.set_index('file') + df = df.set_index("file") # Sort columns by the number in column name (best@8, best@4, best@2, best@1) - pass_columns = [col for col in df.columns if col.startswith('pass@')] + pass_columns = [col for col in df.columns if col.startswith("pass@")] # best_columns = [col for col in df.columns] pass_columns.sort(key=lambda x: x, reverse=False) df = df[pass_columns] @@ -157,4 +157,4 @@ def run_exp_statistic(): if __name__ == "__main__": - run_exp_statistic() \ No newline at end of file + run_exp_statistic() diff --git a/cookbook/bfcl/bfcl_agent.py b/cookbook/bfcl/bfcl_agent.py index 75409b4e..a59e95df 100644 --- a/cookbook/bfcl/bfcl_agent.py +++ b/cookbook/bfcl/bfcl_agent.py @@ -1,3 +1,4 @@ +# flake8: noqa: E402 import os os.environ["BFCL_DATA_PATH"] = "data/multiturn_data_base_val.jsonl" @@ -6,7 +7,6 @@ from dotenv import load_dotenv load_dotenv("../../.env") -import re import time import json import ray @@ -29,7 +29,7 @@ from bfcl_utils import ( extract_single_turn_response, extract_multi_turn_responses, capture_and_print_score_files, - create_error_response + create_error_response, ) from bfcl_eval.model_handler.api_inference.qwen import QwenAPIHandler from bfcl_eval.eval_checker.multi_turn_eval.multi_turn_utils import ( @@ -47,30 +47,33 @@ from bfcl_eval.utils import ( load_file, ) + @ray.remote class BFCLAgent: """A minimal ReAct Agent for BFCL-v3(multi-turn) tasks.""" - def __init__(self, - index: int, - task_ids: List[str], - experiment_name: str, - data_path: str = os.getenv("BFCL_DATA_PATH"), - answer_path: Path = Path(os.getenv("BFCL_ANSWER_PATH")), - model_name: str = "qwen3-8b", - temperature: float = 0.9, - max_interactions: int = 30, - max_response_size: int = 2000, - num_runs: int = 1, - enable_thinking: bool = False, - use_memory: bool = False, - use_memory_addition: bool = False, - use_memory_deletion: bool = False, - delete_freq: int = 10, - freq_threshold: int = 5, - utility_threshold: float = 0.5, - memory_base_url: str = "http://0.0.0.0:8001/", - memory_workspace_id: str = "bfcl_8b_0725"): + def __init__( + self, + index: int, + task_ids: List[str], + experiment_name: str, + data_path: str = os.getenv("BFCL_DATA_PATH"), + answer_path: Path = Path(os.getenv("BFCL_ANSWER_PATH")), + model_name: str = "qwen3-8b", + temperature: float = 0.9, + max_interactions: int = 30, + max_response_size: int = 2000, + num_runs: int = 1, + enable_thinking: bool = False, + use_memory: bool = False, + use_memory_addition: bool = False, + use_memory_deletion: bool = False, + delete_freq: int = 10, + freq_threshold: int = 5, + utility_threshold: float = 0.5, + memory_base_url: str = "http://0.0.0.0:8001/", + memory_workspace_id: str = "bfcl_8b_0725", + ): self.index: int = index self.task_ids: List[str] = task_ids @@ -92,14 +95,14 @@ class BFCLAgent: self.utility_threshold: float = utility_threshold self.memory_base_url: str = memory_base_url self.memory_workspace_id: str = memory_workspace_id - + self.history: List[List[List[dict]]] = [[] for _ in range(num_runs)] self.retrieved_memory_list: List[List[List[Any]]] = [[] for _ in range(num_runs)] self.test_entry: List[List[Dict[str, Any]]] = [[] for _ in range(num_runs)] self.original_test_entry: List[List[Dict[str, Any]]] = [[] for _ in range(num_runs)] self.tool_schema: List[List[List[dict]]] = [[] for _ in range(num_runs)] self.current_turn = [[0 for _ in range(len(task_ids))] for _ in range(num_runs)] - + for run_id in range(num_runs): for task_index in range(len(task_ids)): self.init_state(run_id, task_index) @@ -113,7 +116,7 @@ class BFCLAgent: if self.use_memory: query = msg["content"] response = self.get_memory(query) - + if len(response["metadata"]["memory_list"]): self.retrieved_memory_list[run_id].append(response["metadata"]["memory_list"]) exp: str = response["answer"] @@ -129,22 +132,25 @@ class BFCLAgent: def get_query_with_memory(self, query: str, memory: str): return { "role": "user", - "content": "Task:\n" + query + "\n\nSome Related Experience to help you to complete the task:\n" + memory + "content": "Task:\n" + query + "\n\nSome Related Experience to help you to complete the task:\n" + memory, } def get_traj_from_task_history(self, task_id: str, task_history: list, reward: float): return { "task_id": task_id, "messages": task_history, - "score": reward + "score": reward, } def get_memory(self, query: str): - response = requests.post(url=self.memory_base_url + "retrieve_task_memory", json={ - "workspace_id": self.memory_workspace_id, - "query": query, - "top_k": 5 - }) + response = requests.post( + url=self.memory_base_url + "retrieve_task_memory", + json={ + "workspace_id": self.memory_workspace_id, + "query": query, + "top_k": 5, + }, + ) if response.status_code != 200: logger.info(response.text) @@ -153,33 +159,42 @@ class BFCLAgent: response = response.json() logger.info(f"query: {query}, response: {response}") return response - + def add_memory(self, trajectories): - response = requests.post(url=self.memory_base_url + "summary_task_memory", json={ - "workspace_id": self.memory_workspace_id, - "trajectories": trajectories, - }) + response = requests.post( + url=self.memory_base_url + "summary_task_memory", + json={ + "workspace_id": self.memory_workspace_id, + "trajectories": trajectories, + }, + ) response.raise_for_status() response = response.json() - logger.info(f"add new memorys: {response["metadata"]["memory_list"]}") - - def update_memory_information(self, memory_list, update_utility: bool=False): - response = requests.post(url=self.memory_base_url + "record_task_memory", json={ - "workspace_id": self.memory_workspace_id, - "memory_dicts": memory_list, - "update_utility": update_utility - }) - response.raise_for_status() - logger.info(response.json()) - - def delete_memory(self): - response = requests.post(url=self.memory_base_url + "delete_task_memory", json={ - "workspace_id": self.memory_workspace_id, - "freq_threshold": self.freq_threshold, - "utility_threshold": self.utility_threshold - }) + logger.info(f'add new memorys: {response["metadata"]["memory_list"]}') + + def update_memory_information(self, memory_list, update_utility: bool = False): + response = requests.post( + url=self.memory_base_url + "record_task_memory", + json={ + "workspace_id": self.memory_workspace_id, + "memory_dicts": memory_list, + "update_utility": update_utility, + }, + ) response.raise_for_status() - + logger.info(response.json()) + + def delete_memory(self): + response = requests.post( + url=self.memory_base_url + "delete_task_memory", + json={ + "workspace_id": self.memory_workspace_id, + "freq_threshold": self.freq_threshold, + "utility_threshold": self.utility_threshold, + }, + ) + response.raise_for_status() + def call_llm(self, messages: list, tool_schemas: list[dict]) -> str: for i in range(100): try: @@ -200,10 +215,12 @@ class BFCLAgent: return out_msg.model_dump(exclude_unset=True, exclude_none=True) else: reasoning_content = "" # Complete reasoning process - answer_content = "" # Define complete response - tool_info = [] # Store tool invocation information - is_answering = False # Determine whether the reasoning process has finished and response has started - + answer_content = "" # Define complete response + tool_info = [] # Store tool invocation information + is_answering = ( + False # Determine whether the reasoning process has finished and response has started + ) + for chunk in response: if not chunk.choices: # Handle usage information @@ -211,36 +228,43 @@ class BFCLAgent: else: delta = chunk.choices[0].delta # Handle AI's thought process (chain reasoning) - if hasattr(delta, 'reasoning_content') and delta.reasoning_content is not None: + if hasattr(delta, "reasoning_content") and delta.reasoning_content is not None: reasoning_content += delta.reasoning_content - + # Handle final response content else: if not is_answering: # Print title when entering the response phase for the first time is_answering = True if delta.content is not None: answer_content += delta.content - + # Handle tool invocation information (support parallel tool calls) if delta.tool_calls is not None: for tool_call in delta.tool_calls: index = tool_call.index # Tool call index, used for parallel calls - + # Dynamically expand tool information storage list while len(tool_info) <= index: - tool_info.append({"id": "", "type": "function", "index": index, "function": { "name": "", "arguments": "" }}) - + tool_info.append( + { + "id": "", + "type": "function", + "index": index, + "function": {"name": "", "arguments": ""}, + }, + ) + # Collect tool call ID (used for subsequent function calls) if tool_call.id: - tool_info[index]['id'] += tool_call.id - + tool_info[index]["id"] += tool_call.id + # Collect function name (used for subsequent routing to specific functions) if tool_call.function and tool_call.function.name: - tool_info[index]['function']['name'] += tool_call.function.name - + tool_info[index]["function"]["name"] += tool_call.function.name + # Collect function parameters (in JSON string format, need subsequent parsing) if tool_call.function and tool_call.function.arguments: - tool_info[index]['function']['arguments'] += tool_call.function.arguments + tool_info[index]["function"]["arguments"] += tool_call.function.arguments msg = { "role": "assistant", "content": answer_content, @@ -269,29 +293,35 @@ class BFCLAgent: Dict containing next message and tools if applicable """ try: - if not messages: + if not messages: return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index]) if messages[-1]["role"] != "assistant": return create_error_response( - "Last message must be from assistant" + "Last message must be from assistant", ) if "tool_calls" in messages[-1] and len(messages[-1]["tool_calls"]) > 0: try: tool_calls = messages[-1]["tool_calls"] decoded_calls = self._convert_tool_calls_to_execution_format( - tool_calls + tool_calls, ) # decoded_calls:[function(param=xxx)] print(f"decoded_calls: {decoded_calls}") if is_empty_execute_response(decoded_calls): warnings.warn( - f"is_empty_execute_response: {is_empty_execute_response(decoded_calls)}" + f"is_empty_execute_response: {is_empty_execute_response(decoded_calls)}", + ) + return handle_user_turn( + self.original_test_entry[run_id][index], + self.current_turn[run_id][index], ) - return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index]) return handle_tool_calls( - tool_calls, decoded_calls, self.original_test_entry[run_id][index], self.current_turn[run_id][index] + tool_calls, + decoded_calls, + self.original_test_entry[run_id][index], + self.current_turn[run_id][index], ) except Exception as e: warnings.warn(f"Errors during tool invocation: {str(e)}") @@ -303,7 +333,8 @@ class BFCLAgent: return create_error_response(f"Failed to process request: {str(e)}") def _convert_tool_calls_to_execution_format( - self, tool_calls: List[Dict[str, Any]] + self, + tool_calls: List[Dict[str, Any]], ) -> List[str]: """ Convert OpenAI format tool calls to execution format. @@ -334,7 +365,7 @@ class BFCLAgent: execution_list.append(f"{function_name}()") return execution_list - + def get_reward(self, run_id, index) -> float: try: if not self.history[run_id][index] or not self.original_test_entry[run_id][index]: @@ -342,7 +373,8 @@ class BFCLAgent: model_name = "env_handler" handler = QwenAPIHandler( - model_name, temperature=1.0 + model_name, + temperature=1.0, ) # FIXME: magic number model_result_data = self._convert_conversation_to_eval_format(run_id, index) @@ -351,23 +383,28 @@ class BFCLAgent: state = {"leaderboard_table": {}} record_cost_latency( - state["leaderboard_table"], model_name, [model_result_data] + state["leaderboard_table"], + model_name, + [model_result_data], ) if is_relevance_or_irrelevance(self.categories[index]): accuracy, _ = self._eval_relevance_test( - handler, model_result_data, prompt_data, model_name, self.category + handler, + model_result_data, + prompt_data, + model_name, + self.category, ) else: # Find the corresponding possible answer file possible_answer_file = find_file_with_suffix( - self.answer_path, self.categories[index] + self.answer_path, + self.categories[index], ) possible_answer = load_file(possible_answer_file, sort_by_id=True) - possible_answer = [ - item for item in possible_answer if item["id"] == self.task_ids[index] - ] + possible_answer = [item for item in possible_answer if item["id"] == self.task_ids[index]] if is_multi_turn(self.categories[index]): accuracy, _ = self._eval_multi_turn_test( handler, @@ -396,7 +433,7 @@ class BFCLAgent: traceback.print_exc() return 0 - + def _convert_conversation_to_eval_format(self, run_id, index) -> Dict[str, Any]: """ Convert conversation history to evaluation format. @@ -422,7 +459,7 @@ class BFCLAgent: } return model_result_data - + def _eval_multi_turn_test( self, handler, @@ -458,10 +495,13 @@ class BFCLAgent: score_dir=score_dir, ) capture_and_print_score_files( - score_dir, model_name, test_category, "multi_turn" + score_dir, + model_name, + test_category, + "multi_turn", ) return accuracy, total_count - + def _eval_single_turn_test( self, handler, @@ -504,7 +544,10 @@ class BFCLAgent: score_dir=score_dir, ) capture_and_print_score_files( - score_dir, model_name, test_category, "single_turn" + score_dir, + model_name, + test_category, + "single_turn", ) return accuracy, total_count @@ -516,37 +559,42 @@ class BFCLAgent: try: start_time = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") for i in range(self.max_interactions): - llm_output = self.call_llm(self.history[run_id][task_index], self.tool_schema[run_id][task_index]) + llm_output = self.call_llm( + self.history[run_id][task_index], + self.tool_schema[run_id][task_index], + ) self.history[run_id][task_index].append(llm_output) env_output = self.env_step(run_id, task_index, self.history[run_id][task_index]) # Possible env_output returns after environment interaction: - # 1. Triggers a query with available tools list: {"messages": [{"role": "user", "content": user_query}], "tools": tools} + # 1. Triggers a query with available tools list: {"messages": [{"role": "user", "content": user_query}], "tools": tools} # 2. Returns tool invocation result: {"messages": [{"role": "tool", "content": {}, 'tool_call_id': 'chatcmpl-tool-xxx'}]} # : when success, returns result dicts, e.g., {"travel_cost_list": [1140.0]}, when error, returns error message, e.g., {"error": "cd: temporary: No such directory. You cannot use path to change directory."} # 3. Conversation completion: {"messages": [{"role": "env", "content": "[CONVERSATION_COMPLETED]"}]} # 4. Program error: {"messages": [{"role": "env", "content": f"[ERROR] {error_message}"}]} - + # tool_list update if "tools" in env_output: self.tool_schema[run_id][task_index] = extract_tool_schema(env_output["tools"]) - - new_tool_calls=[] - new_tool_call_ids=[] + + new_tool_calls = [] + new_tool_call_ids = [] next_user_msg = "" for idx, msg in enumerate(env_output.get("messages", [])): - if msg["role"] == "tool" and len(msg["content"])>0: + if msg["role"] == "tool" and len(msg["content"]) > 0: new_tool_calls.append(msg.get("content", "")) new_tool_call_ids.append(msg.get("tool_call_id", "")) elif msg["role"] == "user": next_user_msg = msg.get("content", "") self.current_turn[run_id][task_index] += 1 - else: # for env role messages + else: # for env role messages next_user_msg = msg.get("content", "") - + if new_tool_calls: for idx, call in enumerate(new_tool_calls): - self.history[run_id][task_index].append({"role": "tool", "content": str(call), "tool_call_id": new_tool_call_ids[idx]}) + self.history[run_id][task_index].append( + {"role": "tool", "content": str(call), "tool_call_id": new_tool_call_ids[idx]}, + ) else: self.history[run_id][task_index].append({"role": "user", "content": next_user_msg}) @@ -557,14 +605,16 @@ class BFCLAgent: reward = self.get_reward(run_id, task_index) if self.use_memory: - if reward == 1 and self.use_memory_addition: # selectively add memories when succeed - new_traj_list = [self.get_traj_from_task_history(task_id, self.history[run_id][task_index], reward)] - self.add_memory(new_traj_list) - + if reward == 1 and self.use_memory_addition: # selectively add memories when succeed + new_traj_list = [ + self.get_traj_from_task_history(task_id, self.history[run_id][task_index], reward), + ] + self.add_memory(new_traj_list) + # update the freq & utility attributes of retrieved memories - update_utility: bool = (reward == 1) + update_utility: bool = reward == 1 self.update_memory_information(self.retrieved_memory_list[run_id][task_index], update_utility) - + counter += 1 if self.use_memory_deletion and counter % self.delete_freq == 0: self.delete_memory() @@ -594,14 +644,15 @@ class BFCLAgent: """ return self.history[run_id][index][-1]["content"] == "[CONVERSATION_COMPLETED]" + def main(): with open(os.getenv("BFCL_DATA_PATH"), "r", encoding="utf-8") as f: task_ids = [json.loads(l)["id"] for l in f] dataset_name = "dev" agent = BFCLAgent( - index=0, - task_id=task_ids[0], - experiment_name=f"zouying_{dataset_name}", + index=0, + task_id=task_ids[0], + experiment_name=f"zouying_{dataset_name}", ) result = agent.execute() logger.info(f"result={json.dumps(result)}") diff --git a/cookbook/bfcl/bfcl_utils.py b/cookbook/bfcl/bfcl_utils.py index 04f540ca..5c16d26b 100644 --- a/cookbook/bfcl/bfcl_utils.py +++ b/cookbook/bfcl/bfcl_utils.py @@ -1,26 +1,20 @@ import json -import tempfile from pathlib import Path from typing import Dict, List, Any -from bfcl_eval.constants.type_mappings import GORILLA_TO_OPENAPI from bfcl_eval.constants.default_prompts import ( DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC, ) -from bfcl_eval.model_handler.model_style import ModelStyle -from bfcl_eval.model_handler.utils import ( - convert_to_function_call, - convert_to_tool, - default_decode_ast_prompting, - default_decode_execute_prompting, - format_execution_results_prompting, - func_doc_language_specific_pre_processing, - retry_with_backoff, - system_prompt_pre_processing_chat_model, -) +from bfcl_eval.constants.type_mappings import GORILLA_TO_OPENAPI from bfcl_eval.eval_checker.multi_turn_eval.multi_turn_utils import ( execute_multi_turn_func_call, ) +from bfcl_eval.model_handler.model_style import ModelStyle +from bfcl_eval.model_handler.utils import ( + convert_to_tool, + default_decode_execute_prompting, + func_doc_language_specific_pre_processing, +) def load_test_case(data_path: str, test_id: str | None) -> Dict[str, Any]: @@ -44,8 +38,10 @@ def load_test_case(data_path: str, test_id: str | None) -> Dict[str, Any]: return data raise ValueError(f"Test case id '{test_id}' not found in {data_path}") + def handle_user_turn( - test_entry: Dict[str, Any], current_turn: int + test_entry: Dict[str, Any], + current_turn: int, ) -> Dict[str, Any]: """ Handle user turn by returning appropriate content from test_entry["question"]. @@ -67,19 +63,17 @@ def handle_user_turn( if str(current_turn) in holdout_function: test_entry["function"].extend(holdout_function[str(current_turn)]) tools = compile_tools(test_entry) - assert ( - len(questions[current_turn]) == 0 - ), "Holdout turn should not have user message." + assert len(questions[current_turn]) == 0, "Holdout turn should not have user message." current_turn_message = [ { "role": "user", "content": DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC, - } + }, ] return create_user_response(current_turn_message, tools) if current_turn >= len(questions): return create_completion_response() - + current_turn_message = questions[current_turn] return create_user_response(current_turn_message, tools) @@ -87,6 +81,7 @@ def handle_user_turn( except Exception as e: return create_error_response(f"Failed to process user message: {str(e)}") + def handle_tool_calls( tool_calls: List[Dict[str, Any]], decoded_calls: list[str], @@ -111,9 +106,7 @@ def handle_tool_calls( involved_classes=test_entry["involved_classes"], model_name="env_handler", test_entry_id=test_entry["id"], - long_context=( - "long_context" in test_entry["id"] or "composite" in test_entry["id"] - ), + long_context=("long_context" in test_entry["id"] or "composite" in test_entry["id"]), is_evaL_run=False, ) # print('execution_results in handler_tool_calls:', execution_results) @@ -138,9 +131,11 @@ def compile_tools(test_entry: dict) -> list: tools = convert_to_tool(functions, GORILLA_TO_OPENAPI, ModelStyle.OpenAI_Completions) return tools - + + def create_tool_response( - tool_calls: List[Dict[str, Any]], execution_results: List[str] + tool_calls: List[Dict[str, Any]], + execution_results: List[str], ) -> Dict[str, Any]: """ Create response for tool calls. @@ -159,13 +154,15 @@ def create_tool_response( "role": "tool", "content": result, "tool_call_id": tool_call.get("id", f"call_{i}"), - } + }, ) return {"messages": tool_messages} + def create_user_response( - question_turn: List[Dict[str, Any]], tools: List[Dict[str, Any]] + question_turn: List[Dict[str, Any]], + tools: List[Dict[str, Any]], ) -> Dict[str, Any]: """ Create response containing user message. @@ -173,7 +170,7 @@ def create_user_response( Args: question_turn: List of messages for current turn tools: List of available tools - + Returns: Response containing user message and tools """ @@ -185,6 +182,7 @@ def create_user_response( return {"messages": [{"role": "user", "content": user_content}], "tools": tools} + def create_completion_response() -> Dict[str, Any]: """ Create response indicating conversation completion. @@ -194,6 +192,7 @@ def create_completion_response() -> Dict[str, Any]: """ return {"messages": [{"role": "env", "content": "[CONVERSATION_COMPLETED]"}]} + def create_error_response(error_message: str) -> Dict[str, Any]: """ Create response for error conditions. @@ -206,6 +205,7 @@ def create_error_response(error_message: str) -> Dict[str, Any]: """ return {"messages": [{"role": "env", "content": f"[ERROR] {error_message}"}]} + def decode_execute(result): """ Decode execute results for compatibility with evaluation framework. @@ -218,6 +218,7 @@ def decode_execute(result): """ return default_decode_execute_prompting(result) + def extract_single_turn_response(messages: List[Dict[str, Any]]) -> str: """ Extract single-turn response from conversation messages. @@ -234,7 +235,7 @@ def extract_single_turn_response(messages: List[Dict[str, Any]]) -> str: formatted_calls = [] for tool_call in message["tool_calls"]: formatted_call = format_single_tool_call_for_eval( - tool_call + tool_call, ) if formatted_call: formatted_calls.append(formatted_call) @@ -243,9 +244,10 @@ def extract_single_turn_response(messages: List[Dict[str, Any]]) -> str: return message["content"] return "" - + + def extract_multi_turn_responses( - messages: List[Dict[str, Any]] + messages: List[Dict[str, Any]], ) -> List[List[str]]: """ Extract multi-turn responses from conversation messages. @@ -275,7 +277,7 @@ def extract_multi_turn_responses( if "tool_calls" in assistant_msg and assistant_msg["tool_calls"]: for tool_call in assistant_msg["tool_calls"]: formatted_call = format_single_tool_call_for_eval( - tool_call + tool_call, ) if formatted_call: current_turn_responses.append(formatted_call) @@ -292,10 +294,11 @@ def extract_multi_turn_responses( return turns_data + def format_single_tool_call_for_eval(tool_call: Dict[str, Any]) -> str: """ Format a single tool call into string representation for evaluation. - + Args: tool_call: Single tool call in OpenAI format @@ -315,11 +318,15 @@ def format_single_tool_call_for_eval(tool_call: Dict[str, Any]) -> str: args_str = ", ".join([f"{k}={repr(v)}" for k, v in args_dict.items()]) return f"{function_name}({args_str})" - except Exception as e: + except Exception: return f"{function_name}()" + def capture_and_print_score_files( - score_dir: Path, model_name: str, test_category: str, eval_type: str + score_dir: Path, + model_name: str, + test_category: str, + eval_type: str, ): """ Capture and print contents of score files written to score_dir. @@ -360,8 +367,10 @@ def capture_and_print_score_files( parsed = json.loads(line) formatted_lines.append( json.dumps( - parsed, ensure_ascii=False, indent=2 - ) + parsed, + ensure_ascii=False, + indent=2, + ), ) content = "\n".join(formatted_lines) except json.JSONDecodeError: @@ -379,7 +388,8 @@ def capture_and_print_score_files( except Exception as e: print(f"Error capturing evaluation result files: {str(e)}") + def extract_tool_schema(tools): for i in range(len(tools)): - tools[i]['function'].pop("response") + tools[i]["function"].pop("response") return tools diff --git a/cookbook/bfcl/init_exp_pool.py b/cookbook/bfcl/init_exp_pool.py index 77418766..f9f22b7a 100644 --- a/cookbook/bfcl/init_exp_pool.py +++ b/cookbook/bfcl/init_exp_pool.py @@ -1,10 +1,11 @@ -import json -import requests import argparse -from pathlib import Path -from typing import List, Dict, Any +import json from collections import defaultdict from concurrent.futures import ThreadPoolExecutor, as_completed +from pathlib import Path +from typing import List, Dict, Any + +import requests def load_task_case(data_path: str, task_id: str | None) -> Dict[str, Any]: @@ -34,32 +35,33 @@ def get_tool_prompt(tools): tool_prompt = "\n\n# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within XML tags:\n" for tool in tools: tool_prompt += "\n" + json.dumps(tool) - tool_prompt += "\n\n\nFor each function call, return a json object with function name and arguments within XML tags:\n\n{\"name\": , \"arguments\": }\n" + tool_prompt += '\n\n\nFor each function call, return a json object with function name and arguments within XML tags:\n\n{"name": , "arguments": }\n' return tool_prompt def group_trajectories_by_task_id(jsonl_entries: List[Dict[str, Any]]) -> List[List[Any]]: """ 根据task_id字段对trajectories进行分组 - + Args: jsonl_entries: JSONL条目列表 - + Returns: List[List[Any]]: 按task_id分组的trajectory列表 """ # 按task_id分组 grouped = defaultdict(list) - + for entry in jsonl_entries: task_id = entry.get("task_id", "") taks_case = load_task_case("data/multiturn_data_base.jsonl", task_id) tools = taks_case.get("tools", [{}]) from bfcl_utils import extract_tool_schema + tool_schema = extract_tool_schema(tools) entry["task_history"][0]["content"] += get_tool_prompt(tool_schema) grouped[task_id].append(entry) - + # 对每组只保留最大和最小reward的两个 filtered_groups = [] for key, trajectories in grouped.items(): @@ -75,33 +77,36 @@ def group_trajectories_by_task_id(jsonl_entries: List[Dict[str, Any]]) -> List[L min_reward_traj = trajectories[0] # 最小reward max_reward_traj = trajectories[-1] # 最大reward filtered_groups.append([min_reward_traj, max_reward_traj]) - + return filtered_groups def post_to_summarizer(trajectories: List[Any], service_url: str, workspace_id: str) -> Dict[str, Any]: """ 将trajectories发送到summarizer服务 - + Args: trajectories: trajectory列表 service_url: 服务URL workspace_id: 工作空间ID - + Returns: 响应结果 """ - trajectory_dicts = [{ - "task_id": traj["task_id"], - "messages": traj["task_history"], - "score": traj["reward"] - } for traj in trajectories] + trajectory_dicts = [ + { + "task_id": traj["task_id"], + "messages": traj["task_history"], + "score": traj["reward"], + } + for traj in trajectories + ] request_data = { "traj_list": trajectory_dicts, - "workspace_id": workspace_id + "workspace_id": workspace_id, } - + try: response = requests.post(f"{service_url}/summarizer", json=request_data) response.raise_for_status() @@ -110,31 +115,33 @@ def post_to_summarizer(trajectories: List[Any], service_url: str, workspace_id: return {"error": str(e), "trajectories_count": len(trajectories)} -def process_trajectories_with_threads(grouped_trajectories: List[List[Any]], - service_url: str, - workspace_id: str, - n_threads: int = 4) -> List[Dict[str, Any]]: +def process_trajectories_with_threads( + grouped_trajectories: List[List[Any]], + service_url: str, + workspace_id: str, + n_threads: int = 4, +) -> List[Dict[str, Any]]: """ 使用多线程处理trajectories组 - + Args: grouped_trajectories: 按task_id分组的trajectory列表 service_url: summarizer服务URL workspace_id: 工作空间ID n_threads: 线程数 - + Returns: 所有结果列表 """ results = [] - + with ThreadPoolExecutor(max_workers=n_threads) as executor: # 提交所有任务 future_to_group = { - executor.submit(post_to_summarizer, group, service_url, workspace_id): i + executor.submit(post_to_summarizer, group, service_url, workspace_id): i for i, group in enumerate(grouped_trajectories) } - + # 收集结果 for future in as_completed(future_to_group): group_index = future_to_group[future] @@ -143,16 +150,18 @@ def process_trajectories_with_threads(grouped_trajectories: List[List[Any]], result["group_index"] = group_index result["group_size"] = len(grouped_trajectories[group_index]) results.append(result) - print(f"✅ Group {group_index} processed: {result.get('experience_list', 0) if 'experience_list' in result else 'error'}") + print( + f"✅ Group {group_index} processed: {result.get('experience_list', 0) if 'experience_list' in result else 'error'}", + ) except Exception as e: error_result = { "group_index": group_index, "group_size": len(grouped_trajectories[group_index]), - "error": str(e) + "error": str(e), } results.append(error_result) print(f"❌ Group {group_index} failed: {e}") - + return results @@ -160,20 +169,20 @@ def main(): """ 主函数,支持命令行参数 """ - parser = argparse.ArgumentParser(description='Convert JSONL to experiences using experience maker service') - parser.add_argument('--jsonl_file', type=str, required=True, help='Path to the JSONL file') - parser.add_argument('--service_url', type=str, default='http://localhost:8001', help='Experience maker service URL') - parser.add_argument('--workspace_id', type=str, required=True, help='Workspace ID for the experience') - parser.add_argument('--output_file', type=str, help='Output file to save results (optional)') - parser.add_argument('--n_threads', type=int, default=4, help='Number of threads for processing') - + parser = argparse.ArgumentParser(description="Convert JSONL to experiences using experience maker service") + parser.add_argument("--jsonl_file", type=str, required=True, help="Path to the JSONL file") + parser.add_argument("--service_url", type=str, default="http://localhost:8001", help="Experience maker service URL") + parser.add_argument("--workspace_id", type=str, required=True, help="Workspace ID for the experience") + parser.add_argument("--output_file", type=str, help="Output file to save results (optional)") + parser.add_argument("--n_threads", type=int, default=4, help="Number of threads for processing") + args = parser.parse_args() - + print(f"Processing JSONL file: {args.jsonl_file}") print(f"Service URL: {args.service_url}") print(f"Workspace ID: {args.workspace_id}") print(f"Threads: {args.n_threads}") - + # 读取JSONL文件 try: with open(args.jsonl_file, "r") as f: @@ -182,31 +191,30 @@ def main(): except Exception as e: print(f"Error reading JSONL file: {e}") return - + # 分组处理 grouped_trajectories = group_trajectories_by_task_id(data) print(f"Total groups: {len(grouped_trajectories)}") - + # 多线程处理 results = process_trajectories_with_threads( - grouped_trajectories, - args.service_url, + grouped_trajectories, + args.service_url, args.workspace_id, - n_threads=args.n_threads + n_threads=args.n_threads, ) - - print(f"Processed {len(results)} groups") - - # 统计结果 - success_count = sum(1 for r in results if 'error' not in r) - error_count = len(results) - success_count - total_experiences = sum(len(r.get('experiences', [])) for r in results if 'experiences' in r) - + print(f"Processed {len(results)} groups") + + # 统计结果 + success_count = sum(1 for r in results if "error" not in r) + error_count = len(results) - success_count + total_experiences = sum(len(r.get("experiences", [])) for r in results if "experiences" in r) + print(f"✅ Success: {success_count}") print(f"❌ Errors: {error_count}") print(f"📊 Total experiences created: {total_experiences}") - + # 保存结果到文件 if args.output_file: try: @@ -217,10 +225,10 @@ def main(): "success_count": success_count, "error_count": error_count, "total_experiences": total_experiences, - "results": results + "results": results, } - - with open(args.output_file, 'w') as f: + + with open(args.output_file, "w") as f: json.dump(summary, f, indent=2) print(f"Results saved to: {args.output_file}") except Exception as e: @@ -231,6 +239,7 @@ def main(): if __name__ == "__main__": # 检查是否有命令行参数 import sys + if len(sys.argv) > 1: # 使用新的命令行接口 main() @@ -239,15 +248,15 @@ if __name__ == "__main__": print("Running in compatibility mode...") with open("exp_result/qwen-max-2025-01-25/no_think/bfcl-multi-turn-base-train50_wo-exp.jsonl", "r") as f: data = [json.loads(line) for line in f] - + # 分组 grouped_trajectories = group_trajectories_by_task_id(data) print(f"Total groups: {len(grouped_trajectories)}") - + results = process_trajectories_with_threads( - grouped_trajectories, - "http://localhost:8001", + grouped_trajectories, + "http://localhost:8001", "bfcl_train50_qwen_max_2025_01_25_extract_compare_validate", - n_threads=4 + n_threads=4, ) - print(f"Processed {len(results)} groups") \ No newline at end of file + print(f"Processed {len(results)} groups") diff --git a/cookbook/bfcl/init_task_memory_pool.py b/cookbook/bfcl/init_task_memory_pool.py index 8c00d4a6..01896fe2 100644 --- a/cookbook/bfcl/init_task_memory_pool.py +++ b/cookbook/bfcl/init_task_memory_pool.py @@ -1,10 +1,11 @@ -import json -import requests import argparse -from pathlib import Path -from typing import List, Dict, Any +import json from collections import defaultdict from concurrent.futures import ThreadPoolExecutor, as_completed +from pathlib import Path +from typing import List, Dict, Any + +import requests def load_task_case(data_path: str, task_id: str | None) -> Dict[str, Any]: @@ -36,31 +37,32 @@ def get_tool_prompt(tools): tool_prompt = "\n\n# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within XML tags:\n" for tool in tools: tool_prompt += "\n" + json.dumps(tool) - tool_prompt += "\n\n\nFor each function call, return a json object with function name and arguments within XML tags:\n\n{\"name\": , \"arguments\": }\n" + tool_prompt += '\n\n\nFor each function call, return a json object with function name and arguments within XML tags:\n\n{"name": , "arguments": }\n' return tool_prompt def group_trajectories_by_task_id(jsonl_entries: List[Dict[str, Any]]) -> List[List[Any]]: """ group trajectories by task_id - + Args: jsonl_entries: JSONL entry list - + Returns: List[List[Any]]: trajectory list grouped by task_id """ grouped = defaultdict(list) - + for entry in jsonl_entries: task_id = entry.get("task_id", "") taks_case = load_task_case("data/multiturn_data_base.jsonl", task_id) tools = taks_case.get("tools", [{}]) from bfcl_utils import extract_tool_schema + tool_schema = extract_tool_schema(tools) entry["task_history"][0]["content"] += get_tool_prompt(tool_schema) grouped[task_id].append(entry) - + # retain only the two with the highest and lowest rewards filtered_groups = [] for key, trajectories in grouped.items(): @@ -76,22 +78,25 @@ def group_trajectories_by_task_id(jsonl_entries: List[Dict[str, Any]]) -> List[L min_reward_traj = trajectories[0] # highest reward max_reward_traj = trajectories[-1] # lowest reward filtered_groups.append([min_reward_traj, max_reward_traj]) - + return filtered_groups def post_to_summarizer(trajectories: List[Any], service_url: str, workspace_id: str) -> Dict[str, Any]: - trajectory_dicts = [{ - "task_id": traj["task_id"], - "messages": traj["task_history"], - "score": traj["reward"] - } for traj in trajectories] + trajectory_dicts = [ + { + "task_id": traj["task_id"], + "messages": traj["task_history"], + "score": traj["reward"], + } + for traj in trajectories + ] request_data = { "trajectories": trajectory_dicts, - "workspace_id": workspace_id + "workspace_id": workspace_id, } - + try: response = requests.post(f"{service_url}/summary_task_memory", json=request_data) response.raise_for_status() @@ -100,30 +105,32 @@ def post_to_summarizer(trajectories: List[Any], service_url: str, workspace_id: return {"error": str(e), "trajectories_count": len(trajectories)} -def process_trajectories_with_threads(grouped_trajectories: List[List[Any]], - service_url: str, - workspace_id: str, - n_threads: int = 4) -> List[Dict[str, Any]]: +def process_trajectories_with_threads( + grouped_trajectories: List[List[Any]], + service_url: str, + workspace_id: str, + n_threads: int = 4, +) -> List[Dict[str, Any]]: """ use threads to process trajectories - + Args: grouped_trajectories: group trajectory list by task_id service_url: memory summarizer service URL workspace_id: workspace ID n_threads: number of threads - + Returns: all results """ results = [] - + with ThreadPoolExecutor(max_workers=n_threads) as executor: future_to_group = { - executor.submit(post_to_summarizer, group, service_url, workspace_id): i + executor.submit(post_to_summarizer, group, service_url, workspace_id): i for i, group in enumerate(grouped_trajectories) } - + for future in as_completed(future_to_group): group_index = future_to_group[future] try: @@ -131,58 +138,60 @@ def process_trajectories_with_threads(grouped_trajectories: List[List[Any]], result["group_index"] = group_index result["group_size"] = len(grouped_trajectories[group_index]) results.append(result) - print(f"✅ Group {group_index} processed: {result["metadata"].get('memory_list', 0) if 'memory_list' in result["metadata"] else 'error'}") + print( + f'✅ Group {group_index} processed: {result["metadata"].get("memory_list", 0) if "memory_list" in result["metadata"] else "error"}', + ) except Exception as e: error_result = { "group_index": group_index, "group_size": len(grouped_trajectories[group_index]), - "error": str(e) + "error": str(e), } results.append(error_result) print(f"❌ Group {group_index} failed: {e}") - + return results def main(): - parser = argparse.ArgumentParser(description='Convert JSONL to memories using ReMe service') - parser.add_argument('--jsonl_file', type=str, required=True, help='Path to the JSONL file') - parser.add_argument('--service_url', type=str, default='http://localhost:8001', help='ReMe service URL') - parser.add_argument('--workspace_id', type=str, required=True, help='Workspace ID for the task memory pool') - parser.add_argument('--output_file', type=str, help='Output file to save results (optional)') - parser.add_argument('--n_threads', type=int, default=4, help='Number of threads for processing') - + parser = argparse.ArgumentParser(description="Convert JSONL to memories using ReMe service") + parser.add_argument("--jsonl_file", type=str, required=True, help="Path to the JSONL file") + parser.add_argument("--service_url", type=str, default="http://localhost:8001", help="ReMe service URL") + parser.add_argument("--workspace_id", type=str, required=True, help="Workspace ID for the task memory pool") + parser.add_argument("--output_file", type=str, help="Output file to save results (optional)") + parser.add_argument("--n_threads", type=int, default=4, help="Number of threads for processing") + args = parser.parse_args() - + print(f"Processing JSONL file: {args.jsonl_file}") print(f"Service URL: {args.service_url}") print(f"Workspace ID: {args.workspace_id}") print(f"Threads: {args.n_threads}") - + with open(args.jsonl_file, "r") as f: data = [json.loads(line) for line in f] print(f"Loaded {len(data)} entries from JSONL file") - + grouped_trajectories = group_trajectories_by_task_id(data) print(f"Total groups: {len(grouped_trajectories)}") - + results = process_trajectories_with_threads( - grouped_trajectories, - args.service_url, + grouped_trajectories, + args.service_url, args.workspace_id, - n_threads=args.n_threads + n_threads=args.n_threads, ) - + print(f"Processed {len(results)} groups") - - success_count = sum(1 for r in results if 'error' not in r) + + success_count = sum(1 for r in results if "error" not in r) error_count = len(results) - success_count - total_memories = sum(len(r["metadata"].get('memory_list', [])) for r in results if 'memory_list' in r["metadata"]) + total_memories = sum(len(r["metadata"].get("memory_list", [])) for r in results if "memory_list" in r["metadata"]) print(f"✅ Success: {success_count}") print(f"❌ Errors: {error_count}") print(f"📊 Total task memories created: {total_memories}") - + if args.output_file: try: summary = { @@ -192,10 +201,10 @@ def main(): "success_count": success_count, "error_count": error_count, "total_task_memories": total_memories, - "results": results + "results": results, } - - with open(args.output_file, 'w') as f: + + with open(args.output_file, "w") as f: json.dump(summary, f, indent=2) print(f"Results saved to: {args.output_file}") except Exception as e: @@ -204,20 +213,21 @@ def main(): if __name__ == "__main__": import sys + if len(sys.argv) > 1: main() else: print("Running in compatibility mode...") with open("exp_result/qwen3-14b/no_think/bfcl-multi-turn-base_wo-exp.jsonl", "r") as f: data = [json.loads(line) for line in f] - + grouped_trajectories = group_trajectories_by_task_id(data) print(f"Total groups: {len(grouped_trajectories)}") - + results = process_trajectories_with_threads( - grouped_trajectories, - "http://localhost:8001", + grouped_trajectories, + "http://localhost:8001", "bfcl_test", - n_threads=4 + n_threads=4, ) - print(f"Processed {len(results)} groups") \ No newline at end of file + print(f"Processed {len(results)} groups") diff --git a/cookbook/bfcl/local_file_to_library.py b/cookbook/bfcl/local_file_to_library.py index afee8374..9131b512 100644 --- a/cookbook/bfcl/local_file_to_library.py +++ b/cookbook/bfcl/local_file_to_library.py @@ -1,27 +1,26 @@ import json -with open("../../file_vector_store/bfcl_test.jsonl", 'r') as f: + +with open("../../file_vector_store/bfcl_test.jsonl", "r") as f: bfcl = [json.loads(line) for line in f] - + new_bfcl = [] for exp in bfcl: new_exp = {} new_exp["workspace_id"] = exp["workspace_id"] new_exp["memory_id"] = exp["unique_id"] new_exp["memory_type"] = exp["metadata"]["memory_type"] - + new_exp["when_to_use"] = exp["content"] new_exp["content"] = exp["metadata"]["content"] - new_exp["score"] = exp["metadata"]["score"] - + new_exp["score"] = exp["metadata"]["score"] + new_exp["time_created"] = exp["metadata"]["time_created"] new_exp["time_modified"] = exp["metadata"]["time_modified"] new_exp["author"] = exp["metadata"]["author"] - - new_exp["metadata"]= exp["metadata"]["metadata"] + + new_exp["metadata"] = exp["metadata"]["metadata"] new_bfcl.append(new_exp) - - -with open('../../library/bfcl_test.jsonl', 'w', encoding='utf-8') as f: - f.writelines(json.dumps(item, ensure_ascii=False) + '\n' for item in new_bfcl) \ No newline at end of file +with open("../../library/bfcl_test.jsonl", "w", encoding="utf-8") as f: + f.writelines(json.dumps(item, ensure_ascii=False) + "\n" for item in new_bfcl) diff --git a/cookbook/bfcl/run_bfcl.py b/cookbook/bfcl/run_bfcl.py index 4eb1570b..c3acbdda 100644 --- a/cookbook/bfcl/run_bfcl.py +++ b/cookbook/bfcl/run_bfcl.py @@ -1,11 +1,11 @@ -import os import time + import ray +from dotenv import load_dotenv + # from ray import logger from loguru import logger -from dotenv import load_dotenv - load_dotenv("../../.env") import json @@ -15,26 +15,30 @@ from pathlib import Path from bfcl_agent import BFCLAgent -def run_agent(dataset_name: str, - experiment_suffix: str, - max_workers: int, - num_runs: int = 4, - model_name: str = "qwen3-8b", - data_path: str = "data/multiturn_data_base_val.jsonl", - answer_path: Path = Path("data/possible_answer"), - use_memory: bool = False, - use_memory_addition: bool = True, - use_memory_deletion: bool = False, - delete_freq: int = 10, - freq_threshold: int = 5, - utility_threshold: float = 0.5, - enable_thinking: bool = False, - memory_base_url: str = "http://0.0.0.0:8001/", - memory_workspace_id: str = "bfcl_test"): +def run_agent( + dataset_name: str, + experiment_suffix: str, + max_workers: int, + num_runs: int = 4, + model_name: str = "qwen3-8b", + data_path: str = "data/multiturn_data_base_val.jsonl", + answer_path: Path = Path("data/possible_answer"), + use_memory: bool = False, + use_memory_addition: bool = True, + use_memory_deletion: bool = False, + delete_freq: int = 10, + freq_threshold: int = 5, + utility_threshold: float = 0.5, + enable_thinking: bool = False, + memory_base_url: str = "http://0.0.0.0:8001/", + memory_workspace_id: str = "bfcl_test", +): experiment_name = dataset_name + "_" + experiment_suffix - path: Path = Path(f"./exp_result/{model_name}/with_think" if enable_thinking else f"./exp_result/{model_name}/no_think") + path: Path = Path( + f"./exp_result/{model_name}/with_think" if enable_thinking else f"./exp_result/{model_name}/no_think", + ) path.mkdir(parents=True, exist_ok=True) - + with open(data_path, "r", encoding="utf-8") as f: task_ids = [json.loads(l)["id"] for l in f] @@ -48,7 +52,7 @@ def run_agent(dataset_name: str, future_list: list = [] for i in range(max_workers): actor = BFCLAgent.remote( - index=i, + index=i, task_ids=task_ids[i::max_workers], experiment_name=experiment_name, data_path=data_path, @@ -63,7 +67,7 @@ def run_agent(dataset_name: str, utility_threshold=utility_threshold, enable_thinking=enable_thinking, memory_base_url=memory_base_url, - memory_workspace_id=memory_workspace_id + memory_workspace_id=memory_workspace_id, ) future = actor.execute.remote() future_list.append(future) @@ -80,7 +84,7 @@ def run_agent(dataset_name: str, logger.info(f"{i + 1}/{len(task_ids)} complete") dump_file() - + def main(): max_workers = 4 @@ -94,11 +98,11 @@ def main(): ray.init(num_cpus=max_workers) for run_id in range(num_runs): run_agent( - dataset_name="bfcl-multi-turn-base", + dataset_name="bfcl-multi-turn-base", experiment_suffix=f"wo-exp", model_name="qwen3-8b", - max_workers=max_workers, - num_runs=1, + max_workers=max_workers, + num_runs=1, data_path="data/multiturn_data_base_val.jsonl", answer_path=Path("data/possible_answer"), enable_thinking=False, diff --git a/cookbook/bfcl/run_exp_statistic.py b/cookbook/bfcl/run_exp_statistic.py index 8d689599..aa24ccb8 100644 --- a/cookbook/bfcl/run_exp_statistic.py +++ b/cookbook/bfcl/run_exp_statistic.py @@ -1,8 +1,8 @@ import json -from pathlib import Path from collections import defaultdict -import pandas as pd +from pathlib import Path +import pandas as pd from loguru import logger @@ -24,7 +24,7 @@ def calculate_best_at_k(scores: list, k: int) -> float: group_maxs = [] for i in range(0, len(scores), k): - group = scores[i:i + k] + group = scores[i : i + k] group_maxs.append(max(group)) return sum(group_maxs) / len(group_maxs) @@ -36,8 +36,8 @@ def calculate_pass_at_k(scores: list, k: int) -> float: group_maxs = [] for i in range(0, len(scores), k): - group = scores[i:i + k] - is_pass = 1.0 if max(group) >=1.0 else 0.0 + group = scores[i : i + k] + is_pass = 1.0 if max(group) >= 1.0 else 0.0 group_maxs.append(is_pass) return sum(group_maxs) / len(group_maxs) @@ -133,7 +133,7 @@ def run_exp_statistic(): # Create and display table if all_results: df = pd.DataFrame(list(all_results.values())) - df = df.set_index('file') + df = df.set_index("file") # Sort columns by the number in column name (best@8, best@4, best@2, best@1) # best_columns = [col for col in df.columns if col.startswith('best@')] @@ -156,4 +156,4 @@ def run_exp_statistic(): if __name__ == "__main__": - run_exp_statistic() \ No newline at end of file + run_exp_statistic() diff --git a/cookbook/bfcl/split_into_trainval.py b/cookbook/bfcl/split_into_trainval.py index dee65376..388deb09 100644 --- a/cookbook/bfcl/split_into_trainval.py +++ b/cookbook/bfcl/split_into_trainval.py @@ -1,9 +1,10 @@ +import argparse import json import random -import argparse + def split_jsonl(input_file, train_file, val_file, ratio=0.8): - with open(input_file, 'r', encoding='utf-8') as f: + with open(input_file, "r", encoding="utf-8") as f: data = [json.loads(line) for line in f] random.shuffle(data) @@ -11,19 +12,20 @@ def split_jsonl(input_file, train_file, val_file, ratio=0.8): train_data = data[:split_idx] val_data = data[split_idx:] - with open(train_file, 'w', encoding='utf-8') as f: + with open(train_file, "w", encoding="utf-8") as f: for item in train_data: - f.write(json.dumps(item, ensure_ascii=False) + '\n') - with open(val_file, 'w', encoding='utf-8') as f: + f.write(json.dumps(item, ensure_ascii=False) + "\n") + with open(val_file, "w", encoding="utf-8") as f: for item in val_data: - f.write(json.dumps(item, ensure_ascii=False) + '\n') + f.write(json.dumps(item, ensure_ascii=False) + "\n") + if __name__ == "__main__": - parser = argparse.ArgumentParser(description='Split JSONL file into train and validation sets.') - parser.add_argument('--input', required=True, help='Path to input JSONL file') - parser.add_argument('--train', required=True, help='Path to output train file') - parser.add_argument('--val', required=True, help='Path to output validation file') - parser.add_argument('--ratio', type=float, default=0.5, help='Train ratio (default: 0.8)') + parser = argparse.ArgumentParser(description="Split JSONL file into train and validation sets.") + parser.add_argument("--input", required=True, help="Path to input JSONL file") + parser.add_argument("--train", required=True, help="Path to output train file") + parser.add_argument("--val", required=True, help="Path to output validation file") + parser.add_argument("--ratio", type=float, default=0.5, help="Train ratio (default: 0.8)") args = parser.parse_args() split_jsonl(args.input, args.train, args.val, args.ratio) diff --git a/cookbook/frozenlake/frozenlake_prompts.yaml b/cookbook/frozenlake/frozenlake_prompts.yaml index 4c1f5a70..080cdb17 100644 --- a/cookbook/frozenlake/frozenlake_prompts.yaml +++ b/cookbook/frozenlake/frozenlake_prompts.yaml @@ -1,47 +1,47 @@ frozenlake_sys_prompt_no_slippery: | You are an AI agent playing FrozenLake game. Your goal is to navigate from Start (S) to Goal (G) while avoiding Holes (H). - + Game Rules: - S: Starting position (safe) - F: Frozen surface (safe to walk on) - H: Hole (you fall in and lose) - G: Goal (you win!) - []: Your current position - + Actions: - 0: Move LEFT - - 1: Move DOWN + - 1: Move DOWN - 2: Move RIGHT - 3: Move UP - - Your task: Analyze the current state and choose the best action (0-3) to reach the Goal while avoiding Holes. + + Your task: Analyze the current state and choose the best action (0-3) to reach the Goal while avoiding Holes. While ensuring a safe arrival at the goal, you should aim to complete the task in as few steps as possible. Think step by step, and respond with your thoughts and then clearly state your action as a number (0-3) in format {"action":"(0-3)"}. frozenlake_sys_prompt_slippery: | You are an AI agent playing FrozenLake game. Your goal is to navigate from Start (S) to Goal (G) while avoiding Holes (H). - + Game Rules: - S: Starting position (safe) - F: Frozen surface (safe to walk on) - H: Hole (you fall in and lose) - G: Goal (you win!) - []: Your current position - + Actions: - 0: Move LEFT - - 1: Move DOWN + - 1: Move DOWN - 2: Move RIGHT - 3: Move UP - + The ice is slippery, so you might not always move in the intended direction! you will move in intended direction with probability of 1/3 else will move in either perpendicular direction with equal probability of 1/3 in both directions. - + For example, if action is left, then: - P(move left)=1/3 - P(move up)=1/3 - P(move down)=1/3 - - Your task: Analyze the current state and choose the best action (0-3) to reach the Goal while avoiding Holes. + + Your task: Analyze the current state and choose the best action (0-3) to reach the Goal while avoiding Holes. While ensuring a safe arrival at the goal, you should aim to complete the task in as few steps as possible. Think step by step, and respond with your thoughts and then clearly state your action as a number (0-3) in format {{"action":"(0-3)"}}. diff --git a/cookbook/frozenlake/frozenlake_react_agent.py b/cookbook/frozenlake/frozenlake_react_agent.py index 1dceb0ea..91737d13 100644 --- a/cookbook/frozenlake/frozenlake_react_agent.py +++ b/cookbook/frozenlake/frozenlake_react_agent.py @@ -1,24 +1,22 @@ -import os +import random import re import time -import json +from dataclasses import dataclass +from typing import List, Dict, Any +import gymnasium as gym import ray import requests -import random -from typing import List, Dict, Any, Optional -from dataclasses import dataclass -import numpy as np -import gymnasium as gym -from gymnasium.envs.toy_text.frozen_lake import generate_random_map -from openai import OpenAI -from loguru import logger import yaml from dotenv import load_dotenv +from gymnasium.envs.toy_text.frozen_lake import generate_random_map +from loguru import logger +from openai import OpenAI from tqdm import tqdm load_dotenv("../../.env") + @dataclass class GameResult: task_id: str @@ -30,20 +28,23 @@ class GameResult: trajectory: List[Dict] map_config: Dict[str, Any] + @ray.remote class FrozenLakeReactAgent: """A ReAct Agent for FrozenLake game with task memory learning.""" - def __init__(self, - index: int, - task_configs: List[Dict], - experiment_name: str, - model_name: str = "qwen3-8b", - temperature: float = 0.7, - max_steps: int = 50, - num_runs: int = 1, - use_task_memory: bool = False, - make_task_memory: bool = False): + def __init__( + self, + index: int, + task_configs: List[Dict], + experiment_name: str, + model_name: str = "qwen3-8b", + temperature: float = 0.7, + max_steps: int = 50, + num_runs: int = 1, + use_task_memory: bool = False, + make_task_memory: bool = False, + ): self.index = index self.task_configs = task_configs @@ -64,11 +65,13 @@ class FrozenLakeReactAgent: def _load_prompts(self) -> Dict[str, str]: """Load prompts from yaml file""" try: - with open("frozenlake_prompts.yaml", 'r', encoding='utf-8') as f: + with open("frozenlake_prompts.yaml", "r", encoding="utf-8") as f: return yaml.safe_load(f) except FileNotFoundError: logger.warning("Prompt file not found, using default prompts") - raise FileNotFoundError("Prompt file not found. Please check your current path (should be ./cook/frozenlake) and try again.") + raise FileNotFoundError( + "Prompt file not found. Please check your current path (should be ./cook/frozenlake) and try again.", + ) def call_llm(self, messages: List[Dict]) -> str: """Call LLM with retry logic""" @@ -79,7 +82,7 @@ class FrozenLakeReactAgent: messages=messages, temperature=self.temperature, extra_body={"enable_thinking": False}, - seed=0 + seed=0, ) return response.choices[0].message.content except Exception as e: @@ -93,7 +96,7 @@ class FrozenLakeReactAgent: nrow, ncol = desc.shape # Convert to string grid - grid = [[cell.decode('utf-8') for cell in row] for row in desc] + grid = [[cell.decode("utf-8") for cell in row] for row in desc] # Get current position row, col = observation // ncol, observation % ncol @@ -134,7 +137,7 @@ class FrozenLakeReactAgent: "workspace_id": workspace_id, "query": query, }, - timeout=60 + timeout=60, ) if response.status_code == 200: @@ -155,7 +158,7 @@ class FrozenLakeReactAgent: r'["\']action["\']\s*:\s*["\']([0-3])["\']', r'"action"\s*:\s*"([0-3])"', r"'action'\s*:\s*'([0-3])'", - r'\baction["\']?\s*[:=]\s*["\']?([0-3])' + r'\baction["\']?\s*[:=]\s*["\']?([0-3])', ] for pattern in patterns: @@ -190,8 +193,7 @@ class FrozenLakeReactAgent: env = gym.make("FrozenLake-v1", **env_kwargs) # Get map description for task memory - map_str = '\n'.join([''.join([cell.decode('utf-8') for cell in row]) - for row in env.unwrapped.desc]) + map_str = "\n".join(["".join([cell.decode("utf-8") for cell in row]) for row in env.unwrapped.desc]) # Build messages system_prompt = self.build_system_prompt(is_slippery) @@ -203,7 +205,8 @@ class FrozenLakeReactAgent: memory_content = f"Here are some relevant tips from previous successful games:\n\n{task_memory}\n\nUse these tips to help you succeed." messages.append({"role": "user", "content": memory_content}) messages.append( - {"role": "assistant", "content": "I'll use these tips to navigate the frozen lake successfully."}) + {"role": "assistant", "content": "I'll use these tips to navigate the frozen lake successfully."}, + ) # Initialize game observation, info = env.reset() @@ -230,16 +233,18 @@ class FrozenLakeReactAgent: done = terminated or truncated # Record trajectory step - trajectory.append({ - "step": step, - "state": observation, - "action": action, - "action_name": self.action_map[action], - "reward": reward, - "next_state": next_observation, - "done": done, - "llm_response": response - }) + trajectory.append( + { + "step": step, + "state": observation, + "action": action, + "action_name": self.action_map[action], + "reward": reward, + "next_state": next_observation, + "done": done, + "llm_response": response, + }, + ) if done: if terminated and reward > 0: @@ -275,8 +280,8 @@ class FrozenLakeReactAgent: "map_id": map_id, "is_slippery": is_slippery, "map_size": map_size, - "use_task_memory": self.use_task_memory - } + "use_task_memory": self.use_task_memory, + }, ) return result, messages @@ -311,9 +316,9 @@ class FrozenLakeReactAgent: url=base_url + "summary_task_memory", json={ "workspace_id": workspace_id, - "trajectories": trajectories + "trajectories": trajectories, }, - timeout=300 + timeout=300, ) if response.status_code == 200: @@ -347,7 +352,7 @@ class FrozenLakeReactAgent: "steps": result.steps, "reward": result.reward, "map_config": result.map_config, - "trajectory": result.trajectory + "trajectory": result.trajectory, } all_results[-1] = result_dict @@ -364,10 +369,10 @@ class FrozenLakeReactAgent: steps=result_dict["steps"], reward=result_dict["reward"], trajectory=result_dict["trajectory"], - map_config=result_dict["map_config"] + map_config=result_dict["map_config"], ) game_results.append(game_result) self.save_task_memory(game_results, all_messages) - return all_results \ No newline at end of file + return all_results diff --git a/cookbook/frozenlake/map_manager.py b/cookbook/frozenlake/map_manager.py index 3184c9a7..a37617a0 100644 --- a/cookbook/frozenlake/map_manager.py +++ b/cookbook/frozenlake/map_manager.py @@ -4,118 +4,125 @@ Map Management Tool - Pre-generate and manage test maps """ import json -import numpy as np from pathlib import Path from typing import List, Optional, Dict, Any -from loguru import logger + +import numpy as np from gymnasium.envs.toy_text.frozen_lake import generate_random_map +from loguru import logger class MapManager: - """Map Manager - pre-generating, storing and loading test maps""" + """Map Manager - pre-generating, storing and loading test maps""" - def __init__(self, data_dir: str = "./map/"): - self.data_dir = Path(data_dir) - self.data_dir.mkdir(parents=True, exist_ok=True) + def __init__(self, data_dir: str = "./map/"): + self.data_dir = Path(data_dir) + self.data_dir.mkdir(parents=True, exist_ok=True) - def generate_test_maps(self, num_maps: int, map_size: int = 4, - base_seed: int = 10000) -> str: - """ - Generate test map collection and save + def generate_test_maps( + self, + num_maps: int, + map_size: int = 4, + base_seed: int = 10000, + ) -> str: + """ + Generate test map collection and save - Args: - num_maps: Number of maps to generate - map_size: Map size - base_seed: Base random seed + Args: + num_maps: Number of maps to generate + map_size: Map size + base_seed: Base random seed - Returns: - Path of saved file - """ - logger.info(f"🗺️ Generating {num_maps} test maps (size={map_size})") + Returns: + Path of saved file + """ + logger.info(f"🗺️ Generating {num_maps} test maps (size={map_size})") - maps_data = [] - for i in range(num_maps): - seed = base_seed + i - np.random.seed(seed) - map_desc = generate_random_map(size=map_size) + maps_data = [] + for i in range(num_maps): + seed = base_seed + i + np.random.seed(seed) + map_desc = generate_random_map(size=map_size) - maps_data.append({ - "map_id": i, - "seed": seed, - "map_size": map_size, - "map_desc": map_desc # Convert to list for JSON serialization - }) + maps_data.append( + { + "map_id": i, + "seed": seed, + "map_size": map_size, + "map_desc": map_desc, # Convert to list for JSON serialization + }, + ) - # Save to file - filename = f"test_maps_{num_maps}_{map_size}x{map_size}.jsonl" - filepath = self.data_dir / filename + # Save to file + filename = f"test_maps_{num_maps}_{map_size}x{map_size}.jsonl" + filepath = self.data_dir / filename - with open(filepath, "w", encoding="utf-8") as f: - for map_data in maps_data: - f.write(json.dumps(map_data, ensure_ascii=False) + "\n") + with open(filepath, "w", encoding="utf-8") as f: + for map_data in maps_data: + f.write(json.dumps(map_data, ensure_ascii=False) + "\n") - logger.info(f"✅ Test maps saved to {filepath}") - return str(filepath) + logger.info(f"✅ Test maps saved to {filepath}") + return str(filepath) - def load_test_maps(self, filepath: str) -> List[Dict[str, Any]]: - """ - Load test maps + def load_test_maps(self, filepath: str) -> List[Dict[str, Any]]: + """ + Load test maps - Args: - filepath: Map file path + Args: + filepath: Map file path - Returns: - Map data list - """ - if not Path(filepath).exists(): - raise FileNotFoundError(f"Map file not found: {filepath}") + Returns: + Map data list + """ + if not Path(filepath).exists(): + raise FileNotFoundError(f"Map file not found: {filepath}") - maps_data = [] - with open(filepath, "r", encoding="utf-8") as f: - for line in f: - if line.strip(): - map_data = json.loads(line) - # Convert list back to numpy array - maps_data.append(map_data) + maps_data = [] + with open(filepath, "r", encoding="utf-8") as f: + for line in f: + if line.strip(): + map_data = json.loads(line) + # Convert list back to numpy array + maps_data.append(map_data) - logger.info(f"📖 Loaded {len(maps_data)} test maps from {filepath}") - return maps_data + logger.info(f"📖 Loaded {len(maps_data)} test maps from {filepath}") + return maps_data - def get_map_by_index(self, maps_data: List[Dict], index: int) -> Optional[list]: - """Get map by index""" - if 0 <= index < len(maps_data): - return maps_data[index]["map_desc"] - return None + def get_map_by_index(self, maps_data: List[Dict], index: int) -> Optional[list]: + """Get map by index""" + if 0 <= index < len(maps_data): + return maps_data[index]["map_desc"] + return None - def get_or_create_test_maps(self, num_maps: int, map_size: int = 4) -> List[Dict[str, Any]]: - """ - Get or create test maps - If file exists and has sufficient quantity, load directly; otherwise regenerate - """ - filename = f"test_maps_{num_maps}_{map_size}x{map_size}.jsonl" - filepath = self.data_dir / filename + def get_or_create_test_maps(self, num_maps: int, map_size: int = 4) -> List[Dict[str, Any]]: + """ + Get or create test maps + If file exists and has sufficient quantity, load directly; otherwise regenerate + """ + filename = f"test_maps_{num_maps}_{map_size}x{map_size}.jsonl" + filepath = self.data_dir / filename - if filepath.exists(): - try: - maps_data = self.load_test_maps(str(filepath)) - if len(maps_data) >= num_maps: - logger.info(f"✅ Using existing test maps: {filepath}") - return maps_data[:num_maps] # Return required number of maps - except Exception as e: - logger.warning(f"⚠️ Failed to load existing maps: {e}, regenerating...") + if filepath.exists(): + try: + maps_data = self.load_test_maps(str(filepath)) + if len(maps_data) >= num_maps: + logger.info(f"✅ Using existing test maps: {filepath}") + return maps_data[:num_maps] # Return required number of maps + except Exception as e: + logger.warning(f"⚠️ Failed to load existing maps: {e}, regenerating...") - # File doesn't exist or insufficient quantity, regenerate - self.generate_test_maps(num_maps, map_size) - return self.load_test_maps(str(filepath)) + # File doesn't exist or insufficient quantity, regenerate + self.generate_test_maps(num_maps, map_size) + return self.load_test_maps(str(filepath)) if __name__ == "__main__": - # Usage example - manager = MapManager() + # Usage example + manager = MapManager() - # Generate 100 4x4 test maps - manager.generate_test_maps(num_maps=100, map_size=4) + # Generate 100 4x4 test maps + manager.generate_test_maps(num_maps=100, map_size=4) - # Load and view the first map - maps = manager.load_test_maps("./map/test_maps_100_4x4.jsonl") - print(f"First map:\n{maps[0]['map_desc']}") \ No newline at end of file + # Load and view the first map + maps = manager.load_test_maps("./map/test_maps_100_4x4.jsonl") + print(f"First map:\n{maps[0]['map_desc']}") diff --git a/cookbook/frozenlake/run_exp_statistic.py b/cookbook/frozenlake/run_exp_statistic.py index 3646b896..3d7af652 100644 --- a/cookbook/frozenlake/run_exp_statistic.py +++ b/cookbook/frozenlake/run_exp_statistic.py @@ -1,8 +1,9 @@ import json -import pandas as pd -from pathlib import Path from collections import defaultdict +from pathlib import Path from typing import Dict, List, Tuple + +import pandas as pd from loguru import logger @@ -24,7 +25,7 @@ def calculate_best_at_k(scores: List[float], k: int) -> float: group_maxs = [] for i in range(0, len(scores), k): - group = scores[i:i + k] + group = scores[i : i + k] group_maxs.append(max(group)) return sum(group_maxs) / len(group_maxs) @@ -125,7 +126,7 @@ def process_single_result(data: Dict, condition_results: Dict): if part.startswith("map"): try: # Extract number from "mapXX" - map_num = ''.join(filter(str.isdigit, part)) + map_num = "".join(filter(str.isdigit, part)) if map_num: map_id = int(map_num) break @@ -189,8 +190,10 @@ def calculate_file_metrics(condition_results: Dict, filename: str) -> Dict: file_metrics[f"{condition}_map_details"] = map_success_rates - logger.info(f"{filename} - {condition}: {overall_success:.3f} success rate, " - f"{len(map_results)} maps, {num_runs} runs each") + logger.info( + f"{filename} - {condition}: {overall_success:.3f} success rate, " + f"{len(map_results)} maps, {num_runs} runs each", + ) return file_metrics @@ -213,7 +216,7 @@ def generate_analysis_report(all_results: Dict): if summary_data: df_summary = pd.DataFrame(summary_data) - df_summary = df_summary.set_index('file') + df_summary = df_summary.set_index("file") print("\n" + "=" * 100) print("FROZENLAKE EXPERIMENT RESULTS SUMMARY") @@ -367,4 +370,4 @@ def main(): if __name__ == "__main__": - main() \ No newline at end of file + main() diff --git a/cookbook/frozenlake/run_frozenlake.py b/cookbook/frozenlake/run_frozenlake.py index 66f30fb8..bccad078 100644 --- a/cookbook/frozenlake/run_frozenlake.py +++ b/cookbook/frozenlake/run_frozenlake.py @@ -1,18 +1,18 @@ -import os -import time import json -import ray +import time from pathlib import Path from typing import List, Dict + import numpy as np -from loguru import logger +import ray from gymnasium.envs.toy_text.frozen_lake import generate_random_map +from loguru import logger from frozenlake_react_agent import FrozenLakeReactAgent from map_manager import MapManager -def generate_training_configs(num_maps: int = 20, map_size: int = 4, is_slippery: bool=False) -> List[Dict]: +def generate_training_configs(num_maps: int = 20, map_size: int = 4, is_slippery: bool = False) -> List[Dict]: """Generate random maps for training/task memory generation""" configs = [] @@ -25,7 +25,7 @@ def generate_training_configs(num_maps: int = 20, map_size: int = 4, is_slippery "map_desc": random_map, "map_size": map_size, "is_slippery": is_slippery, - "task_id": f"train_{i}_{is_slippery}" + "task_id": f"train_{i}_{is_slippery}", } configs.append(config) @@ -43,7 +43,7 @@ def generate_test_configs(num_test_maps: int = 100, is_slippery: bool = False) - configs = [] for map_data in maps_data: - map_desc = np.array([list(row) for row in map_data["map_desc"]], dtype='c') + map_desc = np.array([list(row) for row in map_data["map_desc"]], dtype="c") map_id = map_data["map_id"] for use_memory in [True, False]: @@ -54,7 +54,7 @@ def generate_test_configs(num_test_maps: int = 100, is_slippery: bool = False) - "is_slippery": is_slippery, "use_task_memory": use_memory, "map_id": map_id, - "task_id": f"test_map{map_id}_slip{is_slippery}_mem{use_memory}" + "task_id": f"test_map{map_id}_slip{is_slippery}_mem{use_memory}", } configs.append(config) @@ -62,7 +62,13 @@ def generate_test_configs(num_test_maps: int = 100, is_slippery: bool = False) - return configs -def train(experiment_name: str, max_workers: int = 2, num_runs: int = 3, num_training_maps= 15, is_slippery: bool= False) -> None: +def train( + experiment_name: str, + max_workers: int = 2, + num_runs: int = 3, + num_training_maps=15, + is_slippery: bool = False, +) -> None: """Phase 1: Generate task memory from random maps""" logger.info("🎯 Starting Training Phase - Generating Task Memory") logger.info("=" * 60) @@ -116,7 +122,7 @@ def train(experiment_name: str, max_workers: int = 2, num_runs: int = 3, num_tra experiment_name=experiment_name, num_runs=num_runs, use_task_memory=False, - make_task_memory=True + make_task_memory=True, ) results = agent.execute() dump_results() @@ -130,7 +136,13 @@ def train(experiment_name: str, max_workers: int = 2, num_runs: int = 3, num_tra return results -def test(experiment_name: str, max_workers: int = 2, num_runs: int = 5, num_test_maps: int = 100, is_slippery: bool=False) -> None: +def test( + experiment_name: str, + max_workers: int = 2, + num_runs: int = 5, + num_test_maps: int = 100, + is_slippery: bool = False, +) -> None: """Phase 2: Test on fixed maps with/without task memory""" logger.info("🧪 Starting Test Phase - Evaluating Performance") logger.info(f"📊 Testing on {num_test_maps} maps with {num_runs} runs each") @@ -147,8 +159,6 @@ def test(experiment_name: str, max_workers: int = 2, num_runs: int = 5, num_test logger.info(f"📝 Configs without task memory: {len(no_memory_configs)}") logger.info(f"📝 Configs with task memory: {len(memory_configs)}") - - def dump_results(suffix: str): output_file = path / f"{experiment_name}_test_{suffix}.jsonl" with open(output_file, "w") as f: @@ -164,7 +174,7 @@ def test(experiment_name: str, max_workers: int = 2, num_runs: int = 5, num_test experiment_name=experiment_name, max_workers=max_workers, num_runs=num_runs, - use_task_memory=False + use_task_memory=False, ) all_results.extend(results_no_memory) dump_results("no_memory") @@ -177,7 +187,7 @@ def test(experiment_name: str, max_workers: int = 2, num_runs: int = 5, num_test experiment_name=experiment_name, max_workers=max_workers, num_runs=num_runs, - use_task_memory=True + use_task_memory=True, ) all_results.extend(results_with_memory) dump_results("with_memory") @@ -185,8 +195,13 @@ def test(experiment_name: str, max_workers: int = 2, num_runs: int = 5, num_test return all_results -def run_test_configs(configs: List[Dict], experiment_name: str, max_workers: int, - num_runs: int, use_task_memory: bool) -> List[Dict]: +def run_test_configs( + configs: List[Dict], + experiment_name: str, + max_workers: int, + num_runs: int, + use_task_memory: bool, +) -> List[Dict]: """Run a set of test configurations""" results = [] @@ -201,7 +216,7 @@ def run_test_configs(configs: List[Dict], experiment_name: str, max_workers: int experiment_name=experiment_name, num_runs=num_runs, use_task_memory=use_task_memory, - make_task_memory=False + make_task_memory=False, ) future = agent.execute.remote() future_list.append(future) @@ -220,7 +235,7 @@ def run_test_configs(configs: List[Dict], experiment_name: str, max_workers: int experiment_name=experiment_name, num_runs=num_runs, use_task_memory=use_task_memory, - make_task_memory=False + make_task_memory=False, ) results = agent.execute() @@ -255,21 +270,20 @@ def main(): max_workers=max_workers, num_runs=training_runs, num_training_maps=num_training_maps, - is_slippery=is_slippery + is_slippery=is_slippery, ) # Wait a bit for task memory service to process logger.info("⏰ Waiting for task memory service to process data...") time.sleep(10) - # Phase 2: Testing (Performance Evaluation) test_results = test( experiment_name=experiment_name, max_workers=max_workers, num_runs=test_runs, num_test_maps=num_test_maps, - is_slippery=is_slippery + is_slippery=is_slippery, ) # Summary @@ -293,4 +307,4 @@ def main(): if __name__ == "__main__": - main() \ No newline at end of file + main() diff --git a/cookbook/simple_demo/import_usage_demo.py b/cookbook/simple_demo/import_usage_demo.py index 2781b804..5e682b2c 100644 --- a/cookbook/simple_demo/import_usage_demo.py +++ b/cookbook/simple_demo/import_usage_demo.py @@ -7,10 +7,11 @@ from reme_ai import ReMeApp # Task Memory Management Examples # ============================================ + async def summary_task_memory(app: ReMeApp): """ Experience Summarizer: Learn from execution trajectories - + curl -X POST http://localhost:8002/summary_task_memory \ -H "Content-Type: application/json" \ -d '{ @@ -26,11 +27,11 @@ async def summary_task_memory(app: ReMeApp): trajectories=[ { "messages": [ - {"role": "user", "content": "Help me create a project plan"} + {"role": "user", "content": "Help me create a project plan"}, ], - "score": 1.0 - } - ] + "score": 1.0, + }, + ], ) print("Summary Task Memory Result:") print(result["answer"]) @@ -39,7 +40,7 @@ async def summary_task_memory(app: ReMeApp): async def retrieve_task_memory(app: ReMeApp): """ Retriever: Get relevant memories - + curl -X POST http://localhost:8002/retrieve_task_memory \ -H "Content-Type: application/json" \ -d '{ @@ -52,7 +53,7 @@ async def retrieve_task_memory(app: ReMeApp): name="retrieve_task_memory", workspace_id="task_workspace", query="How to efficiently manage project progress?", - top_k=1 + top_k=1, ) print("Retrieve Task Memory Result:") print(result["answer"]) @@ -62,10 +63,11 @@ async def retrieve_task_memory(app: ReMeApp): # Personal Memory Management Examples # ============================================ + async def summary_personal_memory(app: ReMeApp): """ Memory Integration: Learn from user interactions - + curl -X POST http://localhost:8002/summary_personal_memory \ -H "Content-Type: application/json" \ -d '{ @@ -85,11 +87,13 @@ async def summary_personal_memory(app: ReMeApp): { "messages": [ {"role": "user", "content": "I like to drink coffee while working in the morning"}, - {"role": "assistant", - "content": "I understand, you prefer to start your workday with coffee to stay energized"} - ] - } - ] + { + "role": "assistant", + "content": "I understand, you prefer to start your workday with coffee to stay energized", + }, + ], + }, + ], ) print("Summary Personal Memory Result:") print(result["answer"]) @@ -98,7 +102,7 @@ async def summary_personal_memory(app: ReMeApp): async def retrieve_personal_memory(app: ReMeApp): """ Memory Retrieval: Get personal memory fragments - + curl -X POST http://localhost:8002/retrieve_personal_memory \ -H "Content-Type: application/json" \ -d '{ @@ -111,7 +115,7 @@ async def retrieve_personal_memory(app: ReMeApp): name="retrieve_personal_memory", workspace_id="task_workspace", query="What are the user's work habits?", - top_k=5 + top_k=5, ) print("Retrieve Personal Memory Result:") print(result["answer"]) @@ -121,10 +125,11 @@ async def retrieve_personal_memory(app: ReMeApp): # Tool Memory Management Examples # ============================================ + async def add_tool_call_result(app: ReMeApp): """ Record tool execution results - + curl -X POST http://localhost:8002/add_tool_call_result \ -H "Content-Type: application/json" \ -d '{ @@ -153,9 +158,9 @@ async def add_tool_call_result(app: ReMeApp): "output": "Found 10 relevant results...", "token_cost": 150, "success": True, - "time_cost": 2.3 - } - ] + "time_cost": 2.3, + }, + ], ) print("Add Tool Call Result:") print(result["answer"]) @@ -164,7 +169,7 @@ async def add_tool_call_result(app: ReMeApp): async def summary_tool_memory(app: ReMeApp): """ Generate usage guidelines from history - + curl -X POST http://localhost:8002/summary_tool_memory \ -H "Content-Type: application/json" \ -d '{ @@ -175,7 +180,7 @@ async def summary_tool_memory(app: ReMeApp): result = await app.async_execute( name="summary_tool_memory", workspace_id="tool_workspace", - tool_names="web_search" + tool_names="web_search", ) print("Summary Tool Memory Result:") print(result["answer"]) @@ -184,7 +189,7 @@ async def summary_tool_memory(app: ReMeApp): async def retrieve_tool_memory(app: ReMeApp): """ Retrieve tool guidelines before use - + curl -X POST http://localhost:8002/retrieve_tool_memory \ -H "Content-Type: application/json" \ -d '{ @@ -195,7 +200,7 @@ async def retrieve_tool_memory(app: ReMeApp): result = await app.async_execute( name="retrieve_tool_memory", workspace_id="tool_workspace", - tool_names="web_search" + tool_names="web_search", ) print("Retrieve Tool Memory Result:") print(result["answer"]) @@ -205,10 +210,11 @@ async def retrieve_tool_memory(app: ReMeApp): # Vector Store Management Example # ============================================ + async def load_vector_store(app: ReMeApp): """ Load pre-built memories - + curl -X POST http://localhost:8002/vector_store \ -H "Content-Type: application/json" \ -d '{ @@ -221,7 +227,7 @@ async def load_vector_store(app: ReMeApp): name="vector_store", workspace_id="appworld", action="load", - path="./docs/library/" + path="./docs/library/", ) print("Load Vector Store Result:") print(result["answer"]) @@ -231,12 +237,13 @@ async def load_vector_store(app: ReMeApp): # Main Execution # ============================================ + async def main(): """Run all examples""" async with ReMeApp( - "llm.default.model_name=qwen3-30b-a3b-thinking-2507", - "embedding_model.default.model_name=text-embedding-v4", - "vector_store.default.backend=memory" + "llm.default.model_name=qwen3-30b-a3b-thinking-2507", + "embedding_model.default.model_name=text-embedding-v4", + "vector_store.default.backend=memory", ) as app: print("=" * 60) print("Task Memory Examples") diff --git a/cookbook/simple_demo/use_personal_memory_demo.py b/cookbook/simple_demo/use_personal_memory_demo.py index 24a43092..9439e59e 100644 --- a/cookbook/simple_demo/use_personal_memory_demo.py +++ b/cookbook/simple_demo/use_personal_memory_demo.py @@ -6,10 +6,11 @@ import aiohttp # API base URL base_url = "http://0.0.0.0:8002" + async def main(): # Create a unique workspace ID workspace_id = "personal_memory_demo" - + async with aiohttp.ClientSession() as session: # Step 1: Clear existing memories in the workspace print("Clearing existing memories...") @@ -19,39 +20,57 @@ async def main(): "action": "delete", "workspace_id": workspace_id, }, - headers={"Content-Type": "application/json"} + headers={"Content-Type": "application/json"}, ) as response: result = await response.json() print(json.dumps(result, ensure_ascii=False)) - + # Step 2: Create a conversation with rich personal information print("\nCreating conversation with personal information...") messages = [ - {"role": "user", "content": "My name is John Smith, I'm 28 years old, and I work at a tech company in San Francisco"}, + { + "role": "user", + "content": "My name is John Smith, I'm 28 years old, and I work at a tech company in San Francisco", + }, {"role": "assistant", "content": "Nice to meet you, John!"}, - {"role": "user", "content": "I'm a software engineer, mainly doing backend development using Python and Go"}, + { + "role": "user", + "content": "I'm a software engineer, mainly doing backend development using Python and Go", + }, {"role": "assistant", "content": "I see, you're a backend engineer working with Python and Go."}, - {"role": "user", "content": "I enjoy playing basketball and watching sci-fi movies. I recently watched Dune Part 2"}, - {"role": "assistant", "content": "Basketball and sci-fi movies are great hobbies! Dune Part 2 was indeed amazing."}, + { + "role": "user", + "content": "I enjoy playing basketball and watching sci-fi movies. I recently watched Dune Part 2", + }, + { + "role": "assistant", + "content": "Basketball and sci-fi movies are great hobbies! Dune Part 2 was indeed amazing.", + }, {"role": "user", "content": "I have a cat named Shadow who is 3 years old"}, {"role": "assistant", "content": "Shadow sounds adorable! 3-year-old cats are quite playful."}, {"role": "user", "content": "I'm planning a trip to Japan next month, mainly to Tokyo and Kyoto"}, - {"role": "assistant", "content": "Your Japan trip sounds exciting! Tokyo and Kyoto are both wonderful destinations with their own unique charm."}, + { + "role": "assistant", + "content": "Your Japan trip sounds exciting! Tokyo and Kyoto are both wonderful destinations with their own unique charm.", + }, {"role": "user", "content": "I'm really interested in Japanese cuisine, especially sushi and ramen"}, - {"role": "assistant", "content": "Japanese cuisine is delicious! Sushi and ramen are very popular choices."}, + { + "role": "assistant", + "content": "Japanese cuisine is delicious! Sushi and ramen are very popular choices.", + }, ] - + # Step 3: Summarize personal memories from the conversation print("\nSummarizing personal memories...") async with session.post( f"{base_url}/summary_personal_memory", json={ "trajectories": [ - {"messages": messages, "score": 1.0} + {"messages": messages, "score": 1.0}, ], "workspace_id": workspace_id, }, - headers={"Content-Type": "application/json"} + headers={"Content-Type": "application/json"}, ) as response: result = await response.json() result = json.dumps(result, ensure_ascii=False, indent=2) @@ -59,11 +78,11 @@ async def main(): with open("personal_memory.jsonl", "w") as f: f.write(result) - + # Wait for the memories to be processed and stored print("\nWaiting for memories to be processed...") await asyncio.sleep(2) - + # Step 4: Retrieve personal memories with different queries queries = [ "What's my name and age?", @@ -71,9 +90,9 @@ async def main(): "What are my hobbies?", "Do I have any pets?", "What are my travel plans?", - "What foods do I like?" + "What foods do I like?", ] - + print("\nRetrieving personal memories...") for query in queries: print(f"\nQuery: {query}") @@ -83,10 +102,11 @@ async def main(): "query": query, "workspace_id": workspace_id, }, - headers={"Content-Type": "application/json"} + headers={"Content-Type": "application/json"}, ) as response: result = await response.json() print(json.dumps(result, ensure_ascii=False, indent=2)) + if __name__ == "__main__": asyncio.run(main()) diff --git a/cookbook/simple_demo/use_task_memory_demo.py b/cookbook/simple_demo/use_task_memory_demo.py index 6dcd63fa..27a012d6 100644 --- a/cookbook/simple_demo/use_task_memory_demo.py +++ b/cookbook/simple_demo/use_task_memory_demo.py @@ -26,10 +26,10 @@ WORKSPACE_ID = "test_workspace" def handle_api_response(response: requests.Response) -> Optional[Dict[str, Any]]: """ Handle API response with proper error checking - + Args: response: Response object from requests - + Returns: Response JSON if successful, None otherwise """ @@ -44,7 +44,7 @@ def handle_api_response(response: requests.Response) -> Optional[Dict[str, Any]] def delete_workspace() -> None: """ Delete the current workspace from the vector store - + Returns: None """ @@ -53,7 +53,7 @@ def delete_workspace() -> None: json={ "workspace_id": WORKSPACE_ID, "action": "delete", - } + }, ) result = handle_api_response(response) @@ -64,17 +64,17 @@ def delete_workspace() -> None: def run_agent(query: str, dump_messages: bool = False) -> List[Dict[str, Any]]: """ Run the agent with a specific query - + Args: query: The query to send to the agent dump_messages: Whether to save messages to a file - + Returns: List of message objects from the conversation """ response = requests.post( url=f"{BASE_URL}react", - json={"query": query} + json={"query": query}, ) result = handle_api_response(response) @@ -93,18 +93,18 @@ def run_agent(query: str, dump_messages: bool = False) -> List[Dict[str, Any]]: with open("task_messages.jsonl", "w") as f: f.write(json.dumps(messages, indent=2, ensure_ascii=False)) print(f"Messages saved to messages.jsonl") - + return messages def run_summary(messages: List[Dict[str, Any]], enable_dump_memory: bool = True) -> None: """ Generate a summary of conversation messages and create task memories - + Args: messages: List of message objects from a conversation enable_dump_memory: Whether to save memory list to a file - + Returns: None """ @@ -118,9 +118,9 @@ def run_summary(messages: List[Dict[str, Any]], enable_dump_memory: bool = True) json={ "workspace_id": WORKSPACE_ID, "trajectories": [ - {"messages": messages, "score": 1.0} - ] - } + {"messages": messages, "score": 1.0}, + ], + }, ) result = handle_api_response(response) @@ -141,10 +141,10 @@ def run_summary(messages: List[Dict[str, Any]], enable_dump_memory: bool = True) def run_retrieve(query: str) -> str: """ Retrieve relevant task memories based on a query - + Args: query: The query to retrieve relevant memories - + Returns: String containing the retrieved memory answer """ @@ -154,7 +154,7 @@ def run_retrieve(query: str) -> str: json={ "workspace_id": WORKSPACE_ID, "query": query, - } + }, ) result = handle_api_response(response) @@ -170,18 +170,18 @@ def run_retrieve(query: str) -> str: def run_agent_with_memory(query_first: str, query_second: str, enable_dump_memory: bool = True) -> List[Dict[str, Any]]: """ Run the agent with memory augmentation - + This function demonstrates how to use task memory to enhance agent responses: 1. First run the agent with the second query to build memory 2. Then summarize the conversation to create memories 3. Retrieve relevant memories for the first query 4. Run the agent with the first query augmented with retrieved memories - + Args: query_first: The query to run with memory augmentation query_second: The query to build initial memories enable_dump_memory: Whether to save memory list to a file - + Returns: List of message objects from the final conversation """ @@ -203,17 +203,17 @@ def run_agent_with_memory(query_first: str, query_second: str, enable_dump_memor augmented_query = f"{retrieved_memory}\n\nUser Question:\n{query_first}" print(f"Augmented query: {augmented_query}") messages = run_agent(query=augmented_query) - + return messages def dump_memory(path: str = "./") -> None: """ Dump the vector store memories to disk - + Args: path: Directory path to save the memories - + Returns: None """ @@ -223,7 +223,7 @@ def dump_memory(path: str = "./") -> None: "workspace_id": WORKSPACE_ID, "action": "dump", "path": path, - } + }, ) result = handle_api_response(response) @@ -234,10 +234,10 @@ def dump_memory(path: str = "./") -> None: def load_memory(path: str = "./") -> None: """ Load memories from disk into the vector store - + Args: path: Directory path to load the memories from - + Returns: None """ @@ -247,7 +247,7 @@ def load_memory(path: str = "./") -> None: "workspace_id": WORKSPACE_ID, "action": "load", "path": path, - } + }, ) result = handle_api_response(response) diff --git a/cookbook/simple_demo/use_task_memory_mcp_demo.py b/cookbook/simple_demo/use_task_memory_mcp_demo.py index 96fe7780..f3bf5ed8 100644 --- a/cookbook/simple_demo/use_task_memory_mcp_demo.py +++ b/cookbook/simple_demo/use_task_memory_mcp_demo.py @@ -26,10 +26,10 @@ WORKSPACE_ID = "test_workspace" async def delete_workspace(client: Client) -> None: """ Delete the current workspace from the vector store - + Args: client: MCP client instance - + Returns: None """ @@ -38,7 +38,7 @@ async def delete_workspace(client: Client) -> None: arguments={ "workspace_id": WORKSPACE_ID, "action": "delete", - } + }, ) print(f"Workspace '{WORKSPACE_ID}' deleted successfully") @@ -53,12 +53,12 @@ async def run_agent(client: Client, query: str, dump_messages: bool = False) -> async def run_summary(client: Client, messages: List[Dict[str, Any]], enable_dump_memory: bool = True) -> None: """ Generate a summary of conversation messages and create task memories - + Args: client: MCP client instance messages: List of message objects from a conversation enable_dump_memory: Whether to save memory list to a file - + Returns: None """ @@ -71,9 +71,9 @@ async def run_summary(client: Client, messages: List[Dict[str, Any]], enable_dum arguments={ "workspace_id": WORKSPACE_ID, "trajectories": [ - {"messages": messages, "score": 1.0} - ] - } + {"messages": messages, "score": 1.0}, + ], + }, ) answer = result.content[0].text @@ -90,11 +90,11 @@ async def run_summary(client: Client, messages: List[Dict[str, Any]], enable_dum async def run_retrieve(client: Client, query: str) -> str: """ Retrieve relevant task memories based on a query - + Args: client: MCP client instance query: The query to retrieve relevant memories - + Returns: String containing the retrieved memory answer """ @@ -103,7 +103,7 @@ async def run_retrieve(client: Client, query: str) -> str: arguments={ "workspace_id": WORKSPACE_ID, "query": query, - } + }, ) answer = result.content[0].text @@ -111,22 +111,27 @@ async def run_retrieve(client: Client, query: str) -> str: return answer -async def run_agent_with_memory(client: Client, query_first: str, query_second: str, enable_dump_memory: bool = True) -> List[Dict[str, Any]]: +async def run_agent_with_memory( + client: Client, + query_first: str, + query_second: str, + enable_dump_memory: bool = True, +) -> List[Dict[str, Any]]: """ Run the agent with memory augmentation - + This function demonstrates how to use task memory to enhance agent responses: 1. First run the agent with the second query to build memory 2. Then summarize the conversation to create memories 3. Retrieve relevant memories for the first query 4. Run the agent with the first query augmented with retrieved memories - + Args: client: MCP client instance query_first: The query to run with memory augmentation query_second: The query to build initial memories enable_dump_memory: Whether to save memory list to a file - + Returns: List of message objects from the final conversation """ @@ -148,18 +153,18 @@ async def run_agent_with_memory(client: Client, query_first: str, query_second: augmented_query = f"{retrieved_memory}\n\nUser Question:\n{query_first}" print(f"Augmented query: {augmented_query}") messages = await run_agent(client, query=augmented_query) - + return messages async def dump_memory(client: Client, path: str = "./") -> None: """ Dump the vector store memories to disk - + Args: client: MCP client instance path: Directory path to save the memories - + Returns: None """ @@ -169,7 +174,7 @@ async def dump_memory(client: Client, path: str = "./") -> None: "workspace_id": WORKSPACE_ID, "action": "dump", "path": path, - } + }, ) print(f"Memory dumped to {path}") @@ -177,11 +182,11 @@ async def dump_memory(client: Client, path: str = "./") -> None: async def load_memory(client: Client, path: str = "./") -> None: """ Load memories from disk into the vector store - + Args: client: MCP client instance path: Directory path to load the memories from - + Returns: None """ @@ -191,7 +196,7 @@ async def load_memory(client: Client, path: str = "./") -> None: "workspace_id": WORKSPACE_ID, "action": "load", "path": path, - } + }, ) print(f"Memory loaded from {path}") diff --git a/cookbook/simple_demo/use_tool_memory_demo.py b/cookbook/simple_demo/use_tool_memory_demo.py index 7b9f66dd..c84f9e75 100644 --- a/cookbook/simple_demo/use_tool_memory_demo.py +++ b/cookbook/simple_demo/use_tool_memory_demo.py @@ -8,6 +8,7 @@ from typing import List, Dict, Any, Optional import requests from dotenv import load_dotenv + from reme_ai.utils.tool_memory_utils import create_mock_tool_call_results load_dotenv() @@ -27,7 +28,7 @@ def api_call(endpoint: str, data: dict) -> Optional[Dict[str, Any]]: def delete_workspace() -> None: """删除工作空间 - + curl example: curl -X POST http://0.0.0.0:8002/vector_store \ -H "Content-Type: application/json" \ @@ -43,7 +44,7 @@ def delete_workspace() -> None: def add_tool_call_results(tool_call_results: List[Dict[str, Any]]) -> bool: """添加工具调用结果到记忆库 - + curl example: curl -X POST http://0.0.0.0:8002/add_tool_call_result \ -H "Content-Type: application/json" \ @@ -62,11 +63,14 @@ def add_tool_call_results(tool_call_results: List[Dict[str, Any]]) -> bool: # 统计不同的工具 tool_names = set(r.get("tool_name") for r in tool_call_results) print(f"\n[ADD] {len(tool_call_results)} results for {len(tool_names)} tools: {', '.join(sorted(tool_names))}") - - result = api_call("add_tool_call_result", { - "workspace_id": WORKSPACE_ID, - "tool_call_results": tool_call_results - }) + + result = api_call( + "add_tool_call_result", + { + "workspace_id": WORKSPACE_ID, + "tool_call_results": tool_call_results, + }, + ) if result: memory_list = result.get("metadata", {}).get("memory_list", []) print(f"✓ Added successfully, created/updated {len(memory_list)} tool memories") @@ -76,7 +80,7 @@ def add_tool_call_results(tool_call_results: List[Dict[str, Any]]) -> bool: def summarize_tool_memory(tool_names: str) -> Optional[Dict[str, Any]]: """总结工具使用模式 - + curl example: curl -X POST http://0.0.0.0:8002/summary_tool_memory \ -H "Content-Type: application/json" \ @@ -86,11 +90,14 @@ def summarize_tool_memory(tool_names: str) -> Optional[Dict[str, Any]]: }' """ print(f"\n[SUMMARIZE] {tool_names}") - result = api_call("summary_tool_memory", { - "workspace_id": WORKSPACE_ID, - "tool_names": tool_names - }) - + result = api_call( + "summary_tool_memory", + { + "workspace_id": WORKSPACE_ID, + "tool_names": tool_names, + }, + ) + if result: memory_list = result.get("metadata", {}).get("memory_list", []) print(f"✓ Summarized {len(memory_list)} tool memories") @@ -98,13 +105,13 @@ def summarize_tool_memory(tool_names: str) -> Optional[Dict[str, Any]]: print(f"\n{'=' * 60}") print(f"Tool: {memory.get('when_to_use', 'N/A')}") print(f"{'=' * 60}") - print(memory.get('content', 'No content')) + print(memory.get("content", "No content")) return result def retrieve_tool_memory(tool_names: str, save_to_file: bool = False) -> str: """检索工具记忆 - + curl example: curl -X POST http://0.0.0.0:8002/retrieve_tool_memory \ -H "Content-Type: application/json" \ @@ -114,40 +121,45 @@ def retrieve_tool_memory(tool_names: str, save_to_file: bool = False) -> str: }' """ print(f"\n[RETRIEVE] {tool_names}") - result = api_call("retrieve_tool_memory", { - "workspace_id": WORKSPACE_ID, - "tool_names": tool_names - }) - + result = api_call( + "retrieve_tool_memory", + { + "workspace_id": WORKSPACE_ID, + "tool_names": tool_names, + }, + ) + if not result: return "" - + memory_list = result.get("metadata", {}).get("memory_list", []) if not memory_list: print("No memories found") return "" - + print(f"✓ Retrieved {len(memory_list)} memories") - + formatted_memories = [] for memory in memory_list: - content = f"\nTool: {memory.get('when_to_use', 'N/A')}\n" \ - f"Calls: {len(memory.get('tool_call_results', []))}\n" \ - f"{'-' * 60}\n{memory.get('content', 'No content')}\n" + content = ( + f"\nTool: {memory.get('when_to_use', 'N/A')}\n" + f"Calls: {len(memory.get('tool_call_results', []))}\n" + f"{'-' * 60}\n{memory.get('content', 'No content')}\n" + ) formatted_memories.append(content) print(content) - + if save_to_file: with open("tool_memory.json", "w") as f: json.dump(memory_list, f, indent=2, ensure_ascii=False) print("✓ Saved to tool_memory.json") - + return "\n".join(formatted_memories) def dump_memory(path: str = "./") -> None: """导出记忆到磁盘 - + curl example: curl -X POST http://0.0.0.0:8002/vector_store \ -H "Content-Type: application/json" \ @@ -157,18 +169,21 @@ def dump_memory(path: str = "./") -> None: "path": "./" }' """ - result = api_call("vector_store", { - "workspace_id": WORKSPACE_ID, - "action": "dump", - "path": path - }) + result = api_call( + "vector_store", + { + "workspace_id": WORKSPACE_ID, + "action": "dump", + "path": path, + }, + ) if result: print(f"✓ Memory dumped to {path}") def load_memory(path: str = "./") -> None: """从磁盘加载记忆 - + curl example: curl -X POST http://0.0.0.0:8002/vector_store \ -H "Content-Type: application/json" \ @@ -178,11 +193,14 @@ def load_memory(path: str = "./") -> None: "path": "./" }' """ - result = api_call("vector_store", { - "workspace_id": WORKSPACE_ID, - "action": "load", - "path": path - }) + result = api_call( + "vector_store", + { + "workspace_id": WORKSPACE_ID, + "action": "load", + "path": path, + }, + ) if result: print(f"✓ Memory loaded from {path}") @@ -192,43 +210,43 @@ def main() -> None: print("\n[1] Cleaning workspace...") delete_workspace() time.sleep(1) - + # 2. 创建和添加模拟工具调用结果 print("\n[2] Adding mock tool call results...") tools_to_test = [ ("web_search", 30), ("database_query", 22), - ("file_processor", 18) + ("file_processor", 18), ] - + # 收集所有工具的结果,然后一次性添加 all_mock_results = [] for tool_name, count in tools_to_test: mock_results = create_mock_tool_call_results(tool_name, count) all_mock_results.extend(mock_results) - + if not add_tool_call_results(all_mock_results): print("✗ Failed to add results") else: time.sleep(1) - + # 3. 总结工具记忆 print("\n[3] Summarizing tool memories...") all_tool_names = ",".join([tool[0] for tool in tools_to_test]) summarize_tool_memory(all_tool_names) time.sleep(1) - + # 4. 检索工具记忆 print("\n[4] Retrieving tool memories...") for tool_name, _ in tools_to_test: retrieve_tool_memory(tool_name, save_to_file=True) time.sleep(0.5) - + # 5. 测试记忆持久化 print("\n[5] Testing memory persistence...") dump_memory() load_memory() - + print("\n" + "=" * 60) print("DEMO COMPLETE ✓") print("=" * 60) diff --git a/cookbook/tool_memory/run_reme_tool_bench.py b/cookbook/tool_memory/run_reme_tool_bench.py index e9cd33e6..c95e2bdd 100644 --- a/cookbook/tool_memory/run_reme_tool_bench.py +++ b/cookbook/tool_memory/run_reme_tool_bench.py @@ -38,8 +38,8 @@ class BenchmarkStats: def add_result(self, result: Dict[str, Any]): """添加一个工具调用结果 - - Note: + + Note: - score: Quality/relevance of the result (0.0 or 1.0) """ self.total_count += 1 @@ -54,13 +54,13 @@ class BenchmarkStats: return { "name": self.name, "total_calls": 0, - "avg_score": 0.0 + "avg_score": 0.0, } return { "name": self.name, "total_calls": self.total_count, - "avg_score": round(sum(self.scores) / len(self.scores), 3) + "avg_score": round(sum(self.scores) / len(self.scores), 3), } @@ -87,18 +87,18 @@ def delete_workspace(workspace_id: str) -> bool: def load_queries(query_file: str = "query.json") -> Dict[str, Any]: """加载查询数据""" query_path = Path(__file__).parent / query_file - with open(query_path, 'r', encoding='utf-8') as f: + with open(query_path, "r", encoding="utf-8") as f: return json.load(f) def run_use_mock_search(workspace_id: str, queries: List[str], prompt_template: str = "") -> List[ToolCallResult]: """运行use_mock_search并收集结果(支持并发) - + Args: workspace_id: 工作空间ID queries: 查询列表 prompt_template: 提示模板 - + Returns: 工具调用结果列表 """ @@ -112,10 +112,13 @@ def run_use_mock_search(workspace_id: str, queries: List[str], prompt_template: # 提交之前sleep 1秒 time.sleep(1) - result = api_call("use_mock_search", { - "workspace_id": workspace_id, - "query": prompt_template.format(query=query), - }) + result = api_call( + "use_mock_search", + { + "workspace_id": workspace_id, + "query": prompt_template.format(query=query), + }, + ) if result: tool_call_result = result.get("answer") @@ -128,8 +131,7 @@ def run_use_mock_search(workspace_id: str, queries: List[str], prompt_template: with ThreadPoolExecutor(max_workers=4) as executor: # 提交所有任务 future_to_query = { - executor.submit(process_single_query, idx, query): (idx, query) - for idx, query in enumerate(queries) + executor.submit(process_single_query, idx, query): (idx, query) for idx, query in enumerate(queries) } # 收集结果(按完成顺序) @@ -153,11 +155,11 @@ def run_use_mock_search(workspace_id: str, queries: List[str], prompt_template: def add_tool_call_results(workspace_id: str, results: List[ToolCallResult]) -> List[ToolCallResult]: """批量添加工具调用结果到记忆库,并返回带评分的结果 - + Args: workspace_id: 工作空间ID results: 工具调用结果列表 - + Returns: 从API返回的memory_list中提取的带评分的ToolCallResult列表 """ @@ -171,10 +173,13 @@ def add_tool_call_results(workspace_id: str, results: List[ToolCallResult]) -> L logger.info(f"Adding tool call results to {workspace_id}: {len(tool_call_results)} results") # 统一调用API,让后端自动按tool_name分组处理 - api_result = api_call("add_tool_call_result", { - "workspace_id": workspace_id, - "tool_call_results": tool_call_results - }) + api_result = api_call( + "add_tool_call_result", + { + "workspace_id": workspace_id, + "tool_call_results": tool_call_results, + }, + ) if not api_result: logger.error("Failed to add results") @@ -206,10 +211,13 @@ def add_tool_call_results(workspace_id: str, results: List[ToolCallResult]) -> L def summarize_tool_memory(workspace_id: str, tool_names: str) -> bool: """总结工具记忆""" logger.info(f"Summarizing tool memory for {workspace_id}: {tool_names}") - result = api_call("summary_tool_memory", { - "workspace_id": workspace_id, - "tool_names": tool_names - }) + result = api_call( + "summary_tool_memory", + { + "workspace_id": workspace_id, + "tool_names": tool_names, + }, + ) if result: memory_list = result.get("metadata", {}).get("memory_list", []) @@ -221,19 +229,22 @@ def summarize_tool_memory(workspace_id: str, tool_names: str) -> bool: def retrieve_tool_memory(workspace_id: str, tool_names: str) -> str: """检索工具记忆并返回格式化的内容 - + Args: workspace_id: 工作空间ID tool_names: 逗号分隔的工具名称 - + Returns: 格式化的工具记忆内容,每个工具名称作为一级markdown标题 """ logger.info(f"Retrieving tool memory for {workspace_id}: {tool_names}") - result = api_call("retrieve_tool_memory", { - "workspace_id": workspace_id, - "tool_names": tool_names - }) + result = api_call( + "retrieve_tool_memory", + { + "workspace_id": workspace_id, + "tool_names": tool_names, + }, + ) if not result: logger.error("Failed to retrieve tool memory") @@ -251,8 +262,9 @@ def retrieve_tool_memory(workspace_id: str, tool_names: str) -> str: tool_name = tool_memory.when_to_use or "Unknown Tool" formatted_section = f"# {tool_name}\n\n{tool_memory.content}" formatted_contents.append(formatted_section) - logger.info(f"Retrieved content for tool: {tool_name}, " - f"content_length={len(tool_memory.content)}") + logger.info( + f"Retrieved content for tool: {tool_name}, " f"content_length={len(tool_memory.content)}", + ) # 用两个换行符分隔不同工具的记忆 joined_content = "\n\n".join(formatted_contents) @@ -265,7 +277,7 @@ def collect_statistics(results: List[ToolCallResult], stats: BenchmarkStats) -> """从结果列表中收集统计数据""" for result in results: # 转换为字典用于统计 - result_dict = result.model_dump() if hasattr(result, 'model_dump') else result + result_dict = result.model_dump() if hasattr(result, "model_dump") else result stats.add_result(result_dict) @@ -276,11 +288,13 @@ def print_comparison_table(stats_list: List[BenchmarkStats]) -> None: for stats in stats_list: summary = stats.get_summary() - rows.append([ - summary["name"], - summary["total_calls"], - summary["avg_score"] - ]) + rows.append( + [ + summary["name"], + summary["total_calls"], + summary["avg_score"], + ], + ) print("\n" + "=" * 100) print("BENCHMARK RESULTS COMPARISON") @@ -299,8 +313,7 @@ def calculate_improvements(baseline_stats: BenchmarkStats, improved_stats: Bench # 平均分数改进(相对提升百分比) if baseline["avg_score"] > 0: - improvements["avg_score"] = ((improved["avg_score"] - baseline["avg_score"]) - / baseline["avg_score"] * 100) + improvements["avg_score"] = (improved["avg_score"] - baseline["avg_score"]) / baseline["avg_score"] * 100 else: improvements["avg_score"] = 0.0 @@ -314,7 +327,7 @@ def print_improvements(improvements: Dict[str, float]) -> None: print("=" * 100) metric_labels = { - "avg_score": "Average Score" + "avg_score": "Average Score", } for metric, improvement in improvements.items(): @@ -328,19 +341,19 @@ def print_improvements(improvements: Dict[str, float]) -> None: def save_results(results: Dict[str, Any], filename: str = "benchmark_results.json") -> None: """保存结果到文件""" output_path = Path(__file__).parent / filename - with open(output_path, 'w', encoding='utf-8') as f: + with open(output_path, "w", encoding="utf-8") as f: json.dump(results, f, indent=2, ensure_ascii=False) logger.info(f"Results saved to {output_path}") def run_single_epoch(epoch_num: int, train_queries: List[str], test_queries: List[str]) -> Dict[str, Any]: """运行单个epoch的benchmark - + Args: epoch_num: epoch编号(从1开始) train_queries: 训练查询列表 test_queries: 测试查询列表 - + Returns: 包含该epoch统计结果的字典 """ @@ -422,7 +435,7 @@ def run_single_epoch(epoch_num: int, train_queries: List[str], test_queries: Lis tool_names_set = set() results_to_use = train_scored_results if train_scored_results else train_results_no_memory for result in results_to_use: - tool_name = result.tool_name if hasattr(result, 'tool_name') else None + tool_name = result.tool_name if hasattr(result, "tool_name") else None if tool_name: tool_names_set.add(tool_name) @@ -507,15 +520,15 @@ def run_single_epoch(epoch_num: int, train_queries: List[str], test_queries: Lis "statistics": { "train_no_memory": train_no_memory_stats.get_summary(), "test_no_memory": test_no_memory_stats.get_summary(), - "test_with_memory": test_with_memory_stats.get_summary() + "test_with_memory": test_with_memory_stats.get_summary(), }, - "improvements": improvements + "improvements": improvements, } def main(test_mode: bool = False, run_epoch: int = 3): """主函数:运行完整的benchmark流程 - + Args: test_mode: 如果为True,只使用每个难度级别的前3个查询进行快速测试 run_epoch: 运行的epoch数量,默认为3 @@ -567,11 +580,14 @@ def main(test_mode: bool = False, run_epoch: int = 3): # 计算每个场景的平均分数 avg_train_no_memory = sum(e["statistics"]["train_no_memory"]["avg_score"] for e in all_epoch_results) / len( - all_epoch_results) + all_epoch_results, + ) avg_test_no_memory = sum(e["statistics"]["test_no_memory"]["avg_score"] for e in all_epoch_results) / len( - all_epoch_results) + all_epoch_results, + ) avg_test_with_memory = sum(e["statistics"]["test_with_memory"]["avg_score"] for e in all_epoch_results) / len( - all_epoch_results) + all_epoch_results, + ) # 计算平均改进 avg_improvement = sum(e["improvements"]["avg_score"] for e in all_epoch_results) / len(all_epoch_results) @@ -581,7 +597,7 @@ def main(test_mode: bool = False, run_epoch: int = 3): rows = [ ["Train (No Memory)", f"{avg_train_no_memory:.3f}"], ["Test (No Memory)", f"{avg_test_no_memory:.3f}"], - ["Test (With Memory)", f"{avg_test_with_memory:.3f}"] + ["Test (With Memory)", f"{avg_test_with_memory:.3f}"], ] print(tabulate(rows, headers=headers, tablefmt="grid")) @@ -597,13 +613,15 @@ def main(test_mode: bool = False, run_epoch: int = 3): headers = ["Epoch", "Train (No Mem)", "Test (No Mem)", "Test (With Mem)", "Improvement %"] rows = [] for e in all_epoch_results: - rows.append([ - f"Epoch {e['epoch']}", - f"{e['statistics']['train_no_memory']['avg_score']:.3f}", - f"{e['statistics']['test_no_memory']['avg_score']:.3f}", - f"{e['statistics']['test_with_memory']['avg_score']:.3f}", - f"{e['improvements']['avg_score']:+.2f}%" - ]) + rows.append( + [ + f"Epoch {e['epoch']}", + f"{e['statistics']['train_no_memory']['avg_score']:.3f}", + f"{e['statistics']['test_no_memory']['avg_score']:.3f}", + f"{e['statistics']['test_with_memory']['avg_score']:.3f}", + f"{e['improvements']['avg_score']:+.2f}%", + ], + ) print(tabulate(rows, headers=headers, tablefmt="grid")) @@ -616,9 +634,9 @@ def main(test_mode: bool = False, run_epoch: int = 3): "train_no_memory": avg_train_no_memory, "test_no_memory": avg_test_no_memory, "test_with_memory": avg_test_with_memory, - "improvement": avg_improvement + "improvement": avg_improvement, }, - "per_epoch_results": all_epoch_results + "per_epoch_results": all_epoch_results, } save_results(benchmark_summary, "tool_memory_benchmark_results.json") diff --git a/docs/contribution.md b/docs/contribution.md index fd1328b2..079dca68 100644 --- a/docs/contribution.md +++ b/docs/contribution.md @@ -22,6 +22,36 @@ git checkout -b your-feature-branch-name ### Making Changes With your new branch checked out, you can now make your changes to the code. Remember to keep your changes as focused as possible. If you're addressing multiple issues or features, it's better to create separate branches and pull requests for each. +### Set Up Pre-commit Hooks +Before committing your changes, you should set up pre-commit hooks to ensure code quality and consistency. Pre-commit hooks will automatically check your code for common issues and format it according to the project's standards. + +**Install pre-commit:** +```bash +pip install pre-commit +``` + +**Install the git hooks:** +```bash +pre-commit install +``` + +**Run pre-commit manually (optional):** +If you want to run pre-commit checks on all files before committing, you can run: +```bash +pre-commit run --all-files +``` + +After installation, pre-commit will automatically run on `git commit` to check your code. The hooks will check for: +- Code syntax and AST validation +- YAML, XML, TOML, and JSON format validation +- Trailing whitespace +- Code formatting (Black) +- Code style (Flake8) +- Code quality (Pylint) +- Package metadata (Pyroma) + +If any checks fail, please fix the issues before committing. + ### Commit Your Changes Once you've made your changes, it's time to commit them. Write clear and concise commit messages that explain your changes. ```bash diff --git a/docs/cookbook/appworld/quickstart.md b/docs/cookbook/appworld/quickstart.md index f348587e..9fc1f0bf 100644 --- a/docs/cookbook/appworld/quickstart.md +++ b/docs/cookbook/appworld/quickstart.md @@ -1,4 +1,4 @@ -# AppWorld +# AppWorld Experiment Quick Start Guide This guide helps you quickly set up and run AppWorld experiments with ReMe integration. diff --git a/docs/cookbook/bfcl/quickstart.md b/docs/cookbook/bfcl/quickstart.md index 2c9b81a1..93e4de06 100644 --- a/docs/cookbook/bfcl/quickstart.md +++ b/docs/cookbook/bfcl/quickstart.md @@ -1,4 +1,4 @@ -# BFCL +# BFCL Experiment Quick Start Guide This guide helps you quickly set up and run BFCL experiments with ReMe integration. @@ -40,7 +40,7 @@ Run the main experiment script to collect agent trajectories on training data se python run_bfcl.py ``` -**Note**: +**Note**: - `max_workers`: Number of parallel workers (default: `4`) - `num_runs`: Number of times each task is repeated (default: `1`) - `model_name`: LLM model name (default: `qwen3-8b`) diff --git a/docs/cookbook/experiment_overview.md b/docs/cookbook/experiment_overview.md index 4bb1521d..a95b06e4 100644 --- a/docs/cookbook/experiment_overview.md +++ b/docs/cookbook/experiment_overview.md @@ -10,7 +10,7 @@ We tested ReMe on Appworld using qwen3-8b: | with ReMe | 0.109 **(+2.6%)** | 0.175 **(+3.5%)** | 0.281 **(+5.3%)** | Pass@K measures the probability that at least one of the K generated samples successfully completes the task ( -score=1). +score=1). The current experiment uses an internal AppWorld environment, which may have slight differences. You can find more details on reproducing the experiment in [quickstart.md](appworld/quickstart.md). @@ -43,10 +43,10 @@ We tested ReMe on BFCL-V3 multi-turn-base (randomly split 50train/150val) using We evaluated Tool Memory effectiveness using a controlled benchmark with three mock search tools using Qwen3-30B-Instruct: -| Scenario | Avg Score | Improvement | -|-----------------------|-----------|--------------------| -| Train (No Memory) | 0.650 | - | -| Test (No Memory) | 0.672 | Baseline | +| Scenario | Avg Score | Improvement | +|------------------------|-----------|-------------| +| Train (No Memory) | 0.650 | - | +| Test (No Memory) | 0.672 | Baseline | | **Test (With Memory)** | **0.772** | **+14.88%** | **Key Findings:** diff --git a/docs/cookbook/frozenlake/quickstart.md b/docs/cookbook/frozenlake/quickstart.md index 81a35a7b..f6db28b0 100644 --- a/docs/cookbook/frozenlake/quickstart.md +++ b/docs/cookbook/frozenlake/quickstart.md @@ -12,7 +12,7 @@ kernelspec: name: python3 --- -# FrozenLake +# FrozenLake Experiment Quick Start Guide This guide helps you quickly set up and run FrozenLake experiments with ReMe integration. The FrozenLake experiment demonstrates how task memory can improve an agent's performance in a navigation task. diff --git a/docs/index.md b/docs/index.md index f4990b2d..3ee78c4a 100644 --- a/docs/index.md +++ b/docs/index.md @@ -16,8 +16,8 @@ kernelspec: Remember Me, Refine Me.
- Python Version - PyPI Version + Python Version + PyPI Version License GitHub Stars
diff --git a/docs/mcp_quick_start.md b/docs/mcp_quick_start.md index 1fb1dc6c..16591fd4 100644 --- a/docs/mcp_quick_start.md +++ b/docs/mcp_quick_start.md @@ -184,11 +184,11 @@ The `summary_task_memory` tool transforms conversation trajectories into valuabl async def run_summary(client, messages): """ Generate a summary of conversation messages and create task memories - + Args: client: MCP client instance messages: List of message objects from a conversation - + Returns: None """ @@ -227,11 +227,11 @@ The `retrieve_task_memory` tool allows you to retrieve relevant memories based o async def run_retrieve(client, query): """ Retrieve relevant task memories based on a query - + Args: client: MCP client instance query: The query to retrieve relevant memories - + Returns: String containing the retrieved memory answer """ diff --git a/docs/quick_start.md b/docs/quick_start.md index cbe863d6..44e06fd7 100644 --- a/docs/quick_start.md +++ b/docs/quick_start.md @@ -88,7 +88,7 @@ async def main(): ] ) print(result) - + # Retriever: Get relevant memories result = await app.async_execute( name="retrieve_task_memory", @@ -219,7 +219,7 @@ async def main(): ] ) print(result) - + # Memory Retrieval: Get personal memory fragments result = await app.async_execute( name="retrieve_personal_memory", @@ -366,7 +366,7 @@ async def main(): ] ) print(result) - + # Generate usage guidelines from history result = await app.async_execute( name="summary_tool_memory", @@ -374,7 +374,7 @@ async def main(): tool_names="web_search" ) print(result) - + # Retrieve tool guidelines before use result = await app.async_execute( name="retrieve_tool_memory", diff --git a/docs/task_memory/task_memory.md b/docs/task_memory/task_memory.md index c8cd4968..1c96946a 100644 --- a/docs/task_memory/task_memory.md +++ b/docs/task_memory/task_memory.md @@ -139,7 +139,7 @@ Here's a complete example workflow that demonstrates how to use task memory: def run_agent_with_memory(query_first, query_second): # Run agent with second query to build initial memories messages = run_agent(query=query_second) - + # Summarize conversation to create memories requests.post( url=f"{BASE_URL}summary_task_memory", @@ -150,7 +150,7 @@ def run_agent_with_memory(query_first, query_second): ] } ) - + # Retrieve relevant memories for the first query response = requests.post( url=f"{BASE_URL}retrieve_task_memory", @@ -160,7 +160,7 @@ def run_agent_with_memory(query_first, query_second): } ) retrieved_memory = response.json().get("answer", "") - + # Run agent with first query augmented with retrieved memories augmented_query = f"{retrieved_memory}\n\nUser Question:\n{query_first}" return run_agent(query=augmented_query) diff --git a/docs/tool_memory/tool_bench.md b/docs/tool_memory/tool_bench.md index a1b35939..7f0fec94 100644 --- a/docs/tool_memory/tool_bench.md +++ b/docs/tool_memory/tool_bench.md @@ -24,11 +24,11 @@ This benchmark evaluates Tool Memory effectiveness by comparing agent performanc Three LLM-based mock search tools with different performance profiles: -| Tool | Simple Queries | Medium Queries | Complex Queries | -|------|---------------|----------------|-----------------| -| **SearchToolA** | ⭐⭐⭐ Fast, high success (90%) | ❌ Poor (20% success) | ⚠️ Weak (50% success) | -| **SearchToolB** | ⚠️ Over-engineered (30%) | ⭐⭐⭐ Optimal (90% success) | ⚠️ Limited (50% success) | -| **SearchToolC** | ⚠️ Overkill (30%) | ⚠️ Excessive (40%) | ⭐⭐⭐ Best (90% success) | +| Tool | Simple Queries | Medium Queries | Complex Queries | +|-----------------|------------------------------|---------------------------|--------------------------| +| **SearchToolA** | ⭐⭐⭐ Fast, high success (90%) | ❌ Poor (20% success) | ⚠️ Weak (50% success) | +| **SearchToolB** | ⚠️ Over-engineered (30%) | ⭐⭐⭐ Optimal (90% success) | ⚠️ Limited (50% success) | +| **SearchToolC** | ⚠️ Overkill (30%) | ⚠️ Excessive (40%) | ⭐⭐⭐ Best (90% success) | **Performance Characteristics:** - `success_rate`: Probability of successful execution (vs "Service busy" error) diff --git a/docs/tool_memory/tool_memory.md b/docs/tool_memory/tool_memory.md index 077ac859..26a0c047 100644 --- a/docs/tool_memory/tool_memory.md +++ b/docs/tool_memory/tool_memory.md @@ -34,7 +34,7 @@ When an LLM faces numerous MCP tools, it relies heavily on tool descriptions to Imagine an LLM choosing between three search tools: ``` Tool A: "Search the web for information" -Tool B: "Perform web searches with customizable parameters" +Tool B: "Perform web searches with customizable parameters" Tool C: "Query search engines and return results" ``` @@ -108,7 +108,7 @@ LLM: "I have 50 search tools, but Tool A has 95% success for technical queries" ``` Before Tool Memory: - Success rate: 75% -- Average time cost: 5.2s +- Average time cost: 5.2s - Token cost: 200+ per call - Repeated parameter errors - Random tool selection @@ -636,7 +636,7 @@ op: params: max_history_tool_call_cnt: 200 # Keep more history evaluation_sleep_interval: 0.5 # Faster evaluation - + summary_tool_memory_op: params: recent_call_count: 50 # Analyze more calls @@ -649,7 +649,7 @@ op: params: max_history_tool_call_cnt: 50 # Less history needed evaluation_sleep_interval: 1.0 # Standard rate - + summary_tool_memory_op: params: recent_call_count: 20 # Analyze fewer calls @@ -661,10 +661,10 @@ op: ```{code-cell} memory = retrieve_tool_memory("web_search")['metadata']['memory_list'][0] stats = ToolMemory(**memory).statistic(recent_frequency=30) - + print(f"Success Rate: {stats['success_rate']:.2%}") print(f"Avg Score: {stats['avg_score']:.2f}") - + if stats['success_rate'] < 0.7: print("⚠️ Low success rate - investigate tool issues") ``` diff --git a/docs/vector_store_api_guide.md b/docs/vector_store_api_guide.md index 394175a5..e522aaf6 100644 --- a/docs/vector_store_api_guide.md +++ b/docs/vector_store_api_guide.md @@ -12,9 +12,9 @@ kernelspec: name: python3 --- -# Vector Store API Guide +# Vector Store Configuration Guide -This guide covers the vector store implementations available in ReMe, their APIs, and how to use them effectively. +This guide covers how to configure vector store backends in ReMe using the `default.yaml` configuration file. ## 📋 Overview @@ -43,71 +43,28 @@ All vector stores implement the `BaseVectorStore` interface, providing a consist | **Async Support** | ❌ No | ❌ No | ❌ No | ✅ Native | ❌ No | | **Best For** | Development | Local Apps | Production | Production/Cloud | Testing | -## 🔄 Common API Methods +## ⚙️ Configuration in default.yaml -All vector store implementations share these core methods: +All vector stores are configured in the `vector_store` section of `reme_ai/config/default.yaml`. The configuration structure is: -### 🔄 Async Support - -All vector stores provide both synchronous and asynchronous versions of every method: - -```python -# Synchronous methods -store.search(query="example", workspace_id="workspace", top_k=5) -store.insert(nodes, workspace_id="workspace") - -# Asynchronous methods (with async_ prefix) -await store.async_search(query="example", workspace_id="workspace", top_k=5) -await store.async_insert(nodes, workspace_id="workspace") +```yaml +vector_store: + default: + backend: # Required: local, chroma, elasticsearch, qdrant, or memory + embedding_model: default # Required: Name of the embedding model configuration + params: # Optional: Backend-specific parameters + # Backend-specific parameters go here ``` -### Workspace Management +### Configuration Fields -```python -# Check if workspace exists -store.exist_workspace(workspace_id: str) -> bool +- **`backend`** (required): The vector store backend to use. Valid values: `local`, `chroma`, `elasticsearch`, `qdrant`, `memory` +- **`embedding_model`** (required): The name of the embedding model configuration from the `embedding_model` section +- **`params`** (optional): A dictionary of backend-specific parameters that will be passed to the vector store constructor -# Create a new workspace -store.create_workspace(workspace_id: str, **kwargs) +## 📁 Vector Store Backend Configurations -# Delete a workspace -store.delete_workspace(workspace_id: str, **kwargs) - -# Copy a workspace -store.copy_workspace(src_workspace_id: str, dest_workspace_id: str, **kwargs) -``` - -### Data Operations - -```python -# Insert nodes (single or list) -store.insert(nodes: VectorNode | List[VectorNode], workspace_id: str, **kwargs) - -# Delete nodes by ID -store.delete(node_ids: str | List[str], workspace_id: str, **kwargs) - -# Search for similar nodes -store.search(query: str, workspace_id: str, top_k: int = 1, **kwargs) -> List[VectorNode] - -# Iterate through workspace nodes -for node in store.iter_workspace_nodes(workspace_id: str, **kwargs): - # Process each node -``` - -### Import/Export - -```python -# Export workspace to file -store.dump_workspace(workspace_id: str, path: str | Path = "", callback_fn=None, **kwargs) - -# Import workspace from file -store.load_workspace(workspace_id: str, path: str | Path = "", nodes: List[VectorNode] = None, - callback_fn=None, **kwargs) -``` - -## ⚡ Vector Store Implementations - -### 1. 📁 LocalVectorStore (`backend=local`) +### 1. LocalVectorStore (`backend=local`) A simple file-based vector store that saves data to local JSONL files. @@ -118,62 +75,22 @@ A simple file-based vector store that saves data to local JSONL files. #### ⚙️ Configuration -```python -from flowllm.storage.vector_store import LocalVectorStore -from flowllm.embedding_model import OpenAICompatibleEmbeddingModel -from flowllm.utils.common_utils import load_env - -# Load environment variables (for API keys) -load_env() - -# Initialize embedding model -embedding_model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4") - -# Initialize vector store -vector_store = LocalVectorStore( - embedding_model=embedding_model, - store_dir="./file_vector_store", # Directory to store JSONL files - batch_size=1024 # Batch size for operations -) +```yaml +vector_store: + default: + backend: local + embedding_model: default + params: + store_dir: "./local_vector_store" # Directory to store JSONL files (default: "./local_vector_store") + batch_size: 1024 # Batch size for operations (default: 1024) ``` -#### 💻 Example Usage +#### Configuration Parameters -```python -from flowllm.schema.vector_node import VectorNode +- **`store_dir`** (optional): Directory path where workspace files are stored. Default: `"./local_vector_store"` +- **`batch_size`** (optional): Batch size for bulk operations. Default: `1024` -# Create workspace -workspace_id = "my_workspace" -vector_store.create_workspace(workspace_id) - -# Create nodes -nodes = [ - VectorNode( - unique_id="node1", - workspace_id=workspace_id, - content="Artificial intelligence is revolutionizing technology", - metadata={"category": "tech", "source": "article1"} - ), - VectorNode( - unique_id="node2", - workspace_id=workspace_id, - content="Machine learning enables data-driven insights", - metadata={"category": "tech", "source": "article2"} - ) -] - -# Insert nodes -vector_store.insert(nodes, workspace_id) - -# Search -results = vector_store.search("What is AI?", workspace_id, top_k=2) -for result in results: - print(f"Content: {result.content}") - print(f"Metadata: {result.metadata}") - print(f"Score: {result.metadata.get('score', 'N/A')}") -``` - -### 2. 🔮 ChromaVectorStore (`backend=chroma`) +### 2. ChromaVectorStore (`backend=chroma`) An embedded vector database that provides persistent storage with advanced features. @@ -184,71 +101,22 @@ An embedded vector database that provides persistent storage with advanced featu #### ⚙️ Configuration -```python -from flowllm.storage.vector_store import ChromaVectorStore -from flowllm.embedding_model import OpenAICompatibleEmbeddingModel -from flowllm.utils.common_utils import load_env - -# Load environment variables -load_env() - -# Initialize embedding model -embedding_model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4") - -# Initialize vector store -vector_store = ChromaVectorStore( - embedding_model=embedding_model, - store_dir="./chroma_vector_store", # Directory for Chroma database - batch_size=1024 # Batch size for operations -) +```yaml +vector_store: + default: + backend: chroma + embedding_model: default + params: + store_dir: "./chroma_vector_store" # Directory for Chroma database (default: "./chroma_vector_store") + batch_size: 1024 # Batch size for operations (default: 1024) ``` -#### 💻 Example Usage +#### Configuration Parameters -```python -from flowllm.schema.vector_node import VectorNode +- **`store_dir`** (optional): Directory path where ChromaDB data is persisted. Default: `"./chroma_vector_store"` +- **`batch_size`** (optional): Batch size for bulk operations. Default: `1024` -workspace_id = "chroma_workspace" - -# Check if workspace exists and create if needed -if not vector_store.exist_workspace(workspace_id): - vector_store.create_workspace(workspace_id) - -# Create nodes with metadata -nodes = [ - VectorNode( - unique_id="node1", - workspace_id=workspace_id, - content="Deep learning models require large datasets", - metadata={ - "category": "AI", - "difficulty": "advanced", - "topic": "deep_learning" - } - ), - VectorNode( - unique_id="node2", - workspace_id=workspace_id, - content="Transformer architecture revolutionized NLP", - metadata={ - "category": "AI", - "difficulty": "intermediate", - "topic": "transformers" - } - ) -] - -# Insert nodes -vector_store.insert(nodes, workspace_id) - -# Search -results = vector_store.search("deep learning", workspace_id, top_k=5) -for result in results: - print(f"Content: {result.content}") - print(f"Metadata: {result.metadata}") -``` - -### 3. 🔍 EsVectorStore (`backend=elasticsearch`) +### 3. EsVectorStore (`backend=elasticsearch`) Production-grade vector search using Elasticsearch with advanced filtering and scaling capabilities. @@ -282,114 +150,24 @@ export FLOW_ES_HOSTS=http://localhost:9200 #### ⚙️ Configuration -```python -from flowllm.storage.vector_store import EsVectorStore -from flowllm.embedding_model import OpenAICompatibleEmbeddingModel -from flowllm.utils.common_utils import load_env -import os - -# Load environment variables -load_env() - -# Initialize embedding model -embedding_model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4") - -# Initialize vector store -vector_store = EsVectorStore( - embedding_model=embedding_model, - hosts=os.getenv("FLOW_ES_HOSTS", "http://localhost:9200"), # Elasticsearch hosts - basic_auth=None, # ("username", "password") for auth - batch_size=1024 # Batch size for bulk operations -) +```yaml +vector_store: + default: + backend: elasticsearch + embedding_model: default + params: + hosts: "http://localhost:9200" # Elasticsearch host(s) - can be string or list (default: from FLOW_ES_HOSTS env var or "http://localhost:9200") + basic_auth: null # Optional: ("username", "password") tuple for authentication + batch_size: 1024 # Batch size for bulk operations (default: 1024) ``` -#### 🎯 Advanced Filtering +#### Configuration Parameters -EsVectorStore supports advanced filtering capabilities through the `filter_dict` parameter: +- **`hosts`** (optional): Elasticsearch host(s) as a string or list of strings. Defaults to the `FLOW_ES_HOSTS` environment variable or `"http://localhost:9200"` if not set +- **`basic_auth`** (optional): Tuple of `("username", "password")` for basic authentication. Default: `null` (no authentication) +- **`batch_size`** (optional): Batch size for bulk operations. Default: `1024` -```python -# Term filters (exact match) -term_filter = { - "category": "technology", - "author": "research_team" -} - -# Range filters (numeric and date ranges) -range_filter = { - "score": {"gte": 0.8}, # Score >= 0.8 - "confidence": {"gte": 0.5, "lte": 0.9}, # Between 0.5 and 0.9 - "timestamp": {"gte": "2024-01-01", "lte": "2024-12-31"} -} - -# Combined filters (filters are combined with AND logic) -combined_filter = { - "category": "AI", - "confidence": {"gte": 0.9} -} - -# Search with filters applied -results = vector_store.search("machine learning", workspace_id, top_k=10, filter_dict=combined_filter) -``` - -#### ⚡ Performance Optimization - -```python -# Refresh index for immediate availability (useful after bulk inserts) -vector_store.insert(nodes, workspace_id, refresh=True) # Auto-refresh -vector_store.refresh(workspace_id) # Manual refresh - -# Bulk operations with custom batch size -vector_store.insert(large_node_list, workspace_id, refresh=False) # Skip refresh for speed -vector_store.refresh(workspace_id) # Refresh once after all inserts -``` - -#### 💻 Example Usage - -```python -from flowllm.schema.vector_node import VectorNode - -# Define workspace -workspace_id = "production_workspace" - -# Create workspace if needed -if not vector_store.exist_workspace(workspace_id): - vector_store.create_workspace(workspace_id) - -# Create nodes with rich metadata -nodes = [ - VectorNode( - unique_id="doc1", - workspace_id=workspace_id, - content="Transformer architecture revolutionized NLP", - metadata={ - "category": "AI", - "subcategory": "NLP", - "author": "research_team", - "timestamp": "2024-01-15", - "confidence": 0.95, - "tags": ["transformer", "nlp", "attention"] - } - ) -] - -# Insert with refresh for immediate availability -vector_store.insert(nodes, workspace_id, refresh=True) - -# Advanced search with filters -filter_dict = { - "category": "AI", - "confidence": {"gte": 0.9} -} - -results = vector_store.search("transformer models", workspace_id, top_k=5, filter_dict=filter_dict) - -for result in results: - print(f"Score: {result.metadata.get('score', 'N/A')}") - print(f"Content: {result.content}") - print(f"Metadata: {result.metadata}") -``` - -### 4. 🎯 QdrantVectorStore (`backend=qdrant`) +### 4. QdrantVectorStore (`backend=qdrant`) A high-performance vector database designed for production workloads with native async support and advanced filtering. @@ -430,284 +208,40 @@ export FLOW_QDRANT_API_KEY=your-api-key #### ⚙️ Configuration -```python -from flowllm.storage.vector_store import QdrantVectorStore -from flowllm.embedding_model import OpenAICompatibleEmbeddingModel -from flowllm.utils.common_utils import load_env -import os - -# Load environment variables -load_env() - -# Initialize embedding model -embedding_model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4") - -# Option 1: Use localhost with environment variables -vector_store = QdrantVectorStore( - embedding_model=embedding_model, - host=os.getenv("FLOW_QDRANT_HOST", "localhost"), - port=int(os.getenv("FLOW_QDRANT_PORT", "6333")), - batch_size=1024 -) - -# Option 2: Use URL (for Qdrant Cloud or remote servers) -vector_store = QdrantVectorStore( - embedding_model=embedding_model, - url="http://your-qdrant-server:6333", - api_key="your-api-key", # Optional, for cloud - batch_size=1024 -) - -# Option 3: Specify custom distance metric -from qdrant_client.http.models import Distance - -vector_store = QdrantVectorStore( - embedding_model=embedding_model, - host="localhost", - port=6333, - distance=Distance.COSINE, # or Distance.EUCLIDEAN, Distance.DOT - batch_size=1024 -) +##### Local Qdrant Instance +```yaml +vector_store: + default: + backend: qdrant + embedding_model: default + params: + host: "localhost" # Qdrant host (default: from FLOW_QDRANT_HOST env var or "localhost") + port: 6333 # Qdrant port (default: from FLOW_QDRANT_PORT env var or 6333) + batch_size: 1024 # Batch size for operations (default: 1024) + distance: "COSINE" # Distance metric: "COSINE", "EUCLIDEAN", or "DOT" (default: "COSINE") ``` -#### 🎯 Advanced Filtering - -QdrantVectorStore supports advanced filtering capabilities similar to Elasticsearch: - -```python -# Term filters (exact match) -term_filter = { - "category": "AI", - "node_type": "research" -} - -# Range filters (numeric) -range_filter = { - "confidence": {"gte": 0.8, "lte": 1.0}, # Between 0.8 and 1.0 - "score": {"gt": 0.5} # Greater than 0.5 -} - -# Combined filters (all conditions must match - AND logic) -combined_filter = { - "category": "AI", - "confidence": {"gte": 0.9}, - "node_type": "research" -} - -# Search with filters -results = vector_store.search( - query="machine learning", - workspace_id=workspace_id, - top_k=10, - filter_dict=combined_filter -) +##### Qdrant Cloud or Remote Server +```yaml +vector_store: + default: + backend: qdrant + embedding_model: default + params: + url: "https://your-cluster.qdrant.io:6333" # Qdrant server URL (if provided, host and port are ignored) + api_key: "your-api-key" # API key for Qdrant Cloud authentication + batch_size: 1024 # Batch size for operations (default: 1024) + distance: "COSINE" # Distance metric (default: "COSINE") ``` -##### Filter Operations Supported: -- **Exact match**: `{"field": "value"}` -- **Range queries**: - - `gte`: Greater than or equal - - `lte`: Less than or equal - - `gt`: Greater than - - `lt`: Less than +#### Configuration Parameters -#### ⚡ Async Operations - -QdrantVectorStore provides **native async support** for all operations: - -```python -import asyncio - -async def main(): - # All operations have async equivalents - - # Check if workspace exists - exists = await vector_store.async_exist_workspace(workspace_id) - - # Create workspace - if not exists: - await vector_store.async_create_workspace(workspace_id) - - # Insert nodes with async embedding - await vector_store.async_insert(nodes, workspace_id) - - # Search with async embedding - results = await vector_store.async_search( - query="AI research", - workspace_id=workspace_id, - top_k=5, - filter_dict={"category": "AI"} - ) - - # Delete nodes - await vector_store.async_delete(node_ids, workspace_id) - - # Delete workspace - await vector_store.async_delete_workspace(workspace_id) - - # Close client - await vector_store.async_close() - -# Run async operations -asyncio.run(main()) -``` - -#### 💻 Example Usage - -```python -from flowllm.schema.vector_node import VectorNode - -workspace_id = "qdrant_workspace" - -# Check and create workspace -if not vector_store.exist_workspace(workspace_id): - vector_store.create_workspace(workspace_id) - -# Create nodes with rich metadata -nodes = [ - VectorNode( - unique_id="node1", - workspace_id=workspace_id, - content="Artificial intelligence is revolutionizing technology", - metadata={ - "category": "AI", - "node_type": "research", - "confidence": 0.95, - "author": "research_team" - } - ), - VectorNode( - unique_id="node2", - workspace_id=workspace_id, - content="Machine learning models require large datasets", - metadata={ - "category": "AI", - "node_type": "tutorial", - "confidence": 0.85, - "author": "education_team" - } - ), - VectorNode( - unique_id="node3", - workspace_id=workspace_id, - content="Deep learning excels at image recognition", - metadata={ - "category": "AI", - "node_type": "research", - "confidence": 0.92, - "author": "research_team" - } - ) -] - -# Insert nodes (upsert - creates or updates) -vector_store.insert(nodes, workspace_id) - -# Simple search -results = vector_store.search("What is AI?", workspace_id, top_k=3) -for result in results: - print(f"Content: {result.content}") - print(f"Score: {result.metadata.get('score', 'N/A')}") - print(f"Metadata: {result.metadata}") - print("-" * 50) - -# Advanced search with filters -filter_dict = { - "node_type": "research", - "confidence": {"gte": 0.9} -} - -filtered_results = vector_store.search( - query="AI technology", - workspace_id=workspace_id, - top_k=5, - filter_dict=filter_dict -) - -print(f"Found {len(filtered_results)} filtered results") -for result in filtered_results: - print(f"Content: {result.content}") - print(f"Metadata: {result.metadata}") - -# Iterate through all nodes -print("\nAll nodes in workspace:") -for node in vector_store.iter_workspace_nodes(workspace_id, limit=100): - print(f"ID: {node.unique_id}, Content: {node.content[:50]}...") - -# Update a node (delete + insert) -updated_node = VectorNode( - unique_id="node1", - workspace_id=workspace_id, - content="Artificial intelligence is transforming industries worldwide", - metadata={ - "category": "AI", - "node_type": "research", - "confidence": 0.98, - "author": "research_team", - "updated": True - } -) -vector_store.delete("node1", workspace_id) -vector_store.insert(updated_node, workspace_id) - -# Export workspace for backup -vector_store.dump_workspace(workspace_id, path="./qdrant_backup") - -# Clean up -vector_store.close() -``` - -#### 🔄 Async Example - -```python -import asyncio -from flowllm.schema.vector_node import VectorNode - -async def async_example(): - workspace_id = "async_qdrant_workspace" - - # Create workspace - if not await vector_store.async_exist_workspace(workspace_id): - await vector_store.async_create_workspace(workspace_id) - - # Create nodes - nodes = [ - VectorNode( - unique_id="async_node1", - workspace_id=workspace_id, - content="Async operations enable better performance", - metadata={"type": "performance", "async": True} - ), - VectorNode( - unique_id="async_node2", - workspace_id=workspace_id, - content="Concurrent requests improve throughput", - metadata={"type": "performance", "async": True} - ) - ] - - # Insert with async embedding - await vector_store.async_insert(nodes, workspace_id) - - # Search with async embedding - results = await vector_store.async_search( - query="performance optimization", - workspace_id=workspace_id, - top_k=2, - filter_dict={"async": True} - ) - - for result in results: - print(f"Score: {result.metadata['score']:.4f}") - print(f"Content: {result.content}") - - # Cleanup - await vector_store.async_delete_workspace(workspace_id) - await vector_store.async_close() - -# Run async example -asyncio.run(async_example()) -``` +- **`url`** (optional): Complete URL for connecting to Qdrant. If provided, `host` and `port` are ignored. Useful for Qdrant Cloud or custom deployments +- **`host`** (optional): Host address of the Qdrant server. Defaults to the `FLOW_QDRANT_HOST` environment variable or `"localhost"` if not set +- **`port`** (optional): Port number of the Qdrant server. Defaults to the `FLOW_QDRANT_PORT` environment variable or `6333` if not set +- **`api_key`** (optional): API key for authentication (required for Qdrant Cloud). Can also be set via `FLOW_QDRANT_API_KEY` environment variable +- **`distance`** (optional): Distance metric for vector similarity. Valid values: `"COSINE"`, `"EUCLIDEAN"`, `"DOT"`. Default: `"COSINE"` +- **`batch_size`** (optional): Batch size for bulk operations. Default: `1024` #### 🌟 Key Features @@ -720,23 +254,7 @@ asyncio.run(async_example()) - **Persistent Storage** - Data is automatically persisted to disk - **Efficient Iteration** - Scroll through large collections with pagination -#### 🚨 Important Notes - -- **Collection = Workspace** - Qdrant uses "collections" which map to workspace_id -- **Automatic Embedding** - Nodes without vectors are automatically embedded -- **ID-based Upsert** - Using the same unique_id will update existing nodes -- **Metadata Indexing** - All metadata fields are automatically indexed for filtering -- **Connection Management** - Call `close()` or `async_close()` to cleanup connections - -#### 📊 Performance Tips - -1. **Batch Operations** - Insert multiple nodes at once for better performance -2. **Use Async** - For high-concurrency scenarios, use async methods -3. **Optimize Filters** - Use indexed metadata fields for faster filtering -4. **Pagination** - Use `iter_workspace_nodes()` with appropriate `limit` for large collections -5. **Distance Metric** - Choose appropriate distance metric for your use case (COSINE for normalized vectors) - -### 5. ⚡ MemoryVectorStore (`backend=memory`) +### 5. MemoryVectorStore (`backend=memory`) An ultra-fast in-memory vector store that keeps all data in RAM for maximum performance. @@ -748,74 +266,20 @@ An ultra-fast in-memory vector store that keeps all data in RAM for maximum perf #### ⚙️ Configuration -```python -from flowllm.storage.vector_store import MemoryVectorStore -from flowllm.embedding_model import OpenAICompatibleEmbeddingModel -from flowllm.utils.common_utils import load_env - -# Load environment variables -load_env() - -# Initialize embedding model -embedding_model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4") - -# Initialize vector store -vector_store = MemoryVectorStore( - embedding_model=embedding_model, - store_dir="./memory_vector_store", # Directory for backup/restore operations - batch_size=1024 # Batch size for operations -) +```yaml +vector_store: + default: + backend: memory + embedding_model: default + params: + store_dir: "./memory_vector_store" # Directory for backup/restore operations (default: "./memory_vector_store") + batch_size: 1024 # Batch size for operations (default: 1024) ``` -#### 💻 Example Usage +#### Configuration Parameters -```python -from flowllm.schema.vector_node import VectorNode - -workspace_id = "memory_workspace" - -# Create workspace in memory -vector_store.create_workspace(workspace_id) - -# Create nodes -nodes = [ - VectorNode( - unique_id="mem_node1", - workspace_id=workspace_id, - content="Memory stores provide ultra-fast access to data", - metadata={ - "category": "performance", - "type": "memory", - "speed": "ultra_fast" - } - ), - VectorNode( - unique_id="mem_node2", - workspace_id=workspace_id, - content="In-memory databases excel at low-latency operations", - metadata={ - "category": "performance", - "type": "database", - "latency": "low" - } - ) -] - -# Insert nodes (stored in memory) -vector_store.insert(nodes, workspace_id) - -# Ultra-fast search -results = vector_store.search("fast memory access", workspace_id, top_k=2) -for result in results: - print(f"Content: {result.content}") - print(f"Score: {result.metadata.get('score', 'N/A')}") - -# Optional: Save to disk for backup -vector_store.dump_workspace(workspace_id, path="./backup") - -# Optional: Load from disk to memory -vector_store.load_workspace(workspace_id, path="./backup") -``` +- **`store_dir`** (optional): Directory path for backup/restore operations. Default: `"./memory_vector_store"` +- **`batch_size`** (optional): Batch size for bulk operations. Default: `1024` #### ⚡ Performance Benefits @@ -831,69 +295,80 @@ vector_store.load_workspace(workspace_id, path="./backup") - **No persistence** - Use `dump_workspace()` to save to disk - **Single process** - Not suitable for distributed applications -## 📝 Working with VectorNode +## 📝 Example Configurations -The `VectorNode` class is the fundamental data unit for all vector stores: - -```python -from flowllm.schema.vector_node import VectorNode - -# Create a node -node = VectorNode( - unique_id="unique_identifier", # Unique ID for the node (required) - workspace_id="my_workspace", # Workspace ID (required) - content="Text content to embed", # Content to be embedded (required) - metadata={ # Optional metadata - "source": "document1", - "category": "technology", - "timestamp": "2024-08-29" - }, - vector=None # Vector will be generated automatically if None -) +### Minimal Configuration (Memory Store) +```yaml +vector_store: + default: + backend: memory + embedding_model: default ``` -## 🔄 Import/Export Example - -Export and import workspaces for backup or transfer: - -```python -# Export workspace to file -vector_store.dump_workspace( - workspace_id="my_workspace", - path="./backup_data" # Directory to store the exported data -) - -# Import workspace from file -vector_store.load_workspace( - workspace_id="new_workspace", - path="./backup_data" # Directory containing the exported data -) - -# Copy workspace within the same store -vector_store.copy_workspace( - src_workspace_id="original_workspace", - dest_workspace_id="copied_workspace" -) +### Local File Storage +```yaml +vector_store: + default: + backend: local + embedding_model: default + params: + store_dir: "./my_vector_store" + batch_size: 2048 ``` +### Elasticsearch Production Setup +```yaml +vector_store: + default: + backend: elasticsearch + embedding_model: default + params: + hosts: "http://elasticsearch.example.com:9200" + basic_auth: ["username", "password"] + batch_size: 2048 +``` + +### Qdrant Cloud Setup +```yaml +vector_store: + default: + backend: qdrant + embedding_model: default + params: + url: "https://your-cluster.qdrant.io:6333" + api_key: "your-api-key-here" + distance: "COSINE" + batch_size: 1024 +``` + +## 🔄 Environment Variables + +Some vector store backends support environment variables for configuration: + +- **Elasticsearch**: `FLOW_ES_HOSTS` - Elasticsearch host(s) +- **Qdrant**: + - `FLOW_QDRANT_HOST` - Qdrant host (default: "localhost") + - `FLOW_QDRANT_PORT` - Qdrant port (default: 6333) + - `FLOW_QDRANT_API_KEY` - Qdrant API key for authentication + +Environment variables are used as fallbacks when parameters are not explicitly set in the YAML configuration. + ## 🧩 Integration with Embedding Models -All vector stores require an embedding model to function: +All vector stores require an embedding model configuration. The `embedding_model` field in the vector store configuration references a model defined in the `embedding_model` section of `default.yaml`: -```python -from flowllm.embedding_model import OpenAICompatibleEmbeddingModel +```yaml +embedding_model: + default: + backend: openai_compatible + model_name: text-embedding-v4 + params: + dimensions: 1024 -# Initialize embedding model -embedding_model = OpenAICompatibleEmbeddingModel( - dimensions=64, # Embedding dimensions - model_name="text-embedding-v4", # Model name - batch_size=32 # Batch size for embedding generation -) - -# Pass to vector store (example with LocalVectorStore) -# You can also use: ChromaVectorStore, EsVectorStore, QdrantVectorStore, or MemoryVectorStore -vector_store = LocalVectorStore( - embedding_model=embedding_model, - store_dir="./vector_store" -) +vector_store: + default: + backend: memory + embedding_model: default # References the embedding_model.default configuration ``` + +The embedding model configuration provides the model name, backend, and parameters needed for generating vector embeddings. diff --git a/pyproject.toml b/pyproject.toml index 0a29f795..a29ded0e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,8 +4,8 @@ build-backend = "setuptools.build_meta" [project] name = "reme_ai" -version = "0.1.10.8" -description = "Remember me" +dynamic = ["version"] +description = "Remember Me, Refine Me." authors = [ { name = "jinli.yl", email = "jinli.yl@alibaba-inc.com" }, { name = "dengjiaji.djj", email = "dengjiaji.djj@alibaba-inc.com" }, @@ -13,24 +13,32 @@ authors = [ ] license = { file = "LICENSE" } readme = "README.md" -requires-python = ">=3.12" +requires-python = ">=3.10" classifiers = [ - "Programming Language :: Python :: 3", + "Development Status :: 4 - Beta", + "Intended Audience :: Developers", + "Intended Audience :: Science/Research", "License :: OSI Approved :: Apache Software License", "Operating System :: OS Independent", + "Programming Language :: Python :: 3", + "Programming Language :: Python :: 3.10", + "Topic :: Scientific/Engineering :: Artificial Intelligence", + "Topic :: Software Development :: Libraries :: Python Modules", + "Topic :: Software Development :: Libraries :: Application Frameworks", + "Typing :: Typed", ] keywords = ["llm", "memory", "experience", "memoryscope", "ai", "mcp", "http"] dependencies = [ - "flowllm[reme]>=0.1.11.6", + "flowllm[reme]>=0.2.0.0", ] [project.optional-dependencies] dev = ["jupyter-book", "ghp-import", "myst-nb", "sphinxcontrib-bibtex", "furo", "sphinxcontrib-mermaid"] -all = ["reme_ai[dev]"] +full = ["reme_ai[dev]"] [tool.setuptools.packages.find] where = ["."] @@ -44,7 +52,15 @@ reme_ai = [ "**/*.json", ] -[project.scripts] -reme = "reme_ai.app:main" +[tool.setuptools.dynamic] +version = { attr = "reme_ai.__version__" } -# python -m build && twine upload dist/* \ No newline at end of file +[project.urls] +Homepage = "https://github.com/agentscope-ai/ReMe" +Documentation = "https://reme.agentscope.io/" +Repository = "https://github.com/agentscope-ai/ReMe" + +[project.scripts] +reme = "reme_ai.main:main" + +# python -m build && twine upload dist/* diff --git a/reme_ai/__init__.py b/reme_ai/__init__.py index aa4a2b22..62450870 100644 --- a/reme_ai/__init__.py +++ b/reme_ai/__init__.py @@ -1,11 +1,33 @@ +# pylint: disable=wrong-import-position +"""ReMe AI - A memory management framework for AI agents.""" + import os os.environ["FLOW_APP_NAME"] = "ReMe" -__version__ = "0.1.10.8" +from . import agent # noqa: E402 +from . import config # noqa: E402 +from . import constants # noqa: E402 +from . import enumeration # noqa: E402 +from . import retrieve # noqa: E402 +from . import schema # noqa: E402 +from . import service # noqa: E402 +from . import summary # noqa: E402 +from . import utils # noqa: E402 +from . import vector_store # noqa: E402 +from .main import ReMeApp # noqa: E402 F401 -from reme_ai.app import ReMeApp -from . import agent -from . import retrieve -from . import summary -from . import vector_store +__all__ = [ + "agent", + "config", + "constants", + "enumeration", + "retrieve", + "schema", + "service", + "summary", + "utils", + "vector_store", +] + +__version__ = "0.2.0.0" diff --git a/reme_ai/agent/__init__.py b/reme_ai/agent/__init__.py index bb1db659..2094f344 100644 --- a/reme_ai/agent/__init__.py +++ b/reme_ai/agent/__init__.py @@ -1,2 +1,14 @@ +"""Agent module for ReAct and tool-based agent implementations. + +This module provides submodules for different types of agent operations: +- react: ReAct (Reasoning and Acting) agent implementations +- tools: Mock search tools for testing and demonstration +""" + from . import react from . import tools + +__all__ = [ + "react", + "tools", +] diff --git a/reme_ai/agent/react/__init__.py b/reme_ai/agent/react/__init__.py index 84721cb6..afb89a82 100644 --- a/reme_ai/agent/react/__init__.py +++ b/reme_ai/agent/react/__init__.py @@ -1 +1,11 @@ +"""ReAct agent operations module. + +This module provides ReAct (Reasoning and Acting) agent implementations for +answering user queries through iterative reasoning and search actions. +""" + from .simple_react_op import SimpleReactOp + +__all__ = [ + "SimpleReactOp", +] diff --git a/reme_ai/agent/react/simple_react_op.py b/reme_ai/agent/react/simple_react_op.py index 4929949f..ae1157ce 100644 --- a/reme_ai/agent/react/simple_react_op.py +++ b/reme_ai/agent/react/simple_react_op.py @@ -1,16 +1,44 @@ +"""Simple ReAct operation module. + +This module provides a simple ReAct (Reasoning and Acting) agent implementation +that extends the base ReactSearchOp for answering user queries through iterative +reasoning and search actions. +""" + import asyncio -from flowllm import C -from flowllm.context.flow_context import FlowContext -from flowllm.op.gallery.react_llm_op import ReactLLMOp +from flowllm.core.context import C, FlowContext +from flowllm.gallery import ReactSearchOp @C.register_op() -class SimpleReactOp(ReactLLMOp): - ... +class SimpleReactOp(ReactSearchOp): + """A simple ReAct (Reasoning and Acting) agent operation. + + This operation extends ReactSearchOp to provide a straightforward implementation + of a ReAct agent that answers user queries by reasoning about the problem and + taking search actions iteratively until a final answer is reached. + + The agent inherits all functionality from ReactSearchOp, including: + - Iterative reasoning and action cycles + - Search tool integration + - Maximum step limits for preventing infinite loops + """ async def main(): + """Main function to demonstrate SimpleReactOp usage. + + This function initializes the FlowLLM context with ReMe configuration, + creates a SimpleReactOp instance, and processes a sample query about + stock prices for Maotai and Wuliangye. + + Example: + Run this module directly to test the SimpleReactOp: + ```bash + python -m reme_ai.agent.react.simple_react_op + ``` + """ from reme_ai.config.config_parser import ConfigParser C.set_service_config(parser=ConfigParser, config_name="config=default").init_by_service_config() @@ -20,5 +48,6 @@ async def main(): await op.async_call(context=context) print(context.response.answer) + if __name__ == "__main__": asyncio.run(main()) diff --git a/reme_ai/agent/tools/__init__.py b/reme_ai/agent/tools/__init__.py index ec567ab5..014482d2 100644 --- a/reme_ai/agent/tools/__init__.py +++ b/reme_ai/agent/tools/__init__.py @@ -1,3 +1,17 @@ +"""Mock search tools for testing and demonstration purposes. + +This module provides mock search operations that simulate different search tool behaviors, +including LLM-based query classification and result generation. +""" + from .llm_mock_search_op import LLMMockSearchOp from .mock_search_tools import SearchToolA, SearchToolB, SearchToolC from .use_mock_search_op import UseMockSearchOp + +__all__ = [ + "LLMMockSearchOp", + "SearchToolA", + "SearchToolB", + "SearchToolC", + "UseMockSearchOp", +] diff --git a/reme_ai/agent/tools/llm_mock_search_op.py b/reme_ai/agent/tools/llm_mock_search_op.py index a2219064..91313e25 100644 --- a/reme_ai/agent/tools/llm_mock_search_op.py +++ b/reme_ai/agent/tools/llm_mock_search_op.py @@ -1,13 +1,19 @@ +"""LLM-based mock search operation for simulating search tool behavior. + +This module provides a mock search operation that uses LLM to classify queries +and generate realistic search results based on query complexity levels. +""" + import asyncio import json import random from typing import Dict, Any -from flowllm.context import FlowContext, C -from flowllm.enumeration.role import Role -from flowllm.op.base_async_tool_op import BaseAsyncToolOp -from flowllm.schema.message import Message -from flowllm.schema.tool_call import ToolCall +from flowllm.core.context import C, FlowContext +from flowllm.core.enumeration.role import Role +from flowllm.core.op import BaseAsyncToolOp +from flowllm.core.schema import Message +from flowllm.core.schema import ToolCall from loguru import logger @@ -26,15 +32,18 @@ class LLMMockSearchOp(BaseAsyncToolOp): - extra_time: Extra sleep time in seconds to simulate latency - relevance_ratio: Probability of returning relevant results (vs random query results) """ + file_path: str = __file__ - def __init__(self, - llm: str = "qwen3_30b_instruct", - simple_config: Dict[str, Any] = None, - medium_config: Dict[str, Any] = None, - complex_config: Dict[str, Any] = None, - seed: int = 0, - **kwargs): + def __init__( + self, + llm: str = "qwen3_30b_instruct", + simple_config: Dict[str, Any] = None, + medium_config: Dict[str, Any] = None, + complex_config: Dict[str, Any] = None, + seed: int = 0, + **kwargs, + ): """ Initialize the LLM Mock Search Op. @@ -55,7 +64,7 @@ class LLMMockSearchOp(BaseAsyncToolOp): seed: Random seed for deterministic behavior, default 0 """ super().__init__(llm=llm, **kwargs) - + # Set random seed for deterministic behavior self.seed = seed random.seed(self.seed) @@ -65,7 +74,7 @@ class LLMMockSearchOp(BaseAsyncToolOp): "success_rate": 0.95, "extra_time": 0.5, "relevance_ratio": 0.98, - "content_length": "short" + "content_length": "short", } if simple_config: self.simple_config.update(simple_config) @@ -74,7 +83,7 @@ class LLMMockSearchOp(BaseAsyncToolOp): "success_rate": 0.85, "extra_time": 1.0, "relevance_ratio": 0.90, - "content_length": "medium" + "content_length": "medium", } if medium_config: self.medium_config.update(medium_config) @@ -83,22 +92,29 @@ class LLMMockSearchOp(BaseAsyncToolOp): "success_rate": 0.70, "extra_time": 1.5, "relevance_ratio": 0.80, - "content_length": "long" + "content_length": "long", } if complex_config: self.complex_config.update(complex_config) def build_tool_call(self) -> ToolCall: - return ToolCall(**{ - "description": "Use search keywords to retrieve relevant information from the internet.", - "input_schema": { - "query": { - "type": "string", - "description": "search keyword or query", - "required": True - } - } - }) + """Build the tool call schema for the search operation. + + Returns: + ToolCall object defining the search tool interface + """ + return ToolCall( + **{ + "description": "Use search keywords to retrieve relevant information from the internet.", + "input_schema": { + "query": { + "type": "string", + "description": "search keyword or query", + "required": True, + }, + }, + }, + ) async def classify_query(self, query: str) -> str: """ @@ -112,7 +128,7 @@ class LLMMockSearchOp(BaseAsyncToolOp): """ classification_prompt = self.prompt_format( prompt_name="classification_prompt", - query=query + query=query, ) messages = [Message(role=Role.USER, content=classification_prompt)] @@ -146,7 +162,7 @@ class LLMMockSearchOp(BaseAsyncToolOp): prompt_name="generation_prompt", query=query, complexity=complexity, - content_length=content_length + content_length=content_length, ) messages = [Message(role=Role.USER, content=generation_prompt)] @@ -171,7 +187,7 @@ class LLMMockSearchOp(BaseAsyncToolOp): "breakthroughs in medical science", "architectural wonders", "wildlife conservation efforts", - "developments in artificial intelligence" + "developments in artificial intelligence", ] random_query = random.choice(random_topics) @@ -179,7 +195,7 @@ class LLMMockSearchOp(BaseAsyncToolOp): prompt_name="generation_prompt", query=random_query, complexity="simple", - content_length="short" + content_length="short", ) messages = [Message(role=Role.USER, content=generation_prompt)] @@ -188,6 +204,11 @@ class LLMMockSearchOp(BaseAsyncToolOp): return f"[Low Relevance Result]\n{response.content}" async def async_execute(self): + """Execute the mock search operation. + + This method classifies the query, applies the appropriate configuration, + simulates delays, and generates search results based on success and relevance rates. + """ query: str = self.input_dict["query"] logger.info(f"LLMMockSearchOp processing query: {query}") @@ -218,9 +239,9 @@ class LLMMockSearchOp(BaseAsyncToolOp): "success": False, "content": error_message, "query": query, - "complexity": complexity + "complexity": complexity, } - self.set_result(json.dumps(result_dict, ensure_ascii=False)) + self.set_output(json.dumps(result_dict, ensure_ascii=False)) return # Step 5: Check relevance ratio @@ -235,7 +256,7 @@ class LLMMockSearchOp(BaseAsyncToolOp): "content": content, "query": query, "complexity": complexity, - "is_relevant": False # Mark as irrelevant for debugging + "is_relevant": False, # Mark as irrelevant for debugging } else: # Generate relevant result @@ -246,14 +267,15 @@ class LLMMockSearchOp(BaseAsyncToolOp): "content": content, "query": query, "complexity": complexity, - "is_relevant": True # Mark as relevant for debugging + "is_relevant": True, # Mark as relevant for debugging } - self.set_result(json.dumps(result_dict, ensure_ascii=False)) + self.set_output(json.dumps(result_dict, ensure_ascii=False)) async def async_main(): - from reme_ai.app import ReMeApp + """Main function for testing the LLMMockSearchOp with various query types.""" + from reme_ai.main import ReMeApp async with ReMeApp(): # Test with different query types @@ -267,25 +289,25 @@ async def async_main(): custom_simple = { "success_rate": 1, "extra_time": 0, - "relevance_ratio": 1 + "relevance_ratio": 1, } custom_medium = { "success_rate": 1, "extra_time": 0, - "relevance_ratio": 1 + "relevance_ratio": 1, } custom_complex = { "success_rate": 1, "extra_time": 0, - "relevance_ratio": 1 + "relevance_ratio": 1, } op = LLMMockSearchOp( simple_config=custom_simple, medium_config=custom_medium, - complex_config=custom_complex + complex_config=custom_complex, ) for query in test_queries: diff --git a/reme_ai/agent/tools/llm_mock_search_prompt.yaml b/reme_ai/agent/tools/llm_mock_search_prompt.yaml index 15670654..47183407 100644 --- a/reme_ai/agent/tools/llm_mock_search_prompt.yaml +++ b/reme_ai/agent/tools/llm_mock_search_prompt.yaml @@ -1,76 +1,76 @@ classification_prompt: | You are a query complexity classifier. Analyze the following search query and classify it into one of three categories: - + 1. **simple** - Simple factual queries that: - Ask for a single, direct fact - Have a clear, unambiguous answer - Require minimal context or explanation - Examples: "What is the capital of France?", "Who invented the telephone?", "When did World War 2 end?" - + 2. **medium** - Medium complexity queries that: - Require some explanation or context - May involve multiple related facts - Need balanced depth without being exhaustive - Examples: "How does photosynthesis work?", "What are the main causes of climate change?", "Explain blockchain technology" - + 3. **complex** - Complex research queries that: - Require comprehensive, multi-dimensional analysis - Involve multiple subtopics or perspectives - Need in-depth exploration and connections - Examples: "Analyze the impact of AI on the global economy", "Compare different renewable energy solutions", "What are the geopolitical implications of space exploration?" - + Query to classify: {query} - + Respond with ONLY one word: simple, medium, or complex. classification_prompt_zh: | 你是一个查询复杂度分类器。分析以下搜索查询并将其分类为以下三类之一: - + 1. **simple** - 简单事实查询: - 询问单一、直接的事实 - 有明确、无歧义的答案 - 需要最少的上下文或解释 - 示例:"法国的首都是什么?"、"谁发明了电话?"、"第二次世界大战何时结束?" - + 2. **medium** - 中等复杂度查询: - 需要一些解释或上下文 - 可能涉及多个相关事实 - 需要平衡的深度但不需要详尽无遗 - 示例:"光合作用如何工作?"、"气候变化的主要原因是什么?"、"解释区块链技术" - + 3. **complex** - 复杂研究查询: - 需要全面、多维度的分析 - 涉及多个子主题或观点 - 需要深入探索和联系 - 示例:"分析人工智能对全球经济的影响"、"比较不同的可再生能源解决方案"、"太空探索的地缘政治影响是什么?" - + 要分类的查询:{query} - + 只用一个词回答:simple、medium 或 complex。 generation_prompt: | You are a search engine generating mock search results. Generate a {content_length} response for the following query. - + Query: {query} Complexity Level: {complexity} - + Instructions based on content length: - **short**: Provide a concise answer in 1-3 sentences. Be direct and factual. - **medium**: Provide a balanced answer in 2-4 paragraphs. Include key details and some context. - **long**: Provide a comprehensive answer in 4-6 paragraphs. Include multiple perspectives, detailed explanations, and relevant context. - + Generate the search result content now: generation_prompt_zh: | 你是一个搜索引擎,正在生成模拟搜索结果。为以下查询生成一个 {content_length} 的响应。 - + 查询:{query} 复杂度级别:{complexity} - + 根据内容长度的指示: - **short**(短):提供 1-3 句话的简洁答案。要直接和事实性。 - **medium**(中):提供 2-4 段的平衡答案。包括关键细节和一些上下文。 - **long**(长):提供 4-6 段的全面答案。包括多个角度、详细解释和相关上下文。 - + 现在生成搜索结果内容: diff --git a/reme_ai/agent/tools/mock_search_tools.py b/reme_ai/agent/tools/mock_search_tools.py index 140fa895..944011af 100644 --- a/reme_ai/agent/tools/mock_search_tools.py +++ b/reme_ai/agent/tools/mock_search_tools.py @@ -1,77 +1,125 @@ -from flowllm.context import C -from flowllm.schema.tool_call import ToolCall +"""Specialized mock search tools with different performance characteristics. + +This module provides three search tools (SearchToolA, SearchToolB, SearchToolC) +each optimized for different query complexity levels, allowing for realistic +testing of tool selection strategies. +""" + +from flowllm.core.context import C +from flowllm.core.schema import ToolCall from reme_ai.agent.tools.llm_mock_search_op import LLMMockSearchOp @C.register_op() class SearchToolA(LLMMockSearchOp): + """Fast search tool optimized for simple queries. + + This tool is configured for quick responses with high success rates + on simple queries, but performs poorly on medium and complex queries. + Best suited for simple factual queries. + """ + def __init__(self, llm: str = "qwen3_30b_instruct", **kwargs): + """Initialize SearchToolA with fast, simple-query-optimized configuration. + + Args: + llm: LLM model name to use + **kwargs: Additional arguments passed to LLMMockSearchOp + """ # Configure for fast but shallow performance simple_config = { "success_rate": 0.9, # High success rate for simple queries "extra_time": 0, # Very fast (0.2-0.5s range) "relevance_ratio": 0.9, # High relevance - "content_length": "short" # Concise answers + "content_length": "short", # Concise answers } medium_config = { "success_rate": 0.2, # Lower success for medium queries "extra_time": 0, # Still fast "relevance_ratio": 0.2, # Moderate relevance - "content_length": "short" # Limited depth + "content_length": "short", # Limited depth } complex_config = { "success_rate": 0.5, # Poor success rate for complex queries "extra_time": 0, # Fast but insufficient "relevance_ratio": 0.5, # Low relevance (often misses key aspects) - "content_length": "short" # Too shallow for complex topics + "content_length": "short", # Too shallow for complex topics } - super().__init__(llm=llm, - simple_config=simple_config, - medium_config=medium_config, - complex_config=complex_config, - **kwargs) + super().__init__( + llm=llm, + simple_config=simple_config, + medium_config=medium_config, + complex_config=complex_config, + **kwargs, + ) def build_tool_call(self) -> ToolCall: + """Build the tool call schema with description indicating simple query optimization. + + Returns: + ToolCall object with description indicating best use for simple queries + """ tool_call = super().build_tool_call() tool_call.description += " Best suited for simple queries." return tool_call + @C.register_op() class SearchToolB(LLMMockSearchOp): + """Balanced search tool optimized for medium complexity queries. + + This tool provides balanced performance across query types, with + excellent results for medium complexity queries. Best suited for + queries requiring moderate depth and context. + """ + def __init__(self, llm: str = "qwen3_30b_instruct", **kwargs): + """Initialize SearchToolB with balanced, medium-query-optimized configuration. + + Args: + llm: LLM model name to use + **kwargs: Additional arguments passed to LLMMockSearchOp + """ # Configure for balanced performance simple_config = { "success_rate": 0.3, # Very high success rate "extra_time": 0, # Moderate speed (1.0-1.5s range) "relevance_ratio": 0.3, # High relevance - "content_length": "medium" # More detailed than needed for simple + "content_length": "medium", # More detailed than needed for simple } medium_config = { "success_rate": 0.9, # Excellent success rate "extra_time": 0, # Balanced speed "relevance_ratio": 0.9, # High relevance - "content_length": "medium" # Perfect depth for medium queries + "content_length": "medium", # Perfect depth for medium queries } complex_config = { "success_rate": 0.5, # Good success rate "extra_time": 0, # Still reasonable speed "relevance_ratio": 0.5, # Decent relevance but not exhaustive - "content_length": "medium" # Covers main points but lacks depth + "content_length": "medium", # Covers main points but lacks depth } - super().__init__(llm=llm, - simple_config=simple_config, - medium_config=medium_config, - complex_config=complex_config, - **kwargs) + super().__init__( + llm=llm, + simple_config=simple_config, + medium_config=medium_config, + complex_config=complex_config, + **kwargs, + ) def build_tool_call(self) -> ToolCall: + """Build the tool call schema with description indicating medium query optimization. + + Returns: + ToolCall object with description indicating best use for medium complexity queries + """ tool_call = super().build_tool_call() tool_call.description += " Best suited for medium complexity queries." return tool_call @@ -79,37 +127,56 @@ class SearchToolB(LLMMockSearchOp): @C.register_op() class SearchToolC(LLMMockSearchOp): + """Comprehensive search tool optimized for complex queries. + + This tool provides thorough, in-depth results with high success rates + on complex queries, but may be slower and overly detailed for simple queries. + Best suited for complex research queries requiring comprehensive analysis. + """ def __init__(self, llm: str = "qwen3_30b_instruct", **kwargs): + """Initialize SearchToolC with comprehensive, complex-query-optimized configuration. + + Args: + llm: LLM model name to use + **kwargs: Additional arguments passed to LLMMockSearchOp + """ # Configure for comprehensive but costly performance simple_config = { "success_rate": 0.3, # Good but not optimal (over-processing) "extra_time": 0, # Slow (3.0-4.0s range) "relevance_ratio": 0.3, # High relevance but unnecessary depth - "content_length": "long" # Too detailed for simple queries + "content_length": "long", # Too detailed for simple queries } medium_config = { "success_rate": 0.4, # High success rate "extra_time": 0, # Slow but thorough "relevance_ratio": 0.4, # High relevance with extra context - "content_length": "long" # More depth than needed + "content_length": "long", # More depth than needed } complex_config = { "success_rate": 0.9, # Excellent success rate "extra_time": 0, # Slow but comprehensive (3.5-5.0s range) "relevance_ratio": 0.9, # Very high relevance - "content_length": "long" # Perfect depth for complex queries + "content_length": "long", # Perfect depth for complex queries } - super().__init__(llm=llm, - simple_config=simple_config, - medium_config=medium_config, - complex_config=complex_config, - **kwargs) + super().__init__( + llm=llm, + simple_config=simple_config, + medium_config=medium_config, + complex_config=complex_config, + **kwargs, + ) def build_tool_call(self) -> ToolCall: + """Build the tool call schema with description indicating complex query optimization. + + Returns: + ToolCall object with description indicating best use for complex queries + """ tool_call = super().build_tool_call() tool_call.description += " Best suited for complex queries." return tool_call diff --git a/reme_ai/agent/tools/use_mock_search_op.py b/reme_ai/agent/tools/use_mock_search_op.py index 67247f97..64098dc3 100644 --- a/reme_ai/agent/tools/use_mock_search_op.py +++ b/reme_ai/agent/tools/use_mock_search_op.py @@ -1,13 +1,19 @@ +"""Tool selection and execution operation for mock search tools. + +This module provides an operation that intelligently selects and executes +the most appropriate mock search tool based on query complexity. +""" + import asyncio import datetime import json -from flowllm.context import C -from flowllm.enumeration.role import Role -from flowllm.op.base_async_tool_op import BaseAsyncToolOp -from flowllm.schema.message import Message -from flowllm.schema.tool_call import ToolCall -from flowllm.utils.timer import Timer +from flowllm.core.context import C +from flowllm.core.enumeration import Role +from flowllm.core.op import BaseAsyncToolOp +from flowllm.core.schema import Message +from flowllm.core.schema import ToolCall +from flowllm.core.utils import Timer from loguru import logger from reme_ai.agent.tools.mock_search_tools import SearchToolA, SearchToolB, SearchToolC @@ -16,27 +22,61 @@ from reme_ai.schema.memory import ToolCallResult @C.register_op() class UseMockSearchOp(BaseAsyncToolOp): + """Operation that selects and executes the most appropriate mock search tool. + + This operation uses LLM to intelligently select from available search tools + (SearchToolA, SearchToolB, SearchToolC) based on query characteristics, + then executes the selected tool and records performance metrics. + """ + file_path: str = __file__ def __init__(self, llm: str = "qwen3_30b_instruct", **kwargs): + """Initialize the UseMockSearchOp. + + Args: + llm: LLM model name to use for tool selection + **kwargs: Additional arguments passed to BaseAsyncToolOp + """ super().__init__(llm=llm, save_answer=True, **kwargs) def build_tool_call(self) -> ToolCall: - return ToolCall(**{ - "description": "Intelligently selects and executes the most appropriate search tool based on query complexity. " - "Automatically tracks performance metrics and records tool usage for optimization.", - "input_schema": { - "query": { - "type": "string", - "description": "query", - "required": True - } - } - }) + """Build the tool call schema for the search tool selector. + + Returns: + ToolCall object defining the search tool selector interface + """ + return ToolCall( + **{ + "description": ( + "Intelligently selects and executes the most appropriate search tool " + "based on query complexity. Automatically tracks performance metrics " + "and records tool usage for optimization." + ), + "input_schema": { + "query": { + "type": "string", + "description": "query", + "required": True, + }, + }, + }, + ) async def select_tool(self, query: str, tool_ops: list[BaseAsyncToolOp]) -> ToolCall | None: - assistant_message = await self.llm.achat(messages=[Message(role=Role.USER, content=query)], - tools=[x.tool_call for x in tool_ops]) + """Select the most appropriate tool for the given query using LLM. + + Args: + query: The search query to process + tool_ops: List of available tool operations to choose from + + Returns: + Selected ToolCall if a tool was chosen, None otherwise + """ + assistant_message = await self.llm.achat( + messages=[Message(role=Role.USER, content=query)], + tools=[x.tool_call for x in tool_ops], + ) logger.info(f"assistant_message={assistant_message.model_dump_json()}") if assistant_message.tool_calls: return assistant_message.tool_calls[0] @@ -44,6 +84,11 @@ class UseMockSearchOp(BaseAsyncToolOp): return None async def async_execute(self): + """Execute the tool selection and execution workflow. + + This method selects an appropriate tool, executes it, measures performance, + and creates a ToolCallResult with metrics. + """ query: str = self.input_dict["query"] logger.info(f"query={query}") @@ -65,9 +110,9 @@ class UseMockSearchOp(BaseAsyncToolOp): output="No appropriate tool was selected for the query", token_cost=0, success=False, - time_cost=0.0 + time_cost=0.0, ) - self.set_result(error_result.model_dump_json()) + self.set_output(error_result.model_dump_json()) return # Step 2: Execute the selected tool @@ -86,9 +131,9 @@ class UseMockSearchOp(BaseAsyncToolOp): output=f"Tool {tool_call.name} not found in available tools", token_cost=0, success=False, - time_cost=0.0 + time_cost=0.0, ) - self.set_result(error_result.model_dump_json()) + self.set_output(error_result.model_dump_json()) return # Step 3: Execute the tool with timer @@ -110,14 +155,15 @@ class UseMockSearchOp(BaseAsyncToolOp): output=content, token_cost=token_cost, success=success, - time_cost=round(time_cost, 3) + time_cost=round(time_cost, 3), ) - self.set_result(tool_call_result.model_dump_json()) + self.set_output(tool_call_result.model_dump_json()) async def async_main(): - from reme_ai.app import ReMeApp + """Main function for testing the UseMockSearchOp with various queries.""" + from reme_ai.main import ReMeApp async with ReMeApp(): test_queries = [ @@ -125,7 +171,7 @@ async def async_main(): "How does quantum computing work?", "Analyze the impact of artificial intelligence on global economy, employment, and society", "When was Python programming language created?", - "Compare different types of renewable energy sources", + "Compare different types of renewable energy sources", ] for query in test_queries: @@ -136,4 +182,3 @@ async def async_main(): if __name__ == "__main__": asyncio.run(async_main()) - diff --git a/reme_ai/config/__init__.py b/reme_ai/config/__init__.py index e69de29b..68829efa 100644 --- a/reme_ai/config/__init__.py +++ b/reme_ai/config/__init__.py @@ -0,0 +1,14 @@ +"""Configuration module for ReMe. + +This module provides configuration parsing capabilities for the ReMe framework. +It includes: + +- ConfigParser: Configuration parser class that extends PydanticConfigParser + to provide configuration parsing with awareness of the current module location +""" + +from .config_parser import ConfigParser + +__all__ = [ + "ConfigParser", +] diff --git a/reme_ai/config/config_parser.py b/reme_ai/config/config_parser.py index ecc92e0b..c21cff0f 100644 --- a/reme_ai/config/config_parser.py +++ b/reme_ai/config/config_parser.py @@ -1,6 +1,24 @@ -from flowllm.config.pydantic_config_parser import PydanticConfigParser +"""Configuration parser module for ReMe. + +This module provides configuration parsing capabilities for the ReMe framework. +It extends the PydanticConfigParser from FlowLLM to provide configuration parsing +with awareness of the current module location. +""" + +from flowllm.core.utils import PydanticConfigParser class ConfigParser(PydanticConfigParser): + """Configuration parser for ReMe framework. + + Extends PydanticConfigParser to provide configuration parsing capabilities + with awareness of the current module location. Uses the default.yaml + configuration file as the default configuration source. + + Attributes: + current_file: Path to the current file, used for relative config file resolution. + default_config_name: Default configuration file name (without .yaml extension). + """ + current_file: str = __file__ default_config_name: str = "default" diff --git a/reme_ai/config/default.yaml b/reme_ai/config/default.yaml index 159c1568..7d1dcee0 100644 --- a/reme_ai/config/default.yaml +++ b/reme_ai/config/default.yaml @@ -14,7 +14,7 @@ http: flow: retrieve_task_memory: - flow_content: build_query_op >> recall_vector_store_op >> rerank_memory_op >> rewrite_memory_op + flow_content: BuildQueryOp() >> RecallVectorStoreOp() >> RerankMemoryOp(enable_llm_rerank=True, enable_score_filter=False, top_k=5) >> RewriteMemoryOp(enable_llm_rewrite=True) description: "Retrieves the most relevant top-k memory experiences from historical data based on the current query to enhance task-solving capabilities" input_schema: query: @@ -23,7 +23,7 @@ flow: required: true summary_task_memory: - flow_content: trajectory_preprocess_op >> (success_extraction_op|failure_extraction_op|comparative_extraction_op) >> memory_validation_op >> update_vector_store_op + flow_content: TrajectoryPreprocessOp(success_threshold=1.0) >> (SuccessExtractionOp()|FailureExtractionOp()|ComparativeExtractionOp(enable_soft_comparison=True)) >> MemoryValidationOp(validation_threshold=0.5) >> UpdateVectorStoreOp() description: "Summarizes conversation trajectories or messages into structured memory representations for long-term storage" input_schema: trajectories: @@ -32,7 +32,7 @@ flow: required: false retrieve_task_memory_simple: - flow_content: build_query_op >> recall_vector_store_op >> merge_memory_op + flow_content: BuildQueryOp() >> RecallVectorStoreOp() >> MergeMemoryOp() description: "Retrieves the most relevant top-k memory experiences from historical data based on the current query to enhance task-solving capabilities" input_schema: query: @@ -41,7 +41,7 @@ flow: required: true summary_task_memory_simple: - flow_content: simple_summary_op >> update_vector_store_op + flow_content: SimpleSummaryOp() >> UpdateVectorStoreOp() description: "Summarizes conversation trajectories or messages into structured memory representations for long-term storage" input_schema: trajectories: @@ -50,7 +50,7 @@ flow: required: false retrieve_personal_memory: - flow_content: set_query_op >> (extract_time_op | (retrieve_memory_op >> semantic_rank_op)) >> fuse_rerank_op + flow_content: SetQueryOp() >> (ExtractTimeOp() | (RetrieveMemoryOp() >> SemanticRankOp())) >> FuseRerankOp() description: "Retrieves the most relevant personal memories from historical data based on the query to enhance response quality" input_schema: query: @@ -59,7 +59,7 @@ flow: required: true summary_personal_memory: - flow_content: info_filter_op >> (get_observation_op | get_observation_with_time_op | load_today_memory_op) >> contra_repeat_op >> update_vector_store_op + flow_content: InfoFilterOp() >> (GetObservationOp() | GetObservationWithTimeOp() | LoadTodayMemoryOp()) >> ContraRepeatOp() >> UpdateVectorStoreOp() description: "Consolidates user observations and memories by filtering information and removing redundancies for efficient storage" input_schema: trajectories: @@ -68,7 +68,7 @@ flow: required: false retrieve_tool_memory: - flow_content: retrieve_tool_memory_op + flow_content: RetrieveToolMemoryOp() description: "Retrieves tool memories from the vector database based on tool names to provide tool usage patterns and best practices" input_schema: tool_names: @@ -77,7 +77,7 @@ flow: required: true add_tool_call_result: - flow_content: parse_tool_call_result_op >> update_vector_store_op + flow_content: ParseToolCallResultOp(max_history_tool_call_cnt=100, evaluation_sleep_interval=1.0) >> UpdateVectorStoreOp() description: "Evaluates and adds tool call results to the tool memory database, creating new memory or updating existing memory for the specified tool" input_schema: tool_call_results: @@ -86,7 +86,7 @@ flow: required: true summary_tool_memory: - flow_content: summary_tool_memory_op >> update_vector_store_op + flow_content: SummaryToolMemoryOp(recent_call_count=20, summary_sleep_interval=1.0) >> UpdateVectorStoreOp() description: "Analyzes tool call history and generates comprehensive usage patterns, best practices, and recommendations for the specified tools" input_schema: tool_names: @@ -95,7 +95,7 @@ flow: required: true use_mock_search: - flow_content: use_mock_search_op + flow_content: UseMockSearchOp() description: "Simulates intelligent search tool selection and execution based on query complexity, with automatic tool memory recording" input_schema: query: @@ -104,7 +104,7 @@ flow: required: true vector_store: - flow_content: vector_store_action_op + flow_content: VectorStoreActionOp() description: "Directly operates on the vector store with various management actions" input_schema: action: @@ -114,7 +114,7 @@ flow: enum: [ copy, delete, delete_ids, dump, load, list] record_task_memory: - flow_content: update_memory_freq_op >> update_memory_utility_op >> update_vector_store_op + flow_content: UpdateMemoryFreqOp() >> UpdateMemoryUtilityOp() >> UpdateVectorStoreOp() description: "Update the freq & utility attributes of retrieved task memories" input_schema: workspace_id: @@ -131,7 +131,7 @@ flow: required: true delete_task_memory: - flow_content: delete_memory_op >> update_vector_store_op + flow_content: DeleteMemoryOp() >> UpdateVectorStoreOp() description: "Delete task memories when utility/freq < utility_threshold and freq >= freq_threshold" input_schema: workspace_id: @@ -148,7 +148,7 @@ flow: required: true react: - flow_content: simple_react_op + flow_content: SimpleReactOp() description: "React to the current task with an agent" input_schema: query: @@ -156,69 +156,6 @@ flow: description: "user query" required: true -op: - # retriever ops - rerank_memory_op: - backend: rerank_memory_op - llm: default - params: - enable_llm_rerank: true - enable_score_filter: false - top_k: 5 - - rewrite_memory_op: - backend: rewrite_memory_op - llm: default - params: - enable_llm_rewrite: true - - #summarizer ops - trajectory_preprocess_op: - backend: trajectory_preprocess_op - params: - success_threshold: 1.0 - - success_extraction_op: - backend: success_extraction_op - llm: default - - failure_extraction_op: - backend: failure_extraction_op - llm: default - - comparative_extraction_op: - backend: comparative_extraction_op - llm: default - params: - enable_soft_comparison: true - - memory_validation_op: - backend: memory_validation_op - llm: default - params: - validation_threshold: 0.5 - - memory_deduplication_op: - backend: memory_deduplication_op - vector_store: default - params: - similarity_threshold: 0.5 - - # tool memory ops - parse_tool_call_result_op: - backend: parse_tool_call_result_op - llm: default - params: - max_history_tool_call_cnt: 100 - evaluation_sleep_interval: 1.0 - - summary_tool_memory_op: - backend: summary_tool_memory_op - llm: default - params: - recent_call_count: 20 - summary_sleep_interval: 1.0 - llm: default: backend: openai_compatible diff --git a/reme_ai/constants/__init__.py b/reme_ai/constants/__init__.py index d351b271..bc1e2d1e 100644 --- a/reme_ai/constants/__init__.py +++ b/reme_ai/constants/__init__.py @@ -1,7 +1,92 @@ +"""Constants module for ReMe AI. + +This module provides access to all constants used throughout the application, +including common workflow keys and language-specific constants. +""" + from . import common_constants from . import language_constants +# Export all constants from common_constants +from .common_constants import ( + WORKFLOW_NAME, + RESULT, + MEMORIES, + CHAT_MESSAGES, + CHAT_MESSAGES_SCATTER, + CHAT_KWARGS, + USER_NAME, + TARGET_NAME, + MEMORY_MANAGER, + QUERY_WITH_TS, + RETRIEVE_MEMORY_NODES, + RANKED_MEMORY_NODES, + NOT_REFLECTED_NODES, + NOT_UPDATED_NODES, + EXTRACT_TIME_DICT, + NEW_OBS_NODES, + NEW_OBS_WITH_TIME_NODES, + INSIGHT_NODES, + TODAY_NODES, + MERGE_OBS_NODES, + TIME_INFER, +) + +# Export all constants from language_constants +from .language_constants import ( + DATATIME_WORD_LIST, + WEEKDAYS, + MONTH_DICT, + NONE_WORD, + REPEATED_WORD, + CONTRADICTORY_WORD, + CONTAINED_WORD, + COLON_WORD, + COMMA_WORD, + DEFAULT_HUMAN_NAME, + DATATIME_KEY_MAP, + TIME_INFER_WORD, + USER_NAME_EXPRESSION, +) + __all__ = [ + # Module exports "common_constants", - "language_constants" + "language_constants", + # Common constants + "WORKFLOW_NAME", + "RESULT", + "MEMORIES", + "CHAT_MESSAGES", + "CHAT_MESSAGES_SCATTER", + "CHAT_KWARGS", + "USER_NAME", + "TARGET_NAME", + "MEMORY_MANAGER", + "QUERY_WITH_TS", + "RETRIEVE_MEMORY_NODES", + "RANKED_MEMORY_NODES", + "NOT_REFLECTED_NODES", + "NOT_UPDATED_NODES", + "EXTRACT_TIME_DICT", + "NEW_OBS_NODES", + "NEW_OBS_WITH_TIME_NODES", + "INSIGHT_NODES", + "TODAY_NODES", + "MERGE_OBS_NODES", + "TIME_INFER", + # Language constants + "DATATIME_WORD_LIST", + "WEEKDAYS", + "MONTH_DICT", + "NONE_WORD", + "REPEATED_WORD", + "CONTRADICTORY_WORD", + "CONTAINED_WORD", + "COLON_WORD", + "COMMA_WORD", + "DEFAULT_HUMAN_NAME", + "DATATIME_KEY_MAP", + "TIME_INFER_WORD", + "USER_NAME_EXPRESSION", ] diff --git a/reme_ai/constants/common_constants.py b/reme_ai/constants/common_constants.py index 99ce3887..547875b6 100644 --- a/reme_ai/constants/common_constants.py +++ b/reme_ai/constants/common_constants.py @@ -1,7 +1,10 @@ -# common_constants.py -# This module defines constants used as keys throughout the application to maintain a consistent reference -# for data structures related to workflow management, chat interactions, context storage, memory operations, -# node processing, and temporal inference functionalities. +"""Common constants module. + +This module defines constants used as keys throughout the application to maintain +a consistent reference for data structures related to workflow management, chat +interactions, context storage, memory operations, node processing, and temporal +inference functionalities. +""" WORKFLOW_NAME = "workflow_name" diff --git a/reme_ai/constants/language_constants.py b/reme_ai/constants/language_constants.py index 5b93c35a..c24e2cc0 100644 --- a/reme_ai/constants/language_constants.py +++ b/reme_ai/constants/language_constants.py @@ -1,3 +1,11 @@ +"""Language constants module. + +This module provides language-specific constants and mappings for datetime +expressions, weekdays, months, and other linguistic elements used throughout +the application. It supports multiple languages (currently Chinese and English) +and facilitates internationalization of temporal and linguistic processing. +""" + from ..enumeration.language_enum import LanguageEnum # This dictionary maps languages to lists of words related to datetime expressions. @@ -29,64 +37,101 @@ DATATIME_WORD_LIST = { ], LanguageEnum.EN: [ # Units of Time - "year", "yr", - "month", "mo", - "week", "wk", - "day", "d", - "hour", "hr", - "minute", "min", - "second", "sec", - + "year", + "yr", + "month", + "mo", + "week", + "wk", + "day", + "d", + "hour", + "hr", + "minute", + "min", + "second", + "sec", # Days of the Week - "Monday", "Mon", - "Tuesday", "Tue", "Tues", - "Wednesday", "Wed", - "Thursday", "Thu", "Thur", "Thurs", - "Friday", "Fri", - "Saturday", "Sat", - "Sunday", "Sun", - + "Monday", + "Mon", + "Tuesday", + "Tue", + "Tues", + "Wednesday", + "Wed", + "Thursday", + "Thu", + "Thur", + "Thurs", + "Friday", + "Fri", + "Saturday", + "Sat", + "Sunday", + "Sun", # Months of the Year - "January", "Jan", - "February", "Feb", - "March", "Mar", - "April", "Apr", - "May", "May", - "June", "Jun", - "July", "Jul", - "August", "Aug", - "September", "Sep", "Sept", - "October", "Oct", - "November", "Nov", - "December", "Dec", - + "January", + "Jan", + "February", + "Feb", + "March", + "Mar", + "April", + "Apr", + "May", + "May", + "June", + "Jun", + "July", + "Jul", + "August", + "Aug", + "September", + "Sep", + "Sept", + "October", + "Oct", + "November", + "Nov", + "December", + "Dec", # Relative Time References "Today", - "Tomorrow", "Tmrw", - "Yesterday", "Yday", + "Tomorrow", + "Tmrw", + "Yesterday", + "Yday", "Now", - "Morning", "AM", "a.m.", - "Afternoon", "PM", "p.m.", + "Morning", + "AM", + "a.m.", + "Afternoon", + "PM", + "p.m.", "Evening", "Night", "Midnight", "Noon", - # Seasonal References "Spring", "Summer", - "Autumn", "Fall", + "Autumn", + "Fall", "Winter", - # General Time References - "Century", "cent.", + "Century", + "cent.", "Decade", "Millennium", - "Quarter", "Q1", "Q2", "Q3", "Q4", + "Quarter", + "Q1", + "Q2", + "Q3", + "Q4", "Semester", "Fortnight", - "Weekend" - ] + "Weekend", + ], } # A mapping of weekdays for each supported language, facilitating calendar-related operations and understanding @@ -99,7 +144,7 @@ WEEKDAYS = { "周四", "周五", "周六", - "周日" + "周日", ], LanguageEnum.EN: [ "Monday", @@ -109,7 +154,7 @@ WEEKDAYS = { "Friday", "Saturday", "Sunday", - ] + ], } MONTH_DICT = { @@ -140,49 +185,49 @@ MONTH_DICT = { "October", "November", "December", - ] + ], } # Constants for the word 'none' in different languages NONE_WORD = { LanguageEnum.CN: "无", - LanguageEnum.EN: "none" + LanguageEnum.EN: "none", } # Constants for the word 'repeated' in different languages REPEATED_WORD = { LanguageEnum.CN: "重复", - LanguageEnum.EN: "repeated" + LanguageEnum.EN: "repeated", } # Constants for the word 'contradictory' in different languages CONTRADICTORY_WORD = { LanguageEnum.CN: "矛盾", - LanguageEnum.EN: "contradiction" + LanguageEnum.EN: "contradiction", } # Constants for the phrase 'included' in different languages CONTAINED_WORD = { LanguageEnum.CN: "被包含", - LanguageEnum.EN: "contained" + LanguageEnum.EN: "contained", } # Constants for the symbol ':' in different languages' representations COLON_WORD = { LanguageEnum.CN: ":", - LanguageEnum.EN: ":" + LanguageEnum.EN: ":", } # Constants for the symbol ',' in different languages' representations COMMA_WORD = { LanguageEnum.CN: ",", - LanguageEnum.EN: "," + LanguageEnum.EN: ",", } # Default human name placeholders for different languages DEFAULT_HUMAN_NAME = { LanguageEnum.CN: "用户", - LanguageEnum.EN: "user" + LanguageEnum.EN: "user", } # Mapping of datetime terms from natural language to standardized keys for each supported language @@ -200,16 +245,16 @@ DATATIME_KEY_MAP = { "Day": "day", "Week": "week", "Weekday": "weekday", - } + }, } # Phrase for indicating inferred time in different languages TIME_INFER_WORD = { LanguageEnum.CN: "推断时间", - LanguageEnum.EN: "Inference time" + LanguageEnum.EN: "Inference time", } USER_NAME_EXPRESSION = { LanguageEnum.CN: "用户姓名是{name}。", - LanguageEnum.EN: "User's name is {name}." + LanguageEnum.EN: "User's name is {name}.", } diff --git a/reme_ai/enumeration/__init__.py b/reme_ai/enumeration/__init__.py index e69de29b..a6466cf3 100644 --- a/reme_ai/enumeration/__init__.py +++ b/reme_ai/enumeration/__init__.py @@ -0,0 +1,11 @@ +"""Enumeration module for ReMe. + +This module provides enumerations used throughout the ReMe system, +including language enumerations and other type definitions. +""" + +from reme_ai.enumeration.language_enum import LanguageEnum + +__all__ = [ + "LanguageEnum", +] diff --git a/reme_ai/enumeration/language_enum.py b/reme_ai/enumeration/language_enum.py index b59ec74c..d50ae0f7 100644 --- a/reme_ai/enumeration/language_enum.py +++ b/reme_ai/enumeration/language_enum.py @@ -1,3 +1,8 @@ +"""Language enumeration module. + +This module provides enumerations for supported languages in the ReMe system. +""" + from enum import Enum @@ -9,6 +14,7 @@ class LanguageEnum(str, Enum): - CN: Represents the Chinese language. - EN: Represents the English language. """ + CN = "cn" EN = "en" diff --git a/reme_ai/app.py b/reme_ai/main.py similarity index 80% rename from reme_ai/app.py rename to reme_ai/main.py index 62692378..8de69055 100644 --- a/reme_ai/app.py +++ b/reme_ai/main.py @@ -12,50 +12,57 @@ which extends FlowLLM with specialized memory management capabilities including: import asyncio import sys -from flowllm import FlowLLMApp, C -from flowllm.schema.flow_response import FlowResponse +from flowllm.core.application import Application +from flowllm.core.context import C +from flowllm.core.schema import FlowResponse from reme_ai.config.config_parser import ConfigParser -class ReMeApp(FlowLLMApp): +class ReMeApp(Application): """ ReMeApp - Main application class for Reflexive Memory system. - + ReMeApp extends FlowLLMApp to provide enhanced memory capabilities for AI agents. It manages multiple types of memories and provides both synchronous and asynchronous execution interfaces for memory-enhanced workflows. """ - def __init__(self, - *args, - llm_api_key: str = None, - llm_api_base: str = None, - embedding_api_key: str = None, - embedding_api_base: str = None, - config_path: str = None, - **kwargs): + def __init__( + self, + *args, + llm_api_key: str = None, + llm_api_base: str = None, + embedding_api_key: str = None, + embedding_api_base: str = None, + config_path: str = None, + **kwargs, + ): """ Initialize ReMeApp with configuration for LLM, embeddings, and vector stores. - - ⚠️ IMPORTANT: The initialization parameters here are consistent with the command-line + + ⚠️ IMPORTANT: The initialization parameters here are consistent with the command-line startup parameters shown in README.md. You can use the same configuration in both ways: - - Command-line startup (from README): - $ reme \\ - backend=http \\ - http.port=8002 \\ - llm.default.model_name=qwen3-30b-a3b-thinking-2507 \\ - embedding_model.default.model_name=text-embedding-v4 \\ + + Command-line startup: + ```bash + reme \ + backend=http \ + http.port=8002 \ + llm.default.model_name=qwen3-30b-a3b-thinking-2507 \ + embedding_model.default.model_name=text-embedding-v4 \ vector_store.default.backend=memory - + ``` + Python API equivalent: - >>> app = ReMeApp( - ... "llm.default.model_name=qwen3-30b-a3b-thinking-2507", - ... "embedding_model.default.model_name=text-embedding-v4", - ... "vector_store.default.backend=memory" - ... ) - + ```python + app = ReMeApp( + "llm.default.model_name=qwen3-30b-a3b-thinking-2507", + "embedding_model.default.model_name=text-embedding-v4", + "vector_store.default.backend=memory" + ) + ``` + Both approaches accept the same configuration parameters and produce identical results. Args: @@ -101,40 +108,42 @@ class ReMeApp(FlowLLMApp): This overrides the default configuration with your custom settings. **kwargs: Additional keyword arguments passed to parser. Same format as args but as key-value pairs. Example: model_name="gpt-4", temperature=0.7 - + Raises: AssertionError: If required configurations are missing or invalid. - + Note: - Parameters here mirror the command-line options in README.md exactly - API keys can be provided via arguments or environment variables (see example.env) - The parser (ConfigParser) handles merging default configs with custom overrides - Vector store configuration determines where memories are persisted - For detailed startup examples and all available parameters, refer to README.md Quick Start section - + See Also: - README.md "Quick Start" section for command-line startup examples - README.md "Environment Configuration" for environment variable setup - example.env for all available environment variables """ - super().__init__(*args, - llm_api_key=llm_api_key, - llm_api_base=llm_api_base, - embedding_api_key=embedding_api_key, - embedding_api_base=embedding_api_base, - service_config=None, - parser=ConfigParser, - config_path=config_path, - load_default_config=True, - **kwargs) + super().__init__( + *args, + llm_api_key=llm_api_key, + llm_api_base=llm_api_base, + embedding_api_key=embedding_api_key, + embedding_api_base=embedding_api_base, + service_config=None, + parser=ConfigParser, + config_path=config_path, + load_default_config=True, + **kwargs, + ) async def async_execute(self, name: str, **kwargs) -> dict: """ Asynchronously execute a named flow with given parameters. - + This method executes a registered flow (workflow) by name and returns the result. Flows are defined in the configuration and registered during app initialization. - + Args: name: Name of the flow to execute. Must be registered in C.flow_dict. Common flows in ReMe: @@ -148,23 +157,25 @@ class ReMeApp(FlowLLMApp): - context (dict): Additional context for the flow - max_results (int): Maximum number of results to return - threshold (float): Similarity threshold for retrieval - + Returns: dict: Flow execution result as a dictionary containing: - status: Execution status (success/failure) - result: Flow output data - metadata: Additional execution metadata - + Raises: AssertionError: If the flow name is not registered in C.flow_dict. - + Example: - >>> result = await app.async_execute( - ... "task_memory_flow", - ... query="Show me all Python debugging tasks", - ... max_results=10 - ... ) - >>> print(result['result']) + ```python + result = await app.async_execute( + "task_memory_flow", + query="Show me all Python debugging tasks", + max_results=10 + ) + print(result['result']) + ``` """ assert name in C.flow_dict, f"Invalid flow_name={name} !" result: FlowResponse = await self.async_execute_flow(name=name, **kwargs) @@ -173,28 +184,30 @@ class ReMeApp(FlowLLMApp): def execute(self, name: str, **kwargs) -> dict: """ Synchronously execute a named flow with given parameters. - + This is a convenience wrapper around async_execute() for synchronous contexts. It internally uses asyncio.run() to execute the async flow. - + Args: name: Name of the flow to execute. See async_execute() for available flows. **kwargs: Keyword arguments passed to the flow. See async_execute() for details. - + Returns: dict: Flow execution result. Same format as async_execute(). - + Raises: AssertionError: If the flow name is not registered. - + Example: - >>> app = ReMeApp() - >>> result = app.execute( - ... "tool_memory_flow", - ... query="How to use the search tool effectively?" - ... ) - >>> print(result) - + ```python + app = ReMeApp() + result = app.execute( + "tool_memory_flow", + query="How to use the search tool effectively?" + ) + print(result) + ``` + Note: For better performance in async contexts, prefer using async_execute() directly. This method creates a new event loop for each call, which has overhead. @@ -205,24 +218,25 @@ class ReMeApp(FlowLLMApp): def main(): """ Entry point for running ReMeApp as a service. - + This function initializes ReMeApp with command-line arguments and starts the service. It's typically called when running the module directly (python -m reme_ai.app). - + Command-line arguments are passed directly to ReMeApp.__init__(), allowing configuration via command line: - + Example: - $ python -m reme_ai.app --llm_api_key=sk-xxx --config_path=config.yaml - + python -m reme_ai.app --llm_api_key=sk-xxx --config_path=config.yaml + The app runs as a context manager, ensuring proper cleanup of resources (database connections, API clients, etc.) on shutdown. - + Note: Press Ctrl+C to gracefully shutdown the service. """ with ReMeApp(*sys.argv[1:]) as app: app.run_service() + if __name__ == "__main__": main() diff --git a/reme_ai/retrieve/__init__.py b/reme_ai/retrieve/__init__.py index 4f26e94d..9c95d946 100644 --- a/reme_ai/retrieve/__init__.py +++ b/reme_ai/retrieve/__init__.py @@ -1,3 +1,17 @@ +"""Retrieval module for memory operations. + +This module provides submodules for different types of memory retrieval: +- personal: Personal memory retrieval operations +- task: Task memory retrieval operations +- tool: Tool memory retrieval operations +""" + from . import personal from . import task from . import tool + +__all__ = [ + "personal", + "task", + "tool", +] diff --git a/reme_ai/retrieve/personal/__init__.py b/reme_ai/retrieve/personal/__init__.py index 16168138..ee4fe03a 100644 --- a/reme_ai/retrieve/personal/__init__.py +++ b/reme_ai/retrieve/personal/__init__.py @@ -1,3 +1,9 @@ +"""Personal memory retrieval operations module. + +This module provides operations for retrieving, ranking, and formatting personal memories +from a vector store, including time extraction, semantic ranking, and memory formatting. +""" + from .extract_time_op import ExtractTimeOp from .fuse_rerank_op import FuseRerankOp from .print_memory_op import PrintMemoryOp @@ -13,5 +19,5 @@ __all__ = [ "ReadMessageOp", "RetrieveMemoryOp", "SemanticRankOp", - "SetQueryOp" + "SetQueryOp", ] diff --git a/reme_ai/retrieve/personal/extract_time_op.py b/reme_ai/retrieve/personal/extract_time_op.py index 70266b87..a493f5f0 100644 --- a/reme_ai/retrieve/personal/extract_time_op.py +++ b/reme_ai/retrieve/personal/extract_time_op.py @@ -1,9 +1,16 @@ +"""Time extraction operation for personal memories. + +This module provides functionality to extract time-related information from queries +using LLM-based extraction and pattern matching. +""" + import re from typing import Dict -from flowllm import C, BaseAsyncOp -from flowllm.enumeration.role import Role -from flowllm.schema.message import Message +from flowllm.core.context import C +from flowllm.core.enumeration import Role +from flowllm.core.op import BaseAsyncOp +from flowllm.core.schema import Message from loguru import logger from reme_ai.constants.common_constants import QUERY_WITH_TS, EXTRACT_TIME_DICT @@ -13,16 +20,25 @@ from reme_ai.utils.datetime_handler import DatetimeHandler @C.register_op() class ExtractTimeOp(BaseAsyncOp): - file_path: str = __file__ - EXTRACT_TIME_PATTERN = r"-\s*(\S+)[::]\s*(\S+)" - """ A specialized worker class designed to identify and extract time-related information from text generated by an LLM, translating date-time keywords based on the set language, and storing this extracted data within a shared context. """ + file_path: str = __file__ + EXTRACT_TIME_PATTERN = r"-\s*(\S+)[::]\s*(\S+)" + def get_language_value(self, value_dict: dict): + """ + Get value from dictionary based on current language setting. + + Args: + value_dict: Dictionary with language keys + + Returns: + Value for current language or English fallback + """ return value_dict.get(self.language, value_dict.get("en")) async def async_execute(self): @@ -51,8 +67,11 @@ class ExtractTimeOp(BaseAsyncOp): # Create message with system and few-shot examples system_prompt = self.prompt_format(prompt_name="extract_time_system") few_shot = self.prompt_format(prompt_name="extract_time_few_shot") - user_prompt = self.prompt_format(prompt_name="extract_time_user_query", - query=query, query_time_str=query_time_str) + user_prompt = self.prompt_format( + prompt_name="extract_time_user_query", + query=query, + query_time_str=query_time_str, + ) full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_prompt}" logger.info(f"Extracting time from query: {query[:100]}...") @@ -77,10 +96,10 @@ class ExtractTimeOp(BaseAsyncOp): def _parse_time_from_response(self, response_text: str) -> Dict[str, str]: """ Parse time information from LLM response using regex. - + Args: response_text: Raw LLM response content - + Returns: Dictionary of extracted time information """ diff --git a/reme_ai/retrieve/personal/fuse_rerank_op.py b/reme_ai/retrieve/personal/fuse_rerank_op.py index 28e3dc1b..766c66c2 100644 --- a/reme_ai/retrieve/personal/fuse_rerank_op.py +++ b/reme_ai/retrieve/personal/fuse_rerank_op.py @@ -1,6 +1,13 @@ +"""Fuse reranking operation for personal memories. + +This module provides functionality to rerank memory nodes by combining scores, +memory types, and temporal relevance to improve retrieval quality. +""" + from typing import Dict, List -from flowllm import C, BaseAsyncOp +from flowllm.core.context import C +from flowllm.core.op import BaseAsyncOp from loguru import logger from reme_ai.constants.common_constants import EXTRACT_TIME_DICT @@ -12,12 +19,20 @@ class FuseRerankOp(BaseAsyncOp): """ Reranks the memory nodes by scores, types, and temporal relevance. Formats the top-K reranked nodes to print. """ + file_path: str = __file__ @staticmethod def match_memory_time(extract_time_dict: Dict[str, str], memory: BaseMemory): """ Determines whether the memory is relevant based on time matching. + + Args: + extract_time_dict: Dictionary containing extracted time information + memory: Memory object to check for time relevance + + Returns: + Tuple of (match_event_flag, match_msg_flag) indicating temporal matches """ if extract_time_dict: match_event_flag = True @@ -25,18 +40,16 @@ class FuseRerankOp(BaseAsyncOp): event_value = memory.metadata.get(f"event_{k}", "") if event_value in ["-1", v]: continue - else: - match_event_flag = False - break + match_event_flag = False + break match_msg_flag = True for k, v in extract_time_dict.items(): msg_value = memory.metadata.get(f"msg_{k}", "") if msg_value == v: continue - else: - match_msg_flag = False - break + match_msg_flag = False + break else: match_event_flag = False match_msg_flag = False @@ -59,12 +72,15 @@ class FuseRerankOp(BaseAsyncOp): """ # Get operation parameters fuse_score_threshold = self.op_params.get("fuse_score_threshold", 0.1) - fuse_ratio_dict = self.op_params.get("fuse_ratio_dict", { - "conversation": 0.5, - "observation": 1, - "obs_customized": 1.2, - "insight": 2.0 - }) + fuse_ratio_dict = self.op_params.get( + "fuse_ratio_dict", + { + "conversation": 0.5, + "observation": 1, + "obs_customized": 1.2, + "insight": 2.0, + }, + ) fuse_time_ratio = self.op_params.get("fuse_time_ratio", 2.0) output_memory_max_count = self.op_params.get("output_memory_max_count", 5) @@ -83,14 +99,19 @@ class FuseRerankOp(BaseAsyncOp): # Perform reranking based on score, type, and time relevance reranked_memories = self._apply_fuse_reranking( - memory_list, extract_time_dict, fuse_score_threshold, - fuse_ratio_dict, fuse_time_ratio + memory_list, + extract_time_dict, + fuse_score_threshold, + fuse_ratio_dict, + fuse_time_ratio, ) # Sort and select top-k memories - reranked_memories = sorted(reranked_memories, - key=lambda x: x.score or 0.0, - reverse=True)[:output_memory_max_count] + reranked_memories = sorted( + reranked_memories, + key=lambda x: x.score or 0.0, + reverse=True, + )[:output_memory_max_count] logger.info(f"Final reranked memories: {len(reranked_memories)}") @@ -101,13 +122,27 @@ class FuseRerankOp(BaseAsyncOp): self.context.response.metadata["memory_list"] = reranked_memories self.context.response.answer = "\n".join(formatted_memories) - def _apply_fuse_reranking(self, - memory_list: List[BaseMemory], - extract_time_dict: Dict[str, str], - fuse_score_threshold: float, - fuse_ratio_dict: Dict[str, float], - fuse_time_ratio: float) -> List[BaseMemory]: - """Apply fuse reranking logic to memories""" + def _apply_fuse_reranking( + self, + memory_list: List[BaseMemory], + extract_time_dict: Dict[str, str], + fuse_score_threshold: float, + fuse_ratio_dict: Dict[str, float], + fuse_time_ratio: float, + ) -> List[BaseMemory]: + """ + Apply fuse reranking logic to memories. + + Args: + memory_list: List of memories to rerank + extract_time_dict: Dictionary containing extracted time information + fuse_score_threshold: Minimum score threshold for memories + fuse_ratio_dict: Dictionary mapping memory types to score multipliers + fuse_time_ratio: Multiplier for time-relevant memories + + Returns: + List of reranked memories with updated scores + """ reranked_memories = [] for memory in memory_list: @@ -130,23 +165,35 @@ class FuseRerankOp(BaseAsyncOp): original_score = memory_score memory.score = memory_score * type_ratio * time_ratio - logger.debug(f"Memory reranked: {original_score:.3f} -> {memory.score:.3f} " - f"(type={type_ratio}, time={time_ratio})") + logger.debug( + f"Memory reranked: {original_score:.3f} -> {memory.score:.3f} " + f"(type={type_ratio}, time={time_ratio})", + ) reranked_memories.append(memory) return reranked_memories def _format_memories_for_output(self, memories: List[BaseMemory]) -> List[str]: - """Format memories for final output""" + """ + Format memories for final output. + + Args: + memories: List of memories to format + + Returns: + List of formatted memory strings + """ formatted_memories = [] for memory in memories: # Log reranking details - logger.info(f"Final memory: Score={memory.score:.3f}, " - f"Event={memory.metadata.get('match_event_flag', '0')}, " - f"Msg={memory.metadata.get('match_msg_flag', '0')}, " - f"Content={memory.content[:50]}...") + logger.info( + f"Final memory: Score={memory.score:.3f}, " + f"Event={memory.metadata.get('match_event_flag', '0')}, " + f"Msg={memory.metadata.get('match_msg_flag', '0')}, " + f"Content={memory.content[:50]}...", + ) # Format memory with timestamp if available formatted_content = self._format_memory_with_timestamp(memory, self.language) @@ -167,8 +214,9 @@ class FuseRerankOp(BaseAsyncOp): Formatted memory content string """ try: - if hasattr(memory, 'timestamp') and memory.timestamp: + if hasattr(memory, "timestamp") and memory.timestamp: from reme_ai.utils.datetime_handler import DatetimeHandler + dt_handler = DatetimeHandler(memory.timestamp) datetime_str = dt_handler.datetime_format("%Y-%m-%d %H:%M:%S") weekday = dt_handler.get_dt_info_dict(language)["weekday"] diff --git a/reme_ai/retrieve/personal/print_memory_op.py b/reme_ai/retrieve/personal/print_memory_op.py index f7296f2b..8cef06bd 100644 --- a/reme_ai/retrieve/personal/print_memory_op.py +++ b/reme_ai/retrieve/personal/print_memory_op.py @@ -1,6 +1,13 @@ +"""Memory printing operation for personal memories. + +This module provides functionality to format and print memories in various formats +for display or output purposes. +""" + from typing import List -from flowllm import C, BaseAsyncOp +from flowllm.core.context import C +from flowllm.core.op import BaseAsyncOp from loguru import logger from reme_ai.schema.memory import BaseMemory @@ -11,6 +18,7 @@ class PrintMemoryOp(BaseAsyncOp): """ Formats the memories to print. """ + file_path: str = __file__ async def async_execute(self): @@ -39,7 +47,15 @@ class PrintMemoryOp(BaseAsyncOp): @staticmethod def _format_memories_for_print(memories: List[BaseMemory]) -> str: - """Format memories for printing""" + """ + Format memories for printing. + + Args: + memories: List of memory objects to format + + Returns: + Formatted string representation of memories + """ if not memories: return "No memories available." @@ -51,10 +67,10 @@ class PrintMemoryOp(BaseAsyncOp): memory_text += f" Content: {memory.content}\n" # Add additional metadata if available - if hasattr(memory, 'metadata') and memory.metadata: + if hasattr(memory, "metadata") and memory.metadata: metadata_items = [] for key, value in memory.metadata.items(): - if key not in ['when_to_use', 'content']: + if key not in ["when_to_use", "content"]: metadata_items.append(f"{key}: {value}") if metadata_items: memory_text += f" Metadata: {', '.join(metadata_items)}\n" @@ -67,10 +83,10 @@ class PrintMemoryOp(BaseAsyncOp): def format_memories_for_output(memories: List) -> str: """ Format memory list for output string. - + Args: memories: List of memory objects - + Returns: Formatted string """ @@ -79,8 +95,8 @@ class PrintMemoryOp(BaseAsyncOp): formatted_parts = [] for i, memory in enumerate(memories, 1): - when_to_use = getattr(memory, 'when_to_use', '') or memory.get('when_to_use', '') - content = getattr(memory, 'content', '') or memory.get('content', '') + when_to_use = getattr(memory, "when_to_use", "") or memory.get("when_to_use", "") + content = getattr(memory, "content", "") or memory.get("content", "") part = f"Memory {i}:\n" if when_to_use: @@ -96,10 +112,10 @@ class PrintMemoryOp(BaseAsyncOp): def format_memories_for_simple_output(memories: List) -> str: """ Format memory list for simple flow output. - + Args: memories: List of memory objects - + Returns: Formatted string suitable for response.answer """ @@ -110,8 +126,8 @@ class PrintMemoryOp(BaseAsyncOp): for memory in memories: # Safely get field values - when_to_use = getattr(memory, 'when_to_use', '') or memory.get('when_to_use', '') - content = getattr(memory, 'content', '') or memory.get('content', '') + when_to_use = getattr(memory, "when_to_use", "") or memory.get("when_to_use", "") + content = getattr(memory, "content", "") or memory.get("content", "") # Skip memories with empty content if not content: @@ -125,7 +141,9 @@ class PrintMemoryOp(BaseAsyncOp): if len(content_parts) == 1: # Only title return "No relevant memories with valid content found." - content_parts.append("\nPlease consider the helpful parts from these in answering the question, " - "to make the response more comprehensive and substantial.") + content_parts.append( + "\nPlease consider the helpful parts from these in answering the question, " + "to make the response more comprehensive and substantial.", + ) return "\n".join(content_parts) diff --git a/reme_ai/retrieve/personal/read_message_op.py b/reme_ai/retrieve/personal/read_message_op.py index 307e9511..1f901f42 100644 --- a/reme_ai/retrieve/personal/read_message_op.py +++ b/reme_ai/retrieve/personal/read_message_op.py @@ -1,7 +1,14 @@ +"""Message reading operation for personal memories. + +This module provides functionality to read and filter unmemorized chat messages +from the context for processing. +""" + from typing import List -from flowllm import C, BaseAsyncOp -from flowllm.schema.message import Message +from flowllm.core.context import C +from flowllm.core.op import BaseAsyncOp +from flowllm.core.schema import Message from loguru import logger @@ -10,6 +17,7 @@ class ReadMessageOp(BaseAsyncOp): """ Fetches unmemorized chat messages. """ + file_path: str = __file__ async def async_execute(self): @@ -19,20 +27,20 @@ class ReadMessageOp(BaseAsyncOp): # Get chat messages from context chat_messages = self.context.chat_messages target_name = self.context.target_name - contextual_msg_max_count = self.op_params.get('contextual_msg_max_count', 10) + contextual_msg_max_count = self.op_params.get("contextual_msg_max_count", 10) chat_messages_not_memorized: List[List[Message]] = [] for messages in chat_messages: if not messages: continue - if hasattr(messages[0], 'memorized') and messages[0].memorized: + if hasattr(messages[0], "memorized") and messages[0].memorized: continue contain_flag = False for msg in messages: - if hasattr(msg, 'role_name') and msg.role_name == target_name: + if hasattr(msg, "role_name") and msg.role_name == target_name: contain_flag = True break @@ -44,7 +52,7 @@ class ReadMessageOp(BaseAsyncOp): chat_message_scatter.extend(messages) # Sort by time_created if available - if chat_message_scatter and hasattr(chat_message_scatter[0], 'time_created'): + if chat_message_scatter and hasattr(chat_message_scatter[0], "time_created"): chat_message_scatter.sort(key=lambda _: _.time_created) # Store result in context diff --git a/reme_ai/retrieve/personal/retrieve_memory_op.py b/reme_ai/retrieve/personal/retrieve_memory_op.py index a02a125e..e0d0fb2b 100644 --- a/reme_ai/retrieve/personal/retrieve_memory_op.py +++ b/reme_ai/retrieve/personal/retrieve_memory_op.py @@ -1,7 +1,14 @@ +"""Memory retrieval operation for personal memories. + +This module provides functionality to retrieve memories from a vector store +based on query similarity and score thresholds. +""" + from typing import List -from flowllm import C, BaseAsyncOp -from flowllm.schema.vector_node import VectorNode +from flowllm.core.context import C +from flowllm.core.op import BaseAsyncOp +from flowllm.core.schema import VectorNode from loguru import logger from reme_ai.schema.memory import BaseMemory, vector_node_to_memory @@ -16,6 +23,15 @@ class RetrieveMemoryOp(BaseAsyncOp): """ async def async_execute(self): + """ + Executes the memory retrieval operation. + + This method: + 1. Retrieves memories from vector store based on query similarity + 2. Removes duplicate memories based on content + 3. Filters memories by score threshold if specified + 4. Stores the retrieved memories in context metadata + """ recall_key: str = self.op_params.get("recall_key", "query") top_k: int = self.context.get("top_k", 3) @@ -23,9 +39,11 @@ class RetrieveMemoryOp(BaseAsyncOp): assert query, "query should be not empty!" workspace_id: str = self.context.workspace_id - nodes: List[VectorNode] = await self.vector_store.async_search(query=query, - workspace_id=workspace_id, - top_k=top_k) + nodes: List[VectorNode] = await self.vector_store.async_search( + query=query, + workspace_id=workspace_id, + top_k=top_k, + ) memory_list: List[BaseMemory] = [] memory_content_list: List[str] = [] for node in nodes: diff --git a/reme_ai/retrieve/personal/semantic_rank_op.py b/reme_ai/retrieve/personal/semantic_rank_op.py index 2a267d4d..2b1a0b88 100644 --- a/reme_ai/retrieve/personal/semantic_rank_op.py +++ b/reme_ai/retrieve/personal/semantic_rank_op.py @@ -1,11 +1,19 @@ +"""Semantic ranking operation for personal memories. + +This module provides functionality to rank memories semantically using LLM-based +relevance scoring to improve retrieval quality. +""" + import json import re from typing import List -from flowllm import C, BaseAsyncOp +from flowllm.core.context import C +from flowllm.core.enumeration import Role +from flowllm.core.op import BaseAsyncOp +from flowllm.core.schema import Message from loguru import logger -from reme_ai.schema import Message, Role from reme_ai.schema.memory import BaseMemory @@ -17,6 +25,7 @@ class SemanticRankOp(BaseAsyncOp): assigning scores, sorting the nodes, and storing the ranked nodes back, while logging relevant information. """ + file_path: str = __file__ async def async_execute(self): @@ -73,7 +82,14 @@ class SemanticRankOp(BaseAsyncOp): async def _semantic_rank_memories(self, query: str, memories: List[BaseMemory]) -> List[BaseMemory]: """ - Use LLM to semantically rank memories based on relevance to the query + Use LLM to semantically rank memories based on relevance to the query. + + Args: + query: User query to rank memories against + memories: List of memories to rank + + Returns: + List of memories with updated semantic scores """ if not memories: return memories @@ -84,7 +100,7 @@ class SemanticRankOp(BaseAsyncOp): # Create prompt for semantic ranking prompt = f"""Given the query: "{query}" -Please rank the following memories by their semantic relevance to the query. +Please rank the following memories by their semantic relevance to the query. Rate each memory on a scale of 0.0 to 1.0 where 1.0 is most relevant. Memories: @@ -112,10 +128,18 @@ Please respond in JSON format: @staticmethod def parse_llm_ranking_response(response: str) -> List[dict]: - """Parse LLM ranking response to extract rankings.""" + """ + Parse LLM ranking response to extract rankings. + + Args: + response: Raw LLM response string containing ranking JSON + + Returns: + List of ranking dictionaries with index and score + """ try: # Try to extract JSON blocks - json_pattern = r'```json\s*([\s\S]*?)\s*```' + json_pattern = r"```json\s*([\s\S]*?)\s*```" json_blocks = re.findall(json_pattern, response) if json_blocks: @@ -135,7 +159,16 @@ Please respond in JSON format: @staticmethod def apply_semantic_scores_to_memories(memories: List, rankings: List[dict]) -> int: - """Apply semantic ranking scores to memory objects.""" + """ + Apply semantic ranking scores to memory objects. + + Args: + memories: List of memory objects to update + rankings: List of ranking dictionaries with index and score + + Returns: + Number of memories successfully updated with scores + """ applied_count = 0 for ranking in rankings: @@ -144,21 +177,29 @@ Please respond in JSON format: if 0 <= idx < len(memories): # Set score on memory object - if hasattr(memories[idx], 'score'): + if hasattr(memories[idx], "score"): memories[idx].score = score applied_count += 1 else: # Add score as metadata if score attribute doesn't exist - if not hasattr(memories[idx], 'metadata'): + if not hasattr(memories[idx], "metadata"): memories[idx].metadata = {} - memories[idx].metadata['semantic_score'] = score + memories[idx].metadata["semantic_score"] = score applied_count += 1 return applied_count @staticmethod def format_memories_for_llm_ranking(memories: List) -> str: - """Format memories for LLM ranking input.""" + """ + Format memories for LLM ranking input. + + Args: + memories: List of memory objects to format + + Returns: + Formatted string representation of memories for LLM input + """ formatted_memories = [] for i, memory in enumerate(memories): diff --git a/reme_ai/retrieve/personal/set_query_op.py b/reme_ai/retrieve/personal/set_query_op.py index aa9617dc..f382f4a5 100644 --- a/reme_ai/retrieve/personal/set_query_op.py +++ b/reme_ai/retrieve/personal/set_query_op.py @@ -1,7 +1,14 @@ +"""Query setting operation for personal memories. + +This module provides functionality to set query and timestamp in the context +for downstream memory retrieval operations. +""" + import datetime from typing import Tuple -from flowllm import C, BaseAsyncOp +from flowllm.core.context import C +from flowllm.core.op import BaseAsyncOp from loguru import logger from reme_ai.constants.common_constants import QUERY_WITH_TS diff --git a/reme_ai/retrieve/task/__init__.py b/reme_ai/retrieve/task/__init__.py index cf6c1ded..fb4e3dd7 100644 --- a/reme_ai/retrieve/task/__init__.py +++ b/reme_ai/retrieve/task/__init__.py @@ -1,4 +1,17 @@ +"""Task memory retrieval operations module. + +This module provides operations for building queries, reranking memories, +rewriting memory context, and merging memories for task-related retrieval. +""" + from .build_query_op import BuildQueryOp from .merge_memory_op import MergeMemoryOp from .rerank_memory_op import RerankMemoryOp from .rewrite_memory_op import RewriteMemoryOp + +__all__ = [ + "BuildQueryOp", + "MergeMemoryOp", + "RerankMemoryOp", + "RewriteMemoryOp", +] diff --git a/reme_ai/retrieve/task/build_query_op.py b/reme_ai/retrieve/task/build_query_op.py index e2463b3d..d0339c6f 100644 --- a/reme_ai/retrieve/task/build_query_op.py +++ b/reme_ai/retrieve/task/build_query_op.py @@ -1,15 +1,38 @@ -from flowllm import C, BaseAsyncOp -from flowllm.utils.llm_utils import merge_messages_content -from loguru import logger +"""Query building operation module. -from reme_ai.schema import Message, Role +This module provides functionality to build retrieval queries from either +explicit query strings or conversation messages, optionally using LLM to +generate optimized queries. +""" + +from flowllm.core.context import C +from flowllm.core.enumeration import Role +from flowllm.core.op import BaseAsyncOp +from flowllm.core.schema import Message +from flowllm.core.utils import merge_messages_content +from loguru import logger @C.register_op() class BuildQueryOp(BaseAsyncOp): + """Build retrieval query from context or messages. + + This operation constructs a query string for memory retrieval. It can use + an explicit query from context, or generate one from conversation messages + using either LLM-based generation or simple message concatenation. + """ + file_path: str = __file__ async def async_execute(self): + """Execute the query building operation. + + Builds a query string from either: + 1. An explicit query in the context + 2. Conversation messages (using LLM or simple concatenation) + + Stores the built query in context.query. + """ if "query" in self.context: query = self.context.query diff --git a/reme_ai/retrieve/task/build_query_prompt.yaml b/reme_ai/retrieve/task/build_query_prompt.yaml index 73d025d3..7017cec7 100644 --- a/reme_ai/retrieve/task/build_query_prompt.yaml +++ b/reme_ai/retrieve/task/build_query_prompt.yaml @@ -1,6 +1,6 @@ query_build: | # Execution Process {execution_process} - - Read through the entire execution process to understand which part is currently being executed. + + Read through the entire execution process to understand which part is currently being executed. Generate a `query` that reflects the current state, which will later be used to search for similar problems in the database and help resolve the issue at hand. diff --git a/reme_ai/retrieve/task/merge_memory_op.py b/reme_ai/retrieve/task/merge_memory_op.py index 6d67d203..c56e4106 100644 --- a/reme_ai/retrieve/task/merge_memory_op.py +++ b/reme_ai/retrieve/task/merge_memory_op.py @@ -1,6 +1,13 @@ +"""Memory merging operation module. + +This module provides functionality to merge multiple retrieved memories +into a single formatted context string for use in LLM responses. +""" + from typing import List -from flowllm import C, BaseAsyncOp +from flowllm.core.context import C +from flowllm.core.op import BaseAsyncOp from loguru import logger from reme_ai.schema.memory import BaseMemory @@ -8,8 +15,19 @@ from reme_ai.schema.memory import BaseMemory @C.register_op() class MergeMemoryOp(BaseAsyncOp): + """Merge multiple memories into a single formatted context. + + This operation takes a list of retrieved memories and formats them into + a single context string that can be used to guide LLM responses. It includes + instructions for the LLM to consider the helpful parts from these memories. + """ async def async_execute(self): + """Execute the memory merging operation. + + Merges memories from context metadata into a formatted string with + instructions for the LLM. Stores the merged result in response.answer. + """ memory_list: List[BaseMemory] = self.context.response.metadata["memory_list"] if not memory_list: @@ -21,7 +39,9 @@ class MergeMemoryOp(BaseAsyncOp): continue content_collector.append(f"- {memory.when_to_use} {memory.content}\n") - content_collector.append("Please consider the helpful parts from these in answering the question, " - "to make the response more comprehensive and substantial.") + content_collector.append( + "Please consider the helpful parts from these in answering the question, " + "to make the response more comprehensive and substantial.", + ) self.context.response.answer = "\n".join(content_collector) logger.info(f"response.answer={self.context.response.answer}") diff --git a/reme_ai/retrieve/task/rerank_memory_op.py b/reme_ai/retrieve/task/rerank_memory_op.py index 298300b6..beb464e3 100644 --- a/reme_ai/retrieve/task/rerank_memory_op.py +++ b/reme_ai/retrieve/task/rerank_memory_op.py @@ -1,10 +1,18 @@ +"""Memory reranking operation module. + +This module provides functionality to rerank and filter retrieved memories +using LLM-based reranking and score-based filtering to select the most relevant +memories for the current task. +""" + import json import re from typing import List -from flowllm import C, BaseAsyncOp -from flowllm.enumeration.role import Role -from flowllm.schema.message import Message +from flowllm.core.context import C +from flowllm.core.enumeration import Role +from flowllm.core.op import BaseAsyncOp +from flowllm.core.schema import Message from loguru import logger from reme_ai.schema.memory import BaseMemory @@ -12,13 +20,22 @@ from reme_ai.schema.memory import BaseMemory @C.register_op() class RerankMemoryOp(BaseAsyncOp): + """Rerank and filter recalled experiences using LLM and score-based filtering. + + This operation takes recalled memories and applies multiple filtering and + ranking strategies to select the most relevant memories for the current task. + It supports LLM-based reranking and score-based filtering. """ - Rerank and filter recalled experiences using LLM and score-based filtering - """ + file_path: str = __file__ async def async_execute(self): - """Execute rerank operation""" + """Execute the memory reranking operation. + + Applies LLM-based reranking (optional) and score-based filtering (optional) + to select the top-k most relevant memories. Stores the reranked results + in the context response metadata. + """ memory_list: List[BaseMemory] = self.context.response.metadata["memory_list"] retrieval_query: str = self.context.query enable_llm_rerank = self.op_params.get("enable_llm_rerank", True) @@ -52,7 +69,15 @@ class RerankMemoryOp(BaseAsyncOp): self.context.response.metadata["memory_list"] = reranked_memories async def _llm_rerank(self, query: str, candidates: List[BaseMemory]) -> List[BaseMemory]: - """LLM-based reranking of candidate experiences""" + """LLM-based reranking of candidate experiences. + + Args: + query: The retrieval query used to rank candidates. + candidates: List of memory candidates to rerank. + + Returns: + List of memories reranked by relevance to the query. + """ if not candidates: return candidates @@ -63,7 +88,8 @@ class RerankMemoryOp(BaseAsyncOp): prompt_name="memory_rerank_prompt", query=query, candidates=candidates_text, - num_candidates=len(candidates)) + num_candidates=len(candidates), + ) response = await self.llm.achat([Message(role=Role.USER, content=prompt)]) @@ -89,7 +115,15 @@ class RerankMemoryOp(BaseAsyncOp): @staticmethod def _score_based_filter(memories: List[BaseMemory], min_score: float) -> List[BaseMemory]: - """Filter memories based on quality scores""" + """Filter memories based on quality scores. + + Args: + memories: List of memories to filter. + min_score: Minimum combined score threshold for filtering. + + Returns: + List of memories that meet the minimum score threshold. + """ filtered_memories = [] for memory in memories: @@ -110,7 +144,14 @@ class RerankMemoryOp(BaseAsyncOp): @staticmethod def _format_candidates_for_rerank(candidates: List[BaseMemory]) -> str: - """Format candidates for LLM reranking""" + """Format candidates for LLM reranking. + + Args: + candidates: List of memory candidates to format. + + Returns: + Formatted string representation of candidates for LLM evaluation. + """ formatted_candidates = [] for i, candidate in enumerate(candidates): @@ -127,10 +168,17 @@ class RerankMemoryOp(BaseAsyncOp): @staticmethod def _parse_rerank_response(response: str) -> List[int]: - """Parse LLM reranking response to extract ranked indices""" + """Parse LLM reranking response to extract ranked indices. + + Args: + response: The LLM response containing ranked indices. + + Returns: + List of indices representing the reranked order. + """ try: # Try to extract JSON format - json_pattern = r'```json\s*([\s\S]*?)\s*```' + json_pattern = r"```json\s*([\s\S]*?)\s*```" json_blocks = re.findall(json_pattern, response) if json_blocks: @@ -141,7 +189,7 @@ class RerankMemoryOp(BaseAsyncOp): return parsed # Try to extract numbers from text - numbers = re.findall(r'\b\d+\b', response) + numbers = re.findall(r"\b\d+\b", response) return [int(num) for num in numbers if int(num) < 100] # Reasonable upper bound except Exception as e: diff --git a/reme_ai/retrieve/task/rerank_memory_prompt.yaml b/reme_ai/retrieve/task/rerank_memory_prompt.yaml index 8cba9147..e5e53337 100644 --- a/reme_ai/retrieve/task/rerank_memory_prompt.yaml +++ b/reme_ai/retrieve/task/rerank_memory_prompt.yaml @@ -1,18 +1,18 @@ memory_rerank_prompt: | You are an expert AI analyst tasked with reranking retrieved experiences based on their relevance to a specific query. - + Your task is to analyze the candidates and rank them by relevance, considering: ● DIRECT RELEVANCE: How directly applicable the experience is to the current query ● SITUATION SIMILARITY: How similar the experience context is to the current situation ● ACTIONABILITY: How actionable and specific the experience is ● QUALITY: The overall quality and clarity of the experience - + # Current Query {query} - + # Candidate Experiences (Total: {num_candidates}) {candidates} - + OUTPUT FORMAT: Provide a ranked list of candidate indices (0-based) from most relevant to least relevant: ```json @@ -21,5 +21,5 @@ memory_rerank_prompt: | "reasoning": "Brief explanation of ranking rationale" }} ``` - + Note: Include ALL candidate indices in the ranking, even if some are less relevant. \ No newline at end of file diff --git a/reme_ai/retrieve/task/rewrite_memory_op.py b/reme_ai/retrieve/task/rewrite_memory_op.py index 37ad179b..708fbf0e 100644 --- a/reme_ai/retrieve/task/rewrite_memory_op.py +++ b/reme_ai/retrieve/task/rewrite_memory_op.py @@ -1,10 +1,17 @@ +"""Memory rewriting operation module. + +This module provides functionality to rewrite and format retrieved memories +into context messages that can be used by LLMs for task completion. +""" + import json import re from typing import List -from flowllm import C, BaseAsyncOp -from flowllm.enumeration.role import Role -from flowllm.schema.message import Message +from flowllm.core.context import C +from flowllm.core.enumeration import Role +from flowllm.core.op import BaseAsyncOp +from flowllm.core.schema import Message from loguru import logger from reme_ai.schema.memory import BaseMemory @@ -12,17 +19,25 @@ from reme_ai.schema.memory import BaseMemory @C.register_op() class RewriteMemoryOp(BaseAsyncOp): + """Generate and rewrite context messages from reranked experiences. + + This operation takes reranked memories and formats them into context messages + that can be used by LLMs. It optionally uses LLM-based rewriting to make + the context more relevant and actionable for the current task. """ - Generate and rewrite context messages from reranked experiences - """ + file_path: str = __file__ async def async_execute(self): - """Execute rewrite operation""" + """Execute the memory rewrite operation. + + Retrieves memories from context metadata, formats them, and optionally + rewrites them using LLM to make them more relevant for the current query. + Stores the rewritten context in the response answer field. + """ memory_list: List[BaseMemory] = self.context.response.metadata["memory_list"] query: str = self.context.query - messages: List[Message] = \ - [Message(**x) if isinstance(x, dict) else x for x in self.context.get('messages', [])] + messages: List[Message] = [Message(**x) if isinstance(x, dict) else x for x in self.context.get("messages", [])] if not memory_list: logger.info("No reranked memories to rewrite") @@ -39,7 +54,16 @@ class RewriteMemoryOp(BaseAsyncOp): self.context.response.metadata["memory_list"] = [memory.model_dump() for memory in memory_list] async def _generate_context_message(self, query: str, messages: List[Message], memories: List[BaseMemory]) -> str: - """Generate context message from retrieved memories""" + """Generate context message from retrieved memories. + + Args: + query: The current query string. + messages: List of conversation messages for context. + memories: List of retrieved memories to format. + + Returns: + Formatted context string, optionally rewritten by LLM. + """ if not memories: return "" @@ -60,7 +84,16 @@ class RewriteMemoryOp(BaseAsyncOp): return self._format_memories_for_context(memories) async def _rewrite_context(self, query: str, context_content: str, messages: List[Message]) -> str: - """LLM-based context rewriting to make experiences more relevant and actionable""" + """LLM-based context rewriting to make experiences more relevant and actionable. + + Args: + query: The current query string. + context_content: The formatted context content to rewrite. + messages: List of conversation messages for additional context. + + Returns: + Rewritten context string optimized for the current task. + """ if not context_content: return context_content @@ -72,7 +105,8 @@ class RewriteMemoryOp(BaseAsyncOp): prompt_name="memory_rewrite_prompt", current_query=query, current_context=current_context, - original_context=context_content) + original_context=context_content, + ) response = await self.llm.achat([Message(role=Role.USER, content=prompt)]) @@ -91,7 +125,14 @@ class RewriteMemoryOp(BaseAsyncOp): @staticmethod def _format_memories_for_context(memories: List[BaseMemory]) -> str: - """Format memories for context generation""" + """Format memories for context generation. + + Args: + memories: List of memories to format. + + Returns: + Formatted string containing all memories with their conditions and content. + """ formatted_memories = [] for i, memory in enumerate(memories, 1): @@ -105,7 +146,14 @@ class RewriteMemoryOp(BaseAsyncOp): @staticmethod def _extract_context(messages: List[Message]) -> str: - """Extract relevant context from messages""" + """Extract relevant context from messages. + + Args: + messages: List of conversation messages. + + Returns: + Formatted string containing recent conversation context. + """ if not messages: return "" @@ -125,10 +173,18 @@ class RewriteMemoryOp(BaseAsyncOp): @staticmethod def _parse_json_response(response: str, key: str) -> str: - """Parse JSON response to extract specific key""" + """Parse JSON response to extract specific key. + + Args: + response: The response string that may contain JSON. + key: The key to extract from the JSON object. + + Returns: + The value associated with the key, or the response string if parsing fails. + """ try: # Try to extract JSON blocks - json_pattern = r'```json\s*([\s\S]*?)\s*```' + json_pattern = r"```json\s*([\s\S]*?)\s*```" json_blocks = re.findall(json_pattern, response) if json_blocks: diff --git a/reme_ai/retrieve/task/rewrite_memory_prompt.yaml b/reme_ai/retrieve/task/rewrite_memory_prompt.yaml index d8804474..93e899b1 100644 --- a/reme_ai/retrieve/task/rewrite_memory_prompt.yaml +++ b/reme_ai/retrieve/task/rewrite_memory_prompt.yaml @@ -1,23 +1,23 @@ memory_rewrite_prompt: | You are an expert AI assistant tasked with rewriting and reorganizing context content to make it more relevant and actionable for the current task. - + Your task is to take the original context (containing multiple experiences) and rewrite it as a cohesive, task-specific guidance that directly addresses the current situation. - + REWRITING GUIDELINES: ● RELEVANCE FOCUS: Emphasize the most relevant aspects of each experience. Prioritize the most relevant experiences. Use clear, direct language. ● ACTIONABLE INSIGHTS: Extract specific, actionable guidance. Make the context immediately actionable ● COHERENT NARRATIVE: Create a flowing narrative rather than disconnected tips ● SITUATIONAL AWARENESS: Adapt the guidance to the current situation - + # Current Task/Query {current_query} - + # Current Trajectory {current_context} - + # Original Context Content (Multiple Experiences) {original_context} - + OUTPUT FORMAT: Provide the rewritten context: ```json @@ -25,7 +25,7 @@ memory_rewrite_prompt: | "rewritten_context": "A cohesive, task-specific context message that reorganizes and adapts the original experiences for the current task. This should be written as a unified guidance rather than separate experience items.", }} ``` - + Guidelines: - Rewrite as a unified, flowing guidance - Adapt terminology and examples to match the current task domain diff --git a/reme_ai/retrieve/tool/__init__.py b/reme_ai/retrieve/tool/__init__.py index c29f58ae..05af6e6e 100644 --- a/reme_ai/retrieve/tool/__init__.py +++ b/reme_ai/retrieve/tool/__init__.py @@ -1 +1,11 @@ +"""Tool memory retrieval operations module. + +This module provides operations for retrieving tool memories from a vector store +based on tool names, including formatting and matching tool memory results. +""" + from .retrieve_tool_memory_op import RetrieveToolMemoryOp + +__all__ = [ + "RetrieveToolMemoryOp", +] diff --git a/reme_ai/retrieve/tool/retrieve_tool_memory_op.py b/reme_ai/retrieve/tool/retrieve_tool_memory_op.py index 558284a9..2c27728e 100644 --- a/reme_ai/retrieve/tool/retrieve_tool_memory_op.py +++ b/reme_ai/retrieve/tool/retrieve_tool_memory_op.py @@ -1,7 +1,15 @@ +"""Tool memory retrieval operation module. + +This module provides functionality to retrieve tool memories from a vector store +based on tool names, format them into structured documents, and match them with +the requested tools. +""" + from typing import List -from flowllm import C, BaseAsyncOp -from flowllm.schema.vector_node import VectorNode +from flowllm.core.context import C +from flowllm.core.op import BaseAsyncOp +from flowllm.core.schema import VectorNode from loguru import logger from reme_ai.schema.memory import ToolMemory, vector_node_to_memory @@ -9,26 +17,51 @@ from reme_ai.schema.memory import ToolMemory, vector_node_to_memory @C.register_op() class RetrieveToolMemoryOp(BaseAsyncOp): + """Retrieves tool memories from vector store based on tool names. + + This operation searches for tool memories in the vector store using tool names, + validates that the retrieved memories match the requested tools, and formats + them into a structured document format for use in the context. + """ + file_path: str = __file__ - def __init__(self, **kwargs): - super().__init__(**kwargs) + @staticmethod + def _format_tool_memories(memories: List[ToolMemory]) -> str: + """Format tool memories into a structured document format. - def _format_tool_memories(self, memories: List[ToolMemory]) -> str: - """Format tool memories into a structured document format""" - lines = [] - lines.append(f"Retrieved {len(memories)} tool memory(ies):\n") + Args: + memories: List of ToolMemory objects to format. + + Returns: + A formatted string containing all tool memories with separators. + """ + lines = [f"Retrieved {len(memories)} tool memory(ies):\n"] for idx, memory in enumerate(memories, 1): lines.append(f"Tool: {memory.when_to_use}") lines.append(memory.content) - + if idx < len(memories): lines.append("\n---\n") return "\n".join(lines) async def async_execute(self): + """Execute the tool memory retrieval operation. + + This method: + 1. Extracts tool names from context + 2. Searches for each tool in the vector store + 3. Validates that retrieved memories match the requested tools + 4. Formats the memories into a structured document + 5. Stores the results in context response + + The operation expects 'tool_names' in the context, which should be a + comma-separated string of tool names. For each tool name, it retrieves + the top matching memory from the vector store and validates that it + matches exactly. + """ tool_names: str = self.context.get("tool_names", "") workspace_id: str = self.context.workspace_id @@ -49,7 +82,7 @@ class RetrieveToolMemoryOp(BaseAsyncOp): nodes: List[VectorNode] = await self.vector_store.async_search( query=tool_name, workspace_id=workspace_id, - top_k=1 + top_k=1, ) if nodes: @@ -59,9 +92,11 @@ class RetrieveToolMemoryOp(BaseAsyncOp): # Ensure it's a ToolMemory and when_to_use matches if isinstance(memory, ToolMemory) and memory.when_to_use == tool_name: matched_tool_memories.append(memory) - logger.info(f"Found tool_memory for tool_name={tool_name}, " - f"memory_id={memory.memory_id}, " - f"total_calls={len(memory.tool_call_results)}") + logger.info( + f"Found tool_memory for tool_name={tool_name}, " + f"memory_id={memory.memory_id}, " + f"total_calls={len(memory.tool_call_results)}", + ) else: logger.warning(f"No exact match found for tool_name={tool_name}") else: @@ -83,6 +118,8 @@ class RetrieveToolMemoryOp(BaseAsyncOp): # Log retrieval results for memory in matched_tool_memories: - logger.info(f"Retrieved tool: {memory.when_to_use}, " - f"total_calls={len(memory.tool_call_results)}, " - f"content_length={len(memory.content)}") + logger.info( + f"Retrieved tool: {memory.when_to_use}, " + f"total_calls={len(memory.tool_call_results)}, " + f"content_length={len(memory.content)}", + ) diff --git a/reme_ai/schema/__init__.py b/reme_ai/schema/__init__.py index 271b1234..70255638 100644 --- a/reme_ai/schema/__init__.py +++ b/reme_ai/schema/__init__.py @@ -1 +1,34 @@ -from flowllm.schema.message import Message, Role, Trajectory # noqa +"""Schema module for ReMe. + +This module provides data structures and schemas for memory management, +including memory types, tool call results, and conversion utilities. +""" + +from flowllm.core.enumeration import Role # noqa +from flowllm.core.schema import Message, Trajectory # noqa + +from reme_ai.schema.memory import ( + BaseMemory, + PersonalMemory, + TaskMemory, + ToolCallResult, + ToolMemory, + dict_to_memory, + vector_node_to_memory, +) + +__all__ = [ + # FlowLLM schema imports + "Message", + "Role", + "Trajectory", + # Memory classes + "BaseMemory", + "TaskMemory", + "PersonalMemory", + "ToolMemory", + "ToolCallResult", + # Utility functions + "vector_node_to_memory", + "dict_to_memory", +] diff --git a/reme_ai/schema/memory.py b/reme_ai/schema/memory.py index 27745f7b..33cfdbfa 100644 --- a/reme_ai/schema/memory.py +++ b/reme_ai/schema/memory.py @@ -1,3 +1,10 @@ +"""Memory schema definitions for ReMe. + +This module defines the core memory data structures used in the ReMe system, +including base memory classes and specialized memory types for tasks, personal +information, and tool call results. +""" + import datetime import hashlib import json @@ -5,12 +12,31 @@ from abc import ABC from typing import List from uuid import uuid4 -from flowllm.schema.vector_node import VectorNode +from flowllm.core.schema import VectorNode from mcp.types import CallToolResult, TextContent from pydantic import BaseModel, Field class BaseMemory(BaseModel, ABC): + """Base class for all memory types in the ReMe system. + + This abstract base class provides common fields and methods for all memory + types, including workspace identification, content storage, timestamps, + and conversion to/from vector nodes for storage and retrieval. + + Attributes: + workspace_id: Identifier for the workspace this memory belongs to. + memory_id: Unique identifier for this memory instance. + memory_type: Type of memory (task, personal, tool, etc.). + when_to_use: Description of when this memory should be retrieved. + content: The actual content of the memory (string or bytes). + score: Relevance score for this memory (0.0 to 1.0). + time_created: Timestamp when the memory was created. + time_modified: Timestamp when the memory was last modified. + author: Identifier of the entity that created this memory. + metadata: Additional metadata dictionary for extensibility. + """ + workspace_id: str = Field(default="") memory_id: str = Field(default_factory=lambda: uuid4().hex) memory_type: str = Field(default=...) @@ -26,90 +52,195 @@ class BaseMemory(BaseModel, ABC): metadata: dict = Field(default_factory=dict) def update_modified_time(self): + """Update the time_modified field to the current timestamp.""" self.time_modified = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") def update_metadata(self, new_metadata): + """Update the metadata dictionary with new values. + + Args: + new_metadata: Dictionary containing new metadata to replace existing metadata. + """ self.metadata = new_metadata def to_vector_node(self) -> VectorNode: + """Convert this memory instance to a VectorNode for storage. + + Returns: + VectorNode: A vector node representation of this memory. + + Raises: + NotImplementedError: Must be implemented by subclasses. + """ raise NotImplementedError @classmethod def from_vector_node(cls, node: VectorNode): + """Create a memory instance from a VectorNode. + + Args: + node: VectorNode containing memory data. + + Returns: + BaseMemory: A memory instance reconstructed from the vector node. + + Raises: + NotImplementedError: Must be implemented by subclasses. + """ raise NotImplementedError class TaskMemory(BaseMemory): + """Memory type for storing task-related information. + + TaskMemory is used to store information about tasks, including when to use + the memory and the task content itself. It extends BaseMemory with + task-specific behavior. + + Attributes: + memory_type: Always set to "task" for task memories. + """ + memory_type: str = Field(default="task") def to_vector_node(self) -> VectorNode: - return VectorNode(unique_id=self.memory_id, - workspace_id=self.workspace_id, - content=self.when_to_use, - metadata={ - "memory_type": self.memory_type, - "content": self.content, - "score": self.score, - "time_created": self.time_created, - "time_modified": self.time_modified, - "author": self.author, - "metadata": self.metadata, - }) + """Convert this TaskMemory to a VectorNode. + + Returns: + VectorNode: Vector node representation with when_to_use as content + and all other fields stored in metadata. + """ + return VectorNode( + unique_id=self.memory_id, + workspace_id=self.workspace_id, + content=self.when_to_use, + metadata={ + "memory_type": self.memory_type, + "content": self.content, + "score": self.score, + "time_created": self.time_created, + "time_modified": self.time_modified, + "author": self.author, + "metadata": self.metadata, + }, + ) @classmethod def from_vector_node(cls, node: VectorNode) -> "TaskMemory": + """Create a TaskMemory instance from a VectorNode. + + Args: + node: VectorNode containing task memory data. + + Returns: + TaskMemory: Reconstructed TaskMemory instance. + """ metadata = node.metadata.copy() - return cls(workspace_id=node.workspace_id, - memory_id=node.unique_id, - memory_type=metadata.pop("memory_type"), - when_to_use=node.content, - content=metadata.pop("content"), - score=metadata.pop("score"), - time_created=metadata.pop("time_created"), - time_modified=metadata.pop("time_modified"), - author=metadata.pop("author"), - metadata=metadata.pop("metadata", {})) + return cls( + workspace_id=node.workspace_id, + memory_id=node.unique_id, + memory_type=metadata.pop("memory_type"), + when_to_use=node.content, + content=metadata.pop("content"), + score=metadata.pop("score"), + time_created=metadata.pop("time_created"), + time_modified=metadata.pop("time_modified"), + author=metadata.pop("author"), + metadata=metadata.pop("metadata", {}), + ) class PersonalMemory(BaseMemory): + """Memory type for storing personal information and user preferences. + + PersonalMemory extends BaseMemory with fields specific to personal data, + including target information and reflection subject attributes. This is + used for storing user preferences, personal insights, and reflection data. + + Attributes: + memory_type: Always set to "personal" for personal memories. + target: Target identifier or category for this personal memory. + reflection_subject: Subject of reflection for storing reflection attributes. + """ + memory_type: str = Field(default="personal") target: str = Field(default="") reflection_subject: str = Field(default="") # For storing reflection subject attributes def to_vector_node(self) -> VectorNode: - return VectorNode(unique_id=self.memory_id, - workspace_id=self.workspace_id, - content=self.when_to_use, - metadata={ - "memory_type": self.memory_type, - "content": self.content, - "target": self.target, - "reflection_subject": self.reflection_subject, - "score": self.score, - "time_created": self.time_created, - "time_modified": self.time_modified, - "author": self.author, - "metadata": self.metadata, - }) + """Convert this PersonalMemory to a VectorNode. + + Returns: + VectorNode: Vector node representation with when_to_use as content + and all other fields including target and reflection_subject + stored in metadata. + """ + return VectorNode( + unique_id=self.memory_id, + workspace_id=self.workspace_id, + content=self.when_to_use, + metadata={ + "memory_type": self.memory_type, + "content": self.content, + "target": self.target, + "reflection_subject": self.reflection_subject, + "score": self.score, + "time_created": self.time_created, + "time_modified": self.time_modified, + "author": self.author, + "metadata": self.metadata, + }, + ) @classmethod def from_vector_node(cls, node: VectorNode) -> "PersonalMemory": + """Create a PersonalMemory instance from a VectorNode. + + Args: + node: VectorNode containing personal memory data. + + Returns: + PersonalMemory: Reconstructed PersonalMemory instance. + """ metadata = node.metadata.copy() - return cls(workspace_id=node.workspace_id, - memory_id=node.unique_id, - memory_type=metadata.pop("memory_type"), - when_to_use=node.content, - content=metadata.pop("content"), - target=metadata.pop("target", ""), - reflection_subject=metadata.pop("reflection_subject", ""), - score=metadata.pop("score"), - time_created=metadata.pop("time_created"), - time_modified=metadata.pop("time_modified"), - author=metadata.pop("author"), - metadata=metadata.pop("metadata", {})) + return cls( + workspace_id=node.workspace_id, + memory_id=node.unique_id, + memory_type=metadata.pop("memory_type"), + when_to_use=node.content, + content=metadata.pop("content"), + target=metadata.pop("target", ""), + reflection_subject=metadata.pop("reflection_subject", ""), + score=metadata.pop("score"), + time_created=metadata.pop("time_created"), + time_modified=metadata.pop("time_modified"), + author=metadata.pop("author"), + metadata=metadata.pop("metadata", {}), + ) class ToolCallResult(BaseModel): + """Represents the result of a tool invocation. + + This class stores comprehensive information about a tool call, including + inputs, outputs, performance metrics, evaluation, and deduplication hash. + + Attributes: + create_time: Timestamp when the tool was invoked. + tool_name: Name of the tool that was called. + input: Input parameters passed to the tool (dict or string). + output: Output result from the tool execution. + token_cost: Number of tokens consumed by the tool call (-1 if unknown). + success: Whether the tool invocation completed successfully. + time_cost: Time taken for the tool invocation in seconds. + summary: Brief summary of the tool call result. + evaluation: Detailed evaluation of the tool invocation. + score: Quality score from 0.0 (failure) to 1.0 (complete success). + is_summarized: Whether this tool call has been included in a summary. + call_hash: MD5 hash of input and output for deduplication. + metadata: Additional metadata dictionary. + """ + create_time: str = Field(default="", description="Time of tool invocation") tool_name: str = Field(default=..., description="Name of the tool") input: dict | str = Field(default="", description="Tool input") @@ -126,24 +257,38 @@ class ToolCallResult(BaseModel): metadata: dict = Field(default_factory=dict) def generate_hash(self) -> str: - """Generate hash value from tool input and output for deduplication""" + """Generate hash value from tool input and output for deduplication. + + Creates an MD5 hash from the combined input and output strings. + This hash is used to identify duplicate tool calls. + + Returns: + str: MD5 hash hexdigest of the combined input and output. + """ # Convert input to string if it's a dict input_str = json.dumps(self.input, sort_keys=True) if isinstance(self.input, dict) else str(self.input) - + # Combine input and output combined = f"{input_str}|{self.output}" - + # Generate MD5 hash - hash_value = hashlib.md5(combined.encode('utf-8')).hexdigest() - + hash_value = hashlib.md5(combined.encode("utf-8")).hexdigest() + return hash_value - + def ensure_hash(self): - """Ensure call_hash is set, generate if empty""" + """Ensure call_hash is set, generate if empty.""" if not self.call_hash: self.call_hash = self.generate_hash() def from_mcp_tool_result(self, tool_result: CallToolResult, max_char_len: int = None): + """Populate this instance from an MCP CallToolResult. + + Args: + tool_result: MCP CallToolResult to extract data from. + max_char_len: Optional maximum character length for output content. + If provided, output will be truncated to this length. + """ text_list = [] for content in tool_result.content: if isinstance(content, TextContent): @@ -162,28 +307,60 @@ class ToolCallResult(BaseModel): class ToolMemory(BaseMemory): + """Memory type for storing tool call execution history. + + ToolMemory extends BaseMemory to store a collection of tool call results, + allowing tracking of tool usage patterns, performance metrics, and + execution history for analysis and summarization. + + Attributes: + memory_type: Always set to "tool" for tool memories. + tool_call_results: List of ToolCallResult instances representing + historical tool invocations. + """ + memory_type: str = Field(default="tool") tool_call_results: List[ToolCallResult] = Field(default_factory=list) def to_vector_node(self) -> VectorNode: - return VectorNode(unique_id=self.memory_id, - workspace_id=self.workspace_id, - content=self.when_to_use, - metadata={ - "memory_type": self.memory_type, - "content": self.content, - "score": self.score, - "time_created": self.time_created, - "time_modified": self.time_modified, - "author": self.author, - "tool_call_results": [x.model_dump() for x in self.tool_call_results], - "metadata": self.metadata, - }) + """Convert this ToolMemory to a VectorNode. + + Returns: + VectorNode: Vector node representation with when_to_use as content + and all tool_call_results serialized in metadata. + """ + return VectorNode( + unique_id=self.memory_id, + workspace_id=self.workspace_id, + content=self.when_to_use, + metadata={ + "memory_type": self.memory_type, + "content": self.content, + "score": self.score, + "time_created": self.time_created, + "time_modified": self.time_modified, + "author": self.author, + "tool_call_results": [x.model_dump() for x in self.tool_call_results], + "metadata": self.metadata, + }, + ) def statistic(self, recent_frequency: int = 20) -> dict: - """ - Calculate statistical information for the most recent N tool calls. - Returns avg token_cost, success rate, avg time_cost, and avg score. + """Calculate statistical information for the most recent N tool calls. + + Analyzes the most recent tool calls and computes average metrics including + token cost, success rate, time cost, and quality scores. + + Args: + recent_frequency: Number of most recent tool calls to analyze. + Defaults to 20. + + Returns: + dict: Dictionary containing: + - avg_token_cost: Average token consumption (rounded to 2 decimals) + - avg_time_cost: Average execution time in seconds (rounded to 3 decimals) + - success_rate: Ratio of successful calls (rounded to 4 decimals) + - avg_score: Average quality score (rounded to 3 decimals) """ if not self.tool_call_results: return { @@ -192,54 +369,80 @@ class ToolMemory(BaseMemory): "avg_token_cost": 0.0, "success_rate": 0.0, "avg_time_cost": 0.0, - "avg_score": 0.0 + "avg_score": 0.0, } - + # Get the most recent N tool calls (or all if less than N) recent_calls = self.tool_call_results[-recent_frequency:] - total_calls = len(self.tool_call_results) + # total_calls = len(self.tool_call_results) recent_calls_count = len(recent_calls) - + # Calculate statistics total_token_cost = sum(call.token_cost for call in recent_calls if call.token_cost >= 0) valid_token_calls = [call for call in recent_calls if call.token_cost >= 0] avg_token_cost = total_token_cost / len(valid_token_calls) if valid_token_calls else 0.0 - + successful_calls = sum(1 for call in recent_calls if call.success) success_rate = successful_calls / recent_calls_count if recent_calls_count > 0 else 0.0 - + total_time_cost = sum(call.time_cost for call in recent_calls) avg_time_cost = total_time_cost / recent_calls_count if recent_calls_count > 0 else 0.0 - + total_score = sum(call.score for call in recent_calls) avg_score = total_score / recent_calls_count if recent_calls_count > 0 else 0.0 - + return { "avg_token_cost": round(avg_token_cost, 2), "avg_time_cost": round(avg_time_cost, 3), "success_rate": round(success_rate, 4), - "avg_score": round(avg_score, 3) + "avg_score": round(avg_score, 3), } @classmethod def from_vector_node(cls, node: VectorNode) -> "ToolMemory": + """Create a ToolMemory instance from a VectorNode. + + Args: + node: VectorNode containing tool memory data. + + Returns: + ToolMemory: Reconstructed ToolMemory instance with tool_call_results + deserialized from metadata. + """ metadata = node.metadata.copy() tool_call_results = [ToolCallResult(**result) for result in metadata.pop("tool_call_results", [])] - return cls(workspace_id=node.workspace_id, - memory_id=node.unique_id, - when_to_use=node.content, - memory_type=metadata.pop("memory_type"), - content=metadata.pop("content"), - score=metadata.pop("score"), - time_created=metadata.pop("time_created"), - time_modified=metadata.pop("time_modified"), - author=metadata.pop("author"), - tool_call_results=tool_call_results, - metadata=metadata.pop("metadata", {})) - + return cls( + workspace_id=node.workspace_id, + memory_id=node.unique_id, + when_to_use=node.content, + memory_type=metadata.pop("memory_type"), + content=metadata.pop("content"), + score=metadata.pop("score"), + time_created=metadata.pop("time_created"), + time_modified=metadata.pop("time_modified"), + author=metadata.pop("author"), + tool_call_results=tool_call_results, + metadata=metadata.pop("metadata", {}), + ) def vector_node_to_memory(node: VectorNode): + """Convert a VectorNode to the appropriate memory type. + + This function inspects the memory_type in the node's metadata and + reconstructs the appropriate memory subclass (TaskMemory, PersonalMemory, + or ToolMemory). + + Args: + node: VectorNode containing memory data with memory_type in metadata. + + Returns: + BaseMemory: Instance of the appropriate memory subclass based on + memory_type. + + Raises: + RuntimeError: If memory_type is not recognized or not present. + """ memory_type = node.metadata.get("memory_type") if memory_type == "task": return TaskMemory.from_vector_node(node) @@ -255,6 +458,23 @@ def vector_node_to_memory(node: VectorNode): def dict_to_memory(memory_dict: dict): + """Create a memory instance from a dictionary. + + This function creates the appropriate memory subclass based on the + memory_type field in the dictionary. Defaults to TaskMemory if + memory_type is not specified. + + Args: + memory_dict: Dictionary containing memory data with optional + memory_type field. + + Returns: + BaseMemory: Instance of the appropriate memory subclass based on + memory_type. + + Raises: + RuntimeError: If memory_type is not recognized. + """ memory_type = memory_dict.get("memory_type", "task") if memory_type == "task": return TaskMemory(**memory_dict) @@ -270,13 +490,15 @@ def dict_to_memory(memory_dict: dict): def task_main(): + """Test function for TaskMemory serialization and deserialization.""" e1 = TaskMemory( workspace_id="w_1024", memory_id="123", when_to_use="test case use", content="test content", score=0.99, - metadata={}) + metadata={}, + ) print(e1.model_dump_json(indent=2)) v1 = e1.to_vector_node() print(v1.model_dump_json(indent=2)) @@ -285,6 +507,7 @@ def task_main(): def personal_main(): + """Test function for PersonalMemory serialization and deserialization.""" p1 = PersonalMemory( workspace_id="w_2048", memory_id="456", @@ -293,7 +516,8 @@ def personal_main(): target="user_preferences", reflection_subject="learning_style", score=0.85, - metadata={"category": "user_profile"}) + metadata={"category": "user_profile"}, + ) print("PersonalMemory test:") print(p1.model_dump_json(indent=2)) v1 = p1.to_vector_node() @@ -305,6 +529,7 @@ def personal_main(): def tool_main(): + """Test function for ToolMemory serialization and deserialization.""" # Create sample tool call results tool_result1 = ToolCallResult( create_time="2025-10-15 10:30:00", @@ -315,9 +540,9 @@ def tool_main(): success=True, time_cost=0.5, evaluation="Successfully executed", - score=0.95 + score=0.95, ) - + tool_result2 = ToolCallResult( create_time="2025-10-15 10:31:00", tool_name="data_processor", @@ -327,9 +552,9 @@ def tool_main(): success=True, time_cost=1.2, evaluation="Good performance", - score=0.88 + score=0.88, ) - + t1 = ToolMemory( workspace_id="w_4096", memory_id="789", @@ -338,8 +563,9 @@ def tool_main(): content="tool execution test content", score=0.92, tool_call_results=[tool_result1, tool_result2], - metadata={"execution_context": "test_environment"}) - + metadata={"execution_context": "test_environment"}, + ) + print("ToolMemory test:") print(t1.model_dump_json(indent=2)) v1 = t1.to_vector_node() diff --git a/reme_ai/service/__init__.py b/reme_ai/service/__init__.py index e69de29b..64eea5bb 100644 --- a/reme_ai/service/__init__.py +++ b/reme_ai/service/__init__.py @@ -0,0 +1,15 @@ +"""Memory service modules for ReMe. + +This package provides memory service implementations for managing different +types of memories including task memories and personal memories. +""" + +from reme_ai.service.agentscope_runtime_memory_service import AgentscopeRuntimeMemoryService +from reme_ai.service.personal_memory_service import PersonalMemoryService +from reme_ai.service.task_memory_service import TaskMemoryService + +__all__ = [ + "AgentscopeRuntimeMemoryService", + "PersonalMemoryService", + "TaskMemoryService", +] diff --git a/reme_ai/service/agentscope_runtime_memory_service.py b/reme_ai/service/agentscope_runtime_memory_service.py index c9c3f3d5..76b50c61 100644 --- a/reme_ai/service/agentscope_runtime_memory_service.py +++ b/reme_ai/service/agentscope_runtime_memory_service.py @@ -1,18 +1,48 @@ +"""Base memory service for Agentscope runtime integration. + +This module provides the abstract base class AgentscopeRuntimeMemoryService +which defines the interface for memory services that integrate with +Agentscope runtime. Concrete implementations should inherit from this class +and implement the abstract methods. +""" + from abc import abstractmethod, ABC from typing import Optional, Dict, Any from pydantic import Field -from reme_ai.app import ReMeApp +from reme_ai.main import ReMeApp class AgentscopeRuntimeMemoryService(ABC): + """Abstract base class for memory services integrated with Agentscope runtime. + + This class provides a common interface for memory services and manages + the underlying ReMeApp instance and session-to-memory-id mappings. + Subclasses must implement the abstract methods to provide specific + memory management functionality. + """ def __init__(self): + """Initialize the memory service. + + Creates a new ReMeApp instance and initializes the session-to-memory-id + mapping dictionary. + """ self.app = ReMeApp() self.session_id_dict: dict = {} def add_session_memory_id(self, session_id: str, memory_id): + """Add a memory ID to a session's memory list. + + Associates a memory_id with a session_id by adding it to the + session's memory list. If the session doesn't exist, it will + be created. + + Args: + session_id: The session identifier. + memory_id: The memory identifier to associate with the session. + """ if session_id not in self.session_id_dict: self.session_id_dict[session_id] = [] @@ -37,21 +67,38 @@ class AgentscopeRuntimeMemoryService(ABC): """ async def __aenter__(self): - """Async context manager entry.""" + """Async context manager entry. + + Starts the service when entering an async context. + + Returns: + The service instance. + """ await self.start() return self async def __aexit__(self, exc_type, exc_val, exc_tb): - """Async context manager exit.""" + """Async context manager exit. + + Stops the service when exiting an async context. + + Args: + exc_type: Exception type if an exception occurred. + exc_val: Exception value if an exception occurred. + exc_tb: Exception traceback if an exception occurred. + + Returns: + False to propagate exceptions, True to suppress them. + """ await self.stop() return False @abstractmethod async def add_memory( - self, - user_id: str, - messages: list, - session_id: Optional[str] = None, + self, + user_id: str, + messages: list, + session_id: Optional[str] = None, ) -> None: """ Adds messages to the memory service. @@ -64,14 +111,13 @@ class AgentscopeRuntimeMemoryService(ABC): @abstractmethod async def search_memory( - self, - user_id: str, - messages: list, - filters: Optional[Dict[str, Any]] = Field( - description="Associated filters for the messages, " - "such as top_k, score etc.", - default=None, - ), + self, + user_id: str, + messages: list, + filters: Optional[Dict[str, Any]] = Field( + description="Associated filters for the messages, " "such as top_k, score etc.", + default=None, + ), ) -> list: """ Searches messages from the memory service. @@ -86,13 +132,12 @@ class AgentscopeRuntimeMemoryService(ABC): @abstractmethod async def list_memory( - self, - user_id: str, - filters: Optional[Dict[str, Any]] = Field( - description="Associated filters for the messages, " - "such as top_k, score etc.", - default=None, - ), + self, + user_id: str, + filters: Optional[Dict[str, Any]] = Field( + description="Associated filters for the messages, " "such as top_k, score etc.", + default=None, + ), ) -> list: """ Lists the memory items for a given user with filters, such as @@ -105,9 +150,9 @@ class AgentscopeRuntimeMemoryService(ABC): @abstractmethod async def delete_memory( - self, - user_id: str, - session_id: Optional[str] = None, + self, + user_id: str, + session_id: Optional[str] = None, ) -> None: """ Deletes the memory items for a given user with certain session id, diff --git a/reme_ai/service/personal_memory_service.py b/reme_ai/service/personal_memory_service.py index 30e8bc95..c97dd465 100644 --- a/reme_ai/service/personal_memory_service.py +++ b/reme_ai/service/personal_memory_service.py @@ -1,7 +1,15 @@ +"""Personal memory service for managing personalized memories. + +This module provides the PersonalMemoryService class which extends the base +AgentscopeRuntimeMemoryService to handle personal memory operations. +It supports creating, retrieving, listing, and deleting personal memories +using flow-based execution. +""" + import asyncio from typing import Optional, Dict, Any, List -from flowllm.schema.flow_response import FlowResponse +from flowllm.core.schema import FlowResponse from loguru import logger from pydantic import Field, BaseModel @@ -10,17 +18,49 @@ from reme_ai.service.agentscope_runtime_memory_service import AgentscopeRuntimeM class PersonalMemoryService(AgentscopeRuntimeMemoryService): + """Service for managing personalized memories. + + PersonalMemoryService empowers you to generate, retrieve, and share + customized memories. Leveraging advanced LLM, embedding, and vector store + technologies, it builds a comprehensive memory system with intelligent, + context- and time-aware memory management. + """ async def start(self): + """Start the personal memory service. + + Returns: + The result of starting the underlying application. + """ return await self.app.async_start() async def stop(self) -> None: + """Stop the personal memory service. + + Releases resources and stops the underlying application. + """ return await self.app.async_stop() async def health(self) -> bool: + """Check the health status of the service. + + Returns: + True if the service is healthy, False otherwise. + """ return True async def add_memory(self, user_id: str, messages: list, session_id: Optional[str] = None) -> None: + """Add personal memory from messages. + + Processes the provided messages and creates personal memories using + the summary_personal_memory flow. The created memories are associated + with the given session_id. + + Args: + user_id: The user identifier. + messages: List of messages (dict or BaseModel instances) to process. + session_id: Optional session identifier to associate with the memory. + """ new_messages: List[dict] = [] for message in messages: if isinstance(message, dict): @@ -33,8 +73,8 @@ class PersonalMemoryService(AgentscopeRuntimeMemoryService): kwargs = { "workspace_id": user_id, "trajectories": [ - {"messages": new_messages, "score": 1.0} - ] + {"messages": new_messages, "score": 1.0}, + ], } result: FlowResponse = await self.app.async_execute_flow(name="summary_personal_memory", **kwargs) @@ -44,11 +84,29 @@ class PersonalMemoryService(AgentscopeRuntimeMemoryService): self.add_session_memory_id(session_id, memory_id) logger.info(f"[personal_memory_service] user_id={user_id} session_id={session_id} add memory: {memory}") - async def search_memory(self, user_id: str, messages: list, filters: Optional[Dict[str, Any]] = Field( - description="Associated filters for the messages, " - "such as top_k, score etc.", - default=None, - )) -> list: + async def search_memory( + self, + user_id: str, + messages: list, + filters: Optional[Dict[str, Any]] = Field( + description="Associated filters for the messages, " "such as top_k, score etc.", + default=None, + ), + ) -> list: + """Search for personal memories matching the given messages. + + Searches the memory store for personal memories relevant to the provided + messages using the retrieve_personal_memory flow. The query is extracted + from the last message in the messages list. + + Args: + user_id: The user identifier. + messages: List of messages (dict or BaseModel instances) to search with. + filters: Optional filters including top_k for controlling search results. + + Returns: + List containing the search result answer. + """ new_messages: List[dict] = [] for message in messages: if isinstance(message, dict): @@ -64,7 +122,7 @@ class PersonalMemoryService(AgentscopeRuntimeMemoryService): kwargs = { "workspace_id": user_id, "query": query, - "top_k": filters.get("top_k", 1) if filters else 1 + "top_k": filters.get("top_k", 1) if filters else 1, } result: FlowResponse = await self.app.async_execute_flow(name="retrieve_personal_memory", **kwargs) @@ -72,11 +130,26 @@ class PersonalMemoryService(AgentscopeRuntimeMemoryService): return [result.answer] - async def list_memory(self, user_id: str, filters: Optional[Dict[str, Any]] = Field( - description="Associated filters for the messages, " - "such as top_k, score etc.", - default=None, - )) -> list: + async def list_memory( + self, + user_id: str, + filters: Optional[Dict[str, Any]] = Field( + description="Associated filters for the messages, " "such as top_k, score etc.", + default=None, + ), + ) -> list: + """List all personal memories for a user. + + Retrieves all personal memories associated with the given user_id + from the vector store. + + Args: + user_id: The user identifier. + filters: Optional filters (currently not used but kept for API consistency). + + Returns: + List of memory items for the user. + """ result = await self.app.async_execute_flow(name="vector_store", workspace_id=user_id, action="list") logger.info(f"[personal_memory_service] list_memory result: {result}") @@ -86,29 +159,51 @@ class PersonalMemoryService(AgentscopeRuntimeMemoryService): return result async def delete_memory(self, user_id: str, session_id: Optional[str] = None) -> None: + """Delete personal memories for a user session. + + Deletes all memories associated with the given session_id for the user. + If no session_id is provided or no memories exist for the session, + no deletion is performed. + + Args: + user_id: The user identifier. + session_id: Optional session identifier. If provided, only memories + associated with this session will be deleted. + """ delete_ids = self.session_id_dict.get(session_id, []) if not delete_ids: return - result = await self.app.async_execute_flow(name="vector_store", - workspace_id=user_id, - action="delete_ids", - memory_ids=delete_ids) + result = await self.app.async_execute_flow( + name="vector_store", + workspace_id=user_id, + action="delete_ids", + memory_ids=delete_ids, + ) result = result.metadata["action_result"] logger.info(f"[personal_memory_service] delete memory result={result}") async def main(): + """Main function for testing the PersonalMemoryService. + + Demonstrates the usage of PersonalMemoryService by adding, searching, + listing, and deleting personal memories. + """ async with PersonalMemoryService() as service: logger.info("========== start personal memory service ==========") - await service.add_memory(user_id="u_12345", - messages=[{"content": "I really enjoy playing tennis on weekends"}], - session_id="s_123456") + await service.add_memory( + user_id="u_12345", + messages=[{"content": "I really enjoy playing tennis on weekends"}], + session_id="s_123456", + ) - await service.search_memory(user_id="u_12345", - messages=[{"content": "What do I like to do for fun?"}], - filters={"top_k": 1}) + await service.search_memory( + user_id="u_12345", + messages=[{"content": "What do I like to do for fun?"}], + filters={"top_k": 1}, + ) await service.list_memory(user_id="u_12345") await service.delete_memory(user_id="u_12345", session_id="s_123456") diff --git a/reme_ai/service/task_memory_service.py b/reme_ai/service/task_memory_service.py index 319334de..a8088cb8 100644 --- a/reme_ai/service/task_memory_service.py +++ b/reme_ai/service/task_memory_service.py @@ -1,7 +1,15 @@ +"""Task memory service for managing task-oriented memories. + +This module provides the TaskMemoryService class which extends the base +AgentscopeRuntimeMemoryService to handle task-related memory operations. +It supports creating, retrieving, listing, and deleting task memories +using flow-based execution. +""" + import asyncio from typing import Optional, Dict, Any, List -from flowllm.schema.flow_response import FlowResponse +from flowllm.core.schema import FlowResponse from loguru import logger from pydantic import Field, BaseModel @@ -10,17 +18,49 @@ from reme_ai.service.agentscope_runtime_memory_service import AgentscopeRuntimeM class TaskMemoryService(AgentscopeRuntimeMemoryService): + """Service for managing task-oriented memories. + + TaskMemoryService helps efficiently manage and schedule task-related memories, + enhancing both the accuracy and efficiency of task execution. Powered by LLM + capabilities, it supports flexible creation, retrieval, update, and deletion + of memories across diverse task scenarios. + """ async def start(self): + """Start the task memory service. + + Returns: + The result of starting the underlying application. + """ return await self.app.async_start() async def stop(self) -> None: + """Stop the task memory service. + + Releases resources and stops the underlying application. + """ return await self.app.async_stop() async def health(self) -> bool: + """Check the health status of the service. + + Returns: + True if the service is healthy, False otherwise. + """ return True async def add_memory(self, user_id: str, messages: list, session_id: Optional[str] = None) -> None: + """Add task memory from messages. + + Processes the provided messages and creates task memories using + the summary_task_memory flow. The created memories are associated + with the given session_id. + + Args: + user_id: The user identifier. + messages: List of messages (dict or BaseModel instances) to process. + session_id: Optional session identifier to associate with the memory. + """ new_messages: List[dict] = [] for message in messages: if isinstance(message, dict): @@ -33,8 +73,8 @@ class TaskMemoryService(AgentscopeRuntimeMemoryService): kwargs = { "workspace_id": user_id, "trajectories": [ - {"messages": new_messages, "score": 1.0} - ] + {"messages": new_messages, "score": 1.0}, + ], } result: FlowResponse = await self.app.async_execute_flow(name="summary_task_memory", **kwargs) @@ -44,11 +84,28 @@ class TaskMemoryService(AgentscopeRuntimeMemoryService): self.add_session_memory_id(session_id, memory_id) logger.info(f"[task_memory_service] user_id={user_id} session_id={session_id} add memory: {memory}") - async def search_memory(self, user_id: str, messages: list, filters: Optional[Dict[str, Any]] = Field( - description="Associated filters for the messages, " - "such as top_k, score etc.", - default=None, - )) -> list: + async def search_memory( + self, + user_id: str, + messages: list, + filters: Optional[Dict[str, Any]] = Field( + description="Associated filters for the messages, " "such as top_k, score etc.", + default=None, + ), + ) -> list: + """Search for task memories matching the given messages. + + Searches the memory store for task memories relevant to the provided + messages using the retrieve_task_memory flow. + + Args: + user_id: The user identifier. + messages: List of messages (dict or BaseModel instances) to search with. + filters: Optional filters including top_k for controlling search results. + + Returns: + List containing the search result answer. + """ new_messages: List[dict] = [] for message in messages: if isinstance(message, dict): @@ -61,7 +118,7 @@ class TaskMemoryService(AgentscopeRuntimeMemoryService): kwargs = { "workspace_id": user_id, "messages": new_messages, - "top_k": filters.get("top_k", 1) if filters else 1 + "top_k": filters.get("top_k", 1) if filters else 1, } result: FlowResponse = await self.app.async_execute_flow(name="retrieve_task_memory", **kwargs) @@ -69,11 +126,26 @@ class TaskMemoryService(AgentscopeRuntimeMemoryService): return [result.answer] - async def list_memory(self, user_id: str, filters: Optional[Dict[str, Any]] = Field( - description="Associated filters for the messages, " - "such as top_k, score etc.", - default=None, - )) -> list: + async def list_memory( + self, + user_id: str, + filters: Optional[Dict[str, Any]] = Field( + description="Associated filters for the messages, " "such as top_k, score etc.", + default=None, + ), + ) -> list: + """List all task memories for a user. + + Retrieves all task memories associated with the given user_id + from the vector store. + + Args: + user_id: The user identifier. + filters: Optional filters (currently not used but kept for API consistency). + + Returns: + List of memory items for the user. + """ result = await self.app.async_execute_flow(name="vector_store", workspace_id=user_id, action="list") print("list_memory result:", result) @@ -83,29 +155,51 @@ class TaskMemoryService(AgentscopeRuntimeMemoryService): return result async def delete_memory(self, user_id: str, session_id: Optional[str] = None) -> None: + """Delete task memories for a user session. + + Deletes all memories associated with the given session_id for the user. + If no session_id is provided or no memories exist for the session, + no deletion is performed. + + Args: + user_id: The user identifier. + session_id: Optional session identifier. If provided, only memories + associated with this session will be deleted. + """ delete_ids = self.session_id_dict.get(session_id, []) if not delete_ids: return - result = await self.app.async_execute_flow(name="vector_store", - workspace_id=user_id, - action="delete_ids", - memory_ids=delete_ids) + result = await self.app.async_execute_flow( + name="vector_store", + workspace_id=user_id, + action="delete_ids", + memory_ids=delete_ids, + ) result = result.metadata["action_result"] logger.info(f"[task_memory_service] delete memory result={result}") async def main(): + """Main function for testing the TaskMemoryService. + + Demonstrates the usage of TaskMemoryService by adding, searching, + listing, and deleting task memories. + """ async with TaskMemoryService() as service: logger.info("========== start task memory service ==========") - await service.add_memory(user_id="u_123456", - messages=[{"content": "please use web search tool to search financial news:"}], - session_id="s_123456") + await service.add_memory( + user_id="u_123456", + messages=[{"content": "please use web search tool to search financial news:"}], + session_id="s_123456", + ) - await service.search_memory(user_id="u_123456", - messages=[{"content": "please use web search tool to search financial news"}], - filters={"top_k": 1}) + await service.search_memory( + user_id="u_123456", + messages=[{"content": "please use web search tool to search financial news"}], + filters={"top_k": 1}, + ) await service.list_memory(user_id="u_123456") await service.delete_memory(user_id="u_123456", session_id="s_123456") diff --git a/reme_ai/summary/__init__.py b/reme_ai/summary/__init__.py index 4f26e94d..0dd27ce5 100644 --- a/reme_ai/summary/__init__.py +++ b/reme_ai/summary/__init__.py @@ -1,3 +1,17 @@ +"""Summary operations module. + +This module provides summary operations for different types of memories: +- Personal memory summary operations +- Task memory summary operations +- Tool memory summary operations +""" + from . import personal from . import task from . import tool + +__all__ = [ + "personal", + "task", + "tool", +] diff --git a/reme_ai/summary/personal/__init__.py b/reme_ai/summary/personal/__init__.py index 1e2a444f..2c5533b5 100644 --- a/reme_ai/summary/personal/__init__.py +++ b/reme_ai/summary/personal/__init__.py @@ -1,8 +1,30 @@ -from .contra_repeat_op import ContraRepeatOp -from .get_observation_op import GetObservationOp -from .get_observation_with_time_op import GetObservationWithTimeOp -from .get_reflection_subject_op import GetReflectionSubjectOp -from .info_filter_op import InfoFilterOp -from .load_today_memory_op import LoadTodayMemoryOp -from .long_contra_repeat_op import LongContraRepeatOp -from .update_insight_op import UpdateInsightOp +"""Personal memory summary operations module. + +This module provides operations for processing personal memories, including: +- Filtering messages based on information content +- Extracting observations from chat messages +- Generating reflection subjects +- Updating insights based on new observations +- Detecting and handling contradictions and redundancies +- Loading today's memories for deduplication +""" + +from reme_ai.summary.personal.contra_repeat_op import ContraRepeatOp +from reme_ai.summary.personal.get_observation_op import GetObservationOp +from reme_ai.summary.personal.get_observation_with_time_op import GetObservationWithTimeOp +from reme_ai.summary.personal.get_reflection_subject_op import GetReflectionSubjectOp +from reme_ai.summary.personal.info_filter_op import InfoFilterOp +from reme_ai.summary.personal.load_today_memory_op import LoadTodayMemoryOp +from reme_ai.summary.personal.long_contra_repeat_op import LongContraRepeatOp +from reme_ai.summary.personal.update_insight_op import UpdateInsightOp + +__all__ = [ + "ContraRepeatOp", + "GetObservationOp", + "GetObservationWithTimeOp", + "GetReflectionSubjectOp", + "InfoFilterOp", + "LoadTodayMemoryOp", + "LongContraRepeatOp", + "UpdateInsightOp", +] diff --git a/reme_ai/summary/personal/contra_repeat_op.py b/reme_ai/summary/personal/contra_repeat_op.py index fc85f436..aec839b2 100644 --- a/reme_ai/summary/personal/contra_repeat_op.py +++ b/reme_ai/summary/personal/contra_repeat_op.py @@ -1,10 +1,20 @@ +"""Module for detecting and handling contradictory and repetitive memories. + +This module provides the ContraRepeatOp class which processes memory nodes +to identify and handle contradictory and repetitive information. It collects +observation memories from context, constructs prompts for language model +analysis, parses responses to detect contradictions or redundancies, and +filters the processed memories accordingly. +""" + import json import re from typing import List, Tuple -from flowllm import C, BaseAsyncOp -from flowllm.enumeration.role import Role -from flowllm.schema.message import Message +from flowllm.core.context import C +from flowllm.core.enumeration import Role +from flowllm.core.op import BaseAsyncOp +from flowllm.core.schema import Message from loguru import logger from reme_ai.schema.memory import BaseMemory @@ -22,6 +32,7 @@ class ContraRepeatOp(BaseAsyncOp): - Parses the model's response to detect contradictions or redundancies. - Filters and returns the processed memories. """ + file_path: str = __file__ async def async_execute(self): @@ -71,12 +82,16 @@ class ContraRepeatOp(BaseAsyncOp): user_name = self.context.get("user_name", "user") # Create prompt using the new pattern - system_prompt = self.prompt_format(prompt_name="contra_repeat_system", - num_obs=len(user_query_list), - user_name=user_name) + system_prompt = self.prompt_format( + prompt_name="contra_repeat_system", + num_obs=len(user_query_list), + user_name=user_name, + ) few_shot = self.prompt_format(prompt_name="contra_repeat_few_shot", user_name=user_name) - user_query = self.prompt_format(prompt_name="contra_repeat_user_query", - user_query="\n".join(user_query_list)) + user_query = self.prompt_format( + prompt_name="contra_repeat_user_query", + user_query="\n".join(user_query_list), + ) full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}" logger.info(f"contra_repeat_prompt={full_prompt}") @@ -105,7 +120,9 @@ class ContraRepeatOp(BaseAsyncOp): @staticmethod def _parse_and_filter_memories(response_text: str, memories: List[BaseMemory]) -> Tuple[ - List[BaseMemory], List[str]]: + List[BaseMemory], + List[str], + ]: """Parse LLM response and filter memories based on contradiction/containment analysis""" # Parse the response to extract judgments @@ -128,7 +145,7 @@ class ContraRepeatOp(BaseAsyncOp): continue judgment_lower = judgment.lower() - if judgment_lower in ['矛盾', 'contradiction', '被包含', 'contained']: + if judgment_lower in ["矛盾", "contradiction", "被包含", "contained"]: indices_to_remove.add(idx) deleted_memory_ids.append(memories[idx].memory_id) logger.info(f"Marking memory {idx + 1} for removal: {judgment} - {memories[idx].content[:100]}...") diff --git a/reme_ai/summary/personal/get_observation_op.py b/reme_ai/summary/personal/get_observation_op.py index 14ead179..e5d7ca85 100644 --- a/reme_ai/summary/personal/get_observation_op.py +++ b/reme_ai/summary/personal/get_observation_op.py @@ -1,8 +1,17 @@ +"""Module for generating observations from chat messages. + +This module provides the GetObservationOp class which extracts personal +observations from chat messages. It filters messages to exclude those with +time-related keywords and uses LLM-based extraction to generate structured +observation memories from the filtered messages. +""" + import re from typing import List -from flowllm import C, BaseAsyncOp -from flowllm.schema.message import Message +from flowllm.core.context import C +from flowllm.core.op import BaseAsyncOp +from flowllm.core.schema import Message from loguru import logger from reme_ai.schema.memory import BaseMemory, PersonalMemory @@ -14,6 +23,7 @@ class GetObservationOp(BaseAsyncOp): """ A specialized operation class to generate observations from chat messages using BaseAsyncOp. """ + file_path: str = __file__ async def async_execute(self): @@ -68,13 +78,17 @@ class GetObservationOp(BaseAsyncOp): user_query_list.append(f"{i + 1} {user_name}: {msg.content}") # Create prompt using the prompt format method - system_prompt = self.prompt_format(prompt_name="get_observation_system", - num_obs=len(user_query_list), - user_name=user_name) + system_prompt = self.prompt_format( + prompt_name="get_observation_system", + num_obs=len(user_query_list), + user_name=user_name, + ) few_shot = self.prompt_format(prompt_name="get_observation_few_shot", user_name=user_name) - user_query = self.prompt_format(prompt_name="get_observation_user_query", - user_query="\n".join(user_query_list), - user_name=user_name) + user_query = self.prompt_format( + prompt_name="get_observation_user_query", + user_query="\n".join(user_query_list), + user_name=user_name, + ) full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}" logger.info(f"get_observation_prompt={full_prompt}") @@ -104,8 +118,8 @@ class GetObservationOp(BaseAsyncOp): metadata={ "keywords": obs["keywords"], "source_message": filtered_messages[idx].content, - "observation_type": "personal_info" - } + "observation_type": "personal_info", + }, ) observation_memories.append(observation) logger.info(f"Created observation: {obs['content'][:50]}...") @@ -127,19 +141,21 @@ class GetObservationOp(BaseAsyncOp): # Handle both Chinese and English patterns if match[0]: # Chinese pattern idx_str, content, keywords = match[0], match[1], match[2] - else: # English pattern + else: # English pattern idx_str, content, keywords = match[3], match[4], match[5] try: idx = int(idx_str) # Skip if content indicates no meaningful observation content_lower = content.lower().strip() - if content_lower not in ['无', 'none', '', 'repeat']: - observations.append({ - "index": idx, - "content": content.strip(), - "keywords": keywords.strip() if keywords else "" - }) + if content_lower not in ["无", "none", "", "repeat"]: + observations.append( + { + "index": idx, + "content": content.strip(), + "keywords": keywords.strip() if keywords else "", + }, + ) except ValueError: logger.warning(f"Invalid index format: {idx_str}") continue diff --git a/reme_ai/summary/personal/get_observation_with_time_op.py b/reme_ai/summary/personal/get_observation_with_time_op.py index b19dfa06..e7c2b91a 100644 --- a/reme_ai/summary/personal/get_observation_with_time_op.py +++ b/reme_ai/summary/personal/get_observation_with_time_op.py @@ -1,8 +1,18 @@ +"""Module for extracting observations with time information from chat messages. + +This module provides the GetObservationWithTimeOp class which extracts +personal observations with time information from chat messages. It filters +messages to only include those with time-related keywords and uses LLM-based +extraction to generate structured observation memories with time information +from the filtered messages. +""" + import re from typing import List -from flowllm import C, BaseAsyncOp -from flowllm.schema.message import Message +from flowllm.core.context import C +from flowllm.core.op import BaseAsyncOp +from flowllm.core.schema import Message from loguru import logger from reme_ai.schema.memory import BaseMemory, PersonalMemory @@ -14,6 +24,7 @@ class GetObservationWithTimeOp(BaseAsyncOp): """ A specialized operation class to extract observations with time information from chat messages using BaseAsyncOp. """ + file_path: str = __file__ async def async_execute(self): @@ -77,13 +88,17 @@ class GetObservationWithTimeOp(BaseAsyncOp): user_query_list.append(f"{i + 1} {dt} {user_name}{colon}{msg.content}") # Create prompt using the prompt format method - system_prompt = self.prompt_format(prompt_name="get_observation_with_time_system", - num_obs=len(user_query_list), - user_name=user_name) + system_prompt = self.prompt_format( + prompt_name="get_observation_with_time_system", + num_obs=len(user_query_list), + user_name=user_name, + ) few_shot = self.prompt_format(prompt_name="get_observation_with_time_few_shot", user_name=user_name) - user_query = self.prompt_format(prompt_name="get_observation_with_time_user_query", - user_query="\n".join(user_query_list), - user_name=user_name) + user_query = self.prompt_format( + prompt_name="get_observation_with_time_user_query", + user_query="\n".join(user_query_list), + user_name=user_name, + ) full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}" logger.info(f"get_observation_with_time_prompt={full_prompt}") @@ -114,8 +129,8 @@ class GetObservationWithTimeOp(BaseAsyncOp): "keywords": obs["keywords"], "time_info": obs.get("time_info", ""), "source_message": filtered_messages[idx].content, - "observation_type": "personal_info_with_time" - } + "observation_type": "personal_info_with_time", + }, ) observation_memories.append(observation) logger.info(f"Created observation with time: {obs['content'][:50]}...") @@ -135,8 +150,12 @@ class GetObservationWithTimeOp(BaseAsyncOp): """Parse observation with time response to extract structured data""" # Pattern to match both Chinese and English observation formats with time information # Chinese: 信息:<1> <时间信息或不输出> <明确的重要信息或"无"> <关键词> - # English: Information: <1>