Merge pull request #32 from agentscope-ai/dev_2.0

update to 0.2.0.0
This commit is contained in:
jinliyl 2025-11-11 00:11:59 +08:00 • committed by GitHub
commit fba0a00802
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
133 changed files with 5352 additions and 3158 deletions

85
.pre-commit-config.yaml Normal file
View file

@ -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, .]

View file

@ -3,8 +3,8 @@
</p>
<p align="center">
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/python-3.12+-blue" alt="Python Version"></a>
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/pypi-v0.1.10.x-blue?logo=pypi" alt="PyPI Version"></a>
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/python-3.10+-blue" alt="Python Version"></a>
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/pypi-0.2.0.0-blue?logo=pypi" alt="PyPI Version"></a>
<a href="./LICENSE"><img src="https://img.shields.io/badge/license-Apache--2.0-black" alt="License"></a>
<a href="https://github.com/agentscope-ai/ReMe"><img src="https://img.shields.io/github/stars/modelscope/ReMe?style=social" alt="GitHub Stars"></a>
</p>
@ -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).

View file

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

View file

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

View file

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

View file

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

View file

@ -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": {<execution_results>}, 'tool_call_id': 'chatcmpl-tool-xxx'}]}
# <execution_results>: 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)}")

View file

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

View file

@ -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 <tools></tools> XML tags:\n<tools>"
for tool in tools:
tool_prompt += "\n" + json.dumps(tool)
tool_prompt += "\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\n</tool_call>"
tool_prompt += '\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{"name": <function-name>, "arguments": <args-json-object>}\n</tool_call>'
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")
print(f"Processed {len(results)} groups")

View file

@ -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 <tools></tools> XML tags:\n<tools>"
for tool in tools:
tool_prompt += "\n" + json.dumps(tool)
tool_prompt += "\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\n</tool_call>"
tool_prompt += '\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{"name": <function-name>, "arguments": <args-json-object>}\n</tool_call>'
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")
print(f"Processed {len(results)} groups")

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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']}")
# 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']}")

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -1,4 +1,4 @@
# AppWorld
# AppWorld
Experiment Quick Start Guide
This guide helps you quickly set up and run AppWorld experiments with ReMe integration.

View file

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

View file

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

View file

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

View file

@ -16,8 +16,8 @@ kernelspec:
<em>Remember Me, Refine Me.</em>
<div class="flex justify-center space-x-3">
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/python-3.12+-blue" alt="Python Version"></a>
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/pypi-v0.1.10.7-blue?logo=pypi" alt="PyPI Version"></a>
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/python-3.10+-blue" alt="Python Version"></a>
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/pypi-0.2.0.0-blue?logo=pypi" alt="PyPI Version"></a>
<a href="./LICENSE"><img src="https://img.shields.io/badge/license-Apache--2.0-black" alt="License"></a>
<a href="https://github.com/agentscope-ai/ReMe"><img src="https://img.shields.io/github/stars/modelscope/ReMe?style=social" alt="GitHub Stars"></a>
</div>

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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: <backend_name> # 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.

View file

@ -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/*
[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/*

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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 段的全面答案。包括多个角度、详细解释和相关上下文。
现在生成搜索结果内容:

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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}.",
}

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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> <Time information or do not output> <Clear important information or "None"> <Keywords>
pattern = r"信息:<(\d+)>\s*<([^<>]*)>\s*<([^<>]+)>\s*<([^<>]*)>|Information:\s*<(\d+)>\s*<([^<>]*)>\s*<([^<>]+)>\s*<([^<>]*)>"
# English: Information: <1> <Time information or do not output>
# <Clear important information or "None"> <Keywords>
pattern = (
r"信息:<(\d+)>\s*<([^<>]*)>\s*<([^<>]+)>\s*<([^<>]*)>|"
r"Information:\s*<(\d+)>\s*<([^<>]*)>\s*<([^<>]+)>\s*<([^<>]*)>"
)
matches = re.findall(pattern, response_text, re.IGNORECASE | re.MULTILINE)
observations = []
@ -144,20 +163,22 @@ class GetObservationWithTimeOp(BaseAsyncOp):
# Handle both Chinese and English patterns
if match[0]: # Chinese pattern
idx_str, time_info, content, keywords = match[0], match[1], match[2], match[3]
else: # English pattern
else: # English pattern
idx_str, time_info, content, keywords = match[4], match[5], match[6], match[7]
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,
"time_info": time_info.strip() if time_info else "",
"content": content.strip(),
"keywords": keywords.strip() if keywords else ""
})
if content_lower not in ["无", "none", "", "repeat"]:
observations.append(
{
"index": idx,
"time_info": time_info.strip() if time_info else "",
"content": content.strip(),
"keywords": keywords.strip() if keywords else "",
},
)
except ValueError:
logger.warning(f"Invalid index format: {idx_str}")
continue

View file

@ -1,7 +1,16 @@
"""Module for generating reflection subjects from personal memories.
This module provides the GetReflectionSubjectOp class which retrieves
unreflected memory nodes, generates reflection prompts with current insights,
invokes an LLM for fresh insights, parses the LLM responses, forms new
insight nodes, and updates memory statuses accordingly.
"""
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 GetReflectionSubjectOp(BaseAsyncOp):
generating reflection prompts with current insights, invoking an LLM for fresh insights,
parsing the LLM responses, forming new insight nodes, and updating memory statuses accordingly.
"""
file_path: str = __file__
def new_insight_memory(self, insight_content: str, target: str) -> PersonalMemory:
@ -35,8 +45,8 @@ class GetReflectionSubjectOp(BaseAsyncOp):
author=getattr(self.llm, "model_name", "system"),
metadata={
"insight_type": "reflection_subject",
"memory_type": "personal_topic"
}
"memory_type": "personal_topic",
},
)
async def async_execute(self):
@ -68,15 +78,16 @@ class GetReflectionSubjectOp(BaseAsyncOp):
# Extract existing insight subjects to avoid duplication
existing_subjects = []
if existing_insights:
existing_subjects = [memory.content for memory in existing_insights if
hasattr(memory, 'content') and memory.content]
existing_subjects = [
memory.content for memory in existing_insights if hasattr(memory, "content") and memory.content
]
logger.info(f"Found {len(existing_subjects)} existing insight subjects")
# Prepare memory content for LLM analysis
memory_contents = []
for memory in personal_memories:
if hasattr(memory, 'content') and memory.content.strip():
if hasattr(memory, "content") and memory.content.strip():
memory_contents.append(memory.content.strip())
if not memory_contents:
@ -86,24 +97,32 @@ class GetReflectionSubjectOp(BaseAsyncOp):
# Generate reflection subjects using LLM
insight_memories = await self._generate_reflection_subjects(
memory_contents, existing_subjects, user_name, reflect_num_questions
memory_contents,
existing_subjects,
user_name,
reflect_num_questions,
)
# Store results in context
self.context.response.metadata["insight_memories"] = insight_memories
logger.info(f"Generated {len(insight_memories)} new reflection subject memories")
async def _generate_reflection_subjects(self, memory_contents: List[str], existing_subjects: List[str],
user_name: str, num_questions: int) -> List[BaseMemory]:
async def _generate_reflection_subjects(
self,
memory_contents: List[str],
existing_subjects: List[str],
user_name: str,
num_questions: int,
) -> List[BaseMemory]:
"""
Generate new reflection subjects using LLM analysis of memory contents.
Args:
memory_contents: List of memory content strings
existing_subjects: List of already existing subject strings
user_name: Target username
num_questions: Maximum number of new subjects to generate
Returns:
List of PersonalMemory objects representing new reflection subjects
"""
@ -111,17 +130,17 @@ class GetReflectionSubjectOp(BaseAsyncOp):
system_prompt = self.prompt_format(
prompt_name="get_reflection_subject_system",
user_name=user_name,
num_questions=num_questions
num_questions=num_questions,
)
few_shot = self.prompt_format(
prompt_name="get_reflection_subject_few_shot",
user_name=user_name
user_name=user_name,
)
user_query = self.prompt_format(
prompt_name="get_reflection_subject_user_query",
user_name=user_name,
exist_keys=", ".join(existing_subjects) if existing_subjects else "None",
user_query="\n".join(memory_contents)
user_query="\n".join(memory_contents),
)
full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}"
@ -140,7 +159,7 @@ class GetReflectionSubjectOp(BaseAsyncOp):
for subject in new_subjects:
insight_memory = self.new_insight_memory(
insight_content=subject,
target=user_name
target=user_name,
)
insight_memories.append(insight_memory)
logger.info(f"Created reflection subject: {subject}")
@ -151,7 +170,15 @@ class GetReflectionSubjectOp(BaseAsyncOp):
return await self.llm.achat(messages=[Message(content=full_prompt)], callback_fn=parse_reflection_response)
def get_language_value(self, value_dict: dict):
"""Get language-specific value from dictionary"""
"""Get language-specific value from dictionary.
Args:
value_dict: Dictionary mapping language codes to values
Returns:
The value corresponding to the current language, or the English
value as fallback if the current language is not found
"""
return value_dict.get(self.language, value_dict.get("en"))
@staticmethod
@ -161,19 +188,20 @@ class GetReflectionSubjectOp(BaseAsyncOp):
existing_subjects = []
# Split response into lines and clean up
lines = response_text.strip().split('\n')
lines = response_text.strip().split("\n")
subjects = []
for line in lines:
line = line.strip()
# Skip empty lines, "None" responses, and existing subjects
if (line and
line not in ['无', 'None', ''] and
line not in existing_subjects and
not line.startswith('新增') and # Skip Chinese header
not line.startswith('New ') and # Skip English header
len(line) > 1): # Skip single character responses
subjects.append(line)
# Check basic validity first
if not line or line in ["无", "None", ""] or len(line) <= 1:
continue
# Check if it's a header or duplicate
is_header = line.startswith("新增") or line.startswith("New ")
if is_header or line in existing_subjects:
continue
subjects.append(line)
logger.info(f"Parsed {len(subjects)} new reflection subjects from response")
return subjects

View file

@ -1,8 +1,17 @@
"""Module for filtering messages based on information content scores.
This module provides the InfoFilterOp class which filters chat messages by
retaining only those that include significant information about the user.
It uses LLM-based scoring to evaluate the information content of messages
and filters them based on configurable score thresholds.
"""
import re
from typing import List
from flowllm import C, BaseAsyncOp
from flowllm.schema.message import Message, Trajectory
from flowllm.core.context import C
from flowllm.core.op import BaseAsyncOp
from flowllm.core.schema import Message, Trajectory
from loguru import logger
from reme_ai.schema.memory import PersonalMemory
@ -14,6 +23,7 @@ class InfoFilterOp(BaseAsyncOp):
A specialized operation class to filter messages based on information content scores using BaseAsyncOp.
This filters chat messages by retaining only those that include significant information about the user.
"""
file_path: str = __file__
async def async_execute(self):
@ -36,7 +46,7 @@ class InfoFilterOp(BaseAsyncOp):
user_name = self.context.get("user_name", "user")
# Filter and process messages
info_messages = self._filter_and_process_messages(messages, user_name, info_filter_msg_max_size)
info_messages = self._filter_and_process_messages(messages, info_filter_msg_max_size)
if not info_messages:
logger.warning("No messages left after filtering")
self.context.messages = []
@ -52,7 +62,7 @@ class InfoFilterOp(BaseAsyncOp):
logger.info(f"Filtered to {len(filtered_memories)} high-information messages")
@staticmethod
def _filter_and_process_messages(messages: List[Message], user_name: str, max_size: int) -> List[Message]:
def _filter_and_process_messages(messages: List[Message], max_size: int) -> List[Message]:
"""Filter and process messages for information filtering"""
info_messages = []
@ -60,7 +70,7 @@ class InfoFilterOp(BaseAsyncOp):
# Ensure metadata exists
# Skip memorized messages
if msg.metadata.get('memorized', False):
if msg.metadata.get("memorized", False):
continue
# Only process messages from the target user
@ -68,7 +78,7 @@ class InfoFilterOp(BaseAsyncOp):
# if role_name and role_name != user_name:
# continue
elif msg.role.value != "user":
if msg.role.value != "user":
continue
# Truncate long messages
@ -81,8 +91,12 @@ class InfoFilterOp(BaseAsyncOp):
logger.info(f"Filtered messages from {len(messages)} to {len(info_messages)}")
return info_messages
async def _filter_messages_with_llm(self, info_messages: List[Message], user_name: str, preserved_scores: str) -> List[
PersonalMemory]:
async def _filter_messages_with_llm(
self,
info_messages: List[Message],
user_name: str,
preserved_scores: str,
) -> List[PersonalMemory]:
"""Filter messages using LLM to score information content"""
# Build prompt for information filtering
@ -92,12 +106,16 @@ class InfoFilterOp(BaseAsyncOp):
user_query_list.append(f"{i + 1} {user_name}{colon} {msg.content}")
# Create prompt using the prompt format method
system_prompt = self.prompt_format(prompt_name="info_filter_system",
batch_size=len(info_messages),
user_name=user_name)
system_prompt = self.prompt_format(
prompt_name="info_filter_system",
batch_size=len(info_messages),
user_name=user_name,
)
few_shot = self.prompt_format(prompt_name="info_filter_few_shot", user_name=user_name)
user_query = self.prompt_format(prompt_name="info_filter_user_query",
user_query="\n".join(user_query_list))
user_query = self.prompt_format(
prompt_name="info_filter_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"info_filter_prompt={full_prompt}")
@ -134,11 +152,11 @@ class InfoFilterOp(BaseAsyncOp):
metadata={
"info_score": score,
"filter_type": "info_content",
"original_message_time": getattr(message, 'time_created', None),
"original_message_time": getattr(message, "time_created", None),
"role_name": message.metadata.pop("role_name", user_name),
"memorized": True,
**message.metadata # Include all original metadata
}
**message.metadata, # Include all original metadata
},
)
filtered_memories.append(memory)
logger.info(f"Info filter: kept message with score {score}: {message.content[:50]}...")

View file

@ -1,7 +1,16 @@
"""Module for loading today's memories from vector store.
This module provides the LoadTodayMemoryOp class which loads memories from
the current date for deduplication purposes. It focuses specifically on
retrieving and deduplicating memories from the current date using vector
store search with date filtering.
"""
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
@ -14,12 +23,13 @@ class LoadTodayMemoryOp(BaseAsyncOp):
Operation to load today's memories from vector store for deduplication.
Focuses specifically on retrieving and deduplicating memories from the current date.
"""
file_path: str = __file__
async def async_execute(self):
"""
Load today's memories from vector store and perform deduplication.
This operation:
1. Retrieves memories from today using vector store search
2. Converts vector nodes to memory objects
@ -50,12 +60,12 @@ class LoadTodayMemoryOp(BaseAsyncOp):
async def _retrieve_today_memories(self, workspace_id: str, user_name: str, top_k: int) -> List[BaseMemory]:
"""
Retrieve memories from today using vector store with date filtering.
Args:
workspace_id: Workspace identifier
user_name: Target username
top_k: Maximum number of memories to retrieve
Returns:
List of today's memories
"""
@ -70,7 +80,7 @@ class LoadTodayMemoryOp(BaseAsyncOp):
filter_dict = {
"memory_type": "personal",
"target": user_name,
"created_date": today_date
"created_date": today_date,
}
# Search vector store with date filter
@ -78,7 +88,8 @@ class LoadTodayMemoryOp(BaseAsyncOp):
query=" ",
workspace_id=workspace_id,
top_k=top_k,
filter_dict=filter_dict)
filter_dict=filter_dict,
)
logger.info(f"Vector store returned {len(nodes)} nodes for today")
@ -96,10 +107,10 @@ class LoadTodayMemoryOp(BaseAsyncOp):
def _convert_nodes_to_memories(nodes: List[VectorNode]) -> List[BaseMemory]:
"""
Convert vector nodes to memory objects.
Args:
nodes: List of vector nodes from vector store
Returns:
List of converted memory objects
"""

View file

@ -1,9 +1,19 @@
"""Module for handling contradictions and redundancies in long conversations.
This module provides the LongContraRepeatOp class which manages and updates
memory entries within a conversation scope by identifying and handling
contradictions or redundancies. It extends BaseAsyncOp to provide specialized
functionality for long conversations with potential contradictory or
repetitive statements.
"""
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, PersonalMemory
@ -17,12 +27,13 @@ class LongContraRepeatOp(BaseAsyncOp):
specialized functionality for long conversations with potential contradictory
or repetitive statements.
"""
file_path: str = __file__
async def async_execute(self):
"""
Analyze memories for contradictions and redundancies, resolving conflicts.
Process:
1. Get updated insight memories from previous operation
2. Check for contradictions and redundancies among memories
@ -51,7 +62,7 @@ class LongContraRepeatOp(BaseAsyncOp):
sorted_memories = sorted(
updated_insights,
key=lambda x: x.time_created,
reverse=True
reverse=True,
)[:max_memories_to_process]
if len(sorted_memories) <= 1:
@ -71,10 +82,10 @@ class LongContraRepeatOp(BaseAsyncOp):
async def _analyze_and_resolve_conflicts(self, memories: List[BaseMemory]) -> List[BaseMemory]:
"""
Analyze memories for contradictions and redundancies using LLM.
Args:
memories: List of memories to analyze
Returns:
List of filtered memories with conflicts resolved
"""
@ -89,15 +100,15 @@ class LongContraRepeatOp(BaseAsyncOp):
system_prompt = self.prompt_format(
prompt_name="long_contra_repeat_system",
num_obs=len(memory_texts),
user_name=user_name
user_name=user_name,
)
few_shot = self.prompt_format(
prompt_name="long_contra_repeat_few_shot",
user_name=user_name
user_name=user_name,
)
user_query = self.prompt_format(
prompt_name="long_contra_repeat_user_query",
user_query="\n".join(memory_texts)
user_query="\n".join(memory_texts),
)
full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}"
@ -139,7 +150,7 @@ class LongContraRepeatOp(BaseAsyncOp):
memory = memories[memory_idx]
judgment_lower = judgment.lower()
if judgment_lower in ['矛盾', 'contradiction']:
if judgment_lower in ["矛盾", "contradiction"]:
# For contradictory memories, either modify content or mark for removal
if modified_content.strip():
# Create new memory with modified content
@ -147,9 +158,9 @@ class LongContraRepeatOp(BaseAsyncOp):
workspace_id=memory.workspace_id,
memory_id=memory.memory_id,
content=modified_content.strip(),
target=memory.target if hasattr(memory, 'target') else user_name,
target=memory.target if hasattr(memory, "target") else user_name,
author=memory.author,
metadata={**memory.metadata, 'modified_by': 'long_contra_repeat'}
metadata={**memory.metadata, "modified_by": "long_contra_repeat"},
)
modified_memory.update_time_modified()
filtered_memories.append(modified_memory)
@ -158,7 +169,7 @@ class LongContraRepeatOp(BaseAsyncOp):
# Remove contradictory memory without modification
logger.info(f"Removing contradictory memory {idx}: {memory.content[:50]}...")
elif judgment_lower in ['被包含', 'contained']:
elif judgment_lower in ["被包含", "contained"]:
# Remove contained/redundant memories
logger.info(f"Removing contained memory {idx}: {memory.content[:50]}...")
@ -179,7 +190,15 @@ class LongContraRepeatOp(BaseAsyncOp):
return filtered_memories
def get_language_value(self, value_dict: dict):
"""Get language-specific value from dictionary"""
"""Get language-specific value from dictionary.
Args:
value_dict: Dictionary mapping language codes to values
Returns:
The value corresponding to the current language, or the English
value as fallback if the current language is not found
"""
return value_dict.get(self.language, value_dict.get("en"))
@staticmethod
@ -188,7 +207,10 @@ class LongContraRepeatOp(BaseAsyncOp):
# Pattern to match both Chinese and English judgment formats
# Chinese: 判断:<序号> <矛盾|被包含|无> <修改后的内容>
# English: Judgment: <Index> <Contradiction|Contained|None> <Modified content>
pattern = r"判断:<(\d+)>\s*<(矛盾|被包含|无)>\s*<([^<>]*)>|Judgment:\s*<(\d+)>\s*<(Contradiction|Contained|None)>\s*<([^<>]*)>"
pattern = (
r"判断:<(\d+)>\s*<(矛盾|被包含|无)>\s*<([^<>]*)>|"
r"Judgment:\s*<(\d+)>\s*<(Contradiction|Contained|None)>\s*<([^<>]*)>"
)
matches = re.findall(pattern, response_text, re.IGNORECASE | re.MULTILINE)
judgments = []

View file

@ -1,8 +1,17 @@
"""Module for updating personal insight memories based on new observations.
This module provides the UpdateInsightOp class which updates insight values
in a memory system by filtering insight nodes based on their association with
observed nodes, utilizing a ranking model to prioritize them, generating
refreshed insights via an LLM, and managing node statuses and content updates.
"""
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 PersonalMemory
@ -15,12 +24,13 @@ class UpdateInsightOp(BaseAsyncOp):
based on their association with observed nodes, utilizes a ranking model to prioritize them,
generates refreshed insights via an LLM, and manages node statuses and content updates.
"""
file_path: str = __file__
async def async_execute(self):
"""
Update insight values based on new observation memories.
Process:
1. Get insight subjects and personal memories from context
2. Find relevant observations for each insight subject
@ -50,7 +60,9 @@ class UpdateInsightOp(BaseAsyncOp):
# Score and filter insights based on relevance to observations
scored_insights = self._score_insights_by_relevance(
insight_memories, personal_memories, update_insight_threshold
insight_memories,
personal_memories,
update_insight_threshold,
)
if not scored_insights:
@ -64,9 +76,11 @@ class UpdateInsightOp(BaseAsyncOp):
# Update each selected insight
updated_insights = []
for insight_memory, relevance_score, relevant_observations in top_insights:
for insight_memory, _relevance_score, relevant_observations in top_insights:
updated_insight = await self._update_insight_with_observations(
insight_memory, relevant_observations, user_name
insight_memory,
relevant_observations,
user_name,
)
if updated_insight:
updated_insights.append(updated_insight)
@ -75,17 +89,20 @@ class UpdateInsightOp(BaseAsyncOp):
self.context.response.metadata["updated_insight_memories"] = updated_insights
logger.info(f"Successfully updated {len(updated_insights)} insight memories")
def _score_insights_by_relevance(self, insight_memories: List[PersonalMemory],
observation_memories: List[PersonalMemory],
threshold: float) -> List[tuple]:
def _score_insights_by_relevance(
self,
insight_memories: List[PersonalMemory],
observation_memories: List[PersonalMemory],
threshold: float,
) -> List[tuple]:
"""
Score insight memories based on relevance to observation memories.
Args:
insight_memories: List of insight memories to score
observation_memories: List of observation memories for comparison
threshold: Minimum relevance score threshold
Returns:
List[tuple]: List of (insight_memory, relevance_score, relevant_observations)
"""
@ -95,13 +112,15 @@ class UpdateInsightOp(BaseAsyncOp):
relevant_observations = []
max_relevance = 0.0
insight_subject = getattr(insight_memory, 'reflection_subject', '') or insight_memory.content
insight_subject = getattr(insight_memory, "reflection_subject", "") or insight_memory.content
insight_keywords = set(insight_memory.content.lower().split())
# Find observations relevant to this insight
for obs_memory in observation_memories:
relevance_score = self._calculate_relevance_score(
insight_memory, obs_memory, insight_keywords
insight_memory,
obs_memory,
insight_keywords,
)
if relevance_score >= threshold:
@ -112,18 +131,36 @@ class UpdateInsightOp(BaseAsyncOp):
if relevant_observations:
scored_insights.append((insight_memory, max_relevance, relevant_observations))
logger.info(
f"Insight '{insight_subject[:40]}...' scored {max_relevance:.3f} with {len(relevant_observations)} observations"
f"Insight '{insight_subject[:40]}...' scored {max_relevance:.3f} "
f"with {len(relevant_observations)} observations",
)
return scored_insights
@staticmethod
def _calculate_relevance_score(insight_memory: PersonalMemory,
obs_memory: PersonalMemory, insight_keywords: set) -> float:
"""Calculate relevance score between insight and observation memory"""
def _calculate_relevance_score(
insight_memory: PersonalMemory,
obs_memory: PersonalMemory,
insight_keywords: set,
) -> float:
"""Calculate relevance score between insight and observation memory.
The relevance score is calculated based on:
- High relevance (0.9) if both memories share the same reflection subject
- Medium relevance based on keyword overlap using Jaccard similarity
Args:
insight_memory: The insight memory to compare
obs_memory: The observation memory to compare against
insight_keywords: Set of keywords extracted from insight memory content
Returns:
float: Relevance score between 0.0 and 1.0, where higher values
indicate greater relevance
"""
# High relevance for same reflection subject
insight_subject = getattr(insight_memory, 'reflection_subject', '')
obs_subject = getattr(obs_memory, 'reflection_subject', '')
insight_subject = getattr(insight_memory, "reflection_subject", "")
obs_subject = getattr(obs_memory, "reflection_subject", "")
if insight_subject and obs_subject and insight_subject == obs_subject:
return 0.9
@ -135,22 +172,26 @@ class UpdateInsightOp(BaseAsyncOp):
return intersection / union if union > 0 else 0.0
async def _update_insight_with_observations(self, insight_memory: PersonalMemory,
relevant_observations: List[PersonalMemory],
user_name: str) -> PersonalMemory:
async def _update_insight_with_observations(
self,
insight_memory: PersonalMemory,
relevant_observations: List[PersonalMemory],
user_name: str,
) -> PersonalMemory:
"""
Update a single insight memory based on relevant observations using LLM.
Args:
insight_memory: The insight memory to update
relevant_observations: List of relevant observation memories
user_name: The target username
Returns:
PersonalMemory: Updated insight memory or None if update failed
"""
logger.info(
f"Updating insight: {insight_memory.content[:50]}... with {len(relevant_observations)} observations")
f"Updating insight: {insight_memory.content[:50]}... with {len(relevant_observations)} observations",
)
# Build observation context
observation_texts = [obs.content for obs in relevant_observations]
@ -161,10 +202,12 @@ class UpdateInsightOp(BaseAsyncOp):
system_prompt = self.prompt_format(prompt_name="update_insight_system", user_name=user_name)
few_shot = self.prompt_format(prompt_name="update_insight_few_shot", user_name=user_name)
user_query = self.prompt_format(prompt_name="update_insight_user_query",
user_query="\n".join(observation_texts),
insight_key=insight_key,
insight_key_value=insight_key_value)
user_query = self.prompt_format(
prompt_name="update_insight_user_query",
user_query="\n".join(observation_texts),
insight_key=insight_key,
insight_key_value=insight_key_value,
)
full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}"
logger.info(f"update_insight_prompt={full_prompt}")
@ -177,7 +220,7 @@ class UpdateInsightOp(BaseAsyncOp):
# Parse the response to extract updated insight
updated_content = UpdateInsightOp.parse_update_insight_response(response_text, self.language)
if not updated_content or updated_content.lower() in ['无', 'none', '']:
if not updated_content or updated_content.lower() in ["无", "none", ""]:
logger.info(f"No update needed for insight: {insight_memory.content[:50]}...")
return insight_memory
@ -198,8 +241,8 @@ class UpdateInsightOp(BaseAsyncOp):
**insight_memory.metadata,
"updated_by": "update_insight_op",
"original_content": insight_memory.content,
"update_reason": "integrated_new_observations"
}
"update_reason": "integrated_new_observations",
},
)
updated_insight.update_time_modified()

View file

@ -1,3 +1,10 @@
"""Task memory operations module.
This module provides various operations for extracting, processing, and managing
task memories from trajectories, including success/failure extraction, comparative
analysis, deduplication, and validation.
"""
from .comparative_extraction_op import ComparativeExtractionOp
from .failure_extraction_op import FailureExtractionOp
from .memory_deduplication_op import MemoryDeduplicationOp
@ -7,3 +14,15 @@ from .simple_summary_op import SimpleSummaryOp
from .success_extraction_op import SuccessExtractionOp
from .trajectory_preprocess_op import TrajectoryPreprocessOp
from .trajectory_segmentation_op import TrajectorySegmentationOp
__all__ = [
"ComparativeExtractionOp",
"FailureExtractionOp",
"MemoryDeduplicationOp",
"MemoryValidationOp",
"SimpleComparativeSummaryOp",
"SimpleSummaryOp",
"SuccessExtractionOp",
"TrajectoryPreprocessOp",
"TrajectorySegmentationOp",
]

View file

@ -1,8 +1,15 @@
"""Comparative extraction operation for task memory generation.
This module provides operations to extract comparative task memories by comparing
different trajectories with varying scores or success/failure outcomes.
"""
from typing import List, Tuple, Optional
from flowllm import C, BaseAsyncOp
from flowllm.enumeration.role import Role
from flowllm.schema.message import Message as FlowMessage
from flowllm.core.context import C
from flowllm.core.enumeration import Role
from flowllm.core.op import BaseAsyncOp
from flowllm.core.schema import Message as FlowMessage
from loguru import logger
from reme_ai.schema import Message, Trajectory
@ -12,6 +19,16 @@ from reme_ai.utils.op_utils import merge_messages_content, parse_json_experience
@C.register_op()
class ComparativeExtractionOp(BaseAsyncOp):
"""Extract comparative task memories by comparing different scoring trajectories.
This operation performs two types of comparisons:
1. Soft comparison: Compares highest vs lowest scoring trajectories
2. Hard comparison: Compares similar success vs failure step sequences
The extracted memories help identify what makes some trajectories more successful
than others.
"""
file_path: str = __file__
async def async_execute(self):
@ -27,20 +44,24 @@ class ComparativeExtractionOp(BaseAsyncOp):
highest_traj, lowest_traj = self._find_highest_lowest_scoring_trajectories(all_trajectories)
if highest_traj and lowest_traj and highest_traj.score > lowest_traj.score:
logger.info(
f"Extracting soft comparative task memories: highest ({highest_traj.score:.2f}) vs lowest ({lowest_traj.score:.2f})")
f"Extracting soft comparative task memories: "
f"highest ({highest_traj.score:.2f}) vs lowest ({lowest_traj.score:.2f})",
)
soft_task_memories = await self._extract_soft_comparative_task_memory(highest_traj, lowest_traj)
comparative_task_memories.extend(soft_task_memories)
# Hard comparison: success vs failure (if similarity search is enabled)
if (success_trajectories and failure_trajectories and
self.op_params.get("enable_similarity_comparison", False)):
if success_trajectories and failure_trajectories and self.op_params.get("enable_similarity_comparison", False):
similar_pairs = self._find_similar_step_sequences(success_trajectories, failure_trajectories)
logger.info(f"Found {len(similar_pairs)} similar pairs for hard comparison")
for success_steps, failure_steps, similarity_score in similar_pairs:
hard_task_memories = await self._extract_hard_comparative_task_memory(success_steps, failure_steps,
similarity_score)
hard_task_memories = await self._extract_hard_comparative_task_memory(
success_steps,
failure_steps,
similarity_score,
)
comparative_task_memories.extend(hard_task_memories)
logger.info(f"Extracted {len(comparative_task_memories)} comparative task memories")
@ -50,7 +71,9 @@ class ComparativeExtractionOp(BaseAsyncOp):
@staticmethod
def _find_highest_lowest_scoring_trajectories(trajectories: List[Trajectory]) -> Tuple[
Optional[Trajectory], Optional[Trajectory]]:
Optional[Trajectory],
Optional[Trajectory],
]:
"""Find the highest and lowest scoring trajectories"""
if len(trajectories) < 2:
return None, None
@ -75,8 +98,11 @@ class ComparativeExtractionOp(BaseAsyncOp):
"""Get trajectory score"""
return trajectory.score
async def _extract_soft_comparative_task_memory(self, higher_traj: Trajectory, lower_traj: Trajectory) -> List[
BaseMemory]:
async def _extract_soft_comparative_task_memory(
self,
higher_traj: Trajectory,
lower_traj: Trajectory,
) -> List[BaseMemory]:
"""Extract soft comparative task memory (high score vs low score)"""
higher_steps = self._get_trajectory_steps(higher_traj)
lower_steps = self._get_trajectory_steps(lower_traj)
@ -88,7 +114,7 @@ class ComparativeExtractionOp(BaseAsyncOp):
higher_steps=merge_messages_content(higher_steps),
lower_steps=merge_messages_content(lower_steps),
higher_score=f"{higher_score:.2f}",
lower_score=f"{lower_score:.2f}"
lower_score=f"{lower_score:.2f}",
)
def parse_task_memories(message: Message) -> List[BaseMemory]:
@ -100,24 +126,30 @@ class ComparativeExtractionOp(BaseAsyncOp):
workspace_id=self.context.get("workspace_id", ""),
when_to_use=tm_data.get("when_to_use", tm_data.get("condition", "")),
content=tm_data.get("experience", ""),
author=getattr(self.llm, 'model_name', 'system'),
metadata=tm_data
author=getattr(self.llm, "model_name", "system"),
metadata=tm_data,
)
task_memories.append(task_memory)
return task_memories
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_task_memories)
return await self.llm.achat(
messages=[FlowMessage(role=Role.USER, content=prompt)],
callback_fn=parse_task_memories,
)
async def _extract_hard_comparative_task_memory(self, success_steps: List[Message],
failure_steps: List[Message], similarity_score: float) -> List[
BaseMemory]:
async def _extract_hard_comparative_task_memory(
self,
success_steps: List[Message],
failure_steps: List[Message],
similarity_score: float,
) -> List[BaseMemory]:
"""Extract hard comparative task memory (success vs failure)"""
prompt = self.prompt_format(
prompt_name="hard_comparative_step_task_memory_prompt",
success_steps=merge_messages_content(success_steps),
failure_steps=merge_messages_content(failure_steps),
similarity_score=similarity_score
similarity_score=similarity_score,
)
def parse_task_memories(message: Message) -> List[BaseMemory]:
@ -129,19 +161,22 @@ class ComparativeExtractionOp(BaseAsyncOp):
workspace_id=self.context.get("workspace_id", ""),
when_to_use=tm_data.get("when_to_use", tm_data.get("condition", "")),
content=tm_data.get("experience", ""),
author=getattr(self.llm, 'model_name', 'system'),
metadata=tm_data
author=getattr(self.llm, "model_name", "system"),
metadata=tm_data,
)
task_memories.append(task_memory)
return task_memories
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_task_memories)
return await self.llm.achat(
messages=[FlowMessage(role=Role.USER, content=prompt)],
callback_fn=parse_task_memories,
)
@staticmethod
def _get_trajectory_steps(trajectory: Trajectory) -> List[Message]:
"""Get trajectory steps, prioritizing segmented steps"""
if hasattr(trajectory, 'segments') and trajectory.segments:
if hasattr(trajectory, "segments") and trajectory.segments:
# If there are segments, merge all segments
all_steps = []
for segment in trajectory.segments:
@ -150,9 +185,11 @@ class ComparativeExtractionOp(BaseAsyncOp):
else:
return trajectory.messages
def _find_similar_step_sequences(self, success_trajectories: List[Trajectory],
failure_trajectories: List[Trajectory]) -> List[
Tuple[List[Message], List[Message], float]]:
def _find_similar_step_sequences(
self,
success_trajectories: List[Trajectory],
failure_trajectories: List[Trajectory],
) -> List[Tuple[List[Message], List[Message], float]]:
"""Find similar step sequences for comparison"""
if not self.op_params.get("enable_similarity_comparison", False):
return []
@ -163,14 +200,14 @@ class ComparativeExtractionOp(BaseAsyncOp):
# Get step sequences
success_step_sequences = []
for traj in success_trajectories:
if hasattr(traj.metadata, 'segments') and traj.metadata["segments"]:
if hasattr(traj.metadata, "segments") and traj.metadata["segments"]:
success_step_sequences.extend(traj.metadata["segments"])
else:
success_step_sequences.append(traj.messages)
failure_step_sequences = []
for traj in failure_trajectories:
if hasattr(traj.metadata, 'segments') and traj.metadata["segments"]:
if hasattr(traj.metadata, "segments") and traj.metadata["segments"]:
failure_step_sequences.extend(traj.metadata["segments"])
else:
failure_step_sequences.append(traj.messages)
@ -188,8 +225,14 @@ class ComparativeExtractionOp(BaseAsyncOp):
failure_texts = [merge_messages_content(seq) for seq in failure_step_sequences]
# Get embedding vectors
if hasattr(self, 'vector_store') and self.vector_store and hasattr(
self.vector_store, 'embedding_model'):
if (
hasattr(self, "vector_store")
and self.vector_store
and hasattr(
self.vector_store,
"embedding_model",
)
):
success_embeddings = self.vector_store.embedding_model.get_embeddings(success_texts)
failure_embeddings = self.vector_store.embedding_model.get_embeddings(failure_texts)
@ -201,11 +244,13 @@ class ComparativeExtractionOp(BaseAsyncOp):
similarity = self._calculate_cosine_similarity(s_emb, f_emb)
if similarity > similarity_threshold:
similar_pairs.append((
success_step_sequences[i],
failure_step_sequences[j],
similarity
))
similar_pairs.append(
(
success_step_sequences[i],
failure_step_sequences[j],
similarity,
),
)
# Return top most similar pairs
max_pairs = self.op_params.get("max_similarity_pairs", 3)

View file

@ -1,28 +1,28 @@
soft_comparative_step_task_memory_prompt: |
You are an expert AI analyst comparing higher-scoring and lower-scoring step sequences to extract performance insights.
Your task is to identify the key differences between higher and lower performing approaches at the step level.
Focus on what made the higher-scoring approach more effective, even when both approaches may have had partial success.
SOFT COMPARATIVE ANALYSIS FRAMEWORK:
● PERFORMANCE FACTORS: Identify what specifically contributed to the higher score
● APPROACH DIFFERENCES: Compare methodologies and execution strategies
● EFFICIENCY ANALYSIS: Analyze why one approach was more efficient or effective
● OPTIMIZATION INSIGHTS: Extract lessons for improving performance
EXTRACTION PRINCIPLES:
● Focus on INCREMENTAL IMPROVEMENTS and performance optimization
● Extract QUALITY INDICATORS that differentiate better vs good approaches
● Identify REFINEMENT STRATEGIES that lead to higher scores
● Frame insights as PERFORMANCE ENHANCEMENT guidelines
# Higher-Scoring Step Sequence (Score: {higher_score})
{higher_steps}
# Lower-Scoring Step Sequence (Score: {lower_score})
{lower_steps}
OUTPUT FORMAT:
Generate 1-2 performance improvement insights as JSON objects:
```json
@ -40,30 +40,30 @@ soft_comparative_step_task_memory_prompt: |
hard_comparative_step_task_memory_prompt: |
You are an expert AI analyst comparing successful and failed step sequences to extract differential insights.
Your task is to identify the key differences between success and failure patterns at the step level.
Focus on critical decision points, technique variations, and approach differences.
COMPARATIVE ANALYSIS FRAMEWORK:
● DECISION CONTRAST: Compare critical decisions made in success vs failure cases
● TECHNIQUE VARIATIONS: Identify different approaches and their outcomes
● TIMING DIFFERENCES: Analyze when certain actions were taken and their impact
● SUCCESS FACTORS: Extract what specifically made the difference
EXTRACTION PRINCIPLES:
● Frame comparisons as PRINCIPLES as well as case-specific SOLUTIONS
● Identify PATTERNS that differentiate effective vs ineffective approaches
● Extract RULES that can guide future similar situations
● Focus on UNDERLYING MECHANISMS rather than surface-level differences
# Successful Step Sequence
{success_steps}
# Failed Step Sequence
{failure_steps}
# Similarity Score: {similarity_score}
OUTPUT FORMAT:
Generate 1-2 comparative insights as JSON objects:
```json

View file

@ -1,8 +1,15 @@
"""Failure extraction operation for task memory generation.
This module provides operations to extract task memories from failed trajectories,
identifying mistakes, pitfalls, and lessons learned from failures.
"""
from typing import List
from flowllm import C, BaseAsyncOp
from flowllm.enumeration.role import Role
from flowllm.schema.message import Message as FlowMessage
from flowllm.core.context import C
from flowllm.core.enumeration import Role
from flowllm.core.op import BaseAsyncOp
from flowllm.core.schema import Message as FlowMessage
from loguru import logger
from reme_ai.schema import Message, Trajectory
@ -12,6 +19,13 @@ from reme_ai.utils.op_utils import merge_messages_content, parse_json_experience
@C.register_op()
class FailureExtractionOp(BaseAsyncOp):
"""Extract task memories from failed trajectories.
This operation analyzes failed trajectories (or their segments) to extract
lessons learned, common mistakes, and anti-patterns that should be avoided
in similar future tasks.
"""
file_path: str = __file__
async def async_execute(self):
@ -28,7 +42,7 @@ class FailureExtractionOp(BaseAsyncOp):
# Process trajectories
for trajectory in failure_trajectories:
if hasattr(trajectory, 'segments') and trajectory.segments:
if hasattr(trajectory, "segments") and trajectory.segments:
# Process segmented step sequences
for segment in trajectory.segments:
task_memories = await self._extract_failure_task_memory_from_steps(segment, trajectory)
@ -43,17 +57,21 @@ class FailureExtractionOp(BaseAsyncOp):
# Add task memories to context
self.context.failure_task_memories = failure_task_memories
async def _extract_failure_task_memory_from_steps(self, steps: List[Message], trajectory: Trajectory) -> List[BaseMemory]:
async def _extract_failure_task_memory_from_steps(
self,
steps: List[Message],
trajectory: Trajectory,
) -> List[BaseMemory]:
"""Extract task memory from failed step sequences"""
step_content = merge_messages_content(steps)
context = get_trajectory_context(trajectory, steps)
prompt = self.prompt_format(
prompt_name="failure_step_task_memory_prompt",
query=trajectory.metadata.get('query', ''),
query=trajectory.metadata.get("query", ""),
step_sequence=step_content,
context=context,
outcome="failed"
outcome="failed",
)
def parse_task_memories(message: Message) -> List[BaseMemory]:
@ -65,11 +83,14 @@ class FailureExtractionOp(BaseAsyncOp):
workspace_id=self.context.get("workspace_id", ""),
when_to_use=tm_data.get("when_to_use", tm_data.get("condition", "")),
content=tm_data.get("experience", ""),
author=getattr(self.llm, 'model_name', 'system'),
metadata=tm_data
author=getattr(self.llm, "model_name", "system"),
metadata=tm_data,
)
task_memories.append(task_memory)
return task_memories
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_task_memories)
return await self.llm.achat(
messages=[FlowMessage(role=Role.USER, content=prompt)],
callback_fn=parse_task_memories,
)

View file

@ -1,31 +1,31 @@
failure_step_task_memory_prompt: |
You are an expert AI analyst reviewing failed step sequences from an AI agent execution.
Your task is to extract learning task memories from failures to prevent similar mistakes in future executions.
Focus on identifying error patterns, missed opportunities, and alternative approaches.
ANALYSIS FRAMEWORK:
● FAILURE POINT IDENTIFICATION: Pinpoint where and why the steps went wrong
● ERROR PATTERN ANALYSIS: Identify recurring mistakes or problematic approaches
● ALTERNATIVE APPROACHES: Suggest what could have been done differently
● PREVENTION STRATEGIES: Extract actionable insights to avoid similar failures
EXTRACTION PRINCIPLES:
● Extract GENERAL PRINCIPLES as well as SPECIFIC INSTRUCTIONS
● Focus on PATTERNS and RULES as well as particular instances
# Original Query
{query}
# Step Sequence Analysis
{step_sequence}
# Context Information
{context}
# Outcome
This step sequence was part of a {outcome} trajectory.
OUTPUT FORMAT:
Generate 1-3 step-level failure prevention insights as JSON objects:
```json

View file

@ -1,6 +1,13 @@
"""Memory deduplication operation for task memory management.
This module provides operations to remove duplicate or highly similar task
memories by comparing embeddings and calculating similarity scores.
"""
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,6 +15,13 @@ from reme_ai.schema.memory import BaseMemory
@C.register_op()
class MemoryDeduplicationOp(BaseAsyncOp):
"""Remove duplicate task memories using embedding similarity.
This operation identifies and removes duplicate or highly similar task
memories by comparing their embeddings against both existing memories
in the vector store and other memories in the current batch.
"""
file_path: str = __file__
async def async_execute(self):
@ -25,7 +39,9 @@ class MemoryDeduplicationOp(BaseAsyncOp):
deduplicated_task_memories = await self._deduplicate_task_memories(task_memories)
logger.info(
f"Deduplication complete: {len(deduplicated_task_memories)} deduplicated task memories out of {len(task_memories)}")
f"Deduplication complete: {len(deduplicated_task_memories)} deduplicated "
f"task memories out of {len(task_memories)}",
)
# Update context
self.context.response.metadata["memory_list"] = deduplicated_task_memories
@ -70,24 +86,25 @@ class MemoryDeduplicationOp(BaseAsyncOp):
async def _get_existing_task_memory_embeddings(self, workspace_id: str) -> List[List[float]]:
"""Get embeddings of existing task memories"""
try:
if not hasattr(self, 'vector_store') or not self.vector_store or not workspace_id:
if not hasattr(self, "vector_store") or not self.vector_store or not workspace_id:
return []
# Query existing task memory nodes
existing_nodes = await self.vector_store.async_search(
query="...", # Empty query to get all
workspace_id=workspace_id,
top_k=self.op_params.get("max_existing_task_memories", 1000)
top_k=self.op_params.get("max_existing_task_memories", 1000),
)
# Extract embeddings
existing_embeddings = []
for node in existing_nodes:
if hasattr(node, 'embedding') and node.embedding:
if hasattr(node, "embedding") and node.embedding:
existing_embeddings.append(node.embedding)
logger.debug(
f"Retrieved {len(existing_embeddings)} existing task memory embeddings from workspace {workspace_id}")
f"Retrieved {len(existing_embeddings)} existing task memory embeddings from workspace {workspace_id}",
)
return existing_embeddings
except Exception as e:
@ -112,9 +129,12 @@ class MemoryDeduplicationOp(BaseAsyncOp):
logger.error(f"Error generating embedding for task memory: {e}")
return None
def _is_similar_to_existing_task_memories(self, current_embedding: List[float],
existing_embeddings: List[List[float]],
threshold: float) -> bool:
def _is_similar_to_existing_task_memories(
self,
current_embedding: List[float],
existing_embeddings: List[List[float]],
threshold: float,
) -> bool:
"""Check if current embedding is similar to existing embeddings"""
for existing_embedding in existing_embeddings:
similarity = self._calculate_cosine_similarity(current_embedding, existing_embedding)
@ -123,9 +143,13 @@ class MemoryDeduplicationOp(BaseAsyncOp):
return True
return False
def _is_similar_to_current_task_memories(self, current_embedding: List[float],
current_task_memories: List[BaseMemory],
threshold: float) -> bool:
def _is_similar_to_current_task_memories(
self,
current_embedding: List[float],
current_task_memories: List[BaseMemory],
threshold: float,
) -> bool:
"""Check if current embedding is similar to other memories in current batch."""
for existing_task_memory in current_task_memories:
existing_embedding = self._get_task_memory_embedding(existing_task_memory)
if existing_embedding is None:

View file

@ -1,18 +1,32 @@
"""Memory validation operation for task memory quality control.
This module provides operations to validate the quality of extracted task
memories using LLM-based evaluation, ensuring only high-quality memories
are stored.
"""
import json
import re
from typing import List, Dict, Any
from flowllm import C, BaseAsyncOp
from flowllm.enumeration.role import Role
from flowllm.schema.message import Message as FlowMessage
from flowllm.core.context import C
from flowllm.core.enumeration import Role
from flowllm.core.op import BaseAsyncOp
from flowllm.core.schema import Message as FlowMessage
from loguru import logger
from reme_ai.schema import Message
from reme_ai.schema.memory import BaseMemory
@C.register_op()
class MemoryValidationOp(BaseAsyncOp):
"""Validate quality of extracted task memories.
This operation uses LLM-based evaluation to assess the quality of extracted
task memories, filtering out low-quality or invalid memories based on
validation scores and criteria.
"""
file_path: str = __file__
async def async_execute(self):
@ -59,15 +73,16 @@ class MemoryValidationOp(BaseAsyncOp):
prompt = self.prompt_format(
prompt_name="task_memory_validation_prompt",
condition=task_memory.when_to_use,
task_memory_content=task_memory.content)
task_memory_content=task_memory.content,
)
def parse_validation(message: Message) -> Dict[str, Any]:
def parse_validation(message: FlowMessage) -> Dict[str, Any]:
try:
response_content = message.content
# Parse validation result
# 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_content)
if json_blocks:
@ -85,8 +100,11 @@ class MemoryValidationOp(BaseAsyncOp):
"is_valid": is_valid and score >= validation_threshold,
"score": score,
"feedback": response_content,
"reason": "" if (
is_valid and score >= validation_threshold) else f"Low validation score ({score:.2f}) or marked as invalid"
"reason": (
""
if (is_valid and score >= validation_threshold)
else f"Low validation score ({score:.2f}) or marked as invalid"
),
}
except Exception as e_inner:
@ -95,10 +113,13 @@ class MemoryValidationOp(BaseAsyncOp):
"is_valid": False,
"score": 0.0,
"feedback": "",
"reason": f"Parse error: {str(e_inner)}"
"reason": f"Parse error: {str(e_inner)}",
}
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_validation)
return await self.llm.achat(
messages=[FlowMessage(role=Role.USER, content=prompt)],
callback_fn=parse_validation,
)
except Exception as e:
logger.error(f"LLM validation failed: {e}")
@ -106,5 +127,5 @@ class MemoryValidationOp(BaseAsyncOp):
"is_valid": False,
"score": 0.0,
"feedback": "",
"reason": f"LLM validation error: {str(e)}"
"reason": f"LLM validation error: {str(e)}",
}

View file

@ -1,19 +1,19 @@
task_memory_validation_prompt: |
You are an expert AI analyst tasked with validating the quality and usefulness of extracted step-level task memories.
Your task is to assess whether the extracted task memory is actionable, accurate, and valuable for future agent executions.
VALIDATION CRITERIA:
● ACTIONABILITY: Is the task memory specific enough to guide future actions?
● ACCURACY: Does the task memory correctly reflect the patterns observed?
● RELEVANCE: Is the task memory applicable to similar future scenarios?
● CLARITY: Is the task memory clearly articulated and understandable?
● UNIQUENESS: Does the task memory provide novel insights or common knowledge?
# Task Memory to Validate
Condition: {condition}
Task Memory Content: {task_memory_content}
OUTPUT FORMAT:
Provide validation assessment:
```json
@ -24,6 +24,6 @@ task_memory_validation_prompt: |
"recommendations": "Suggestions for improvement if applicable"
}}
```
Score should be between 0.0 (poor quality) and 1.0 (excellent quality).
Mark as invalid if score is below 0.3 or if there are fundamental issues with the task memory.

View file

@ -1,9 +1,16 @@
"""Simple comparative summary operation for task memory generation.
This module provides a simplified operation to extract task memories by
comparing trajectories with different scores for the same task.
"""
import json
from typing import List, Dict
from flowllm import C, BaseAsyncOp
from flowllm.enumeration.role import Role
from flowllm.schema.message import Message as FlowMessage
from flowllm.core.context import C
from flowllm.core.enumeration import Role
from flowllm.core.op import BaseAsyncOp
from flowllm.core.schema import Message as FlowMessage
from loguru import logger
from reme_ai.schema import Message, Trajectory
@ -13,13 +20,30 @@ from reme_ai.utils.op_utils import merge_messages_content
@C.register_op()
class SimpleComparativeSummaryOp(BaseAsyncOp):
"""Extract task memories by comparing trajectories with different scores.
This operation compares the highest and lowest scoring trajectories for
each task to extract comparative insights and best practices.
"""
file_path: str = __file__
async def compare_summary_trajectory(self, trajectory_a: Trajectory, trajectory_b: Trajectory) -> List[BaseMemory]:
summary_prompt = self.prompt_format(prompt_name="summary_prompt",
execution_process_a=merge_messages_content(trajectory_a.messages),
execution_process_b=merge_messages_content(trajectory_b.messages),
summary_example=self.get_prompt("summary_example"))
"""Compare two trajectories and extract comparative task memories.
Args:
trajectory_a: First trajectory to compare (typically higher scoring)
trajectory_b: Second trajectory to compare (typically lower scoring)
Returns:
List of extracted task memories from the comparison
"""
summary_prompt = self.prompt_format(
prompt_name="summary_prompt",
execution_process_a=merge_messages_content(trajectory_a.messages),
execution_process_b=merge_messages_content(trajectory_b.messages),
summary_example=self.get_prompt("summary_example"),
)
def parse_content(message: Message):
content = message.content
@ -33,10 +57,14 @@ class SimpleComparativeSummaryOp(BaseAsyncOp):
when_to_use = tm_dict.get("when_to_use", "").strip()
task_memory_content = tm_dict.get("experience", "").strip()
if when_to_use and task_memory_content:
task_memory_list.append(TaskMemory(workspace_id=self.context.get("workspace_id", ""),
when_to_use=when_to_use,
content=task_memory_content,
author=getattr(self.llm, 'model_name', 'system')))
task_memory_list.append(
TaskMemory(
workspace_id=self.context.get("workspace_id", ""),
when_to_use=when_to_use,
content=task_memory_content,
author=getattr(self.llm, "model_name", "system"),
),
)
return task_memory_list
@ -44,9 +72,17 @@ class SimpleComparativeSummaryOp(BaseAsyncOp):
logger.exception(f"parse content failed!\n{content}")
raise e
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=summary_prompt)], callback_fn=parse_content)
return await self.llm.achat(
messages=[FlowMessage(role=Role.USER, content=summary_prompt)],
callback_fn=parse_content,
)
async def async_execute(self):
"""Execute the comparative summary operation.
Groups trajectories by task_id, compares the highest and lowest scoring
trajectories for each task, and extracts task memories from the comparison.
"""
trajectories: list = self.context.get("trajectories", [])
trajectories: List[Trajectory] = [Trajectory(**x) if isinstance(x, dict) else x for x in trajectories]
@ -57,14 +93,16 @@ class SimpleComparativeSummaryOp(BaseAsyncOp):
task_id_dict[trajectory.task_id].append(trajectory)
memory_list = []
for task_id, task_trajectories in task_id_dict.items():
for _, task_trajectories in task_id_dict.items():
task_trajectories: List[Trajectory] = sorted(task_trajectories, key=lambda x: x.score, reverse=True)
if len(task_trajectories) < 2:
continue
if task_trajectories[0].score > task_trajectories[-1].score:
task_memories = await self.compare_summary_trajectory(trajectory_a=task_trajectories[0],
trajectory_b=task_trajectories[-1])
task_memories = await self.compare_summary_trajectory(
trajectory_a=task_trajectories[0],
trajectory_b=task_trajectories[-1],
)
memory_list.extend(task_memories)
self.context.response.answer = json.dumps([x.model_dump() for x in memory_list])

Some files were not shown because too many files have changed in this diff Show more