rename reme4->reme

This commit is contained in:
jinli.yl 2026-06-19 02:04:28 +08:00
parent 26cb5ca62f
commit 195f97857a
602 changed files with 194 additions and 81331 deletions

View file

@ -1,360 +0,0 @@
# flake8: noqa: E402, E501
# pylint: disable=E0611
"""A minimal ReAct Agent for AppWorld tasks."""
import os
import re
import time
import json
import datetime
from typing import List, Any
import ray
import requests
from tqdm import tqdm
from loguru import logger
from openai import OpenAI
from jinja2 import Template
from dotenv import load_dotenv
from prompt import NEW_PROMPT_TEMPLATE
from appworld import AppWorld, load_task_ids
os.environ["APPWORLD_ROOT"] = "."
load_dotenv("../../.env")
@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 = 129024,
num_trials: int = 1,
use_memory: bool = False,
memory_base_url: str = "http://0.0.0.0:8002/",
use_memory_addition: bool = False,
use_memory_deletion: bool = False,
delete_freq: int = 10,
freq_threshold: int = 5,
utility_threshold: float = 0.5,
):
self.index: int = index
self.task_ids: List[str] = task_ids
self.experiment_name: str = experiment_name
self.model_name: str = model_name
self.temperature: float = temperature
self.max_interactions: int = max_interactions
self.max_response_size: int = max_response_size
self.num_trials: int = num_trials
self.use_memory: bool = use_memory
self.use_memory_addition: bool = use_memory_addition if use_memory else False
self.use_memory_deletion: bool = use_memory_deletion if use_memory else False
self.delete_freq: int = delete_freq
self.freq_threshold: int = freq_threshold
self.utility_threshold: float = utility_threshold
self.llm_client = OpenAI()
self.memory_base_url: str = memory_base_url
self.history: List[List[List[dict]]] = [[] for _ in range(num_trials)]
self.retrieved_memory_list: List[List[List[Any]]] = [[] for _ in range(num_trials)]
for run_id in range(num_trials):
for _ in range(len(task_ids)):
self.retrieved_memory_list[run_id].append([])
self.history[run_id].append([])
def call_llm(self, messages: list) -> str:
"""Call the LLM to generate a response to the messages."""
for i in range(100):
try:
response = self.llm_client.chat.completions.create(
model=self.model_name,
messages=messages,
temperature=self.temperature,
extra_body={"enable_thinking": False},
seed=0,
)
return response.choices[0].message.content
except Exception as e:
logger.exception(f"encounter error with {e.args}")
time.sleep(1 + i * 10)
return "call llm error"
def prompt_messages(self, run_id, task_index, previous_memories: None, world: AppWorld):
"""Prompt the messages to the LLM."""
app_descriptions = json.dumps(
[{"name": k, "description": v} for (k, v) in world.task.app_descriptions.items()],
indent=1,
)
dictionary = {"supervisor": world.task.supervisor, "app_descriptions": app_descriptions}
sys_prompt = Template(NEW_PROMPT_TEMPLATE.lstrip()).render(dictionary)
query = world.task.instruction
if self.use_memory:
if len(previous_memories) == 0:
response = self.get_memory(world.task.instruction)
if response and "memory_list" in response["metadata"]:
self.retrieved_memory_list[run_id][task_index] = response["metadata"]["memory_list"]
task_memory = re.sub(r"\bMemory\s*(\d+)\s*[:]", r"Experience \1:", response["answer"])
logger.info(f"loaded task_memory: {task_memory}")
query = (
"Task:\n"
+ query
+ "\n\nSome Related Experience to help you to complete the task:\n"
+ task_memory
)
else:
formatted_memories = []
for i, memory in enumerate(previous_memories, 1):
condition = memory["when_to_use"]
memory_content = memory["content"]
memory_text = f"Experience {i}:\n When to use: {condition}\n Content: {memory_content}\n"
formatted_memories.append(memory_text)
query = (
"Task:\n"
+ query
+ "\n\nSome Related Experience to help you to complete the task:\n"
+ "\n".join(formatted_memories)
)
messages = [
{"role": "system", "content": sys_prompt},
{"role": "user", "content": query},
]
self.history[run_id][task_index] = messages
@staticmethod
def get_reward(world) -> float:
"""Get the reward for the Appworld world."""
tracker = world.evaluate()
num_passes = len(tracker.passes)
num_failures = len(tracker.failures)
return num_passes / (num_passes + num_failures)
def extract_code_and_fix_content(
self,
text: str,
ignore_multiple_calls=True,
) -> tuple[str, str]:
"""Extract the code and fix the content."""
full_code_regex = r"```python\n(.*?)```"
partial_code_regex = r".*```python\n(.*)"
original_text = text
output_code = ""
match_end = 0
# Handle multiple calls
for re_match in re.finditer(full_code_regex, original_text, flags=re.DOTALL):
code = re_match.group(1).strip()
if ignore_multiple_calls:
text = original_text[: re_match.end()]
return code, text
output_code += code + "\n"
match_end = re_match.end()
# check for partial code match at end (no terminating ```) following the last match
partial_match = re.match(
partial_code_regex,
original_text[match_end:],
flags=re.DOTALL,
)
if partial_match:
output_code += partial_match.group(1).strip()
# terminated due to stop condition. Add stop condition to output.
if not text.endswith("\n"):
text = text + "\n"
text = text + "```"
if len(output_code) == 0:
return text, text
else:
return output_code, text
def execute(self):
"""Execute the Appworld tasks."""
result = []
counter = 0
for task_index, task_id in enumerate(tqdm(self.task_ids, desc=f"run_index={self.index}")):
t_result = None
previous_memories = []
# Run each task num_trials times
for run_id in range(self.num_trials):
start_time = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
with AppWorld(task_id=task_id, experiment_name=f"{self.experiment_name}_run_{run_id}") as world:
before_score = self.get_reward(world)
for i in range(self.max_interactions):
if i == 0:
self.prompt_messages(
run_id=run_id,
task_index=task_index,
previous_memories=previous_memories,
world=world,
)
code_msg = self.call_llm(self.history[run_id][task_index])
code, _ = self.extract_code_and_fix_content(code_msg)
self.history[run_id][task_index].append({"role": "assistant", "content": code})
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]
self.history[run_id][task_index].append(
{"role": "user", "content": "Output:\n```\n" + output + "```\n\n"},
)
if world.task_completed():
break
after_score = self.get_reward(world)
uplift_score = after_score - before_score
if self.use_memory:
if self.use_memory_addition:
new_traj_list = [
self.get_traj_from_task_history(task_id, self.history[run_id][task_index], after_score),
]
previous_memories = self.summary_memory(new_traj_list)
if after_score == 1:
self.add_memory(previous_memories)
# update the freq & utility attributes of retrieved memories
update_utility: bool = after_score == 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()
t_result = {
"task_id": world.task_id,
"run_id": run_id,
"experiment_name": self.experiment_name,
"task_completed": world.task_completed(),
"before_score": before_score,
"after_score": after_score,
"uplift_score": uplift_score,
"task_history": self.history[run_id][task_index],
"task_start_time": start_time,
}
if after_score == 1:
break
result.append(t_result)
return result
def handle_api_response(self, response: requests.Response):
"""Handle API response with proper error checking"""
if response.status_code != 200:
print(f"Error: {response.status_code}")
print(response.text)
return None
return response.json()
def get_memory(self, query: str):
"""Retrieve relevant task memories based on a query"""
response = requests.post(
url=f"{self.memory_base_url}retrieve_task_memory",
json={
"query": query,
"enable_llm_rerank": False,
"enable_score_filter": False,
"top_k": 5,
"enable_llm_rewrite": False,
},
)
result = self.handle_api_response(response)
if not result:
return None
logger.info(f"query: {query}, response: {result}")
return result
def get_traj_from_task_history(self, task_id: str, task_history: list, reward: float):
"""Get the trajectory from the task history."""
pattern = r"\n\nSome Related Experience to help you to complete the task:.*"
task_history[1]["content"] = re.sub(pattern, "", task_history[1]["content"], flags=re.DOTALL)
return {
"task_id": task_id,
"messages": task_history,
"score": reward,
}
def summary_memory(self, trajectories):
"""Generate a summary of conversation messages and create task memories"""
response = requests.post(
url=f"{self.memory_base_url}summary_task_memory",
json={
"trajectories": trajectories,
"success_threshold": 1.0,
"enable_soft_comparison": True,
"validation_threshold": 0.5,
},
)
result = self.handle_api_response(response)
if not result:
return []
# Extract memory list from response
memory_list = result.get("metadata", {}).get("memory_list", [])
print(f"Task memory list created: {len(memory_list)} memories")
return memory_list
def add_memory(self, memory_list):
"""Add the memory to the memory pool."""
response = requests.post(
url=f"{self.memory_base_url}add_task_memory",
json={
"memory_list": memory_list,
},
)
response.raise_for_status()
def update_memory_information(self, memory_list, update_utility: bool = False):
"""Update the memory information."""
response = requests.post(
url=f"{self.memory_base_url}record_task_memory",
json={
"memory_list": memory_list,
"update_utility": update_utility,
},
)
response.raise_for_status()
logger.info(response.json())
def delete_memory(self):
"""Delete the memory from the memory pool."""
response = requests.post(
url=f"{self.memory_base_url}delete_task_memory",
json={
"freq_threshold": self.freq_threshold,
"utility_threshold": self.utility_threshold,
},
)
response.raise_for_status()
def main():
"""Main function to run the Appworld React Agent."""
dataset_name = "train"
task_ids = load_task_ids(dataset_name)
agent = AppworldReactAgent(index=0, task_ids=task_ids[0:1], experiment_name=dataset_name, num_trials=1)
result = agent.execute()
logger.info(f"result={json.dumps(result)}")
if __name__ == "__main__":
main()

View file

@ -1,660 +0,0 @@
# flake8: noqa: E402, E501
# pylint: disable=C0114,C0301
# 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.
PROMPT_TEMPLATE = """
USER:
I am your supervisor and you are a super intelligent AI Assistant whose job is to achieve my day-to-day tasks completely autonomously.
To do this, you will need to interact with app/s (e.g., spotify, venmo, etc) using their associated APIs on my behalf. For this you will undertake a *multi-step conversation* using a python REPL environment. That is, you will write the python code and the environment will execute it and show you the result, based on which, you will write python code for the next step and so on, until you've achieved the goal. This environment will let you interact with app/s using their associated APIs on my behalf.
Here are three key APIs that you need to know to get more information
# To get a list of apps that are available to you.
print(apis.api_docs.show_app_descriptions())
# To get the list of apis under any app listed above, e.g. supervisor
print(apis.api_docs.show_api_descriptions(app_name='supervisor'))
# To get the specification of a particular api, e.g. supervisor app's show_account_passwords
print(apis.api_docs.show_api_doc(app_name='supervisor', api_name='show_account_passwords'))
Each code execution will produce an output that you can use in subsequent calls. Using these APIs, you can now generate code, that the environment will execute, to solve the task.
For example, consider the task:
My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}.
Task:
What is the password for my Spotify account?
ASSISTANT:
# Okay. Lets first find which apps are available to get the password by looking at the app descriptions.
print(apis.api_docs.show_app_descriptions())
USER:
[
{
"name": "api_docs",
"description": "An app to search and explore API documentation."
},
{
"name": "supervisor",
"description": "An app to access supervisor's personal information, account credentials, addresses, payment cards, and manage the assigned task."
},
...
{
"name": "spotify",
"description": "A music streaming app to stream songs and manage song, album and playlist libraries."
},
{
"name": "venmo",
"description": "A social payment app to send, receive and request money to and from others."
},
...
]
ASSISTANT:
# Looks like the supervisor app could help me with that. Lets see what apis are available under this app.
print(apis.api_docs.show_api_descriptions(app_name='supervisor'))
USER:
[
...
"show_account_passwords : Show your supervisor's account passwords."
...
]
ASSISTANT:
# I can use `show_account_passwords` to get the passwords. Let me see its detailed specification to understand its arguments and output structure.
print(apis.api_docs.show_api_doc(app_name='supervisor', api_name='show_account_passwords'))
USER:
{
'app_name': 'supervisor',
'api_name': 'show_account_passwords',
'path': '/account_passwords',
'method': 'GET',
'description': "Show your supervisor's app account passwords.",
'parameters': [],
'response_schemas': {
'success': [{'account_name': 'string', 'password': 'string'}],
'failure': {'message': 'string'}
}
}
ASSISTANT:
# Okay, it requires no arguments. So I can just call it directly.
print(apis.supervisor.show_account_passwords())
USER:
[
{
"account_name": "spotify",
"password": "dummy_spotify_pass"
},
{
"account_name": "file_system",
"password": "dummy_fs_pass"
},
...
]
ASSISTANT:
# So the Spotify password is an entry in the `passwords` list with the account_name=spotify.
spotify_password = [account_password["account_name"] == "spotify" for account_password in passwords][0]["password"]
print(spotify_password)
USER:
dummy_spotify_pass
ASSISTANT:
# When the task is completed, I need to call apis.supervisor.complete_task(). If there is an answer, I need to pass it as an argument `answer`. I will pass the spotify_password as an answer.
apis.supervisor.complete_task(answer=spotify_password)
USER:
Marked the active task complete.
----------------------------------------------
USER:
**Key instructions and disclaimers**:
1. The email addresses, access tokens and variables (e.g. spotify_password) in the example above were only for demonstration. Obtain the correct information by calling relevant APIs yourself.
2. Only generate valid code blocks, i.e., do not put them in ```...``` or add any extra formatting. Any thoughts should be put as code comments.
3. You can use the variables from the previous code blocks in the subsequent code blocks.
4. Write small chunks of code and only one chunk of code in every step. Make sure everything is working correctly before making any irreversible change.
5. The provided Python environment has access to its standard library. But modules and functions that have a risk of affecting the underlying OS, file system or process are disabled. You will get an error if do call them.
6. Any reference to a file system in the task instructions means the file system *app*, operable via given APIs, and not the actual file system the code is running on. So do not write code making calls to os-level modules and functions.
7. To interact with apps, only use the provided APIs, and not the corresponding Python packages. E.g., do NOT use `spotipy` for Spotify. Remember, the environment only has the standard library.
8. The provided API documentation has both the input arguments and the output JSON schemas. All calls to APIs and parsing its outputs must be as per this documentation.
9. For APIs that return results in "pages", make sure to consider all pages.
10. To obtain current date or time, use Python functions like `datetime.now()` or obtain it from the phone app. Do not rely on your existing knowledge of what the current date or time is.
11. For all temporal requests, use proper time boundaries, e.g., if I ask for something that happened yesterday, make sure to consider the time between 00:00:00 and 23:59:59. All requests are concerning a single, default (no) time zone.
12. Any reference to my friends, family or any other person or relation refers to the people in my phone's contacts list.
13. All my personal information, and information about my app account credentials, physical addresses and owned payment cards are stored in the "supervisor" app. You can access them via the APIs provided by the supervisor app.
14. Once you have completed the task, call `apis.supervisor.complete_task()`. If the task asks for some information, return it as the answer argument, i.e. call `apis.supervisor.complete_task(answer=<answer>)`. For tasks that do not require an answer, just skip the answer argument or pass it as None.
15. The answers, when given, should be just entity or number, not full sentences, e.g., `answer=10` for "How many songs are in the Spotify queue?". When an answer is a number, it should be in numbers, not in words, e.g., "10" and not "ten".
16. You can also pass `status="fail"` in the complete_task API if you are sure you cannot solve it and want to exit.
17. You must make all decisions completely autonomously and not ask for any clarifications or confirmations from me or anyone else.
USER:
Using these APIs, now generate code to solve the actual task:
My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}.
Task:
{{ instruction }}
"""
PROMPT_TEMPLATE_WITH_EXPERIENCE = """
USER:
I am your supervisor and you are a super intelligent AI Assistant whose job is to achieve my day-to-day tasks completely autonomously.
To do this, you will need to interact with app/s (e.g., spotify, venmo, etc) using their associated APIs on my behalf. For this you will undertake a *multi-step conversation* using a python REPL environment. That is, you will write the python code and the environment will execute it and show you the result, based on which, you will write python code for the next step and so on, until you've achieved the goal. This environment will let you interact with app/s using their associated APIs on my behalf.
Here are three key APIs that you need to know to get more information
# To get a list of apps that are available to you.
print(apis.api_docs.show_app_descriptions())
# To get the list of apis under any app listed above, e.g. supervisor
print(apis.api_docs.show_api_descriptions(app_name='supervisor'))
# To get the specification of a particular api, e.g. supervisor app's show_account_passwords
print(apis.api_docs.show_api_doc(app_name='supervisor', api_name='show_account_passwords'))
Each code execution will produce an output that you can use in subsequent calls. Using these APIs, you can now generate code, that the environment will execute, to solve the task.
For example, consider the task:
My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}.
Task:
What is the password for my Spotify account?
ASSISTANT:
# Okay. Lets first find which apps are available to get the password by looking at the app descriptions.
print(apis.api_docs.show_app_descriptions())
USER:
[
{
"name": "api_docs",
"description": "An app to search and explore API documentation."
},
{
"name": "supervisor",
"description": "An app to access supervisor's personal information, account credentials, addresses, payment cards, and manage the assigned task."
},
...
{
"name": "spotify",
"description": "A music streaming app to stream songs and manage song, album and playlist libraries."
},
{
"name": "venmo",
"description": "A social payment app to send, receive and request money to and from others."
},
...
]
ASSISTANT:
# Looks like the supervisor app could help me with that. Lets see what apis are available under this app.
print(apis.api_docs.show_api_descriptions(app_name='supervisor'))
USER:
[
...
"show_account_passwords : Show your supervisor's account passwords."
...
]
ASSISTANT:
# I can use `show_account_passwords` to get the passwords. Let me see its detailed specification to understand its arguments and output structure.
print(apis.api_docs.show_api_doc(app_name='supervisor', api_name='show_account_passwords'))
USER:
{
'app_name': 'supervisor',
'api_name': 'show_account_passwords',
'path': '/account_passwords',
'method': 'GET',
'description': "Show your supervisor's app account passwords.",
'parameters': [],
'response_schemas': {
'success': [{'account_name': 'string', 'password': 'string'}],
'failure': {'message': 'string'}
}
}
ASSISTANT:
# Okay, it requires no arguments. So I can just call it directly.
print(apis.supervisor.show_account_passwords())
USER:
[
{
"account_name": "spotify",
"password": "dummy_spotify_pass"
},
{
"account_name": "file_system",
"password": "dummy_fs_pass"
},
...
]
ASSISTANT:
# So the Spotify password is an entry in the `passwords` list with the account_name=spotify.
spotify_password = [account_password["account_name"] == "spotify" for account_password in passwords][0]["password"]
print(spotify_password)
USER:
dummy_spotify_pass
ASSISTANT:
# When the task is completed, I need to call apis.supervisor.complete_task(). If there is an answer, I need to pass it as an argument `answer`. I will pass the spotify_password as an answer.
apis.supervisor.complete_task(answer=spotify_password)
USER:
Marked the active task complete.
----------------------------------------------
USER:
**Key instructions and disclaimers**:
1. The email addresses, access tokens and variables (e.g. spotify_password) in the example above were only for demonstration. Obtain the correct information by calling relevant APIs yourself.
2. Only generate valid code blocks, i.e., do not put them in ```...``` or add any extra formatting. Any thoughts should be put as code comments.
3. You can use the variables from the previous code blocks in the subsequent code blocks.
4. Write small chunks of code and only one chunk of code in every step. Make sure everything is working correctly before making any irreversible change.
5. The provided Python environment has access to its standard library. But modules and functions that have a risk of affecting the underlying OS, file system or process are disabled. You will get an error if do call them.
6. Any reference to a file system in the task instructions means the file system *app*, operable via given APIs, and not the actual file system the code is running on. So do not write code making calls to os-level modules and functions.
7. To interact with apps, only use the provided APIs, and not the corresponding Python packages. E.g., do NOT use `spotipy` for Spotify. Remember, the environment only has the standard library.
8. The provided API documentation has both the input arguments and the output JSON schemas. All calls to APIs and parsing its outputs must be as per this documentation.
9. For APIs that return results in "pages", make sure to consider all pages.
10. To obtain current date or time, use Python functions like `datetime.now()` or obtain it from the phone app. Do not rely on your existing knowledge of what the current date or time is.
11. For all temporal requests, use proper time boundaries, e.g., if I ask for something that happened yesterday, make sure to consider the time between 00:00:00 and 23:59:59. All requests are concerning a single, default (no) time zone.
12. Any reference to my friends, family or any other person or relation refers to the people in my phone's contacts list.
13. All my personal information, and information about my app account credentials, physical addresses and owned payment cards are stored in the "supervisor" app. You can access them via the APIs provided by the supervisor app.
14. Once you have completed the task, call `apis.supervisor.complete_task()`. If the task asks for some information, return it as the answer argument, i.e. call `apis.supervisor.complete_task(answer=<answer>)`. For tasks that do not require an answer, just skip the answer argument or pass it as None.
15. The answers, when given, should be just entity or number, not full sentences, e.g., `answer=10` for "How many songs are in the Spotify queue?". When an answer is a number, it should be in numbers, not in words, e.g., "10" and not "ten".
16. You can also pass `status="fail"` in the complete_task API if you are sure you cannot solve it and want to exit.
17. You must make all decisions completely autonomously and not ask for any clarifications or confirmations from me or anyone else.
18. Some Related Experience to help you to complete the task:
{{experience}}
USER:
Using these APIs, now generate code to solve the actual task:
My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}.
Task:
{{ instruction }}
"""
NEW_PROMPT_TEMPLATE = """
USER:
I am your supervisor and you are a super intelligent AI Assistant whose job is to achieve my day-to-day tasks completely autonomously.
To do this, you will need to interact with app/s (e.g., spotify, venmo etc) using their associated APIs on my behalf. For this you will undertake a *multi-step conversation* using a python REPL environment. That is, you will write the python code and the environment will execute it and show you the result, based on which, you will write python code for the next step and so on, until you've achieved the goal. This environment will let you interact with app/s using their associated APIs on my behalf.
Here are three key APIs that you need to know to get more information
# To get a list of apps that are available to you.
```python
print(apis.api_docs.show_app_descriptions())
```
# To get the list of apis under any app listed above, e.g. spotify
```python
print(apis.api_docs.show_api_descriptions(app_name='spotify'))
```
# To get the specification of a particular api, e.g. spotify app's login api
```python
print(apis.api_docs.show_api_doc(app_name='spotify', api_name='login'))
```
Each code execution will produce an output that you can use in subsequent calls. Using these APIs, you can now generate code, that I will execute, to solve the task. Let's start with the task
My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}.
Task: How many playlists do I have in Spotify?
ASSISTANT:
Okay. Lets first find which APIs are available to use in Spotify.
Code:
```python
print(apis.api_docs.show_api_descriptions(app_name='spotify'))
```
USER:
Output:
```
[
...
"login : Login to your account.",
"logout : Logout from your account.",
...
]
```
ASSISTANT:
Okay. Looks like I can use the `login` api. Lets find its specifications.
Code:
```python
print(apis.api_docs.show_api_doc(app_name='spotify', api_name='login'))
```
USER:
Output:
```
{
"app_name": "spotify",
"api_name": "login",
"path": "/auth/token",
"method": "POST",
"description": "Login to your account.",
"parameters": [
{
"name": "username",
"type": "string",
"required": true,
"description": "Your account email.",
"default": null,
"constraints": []
},
{
"name": "password",
"type": "string",
"required": true,
"description": "Your account password.",
"default": null,
"constraints": []
}
],
"response_schemas": {
"success": {
"token_type": "string",
"access_token": "string"
},
"failure": {
"message": "string"
}
}
}
```
ASSISTANT:
I need the supervisor's username and password. Lets see if any app can help me get that.
Code:
```python
print(apis.api_docs.show_app_descriptions())
```
USER:
Output:
```
{{ app_descriptions }}
```
ASSISTANT:
Looks like the supervisor app could help me with that.
Code:
```python
print(apis.api_docs.show_api_descriptions(app_name='supervisor'))
```
USER:
Output:
```
[
...
"show_account_passwords : Show your supervisor's account passwords."
...
]
```
ASSISTANT:
I can use `show_account_passwords` to get the passwords. So I will call that and save it in a variable for future reference.
Code:
```python
passwords=apis.supervisor.show_account_passwords()
print(passwords)
```
USER:
Output:
```
[
{
"account_name": "spotify",
"password": "dummy_spotify_pass"
},
{
"account_name": "file_system",
"password": "dummy_fs_pass"
}
]
```
ASSISTANT:
So the Spotify password is an entry in the `passwords` list with the account_name=spotify. I can use it to login now with the provided email address: {{ supervisor.email }}.
Code:
```python
spotify_password = [account_password["account_name"] == "spotify" for account_password in passwords][0]["password"]
login_result = apis.spotify.login(username='{{ supervisor.email }}', password=spotify_password)
print(login_result)
```
USER:
Output:
```
{
"token_type": "Bearer",
"access_token": "fake_access_token"
}
```
ASSISTANT:
First, let's check the available APIs for the Spotify app.
Code:
```python
print(apis.api_docs.show_api_descriptions(app_name='spotify'))
```
USER:
Output:
```
[
...
{
"name": "show_playlist_library",
"description": "Get a list of playlists in the user's playlist library."
},
{
"name": "show_playlist",
"description": "Get detailed information about a specific playlist. You can view your own playlists or others' playlists if they are public."
},
...
]
```
ASSISTANT:
The `show_playlist_library` API seems to be the one we need. Let's check its specifications.
Code:
```python
print(apis.api_docs.show_api_doc(app_name='spotify', api_name='show_playlist_library'))
```
USER:
Output:
```
{
"app_name": "spotify",
"api_name": "show_playlist_library",
"path": "/private_playlists",
"method": "GET",
"description": "Get a list of playlists in the user's playlist library.",
"parameters": [
{
"name": "access_token",
"type": "string",
"required": true,
"description": "Access token obtained from spotify app login.",
"default": null,
"constraints": []
},
{
"name": "page_index",
"type": "integer",
"required": false,
"description": "The index of the page to retrieve.",
"default": 0,
"constraints": [
"value >= 0.0"
]
},
{
"name": "page_limit",
"type": "integer",
"required": false,
"description": "The maximum number of results to return per page.",
"default": 5,
"constraints": [
"value >= 1.0, <= 20.0"
]
},
{
"name": "is_public",
"type": "boolean",
"required": false,
"description": "Whether to show public playlists or private playlists.",
"default": null,
"constraints": []
}
],
"response_schema": [
{
"title": "string",
"created_at": "2019-01-01T00:00:00",
"is_public": true,
"rating": 0.0,
"like_count": 1,
"owner_email": "user@example.com",
"playlist_id": 1,
"song_ids": [
1
]
}
]
}
```
ASSISTANT:
I need to page through all the playlists to get the list of playlists and save it in `playlists`.
Code:
```python
page_index = 0
playlists = []
while page_index < 10:
playlist_page = apis.spotify.show_playlist_library(access_token=spotify_access_token, page_index=page_index)
if playlist_page:
playlists.extend(playlist_page)
page_index += 1
else:
break
num_playlists = len(playlists)
print(num_playlists)
```
USER:
Output:
```
23
```
ASSISTANT:
Now that the task is completed, I can call apis.supervisor.complete_task(). Since this task has an answer to be returned, I will pass that as an argument.
Code:
```python
apis.supervisor.complete_task(answer=num_playlists)
```
USER:
Output:
Marked the active task complete.
----------------------------------------------
USER:
**Key instructions**:
(1) Make sure to end code blocks with ``` followed by a newline(\n).
(2) Remember you can use the variables in your code in subsequent code blocks.
(3) Remember that the email addresses, access tokens and variables (e.g. spotify_password) in the example above are not valid anymore.
(4) You can use the "supervisor" app to get information about my accounts and use the "phone" app to get information about friends and family.
(5) Always look at API specifications (using apis.api_docs.show_api_doc) before calling an API.
(6) Write small chunks of code and only one chunk of code in every step. Make sure everything is working correctly before making any irreversible change.
(7) Many APIs return items in "pages". Make sure to run through all the pages by looping over `page_index`.
(8) Once you have completed the task, make sure to call apis.supervisor.complete_task(). If the task asked for some information, return it as the answer argument, i.e. call apis.supervisor.complete_task(answer=<answer>). Many tasks do not require an answer, so in those cases, just call apis.supervisor.complete_task() i.e. do not pass any argument.
USER:
Using these APIs, now generate code to solve the actual task:
My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}.
"""

View file

@ -1,197 +0,0 @@
# pylint: disable=E0611
"""Run the Appworld React Agent."""
import os
import json
import time
from pathlib import Path
import ray
import requests
from loguru import logger
from dotenv import load_dotenv
from appworld import load_task_ids
from appworld_react_agent import AppworldReactAgent
os.environ["APPWORLD_ROOT"] = "."
load_dotenv("../../.env")
def run_agent(
run_index: int,
max_workers: int,
model_name: str,
dataset_name: str,
experiment_suffix: str,
num_trials: int = 1,
use_memory: bool = False,
memory_base_url: str = "http://0.0.0.0:8002/",
use_memory_addition: bool = False,
use_memory_deletion: bool = False,
delete_freq: int = 10,
freq_threshold: int = 5,
utility_threshold: float = 0.5,
batch_size: int = 4,
):
"""Run the Appworld React Agent."""
experiment_name = dataset_name + "_" + experiment_suffix
path: Path = Path(f"./exp_result/{model_name}")
path.mkdir(parents=True, exist_ok=True)
task_ids = load_task_ids(dataset_name)
result: list = []
def dump_file():
with open(path / f"{experiment_name}.jsonl", "a", encoding="utf-8") as f:
for x in result:
f.write(json.dumps(x) + "\n")
if max_workers > 1:
# Process tasks in batches
total_tasks = len(task_ids)
num_batches = (total_tasks + batch_size - 1) // batch_size # Ceiling division
logger.info(f"Total tasks: {total_tasks}, Batch size: {batch_size}, Number of batches: {num_batches}")
for batch_idx in range(num_batches):
# Initialize Ray for this batch
start_idx = batch_idx * batch_size
end_idx = min(start_idx + batch_size, total_tasks)
batch_task_ids = task_ids[start_idx:end_idx]
logger.info(f"Starting batch {batch_idx + 1}/{num_batches} with {len(batch_task_ids)} tasks")
# Initialize Ray with the number of CPUs needed for this batch
ray.init(num_cpus=len(batch_task_ids))
future_list: list = []
for i, task_id in enumerate(batch_task_ids):
actor = AppworldReactAgent.remote(
index=start_idx + i,
model_name=model_name,
task_ids=[task_id],
experiment_name=experiment_name,
num_trials=num_trials,
use_memory=use_memory,
memory_base_url=memory_base_url,
use_memory_addition=use_memory_addition,
use_memory_deletion=use_memory_deletion,
delete_freq=delete_freq,
freq_threshold=freq_threshold,
utility_threshold=utility_threshold,
)
future = actor.execute.remote()
future_list.append(future)
time.sleep(1)
logger.info(f"Batch {batch_idx + 1} submit complete, waiting for results...")
# Collect results from this batch
for i, (task_id, future) in enumerate(zip(batch_task_ids, future_list)):
try:
t_result = ray.get(future)
if t_result:
if isinstance(t_result, list):
result.extend(t_result)
else:
result.append(t_result)
except Exception:
logger.exception(f"run ray error with task_id={task_id}")
logger.info(f"Batch {batch_idx + 1}: task {i + 1}/{len(batch_task_ids)} complete")
# Shutdown Ray to free resources before next batch
ray.shutdown()
logger.info(f"Batch {batch_idx + 1}/{num_batches} complete, Ray resources released")
# Optional: small delay between batches
if batch_idx < num_batches - 1:
time.sleep(2)
dump_file()
else:
agent = AppworldReactAgent(
index=run_index,
model_name=model_name,
task_ids=task_ids,
experiment_name=experiment_name,
num_trials=num_trials,
use_memory=use_memory,
memory_base_url=memory_base_url,
use_memory_addition=use_memory_addition,
use_memory_deletion=use_memory_deletion,
delete_freq=delete_freq,
freq_threshold=freq_threshold,
utility_threshold=utility_threshold,
)
result = agent.execute()
dump_file()
def handle_api_response(response: requests.Response):
"""Handle API response with proper error checking"""
if response.status_code != 200:
print(f"Error: {response.status_code}")
print(response.text)
return None
return response.json()
def load_memory(path: str = "docs/library", api_url: str = "http://0.0.0.0:8002/"):
"""Load memories from disk into the vector store"""
response = requests.post(
url=f"{api_url}load_memory",
json={
"load_file_path": path,
"clear_existing": True,
},
)
result = handle_api_response(response)
if result:
print(f"Memory loaded from {path}")
def main():
"""Main function to run the Appworld React Agent."""
max_workers = 16
batch_size = 8
num_runs = 4 # Number of runs
num_trials = 1 # for self-reflection
model_name = "qwen3-8b"
use_memory = True
use_memory_addition = False
use_memory_deletion = False
memory_base_url = "http://0.0.0.0:8002/"
if use_memory:
load_file_path = "docs/library/paper_data/task/appworld_qwen3_8b.jsonl"
load_memory(load_file_path, memory_base_url)
for i in range(num_runs):
run_agent(
run_index=i,
max_workers=max_workers,
model_name=model_name,
dataset_name="test_normal",
experiment_suffix="with-fixed-memory",
num_trials=num_trials,
use_memory=use_memory,
memory_base_url=memory_base_url,
use_memory_addition=use_memory_addition,
use_memory_deletion=use_memory_deletion,
delete_freq=5,
freq_threshold=5,
utility_threshold=0.5,
batch_size=batch_size,
)
if __name__ == "__main__":
main()

View file

@ -1,164 +0,0 @@
"""Run the experiment statistic."""
import json
from collections import defaultdict
from pathlib import Path
import pandas as pd
from loguru import logger
def calculate_best_at_k(scores: list, k: int) -> float:
"""
Calculate best@k
Divide scores into groups of size k, take the maximum value in each group,
then average these maximum values
Args:
scores: List of after_score values for all runs of a task
k: Group size
Returns:
best@k value
"""
if len(scores) % k != 0:
raise ValueError(f"Length of scores ({len(scores)}) must be divisible by k ({k})")
group_maxs = []
for i in range(0, len(scores), k):
group = scores[i : i + k]
group_maxs.append(max(group))
return sum(group_maxs) / len(group_maxs)
def calculate_pass_at_k(scores: list, k: int) -> float:
"""Calculate pass@k."""
if len(scores) % k != 0:
raise ValueError(f"Length of scores ({len(scores)}) must be divisible by k ({k})")
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_maxs.append(is_pass)
return sum(group_maxs) / len(group_maxs)
def get_possible_k_values(total_runs: int) -> list:
"""
Get all possible k values (factors of total_runs)
Args:
total_runs: Total number of runs
Returns:
List of k values in descending order
"""
k_values = []
for k in range(1, total_runs + 1):
if total_runs % k == 0:
k_values.append(k)
return sorted(k_values, reverse=True) # Sort from large to small
def run_exp_statistic():
"""Run the experiment statistic."""
path: Path = Path("./exp_result/qwen3-8b")
# Store results for all experiments
all_results = {}
for file in path.glob("*.jsonl"): # [f for f in path.glob("*.jsonl") if not f.stem[-1].isdigit()]
# Group results by task_id
task_results = defaultdict(list)
with open(file, "r", encoding="utf-8") as f:
for line in f:
if not line.strip():
continue
data = json.loads(line)
if isinstance(data, list):
for part_data in data:
task_id = part_data["task_id"]
after_score = part_data["after_score"]
task_results[task_id].append(after_score)
else:
task_id = data["task_id"]
after_score = data["after_score"]
task_results[task_id].append(after_score)
if not task_results:
logger.warning(f"No valid data found in file {file}")
continue
# Check if each task has consistent number of runs
run_counts = [len(scores) for scores in task_results.values()]
if len(set(run_counts)) > 1:
logger.warning(f"Inconsistent number of runs for different tasks in file {file}: {set(run_counts)}")
continue
num_runs = run_counts[0]
logger.info(f"File {file}: {len(task_results)} tasks, {num_runs} runs per task")
# Get all possible k values
k_values = get_possible_k_values(num_runs)
logger.info(f"Calculable best@k values: {k_values}")
# Calculate various best@k values
file_results = {"file": file.name}
for k in k_values:
best_at_k_scores = []
pass_at_k_scores = []
for task_id, scores in task_results.items():
try:
best_k_score = calculate_best_at_k(scores, k)
pass_at_k_score = calculate_pass_at_k(scores, k)
pass_at_k_scores.append(pass_at_k_score)
best_at_k_scores.append(best_k_score)
except ValueError as e:
logger.error(f"Error calculating best@{k} for task {task_id}: {e}")
continue
if best_at_k_scores:
avg_best_at_k = sum(best_at_k_scores) / len(best_at_k_scores)
file_results[f"best@{k}"] = avg_best_at_k
logger.info(f"file={file.name} best@{k}={avg_best_at_k:.4f}")
if pass_at_k_scores:
avg_pass_at_k = sum(pass_at_k_scores) / len(pass_at_k_scores)
file_results[f"pass@{k}"] = avg_pass_at_k
logger.info(f"file={file.name} pass@{k}={avg_pass_at_k:.4f}")
all_results[file.name] = file_results
# Create and display table
if all_results:
df = pd.DataFrame(list(all_results.values()))
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@")]
# best_columns = [col for col in df.columns]
pass_columns.sort(key=lambda x: x, reverse=False)
df = df[pass_columns]
print("\n" + "=" * 80)
print("Experiment Results Summary Table")
print("=" * 80)
print(df.round(4))
print("=" * 80)
# Save table to CSV
output_path = path / "experiment_summary.csv"
df.to_csv(output_path)
logger.info(f"Results table saved to: {output_path}")
else:
logger.warning("No valid experiment results found")
if __name__ == "__main__":
run_exp_statistic()

View file

@ -1,726 +0,0 @@
# flake8: noqa: E402
# pylint: disable=too-many-return-statements
"""A minimal ReAct Agent for BFCL-v3(multi-turn) tasks."""
import re
import os
import time
import json
import warnings
import tempfile
import datetime
from pathlib import Path
from typing import Dict, List, Any
import ray
import requests
from tqdm import tqdm
from loguru import logger
from openai import OpenAI
from dotenv import load_dotenv
from bfcl_utils import (
load_test_case,
handle_user_turn,
handle_tool_calls,
extract_tool_schema,
extract_single_turn_response,
extract_multi_turn_responses,
capture_and_print_score_files,
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 (
is_empty_execute_response,
)
from bfcl_eval.eval_checker.eval_runner import (
multi_turn_runner,
ast_file_runner,
)
from bfcl_eval.eval_checker.eval_runner_helper import record_cost_latency
from bfcl_eval.utils import (
is_multi_turn,
is_relevance_or_irrelevance,
find_file_with_suffix,
load_file,
)
os.environ["BFCL_DATA_PATH"] = "data/multiturn_data_base_val.jsonl"
os.environ["BFCL_ANSWER_PATH"] = "data/possible_answer"
load_dotenv("../../.env")
@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_trials: 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:8002/",
):
self.index: int = index
self.task_ids: List[str] = task_ids
self.categories: List[str] = [task_id.rsplit("_", 1)[0] if "_" in task_id else task_id for task_id in task_ids]
self.experiment_name: str = experiment_name
self.data_path: str = data_path
self.answer_path: Path = answer_path
self.model_name: str = model_name
self.temperature: float = temperature
self.max_interactions: int = max_interactions
self.max_response_size: int = max_response_size
self.num_trials: int = num_trials
self.enable_thinking: bool = enable_thinking
self.use_memory: bool = use_memory
self.use_memory_addition: bool = use_memory_addition if use_memory else False
self.use_memory_deletion: bool = use_memory_deletion if use_memory else False
self.delete_freq: int = delete_freq
self.freq_threshold: int = freq_threshold
self.utility_threshold: float = utility_threshold
self.memory_base_url: str = memory_base_url
self.llm_client = OpenAI()
self.history: List[List[List[dict]]] = [[] for _ in range(num_trials)]
self.retrieved_memory_list: List[List[List[Any]]] = [[] for _ in range(num_trials)]
self.test_entry: List[List[Dict[str, Any]]] = [[] for _ in range(num_trials)]
self.original_test_entry: List[List[Dict[str, Any]]] = [[] for _ in range(num_trials)]
self.tool_schema: List[List[List[dict]]] = [[] for _ in range(num_trials)]
self.current_turn = [[0 for _ in range(len(task_ids))] for _ in range(num_trials)]
for run_id in range(num_trials):
for task_index in range(len(task_ids)):
self.init_state(run_id, task_index)
def init_state(self, run_id, i) -> Dict[str, Any]:
"""Initialize the state of the agent."""
self.test_entry[run_id].append(load_test_case(self.data_path, self.task_ids[i]))
self.original_test_entry[run_id].append(self.test_entry[run_id][i].get("extra", {}))
self.tool_schema[run_id].append(extract_tool_schema(self.test_entry[run_id][i].get("tools", [{}])))
msg = self.test_entry[run_id][i].get("messages", [])
self.history[run_id].append(msg)
self.retrieved_memory_list[run_id].append([])
self.current_turn[run_id][i] = 1
def update_task_history_with_memory(self, run_id, task_index, previous_memories: None):
"""Update the task history with memory."""
query = self.history[run_id][task_index][0]["content"]
if len(previous_memories) == 0:
response = self.get_memory(query)
if response and "memory_list" in response["metadata"]:
self.retrieved_memory_list[run_id][task_index] = response["metadata"]["memory_list"]
task_memory = re.sub(r"\bMemory\s*(\d+)\s*[:]", r"Experience \1 :", response["answer"])
logger.info(f"loaded task_memory: {task_memory}")
self.history[run_id][task_index][0] = self.get_query_with_memory(query, task_memory)
else:
formatted_memories = []
for i, memory in enumerate(previous_memories, 1):
condition = memory["when_to_use"]
memory_content = memory["content"]
memory_text = f"Experience {i} :\n When to use: {condition}\n Content: {memory_content}\n"
formatted_memories.append(memory_text)
self.history[run_id][task_index][0] = self.get_query_with_memory(query, "\n".join(formatted_memories))
def get_query_with_memory(self, query: str, memory: str):
"""Get the query with memory."""
return {
"role": "user",
"content": "Task:\n" + query + "\n\nSome Related Experience to help you to complete the task:\n" + memory,
}
def get_query_without_experience(self, query: str):
"""Get the query without experience."""
if "\n\nSome Related Experience" in query:
query = query.split("\n\nSome Related Experience")[0].split("Task:\n")[-1]
return query
def get_traj_from_task_history(self, task_id: str, task_history: list, reward: float):
"""Get the trajectory from the task history."""
return {
"task_id": task_id,
"messages": task_history,
"score": reward,
}
def handle_api_response(self, response: requests.Response):
"""Handle API response with proper error checking"""
if response.status_code != 200:
print(f"Error: {response.status_code}")
print(response.text)
return None
return response.json()
def get_memory(self, query: str):
"""Retrieve relevant task memories based on a query"""
response = requests.post(
url=f"{self.memory_base_url}retrieve_task_memory",
json={
"query": query,
"enable_llm_rerank": False,
"enable_score_filter": False,
"top_k": 5,
"enable_llm_rewrite": False,
},
)
result = self.handle_api_response(response)
if not result:
return None
logger.info(f"query: {query}, response: {result}")
return result
def summary_memory(self, trajectories):
"""Generate a summary of conversation messages and create task memories"""
response = requests.post(
url=f"{self.memory_base_url}summary_task_memory",
json={
"trajectories": trajectories,
"success_threshold": 1.0,
"enable_soft_comparison": True,
"validation_threshold": 0.5,
},
)
result = self.handle_api_response(response)
if not result:
return []
# Extract memory list from response
memory_list = result.get("metadata", {}).get("memory_list", [])
logger.info(f"add new memories: {memory_list}")
return memory_list
def add_memory(self, memory_list):
"""Add the memory to the memory pool."""
response = requests.post(
url=f"{self.memory_base_url}add_task_memory",
json={
"memory_list": memory_list,
},
)
response.raise_for_status()
def update_memory_information(self, memory_list, update_utility: bool = False):
"""Update the memory information."""
response = requests.post(
url=f"{self.memory_base_url}record_task_memory",
json={
"memory_list": memory_list,
"update_utility": update_utility,
},
)
response.raise_for_status()
logger.info(response.json())
def delete_memory(self):
"""Delete the memory from the memory pool."""
response = requests.post(
url=f"{self.memory_base_url}delete_task_memory",
json={
"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:
"""Call the LLM."""
for i in range(100):
try:
response = self.llm_client.chat.completions.create(
model=self.model_name,
messages=messages,
tools=tool_schemas,
temperature=self.temperature,
seed=0,
extra_body={"enable_thinking": self.enable_thinking},
stream=self.enable_thinking,
parallel_tool_calls=True,
)
if not self.enable_thinking:
out_msg = response.choices[0].message
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
)
for chunk in response:
if not chunk.choices:
# Handle usage information
continue
delta = chunk.choices[0].delta
# Handle AI's thought process (chain reasoning)
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": ""},
},
)
# Collect tool call ID (used for subsequent function calls)
if 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
# 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
msg = {
"role": "assistant",
"content": answer_content,
"reasoning_content": reasoning_content,
}
if tool_info:
msg["tool_calls"] = tool_info
return msg
except Exception as e:
logger.exception(f"encounter error with {e.args}")
time.sleep(1 + i * 10)
return "call llm error"
def env_step(self, run_id: int, index: int, messages: str) -> str:
"""
Process one step in the conversation.
Both single turn and multi turn are supported.
Args:
messages: List of conversation messages, with the last one being assistant response
test_entry: Test entry containing initial_config, involved_classes, question etc.
**kwargs: Additional arguments for compatibility
Returns:
Dict containing next message and tools if applicable
"""
try:
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",
)
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,
)
# 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)}",
)
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],
)
except Exception as e:
warnings.warn(f"Errors during tool invocation: {str(e)}")
return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index])
else:
return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index])
except Exception as e:
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]],
) -> List[str]:
"""
Convert OpenAI format tool calls to execution format.
Args:
tool_calls: List of tool calls in OpenAI format
Returns:
List of function calls in string format
"""
execution_list = []
for tool_call in tool_calls:
function = tool_call.get("function", {})
function_name = function.get("name", "")
try:
arguments = function.get("arguments", "{}")
if isinstance(arguments, str):
args_dict = json.loads(arguments)
else:
args_dict = arguments
args_str = ", ".join([f"{k}={repr(v)}" for k, v in args_dict.items()])
execution_list.append(f"{function_name}({args_str})")
except Exception:
execution_list.append(f"{function_name}()")
return execution_list
def get_reward(self, run_id, index) -> float:
"""Get the reward."""
try:
if not self.history[run_id][index] or not self.original_test_entry[run_id][index]:
return 0.0
model_name = "env_handler"
handler = QwenAPIHandler(
model_name,
temperature=1.0,
) # FIXME: magic number
model_result_data = self._convert_conversation_to_eval_format(run_id, index)
prompt_data = [self.original_test_entry[run_id][index]]
state = {"leaderboard_table": {}}
record_cost_latency(
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,
)
else:
# Find the corresponding possible answer file
possible_answer_file = find_file_with_suffix(
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]]
if is_multi_turn(self.categories[index]):
accuracy, _ = self._eval_multi_turn_test(
handler,
model_result_data,
prompt_data,
possible_answer,
model_name,
self.categories[index],
)
else:
accuracy, _ = self._eval_single_turn_test(
handler,
model_result_data,
prompt_data,
possible_answer,
model_name,
self.categories[index],
)
print(f"model_result_data: {model_result_data}")
if possible_answer:
print(f"possible_answer: {possible_answer}")
else:
print("possible_answer: None")
return accuracy
except Exception:
import traceback
traceback.print_exc()
return 0
def _convert_conversation_to_eval_format(self, run_id, index) -> Dict[str, Any]:
"""
Convert conversation history to evaluation format.
Args:
conversation_result: Result from run_conversation
original_test_entry: Original test entry data
Returns:
Data in format expected by multi_turn_runner or other runners
"""
if is_multi_turn(self.categories[index]):
turns_data = extract_multi_turn_responses(self.history[run_id][index])
else:
turns_data = extract_single_turn_response(self.history[run_id][index])
model_result_data = {
"id": self.task_ids[index],
"result": turns_data,
"latency": 0,
"input_token_count": 0,
"output_token_count": 0,
}
return model_result_data
def _eval_multi_turn_test(
self,
handler,
model_result_data,
prompt_data,
possible_answer,
model_name,
test_category,
):
"""
Evaluate multi-turn test.
Args:
handler: Model handler instance
model_result_data: Model result data
prompt_data: Prompt data
possible_answer: Possible answer data
model_name: Name of the model
test_category: Category of the test
Returns:
Tuple of (accuracy, total_count)
"""
with tempfile.TemporaryDirectory() as temp_dir:
score_dir = Path(temp_dir)
accuracy, total_count = multi_turn_runner(
handler=handler,
model_result=[model_result_data],
prompt=prompt_data,
possible_answer=possible_answer,
model_name=model_name,
test_category=test_category,
score_dir=score_dir,
)
capture_and_print_score_files(
score_dir,
model_name,
test_category,
"multi_turn",
)
return accuracy, total_count
def _eval_single_turn_test(
self,
handler,
model_result_data,
prompt_data,
possible_answer,
model_name,
test_category,
):
"""
Evaluate single-turn AST test.
Args:
handler: Model handler instance
model_result_data: Model result data
prompt_data: Prompt data
possible_answer: Possible answer data
model_name: Name of the model
test_category: Category of the test
Returns:
Tuple of (accuracy, total_count)
"""
language = "Python"
if "java" in test_category.lower():
language = "Java"
elif "js" in test_category.lower() or "javascript" in test_category.lower():
language = "JavaScript"
with tempfile.TemporaryDirectory() as temp_dir:
score_dir = Path(temp_dir)
accuracy, total_count = ast_file_runner(
handler=handler,
model_result=[model_result_data],
prompt=prompt_data,
possible_answer=possible_answer,
language=language,
test_category=test_category,
model_name=model_name,
score_dir=score_dir,
)
capture_and_print_score_files(
score_dir,
model_name,
test_category,
"single_turn",
)
return accuracy, total_count
def execute(self):
"""Execute the agent."""
result = []
counter = 0
for task_index, task_id in enumerate(tqdm(self.task_ids, desc=f"ray_index={self.index}")):
t_result = None
previous_memories = []
for run_id in range(self.num_trials):
try:
start_time = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
for i in range(self.max_interactions):
if self.use_memory and i == 0:
self.update_task_history_with_memory(run_id, task_index, previous_memories)
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}
# 2. Returns tool invocation result: {"messages":
# [{"role": "tool", "content": {<exec_results>}, 'tool_call_id': 'chatcmpl-tool-xxx'}]}
# <exec_results>: when success, returns result dicts, e.g., {"travel_cost_list": [x]},
# when error, returns error message,
# e.g., {"error": "cd: temporary: No such directory. You cannot use path ..."}
# 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 = []
next_user_msg = ""
for idx, msg in enumerate(env_output.get("messages", [])):
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
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]},
)
else:
self.history[run_id][task_index].append({"role": "user", "content": next_user_msg})
logger.info(f"index={self.index} task_id={task_id} iteration={i}")
if self.task_completed(run_id, task_index):
break
reward = self.get_reward(run_id, task_index)
if self.use_memory:
if self.use_memory_addition:
new_traj_list = [
self.get_traj_from_task_history(task_id, self.history[run_id][task_index], reward),
]
previous_memories = self.summary_memory(new_traj_list)
if reward == 1:
self.add_memory(previous_memories)
# update the freq & utility attributes of retrieved memories
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()
t_result = {
"run_id": run_id,
"task_id": self.task_ids[task_index],
"experiment_name": self.experiment_name,
"task_completed": self.task_completed(run_id, task_index),
"reward": reward,
"task_history": self.history[run_id][task_index],
"task_start_time": start_time,
}
if reward == 1:
break
except Exception as e:
logger.exception(f"encounter error with {e.args}")
result.append(t_result)
return result
def task_completed(self, run_id, index):
"""
Check if task is completed.
Returns:
True if task is completed, False otherwise
"""
return self.history[run_id][index][-1]["content"] == "[CONVERSATION_COMPLETED]"
def main():
"""Main function to run the BFCLAgent."""
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_ids=[task_ids[0]],
experiment_name=f"qwen3_8b_{dataset_name}",
)
result = agent.execute()
logger.info(f"result={json.dumps(result)}")
if __name__ == "__main__":
main()

View file

@ -1,399 +0,0 @@
"""Utils for evaluation on BFCL tasks"""
import json
from pathlib import Path
from typing import Dict, List, Any
from bfcl_eval.constants.default_prompts import (
DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC,
)
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]:
"""
load test cases by id
"""
if not Path(data_path).exists():
raise FileNotFoundError(f"BFCL data file '{data_path}' not found")
if test_id is None:
raise ValueError("task_id is required")
with open(data_path, "r", encoding="utf-8") as f:
if str(test_id).isdigit(): # pylint: disable=R1720
idx = int(test_id)
for line_no, line in enumerate(f):
if line_no == idx:
return json.loads(line)
raise ValueError(f"Test case index {idx} not found in {data_path}")
else:
for line in f:
data = json.loads(line)
if data.get("id") == test_id:
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,
) -> Dict[str, Any]:
"""
Handle user turn by returning appropriate content from test_entry["question"].
For non-first turns, processes user query and tools.
Args:
test_entry: Test entry containing conversation data
current_turn: Current turn number
Returns:
Response containing next user message and tools
"""
try:
current_turn_message = []
tools = compile_tools(test_entry)
questions = test_entry.get("question", [])
holdout_function = test_entry.get("holdout_function", {})
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."
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)
except Exception as e:
return create_error_response(f"Failed to process user message: {str(e)}")
def handle_tool_calls( # pylint: disable=W0613
tool_calls: List[Dict[str, Any]],
decoded_calls: list[str],
test_entry: Dict[str, Any],
current_turn: int,
) -> Dict[str, Any]:
"""
Handle tool calls from assistant.
Args:
tool_calls: List of tool calls in OpenAI format
decoded_calls: List of decoded function calls
test_entry: Test entry containing environment data
current_turn: Current turn number
Returns:
Response containing tool execution results
"""
execution_results, _ = execute_multi_turn_func_call(
func_call_list=decoded_calls,
initial_config=test_entry["initial_config"],
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"]),
is_evaL_run=False,
)
# print('execution_results in handler_tool_calls:', execution_results)
return create_tool_response(tool_calls, execution_results)
def compile_tools(test_entry: dict) -> list:
"""
Compile functions into tools format.
Args:
test_entry: Test entry containing functions
Returns:
List of tools in OpenAI format
"""
functions: list = test_entry["function"]
test_category: str = test_entry["id"].rsplit("_", 1)[0]
functions = func_doc_language_specific_pre_processing(functions, test_category)
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],
) -> Dict[str, Any]:
"""
Create response for tool calls.
Args:
tool_calls: List of tool calls
execution_results: List of execution results
Returns:
Response containing tool execution results
"""
tool_messages = []
for i, (tool_call, result) in enumerate(zip(tool_calls, execution_results)):
tool_messages.append(
{
"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]],
) -> Dict[str, Any]:
"""
Create response containing user message.
Args:
question_turn: List of messages for current turn
tools: List of available tools
Returns:
Response containing user message and tools
"""
user_content = ""
for msg in question_turn:
if msg["role"] == "user":
user_content = msg["content"]
break
return {"messages": [{"role": "user", "content": user_content}], "tools": tools}
def create_completion_response() -> Dict[str, Any]:
"""
Create response indicating conversation completion.
Returns:
Response with completion message
"""
return {"messages": [{"role": "env", "content": "[CONVERSATION_COMPLETED]"}]}
def create_error_response(error_message: str) -> Dict[str, Any]:
"""
Create response for error conditions.
Args:
error_message: Error message to include
Returns:
Response containing error message
"""
return {"messages": [{"role": "env", "content": f"[ERROR] {error_message}"}]}
def decode_execute(result):
"""
Decode execute results for compatibility with evaluation framework.
Args:
result: Result to decode
Returns:
List of decoded function calls
"""
return default_decode_execute_prompting(result)
def extract_single_turn_response(messages: List[Dict[str, Any]]) -> str:
"""
Extract single-turn response from conversation messages.
Args:
messages: List of conversation messages
Returns:
String representation of the response
"""
for message in reversed(messages):
if message["role"] == "assistant":
if "tool_calls" in message and message["tool_calls"]:
formatted_calls = []
for tool_call in message["tool_calls"]:
formatted_call = format_single_tool_call_for_eval(
tool_call,
)
if formatted_call:
formatted_calls.append(formatted_call)
return "\n".join(formatted_calls) if formatted_calls else ""
elif message.get("content"):
return message["content"]
return ""
def extract_multi_turn_responses(
messages: List[Dict[str, Any]],
) -> List[List[str]]:
"""
Extract multi-turn responses from conversation messages.
Args:
messages: List of conversation messages
Returns:
List of turns, each turn is a list of function call strings
"""
turns_data = []
current_turn_responses = []
i = 0
while i < len(messages):
message = messages[i]
if message["role"] == "user":
if current_turn_responses:
turns_data.append(current_turn_responses)
current_turn_responses = []
i += 1
while i < len(messages) and messages[i]["role"] == "assistant":
assistant_msg = messages[i]
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,
)
if formatted_call:
current_turn_responses.append(formatted_call)
i += 1
while i < len(messages) and messages[i]["role"] == "tool":
i += 1
else:
i += 1
if current_turn_responses:
turns_data.append(current_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
Returns:
Formatted string representation
"""
function = tool_call.get("function", {})
function_name = function.get("name", "")
try:
arguments = function.get("arguments", "{}")
if isinstance(arguments, str):
args_dict = json.loads(arguments)
else:
args_dict = arguments
args_str = ", ".join([f"{k}={repr(v)}" for k, v in args_dict.items()])
return f"{function_name}({args_str})"
except Exception:
return f"{function_name}()"
def capture_and_print_score_files(
score_dir: Path,
model_name: str,
test_category: str,
eval_type: str,
):
"""
Capture and print contents of score files written to score_dir.
Args:
score_dir: Directory containing score files
model_name: Name of the model
test_category: Category of the test
eval_type: Type of evaluation (relevance/multi_turn/single_turn)
"""
try:
print(f"\n=== {eval_type.upper()} Evaluation Result Files ===")
print(f"Model: {model_name}")
print(f"Test Category: {test_category}")
print(f"Evaluation Type: {eval_type}")
for file_path in score_dir.rglob("*"):
if file_path.is_file():
relative_path = file_path.relative_to(score_dir)
print(f"\n--- File: {relative_path} ---")
try:
with open(file_path, "r", encoding="utf-8") as f:
content = f.read()
if (
file_path.suffix == ".json"
or content.strip().startswith("{")
or content.strip().startswith("[")
):
try:
lines = content.strip().split("\n")
formatted_lines = []
for line in lines:
if line.strip():
parsed = json.loads(line)
formatted_lines.append(
json.dumps(
parsed,
ensure_ascii=False,
indent=2,
),
)
content = "\n".join(formatted_lines)
except json.JSONDecodeError:
pass
print(content)
except UnicodeDecodeError:
print(f"[Binary file, size: {file_path.stat().st_size} bytes]")
except Exception as e:
print(f"[Error reading file: {str(e)}]")
print(f"=== {eval_type.upper()} Evaluation Result Files End ===\n")
except Exception as e:
print(f"Error capturing evaluation result files: {str(e)}")
def extract_tool_schema(tools):
"""Reformat tool schema"""
for i in range(len(tools)): # pylint: disable=C0200
tools[i]["function"].pop("response")
return tools

View file

@ -1,206 +0,0 @@
# pylint: disable=C0114
DEFAULT_TRAIN_IDS: set[str] = {
"multi_turn_base_102",
"multi_turn_base_107",
"multi_turn_base_110",
"multi_turn_base_114",
"multi_turn_base_115",
"multi_turn_base_118",
"multi_turn_base_122",
"multi_turn_base_123",
"multi_turn_base_128",
"multi_turn_base_13",
"multi_turn_base_130",
"multi_turn_base_132",
"multi_turn_base_133",
"multi_turn_base_143",
"multi_turn_base_144",
"multi_turn_base_146",
"multi_turn_base_15",
"multi_turn_base_158",
"multi_turn_base_169",
"multi_turn_base_17",
"multi_turn_base_172",
"multi_turn_base_176",
"multi_turn_base_182",
"multi_turn_base_187",
"multi_turn_base_197",
"multi_turn_base_199",
"multi_turn_base_22",
"multi_turn_base_23",
"multi_turn_base_24",
"multi_turn_base_36",
"multi_turn_base_40",
"multi_turn_base_44",
"multi_turn_base_47",
"multi_turn_base_48",
"multi_turn_base_5",
"multi_turn_base_51",
"multi_turn_base_59",
"multi_turn_base_63",
"multi_turn_base_65",
"multi_turn_base_66",
"multi_turn_base_67",
"multi_turn_base_68",
"multi_turn_base_70",
"multi_turn_base_75",
"multi_turn_base_77",
"multi_turn_base_78",
"multi_turn_base_79",
"multi_turn_base_81",
"multi_turn_base_83",
"multi_turn_base_93",
}
DEFAULT_VAL_IDS: set[str] = {
"multi_turn_base_0",
"multi_turn_base_1",
"multi_turn_base_10",
"multi_turn_base_100",
"multi_turn_base_101",
"multi_turn_base_103",
"multi_turn_base_104",
"multi_turn_base_105",
"multi_turn_base_106",
"multi_turn_base_108",
"multi_turn_base_109",
"multi_turn_base_11",
"multi_turn_base_111",
"multi_turn_base_112",
"multi_turn_base_113",
"multi_turn_base_116",
"multi_turn_base_117",
"multi_turn_base_119",
"multi_turn_base_12",
"multi_turn_base_120",
"multi_turn_base_121",
"multi_turn_base_124",
"multi_turn_base_125",
"multi_turn_base_126",
"multi_turn_base_127",
"multi_turn_base_129",
"multi_turn_base_131",
"multi_turn_base_134",
"multi_turn_base_135",
"multi_turn_base_136",
"multi_turn_base_137",
"multi_turn_base_138",
"multi_turn_base_139",
"multi_turn_base_14",
"multi_turn_base_140",
"multi_turn_base_141",
"multi_turn_base_142",
"multi_turn_base_145",
"multi_turn_base_147",
"multi_turn_base_148",
"multi_turn_base_149",
"multi_turn_base_150",
"multi_turn_base_151",
"multi_turn_base_152",
"multi_turn_base_153",
"multi_turn_base_154",
"multi_turn_base_155",
"multi_turn_base_156",
"multi_turn_base_157",
"multi_turn_base_159",
"multi_turn_base_16",
"multi_turn_base_160",
"multi_turn_base_161",
"multi_turn_base_162",
"multi_turn_base_163",
"multi_turn_base_164",
"multi_turn_base_165",
"multi_turn_base_166",
"multi_turn_base_167",
"multi_turn_base_168",
"multi_turn_base_170",
"multi_turn_base_171",
"multi_turn_base_173",
"multi_turn_base_174",
"multi_turn_base_175",
"multi_turn_base_177",
"multi_turn_base_178",
"multi_turn_base_179",
"multi_turn_base_18",
"multi_turn_base_180",
"multi_turn_base_181",
"multi_turn_base_183",
"multi_turn_base_184",
"multi_turn_base_185",
"multi_turn_base_186",
"multi_turn_base_188",
"multi_turn_base_189",
"multi_turn_base_19",
"multi_turn_base_190",
"multi_turn_base_191",
"multi_turn_base_192",
"multi_turn_base_193",
"multi_turn_base_194",
"multi_turn_base_195",
"multi_turn_base_196",
"multi_turn_base_198",
"multi_turn_base_2",
"multi_turn_base_20",
"multi_turn_base_21",
"multi_turn_base_25",
"multi_turn_base_26",
"multi_turn_base_27",
"multi_turn_base_28",
"multi_turn_base_29",
"multi_turn_base_3",
"multi_turn_base_30",
"multi_turn_base_31",
"multi_turn_base_32",
"multi_turn_base_33",
"multi_turn_base_34",
"multi_turn_base_35",
"multi_turn_base_37",
"multi_turn_base_38",
"multi_turn_base_39",
"multi_turn_base_4",
"multi_turn_base_41",
"multi_turn_base_42",
"multi_turn_base_43",
"multi_turn_base_45",
"multi_turn_base_46",
"multi_turn_base_49",
"multi_turn_base_50",
"multi_turn_base_52",
"multi_turn_base_53",
"multi_turn_base_54",
"multi_turn_base_55",
"multi_turn_base_56",
"multi_turn_base_57",
"multi_turn_base_58",
"multi_turn_base_6",
"multi_turn_base_60",
"multi_turn_base_61",
"multi_turn_base_62",
"multi_turn_base_64",
"multi_turn_base_69",
"multi_turn_base_7",
"multi_turn_base_71",
"multi_turn_base_72",
"multi_turn_base_73",
"multi_turn_base_74",
"multi_turn_base_76",
"multi_turn_base_8",
"multi_turn_base_80",
"multi_turn_base_82",
"multi_turn_base_84",
"multi_turn_base_85",
"multi_turn_base_86",
"multi_turn_base_87",
"multi_turn_base_88",
"multi_turn_base_89",
"multi_turn_base_9",
"multi_turn_base_90",
"multi_turn_base_91",
"multi_turn_base_92",
"multi_turn_base_94",
"multi_turn_base_95",
"multi_turn_base_96",
"multi_turn_base_97",
"multi_turn_base_98",
"multi_turn_base_99",
}

View file

@ -1,235 +0,0 @@
# pylint: disable=W0621,W1514
"""Init task memory pool"""
import argparse
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]:
"""
load training cases by id
"""
if not Path(data_path).exists():
raise FileNotFoundError(f"BFCL data file '{data_path}' not found")
if task_id is None:
raise ValueError("task_id is required")
with open(data_path, "r", encoding="utf-8") as f:
if str(task_id).isdigit(): # pylint: disable=R1720
idx = int(task_id)
for line_no, line in enumerate(f):
if line_no == idx:
return json.loads(line)
raise ValueError(f"Task case index {idx} not found in {data_path}")
else:
for line in f:
data = json.loads(line)
if data.get("id") == task_id:
return data
raise ValueError(f"Task case id '{task_id}' not found in {data_path}")
def get_tool_prompt(tools):
"""Construct prompt with provided 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>'
)
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 _, trajectories in grouped.items():
if len(trajectories) == 1:
# when only one trajectory, retain it
filtered_groups.append(trajectories)
elif len(trajectories) == 2:
# when there are two trajectories, retain them
filtered_groups.append(trajectories)
else:
# when there are more than two trajectories, choose the two with the highest and lowest rewards
trajectories.sort(key=lambda t: t["reward"])
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) -> Dict[str, Any]:
"""
post trajectories to summarizer service
Args:
trajectories: trajectory list
service_url: summarizer service URL
Returns:
response json
"""
trajectory_dicts = [
{
"task_id": traj["task_id"],
"messages": traj["task_history"],
"score": traj["reward"],
}
for traj in trajectories
]
request_data = {
"trajectories": trajectory_dicts,
"success_threshold": 1.0,
"enable_soft_comparison": True,
"validation_threshold": 0.5,
}
try:
response = requests.post(f"{service_url}/summary_task_memory", json=request_data)
response.raise_for_status()
return response.json()
except Exception as e:
return {"error": str(e), "trajectories_count": len(trajectories)}
def process_trajectories_with_threads(
grouped_trajectories: List[List[Any]],
service_url: 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
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): i for i, group in enumerate(grouped_trajectories)
}
for future in as_completed(future_to_group):
group_index = future_to_group[future]
try:
result = future.result()
result["group_index"] = group_index
result["group_size"] = len(grouped_trajectories[group_index])
results.append(result)
if "memory_list" in result["metadata"]:
print(f'✅ Group {group_index} processed: {result["metadata"].get("memory_list", 0)}')
memory_list = result["metadata"].get("memory_list", [])
response = requests.post(url=f"{service_url}/add_task_memory", json={"memory_list": memory_list})
response.raise_for_status()
else:
print(f"❌ Group {group_index} processed: error")
except Exception as e:
error_result = {
"group_index": group_index,
"group_size": len(grouped_trajectories[group_index]),
"error": str(e),
}
results.append(error_result)
print(f"❌ Group {group_index} failed: {e}")
return results
def main():
"""Main function to convert JSONL to memories using ReMe service."""
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:8002", help="ReMe service URL")
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"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,
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_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 = {
"jsonl_file": args.jsonl_file,
"total_groups": len(grouped_trajectories),
"success_count": success_count,
"error_count": error_count,
"total_task_memories": total_memories,
"results": results,
}
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:
print(f"Error saving results: {e}")
if __name__ == "__main__":
main()

View file

@ -1,73 +0,0 @@
# pylint: disable=W0621
"""Preprocess multi-turn test cases"""
import json
from pathlib import Path
from bfcl_eval.model_handler.model_style import ModelStyle
from bfcl_eval.eval_checker.eval_runner_helper import load_file
from bfcl_eval.constants.type_mappings import GORILLA_TO_OPENAPI
from bfcl_eval.constants.eval_config import MULTI_TURN_FUNC_DOC_PATH
from bfcl_eval.constants.category_mapping import MULTI_TURN_FUNC_DOC_FILE_MAPPING
from bfcl_eval.model_handler.utils import (
convert_to_tool,
func_doc_language_specific_pre_processing,
)
def process_multi_turn_test_case(file_path, output_path):
"""
Multi-turn test cases don't have the function doc in the prompt. We need to add them here.
"""
test_cases = []
with open(output_path, "w", encoding="utf-8") as outf:
with open(file_path, encoding="utf-8") as f:
file = f.readlines()
for line in file:
entry = json.loads(line)
if "multi_turn" not in entry["id"]:
continue
test_category: str = entry["id"].rsplit("_", 1)[0]
involved_classes = entry["involved_classes"]
entry["function"] = []
for func_collection in involved_classes:
# func_doc is a list of dict
func_doc = load_file(
MULTI_TURN_FUNC_DOC_PATH / MULTI_TURN_FUNC_DOC_FILE_MAPPING[func_collection],
)
entry["function"].extend(func_doc)
# Handle Miss Func category; we need to remove the holdout function doc
if "missed_function" in entry:
for turn_index, missed_func_names in entry["missed_function"].items():
entry["missed_function"][turn_index] = []
for missed_func_name in missed_func_names:
for i, func_doc in enumerate(entry["function"]):
if func_doc["name"] == missed_func_name:
# Add the missed function doc to the missed_function list
entry["missed_function"][turn_index].append(func_doc)
# Remove it from the function list
entry["function"].pop(i)
break
functions = func_doc_language_specific_pre_processing(entry["function"], test_category)
tools = convert_to_tool(functions, GORILLA_TO_OPENAPI, ModelStyle.OpenAI_Completions)
test_cases.append(
{
"id": entry["id"],
"messages": entry["question"][0],
"tools": tools,
"extra": entry,
},
)
outf.write(json.dumps(test_cases[-1], ensure_ascii=False) + "\n")
return test_cases
if __name__ == "__main__":
file_path = Path("./gorilla/berkeley-function-call-leaderboard/bfcl_eval/data/BFCL_v3_multi_turn_base.json")
output_path = "data/multiturn_data_base.jsonl"
preprocessed_test_cases = process_multi_turn_test_case(file_path, output_path)

View file

@ -1,151 +0,0 @@
"""Run evaluation on BFCL-V3-Multi-Turn-Base dataset."""
import time
import json
from pathlib import Path
import ray
import requests
from loguru import logger
from dotenv import load_dotenv
from bfcl_agent import BFCLAgent
load_dotenv("../../.env")
def run_agent(
max_workers: int,
dataset_name: str,
experiment_suffix: str,
model_name: str = "qwen3-8b",
enable_thinking: bool = False,
data_path: str = "data/multiturn_data_base_val.jsonl",
answer_path: Path = Path("data/possible_answer"),
num_trials: int = 1,
use_memory: bool = False,
memory_base_url: str = "http://0.0.0.0:8002/",
use_memory_addition: bool = True,
use_memory_deletion: bool = False,
delete_freq: int = 10,
freq_threshold: int = 5,
utility_threshold: float = 0.5,
):
"""Run the agent"""
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.mkdir(parents=True, exist_ok=True)
with open(data_path, "r", encoding="utf-8") as f:
task_ids = [json.loads(line)["id"] for line in f]
result: list = []
def dump_file():
with open(path / f"{experiment_name}.jsonl", "a", encoding="utf-8") as f:
for x in result:
f.write(json.dumps(x) + "\n")
future_list: list = []
for i in range(max_workers):
actor = BFCLAgent.remote(
index=i,
model_name=model_name,
task_ids=task_ids[i::max_workers],
experiment_name=experiment_name,
data_path=data_path,
answer_path=answer_path,
num_trials=num_trials,
use_memory=use_memory,
memory_base_url=memory_base_url,
use_memory_addition=use_memory_addition,
use_memory_deletion=use_memory_deletion,
delete_freq=delete_freq,
freq_threshold=freq_threshold,
utility_threshold=utility_threshold,
enable_thinking=enable_thinking,
)
future = actor.execute.remote()
future_list.append(future)
time.sleep(1)
logger.info("submit complete")
for i, future in enumerate(future_list):
t_result = ray.get(future)
if t_result:
if isinstance(t_result, list):
result.extend(t_result)
else:
result.append(t_result)
logger.info(f"{i + 1}/{len(task_ids)} complete")
dump_file()
def handle_api_response(response: requests.Response):
"""Handle API response with proper error checking"""
if response.status_code != 200:
print(f"Error: {response.status_code}")
print(response.text)
return None
return response.json()
def load_memory(path: str = "docs/library", api_url: str = "http://0.0.0.0:8002/"):
"""Load memories from disk into the vector store"""
response = requests.post(
url=f"{api_url}load_memory",
json={
"load_file_path": path,
"clear_existing": True,
},
)
result = handle_api_response(response)
if result:
print(f"Memory loaded from {path}")
def main():
"""Main function"""
max_workers = 4
if max_workers > 1:
ray.init(num_cpus=max_workers)
num_runs = 4
num_trials = 1
model_name = "qwen3-8b"
enable_thinking = True
use_memory = True
use_memory_addition = False
use_memory_deletion = False
memory_base_url = "http://0.0.0.0:8003/"
if use_memory:
load_file_path = "docs/library/paper_data/task/bfcl_qwen3_8b.jsonl"
load_memory(load_file_path, memory_base_url)
for _ in range(num_runs):
run_agent(
max_workers=max_workers,
model_name=model_name,
dataset_name="bfcl-multi-turn-base-val",
experiment_suffix="w-fixed-memory",
data_path="data/multiturn_data_base_val.jsonl",
answer_path=Path("data/possible_answer"),
enable_thinking=enable_thinking,
num_trials=num_trials,
use_memory=use_memory,
memory_base_url=memory_base_url,
use_memory_addition=use_memory_addition,
use_memory_deletion=use_memory_deletion,
delete_freq=5,
freq_threshold=5,
utility_threshold=0.5,
)
if __name__ == "__main__":
main()

View file

@ -1,163 +0,0 @@
"""Run the experiment statistic."""
import json
from collections import defaultdict
from pathlib import Path
import pandas as pd
from loguru import logger
def calculate_best_at_k(scores: list, k: int) -> float:
"""
Calculate best@k
Divide scores into groups of size k, take the maximum value in each group,
then average these maximum values
Args:
scores: List of after_score values for all runs of a task
k: Group size
Returns:
best@k value
"""
if len(scores) % k != 0:
raise ValueError(f"Length of scores ({len(scores)}) must be divisible by k ({k})")
group_maxs = []
for i in range(0, len(scores), k):
group = scores[i : i + k]
group_maxs.append(max(group))
return sum(group_maxs) / len(group_maxs)
def calculate_pass_at_k(scores: list, k: int) -> float:
"""Calculate pass@k."""
if len(scores) % k != 0:
raise ValueError(f"Length of scores ({len(scores)}) must be divisible by k ({k})")
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_maxs.append(is_pass)
return sum(group_maxs) / len(group_maxs)
def get_possible_k_values(total_runs: int) -> list:
"""
Get all possible k values (factors of total_runs)
Args:
total_runs: Total number of runs
Returns:
List of k values in descending order
"""
k_values = []
for k in range(1, total_runs + 1):
if total_runs % k == 0:
k_values.append(k)
return sorted(k_values, reverse=True) # Sort from large to small
def run_exp_statistic():
"""Run the experiment statistic."""
path: Path = Path("./exp_result/qwen3-8b/with_think")
# Store results for all experiments
all_results = {}
for file in path.glob("*.jsonl"):
# Group results by task_id
task_results = defaultdict(list)
print(file)
with open(file, "r", encoding="utf-8") as f:
for line in f:
if not line.strip():
continue
data = json.loads(line)
if isinstance(data, list):
for part_data in data:
task_id = part_data["task_id"]
after_score = part_data["reward"]
task_results[task_id].append(after_score)
else:
task_id = data["task_id"]
after_score = data["reward"]
task_results[task_id].append(after_score)
if not task_results:
logger.warning(f"No valid data found in file {file}")
continue
# Check if each task has consistent number of runs
run_counts = [len(scores) for scores in task_results.values()]
if len(set(run_counts)) > 1:
logger.warning(f"Inconsistent number of runs for different tasks in file {file}: {set(run_counts)}")
continue
num_runs = run_counts[0]
logger.info(f"File {file}: {len(task_results)} tasks, {num_runs} runs per task")
# Get all possible k values
k_values = get_possible_k_values(num_runs)
logger.info(f"Calculable best@k values: {k_values}")
# Calculate various best@k values
file_results = {"file": file.name}
for k in k_values:
best_at_k_scores = []
pass_at_k_scores = []
for task_id, scores in task_results.items():
try:
best_k_score = calculate_best_at_k(scores, k)
pass_at_k_score = calculate_pass_at_k(scores, k)
pass_at_k_scores.append(pass_at_k_score)
best_at_k_scores.append(best_k_score)
except ValueError as e:
logger.error(f"Error calculating best@{k} for task {task_id}: {e}")
continue
if best_at_k_scores:
avg_best_at_k = sum(best_at_k_scores) / len(best_at_k_scores)
file_results[f"best@{k}"] = avg_best_at_k
logger.info(f"file={file.name} best@{k}={avg_best_at_k:.4f}")
if pass_at_k_scores:
avg_pass_at_k = sum(pass_at_k_scores) / len(pass_at_k_scores)
file_results[f"pass@{k}"] = avg_pass_at_k
logger.info(f"file={file.name} pass@{k}={avg_pass_at_k:.4f}")
all_results[file.name] = file_results
# Create and display table
if all_results:
df = pd.DataFrame(list(all_results.values()))
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@')]
best_columns = list(df.columns)
best_columns.sort(key=lambda x: x, reverse=False)
df = df[best_columns]
print("\n" + "=" * 80)
print("Experiment Results Summary Table")
print("=" * 80)
print(df.round(4))
print("=" * 80)
# Save table to CSV
output_path = path / "experiment_summary.csv"
df.to_csv(output_path)
logger.info(f"Results table saved to: {output_path}")
else:
logger.warning("No valid experiment results found")
if __name__ == "__main__":
run_exp_statistic()

View file

@ -1,69 +0,0 @@
"""Split the JSONL file into train and validation sets."""
import argparse
import json
import random
from default_ids import DEFAULT_TRAIN_IDS, DEFAULT_VAL_IDS
def split_jsonl(
input_file: str,
train_file: str,
val_file: str,
ratio: float = 0.75,
random_split: bool = False,
) -> None:
"""Split the JSONL file into train and validation sets."""
with open(input_file, "r", encoding="utf-8") as f:
data = [json.loads(line) for line in f]
if random_split:
random.shuffle(data)
split_idx = int(len(data) * ratio)
train_data = data[:split_idx]
val_data = data[split_idx:]
else:
train_data = []
val_data = []
unknown_ids: list[str] = []
for obj in data:
if "id" not in obj:
raise ValueError(f"Missing 'id' field in input file: {input_file}")
obj_id = str(obj["id"])
if obj_id in DEFAULT_TRAIN_IDS:
train_data.append(obj)
elif obj_id in DEFAULT_VAL_IDS:
val_data.append(obj)
else:
unknown_ids.append(obj_id)
if len(train_data) + len(val_data) != len(data):
missing = len(data) - (len(train_data) + len(val_data))
examples = ", ".join(unknown_ids) if unknown_ids else "(none)"
raise ValueError(
f"{missing} samples in {input_file} not found in train_ref/val_ref id sets. Examples: {examples}",
)
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:
for item in val_data:
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.add_argument(
"--random",
action="store_true",
help="Whether to randomly split input into train/val. "
"If false, split strictly by default train/val id sets (see default_ids.py).",
)
args = parser.parse_args()
split_jsonl(args.input, args.train, args.val, args.ratio, args.random)

File diff suppressed because it is too large Load diff

File diff suppressed because it is too large Load diff

View file

@ -1,346 +0,0 @@
"""
LongMemEval Evaluation Statistics Analyzer
Computes detailed statistics from evaluation results including:
- Overall accuracy
- Accuracy by question type
- Timing statistics (summary, retrieval)
- Memory extraction statistics
Usage:
python bench/longmemeval/compute_stats.py \
--results_dir bench/longmemeval/bench_results/longmemeval_reme
"""
import argparse
import json
from collections import defaultdict
from pathlib import Path
from typing import Any
def load_results(results_dir: str) -> list[dict]:
"""Load all question result files from the directory.
Args:
results_dir: Path to the results directory
Returns:
List of result dictionaries
"""
results_path = Path(results_dir)
results = []
# Load individual question files
question_files = sorted(results_path.glob("question_*.json"))
for file_path in question_files:
try:
with open(file_path, "r", encoding="utf-8") as f:
result = json.load(f)
results.append(result)
except Exception as e:
print(f"⚠️ Error loading {file_path}: {e}")
return results
def compute_accuracy_stats(results: list[dict]) -> dict[str, Any]:
"""Compute overall and per-type accuracy statistics.
Args:
results: List of result dictionaries
Returns:
Dictionary with accuracy statistics
"""
total = len(results)
correct = 0
incorrect = 0
error = 0
# Per question type statistics
type_stats = defaultdict(lambda: {"total": 0, "correct": 0, "incorrect": 0, "error": 0})
for r in results:
qtype = r.get("question_type", "unknown")
judgment = r.get("judgment", {})
is_correct = judgment.get("is_correct")
type_stats[qtype]["total"] += 1
if is_correct is True:
correct += 1
type_stats[qtype]["correct"] += 1
elif is_correct is False:
incorrect += 1
type_stats[qtype]["incorrect"] += 1
else:
error += 1
type_stats[qtype]["error"] += 1
# Compute accuracies
overall = {
"total": total,
"correct": correct,
"incorrect": incorrect,
"error": error,
"accuracy": correct / total if total > 0 else 0,
"accuracy_valid": correct / (correct + incorrect) if (correct + incorrect) > 0 else 0,
}
by_type = {}
for qtype, stats in type_stats.items():
valid = stats["correct"] + stats["incorrect"]
by_type[qtype] = {
**stats,
"accuracy": stats["correct"] / stats["total"] if stats["total"] > 0 else 0,
"accuracy_valid": stats["correct"] / valid if valid > 0 else 0,
}
return {
"overall": overall,
"by_question_type": by_type,
}
def compute_timing_stats(results: list[dict]) -> dict[str, Any]:
"""Compute timing statistics.
Args:
results: List of result dictionaries
Returns:
Dictionary with timing statistics
"""
summary_times = []
retrieve_times = []
for r in results:
summary_ms = r.get("summary_duration_ms", 0)
retrieve_ms = r.get("retrieve_duration_ms", 0)
if summary_ms > 0:
summary_times.append(summary_ms)
if retrieve_ms > 0:
retrieve_times.append(retrieve_ms)
def compute_stats(times: list[float]) -> dict:
if not times:
return {"count": 0, "total_ms": 0, "avg_ms": 0, "min_ms": 0, "max_ms": 0}
return {
"count": len(times),
"total_ms": sum(times),
"total_min": sum(times) / 1000 / 60,
"avg_ms": sum(times) / len(times),
"min_ms": min(times),
"max_ms": max(times),
}
return {
"summary": compute_stats(summary_times),
"retrieve": compute_stats(retrieve_times),
"total_time_min": (sum(summary_times) + sum(retrieve_times)) / 1000 / 60,
}
def compute_memory_stats(results: list[dict]) -> dict[str, Any]:
"""Compute memory extraction statistics.
Args:
results: List of result dictionaries
Returns:
Dictionary with memory statistics
"""
memory_counts = []
session_counts = []
for r in results:
memories = r.get("extracted_memories", [])
num_sessions = r.get("num_sessions", 0)
memory_counts.append(len(memories))
session_counts.append(num_sessions)
def compute_stats(counts: list[int]) -> dict:
if not counts:
return {"count": 0, "total": 0, "avg": 0, "min": 0, "max": 0}
return {
"count": len(counts),
"total": sum(counts),
"avg": sum(counts) / len(counts),
"min": min(counts),
"max": max(counts),
}
return {
"memories_per_question": compute_stats(memory_counts),
"sessions_per_question": compute_stats(session_counts),
}
def print_report(
accuracy_stats: dict,
timing_stats: dict,
memory_stats: dict,
results_dir: str,
):
"""Print formatted statistics report.
Args:
accuracy_stats: Accuracy statistics
timing_stats: Timing statistics
memory_stats: Memory statistics
results_dir: Path to results directory
"""
print("\n" + "=" * 80)
print("LONGMEMEVAL EVALUATION STATISTICS")
print(f"Results Directory: {results_dir}")
print("=" * 80)
# Overall accuracy
overall = accuracy_stats["overall"]
print("\n📊 Overall Accuracy:")
print(f" Total Questions: {overall['total']}")
print(f" ✅ Correct: {overall['correct']} ({100 * overall['accuracy']:.2f}%)")
print(
f" ❌ Incorrect: {overall['incorrect']} "
f"({100 * overall['incorrect'] / overall['total'] if overall['total'] > 0 else 0:.2f}%)",
)
if overall["error"] > 0:
print(f" ⚠️ Error: {overall['error']} ({100 * overall['error'] / overall['total']:.2f}%)")
print(f" Accuracy (valid): {100 * overall['accuracy_valid']:.2f}%")
# Accuracy by question type
print("\n📊 Accuracy by Question Type:")
print("-" * 60)
print(f"{'Question Type':<30} {'Correct':<10} {'Total':<10} {'Accuracy':<10}")
print("-" * 60)
by_type = accuracy_stats["by_question_type"]
for qtype in sorted(by_type.keys()):
stats = by_type[qtype]
print(f"{qtype:<30} {stats['correct']:<10} {stats['total']:<10} {100 * stats['accuracy']:.2f}%")
print("-" * 60)
# Timing statistics
print("\n⏱️ Timing Statistics:")
summary = timing_stats["summary"]
retrieve = timing_stats["retrieve"]
print(" Memory Summarization:")
print(f" Total Time: {summary['total_min']:.2f} min")
print(f" Avg per Q: {summary['avg_ms']:.0f} ms")
print(f" Min/Max: {summary['min_ms']:.0f} / {summary['max_ms']:.0f} ms")
print(" Memory Retrieval:")
print(f" Total Time: {retrieve['total_min']:.2f} min")
print(f" Avg per Q: {retrieve['avg_ms']:.0f} ms")
print(f" Min/Max: {retrieve['min_ms']:.0f} / {retrieve['max_ms']:.0f} ms")
print(f" Total Time: {timing_stats['total_time_min']:.2f} min")
# Memory statistics
print("\n📝 Memory Statistics:")
mem = memory_stats["memories_per_question"]
sess = memory_stats["sessions_per_question"]
print(" Extracted Memories per Question:")
print(f" Total: {mem['total']}")
print(f" Average: {mem['avg']:.1f}")
print(f" Min/Max: {mem['min']} / {mem['max']}")
print(" Sessions per Question:")
print(f" Average: {sess['avg']:.1f}")
print(f" Min/Max: {sess['min']} / {sess['max']}")
print("\n" + "=" * 80)
def save_statistics(
accuracy_stats: dict,
timing_stats: dict,
memory_stats: dict,
output_file: str,
):
"""Save statistics to JSON file.
Args:
accuracy_stats: Accuracy statistics
timing_stats: Timing statistics
memory_stats: Memory statistics
output_file: Path to output file
"""
stats = {
"accuracy": accuracy_stats,
"timing": timing_stats,
"memory": memory_stats,
}
with open(output_file, "w", encoding="utf-8") as f:
json.dump(stats, f, indent=4, ensure_ascii=False)
print(f"\n📁 Statistics saved to: {output_file}")
def main(results_dir: str, output_file: str = None):
"""Main function to compute and display statistics.
Args:
results_dir: Path to results directory
output_file: Optional path to save statistics JSON
"""
print(f"\nLoading results from: {results_dir}")
results = load_results(results_dir)
if not results:
print("❌ No results found!")
return
print(f"Loaded {len(results)} question results")
# Compute statistics
accuracy_stats = compute_accuracy_stats(results)
timing_stats = compute_timing_stats(results)
memory_stats = compute_memory_stats(results)
# Print report
print_report(accuracy_stats, timing_stats, memory_stats, results_dir)
# Save to file if specified
if output_file:
save_statistics(accuracy_stats, timing_stats, memory_stats, output_file)
else:
# Default output file in results directory
default_output = Path(results_dir) / "statistics.json"
save_statistics(accuracy_stats, timing_stats, memory_stats, str(default_output))
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description="Compute statistics from LongMemEval evaluation results",
)
parser.add_argument(
"--results_dir",
type=str,
default="bench_results/longmemeval_reme",
help="Path to results directory containing question_*.json files",
)
parser.add_argument(
"--output_file",
type=str,
default=None,
help="Path to save statistics JSON (default: <results_dir>/statistics.json)",
)
args = parser.parse_args()
main(
results_dir=args.results_dir,
output_file=args.output_file,
)

File diff suppressed because it is too large Load diff

View file

@ -1,921 +0,0 @@
"""
LongMemEval Benchmark Evaluator for ReMe - Retrieve Only
A simplified evaluation pipeline that only runs the retrieve and judge phases:
1. Loads LongMemEval benchmark data
2. Skips memory summarization (assumes memories are already in vector store)
3. Uses questions to query memory and generate answers
4. Uses LLM to judge answer correctness
5. Generates comprehensive metrics
This is useful for debugging/tuning the retrieve phase without re-running summary.
Usage:
python benchmark/longmemeval/eval_longmemeval_reme_retrieve.py \
--data_path dataset/longmemeval/longmemeval_s_cleaned.json \
--top_k 20 --start_index 0 --end_index 10
"""
import asyncio
import json
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Optional
from loguru import logger
from reme.reme import ReMe
# ==================== Configuration ====================
@dataclass
class RetrieveEvalConfig:
"""Evaluation configuration parameters for retrieve-only mode."""
data_path: str
top_k: int = 10
start_index: int = 0
end_index: Optional[int] = None
max_concurrency: int = 1
output_dir: str = "cache/bench_results/longmemeval_reme_retrieve"
reme_model_name: str = "qwen-flash"
eval_model_name: str = "qwen3-max"
algo_version: str = "v1"
samples_per_type: int = -1 # Number of samples per question type, -1 for all
enable_thinking_params: bool = False
# Optional: path to previous summary results to reload memories
summary_results_dir: Optional[str] = None
# ==================== Answer Judge Prompts ====================
def get_anscheck_prompt(task: str, question: str, answer: str, response: str, abstention: bool = False) -> str:
"""Generate the answer checking prompt based on question type.
Args:
task: Question type, e.g. 'single-session-user', 'multi-session', 'temporal-reasoning'
question: The question content
answer: The reference answer
response: The model's response
abstention: Whether this is an unanswerable question
Returns:
Prompt for judging answer correctness
"""
if not abstention:
if task in ["single-session-user", "single-session-assistant", "multi-session"]:
template = (
"I will give you a question, a correct answer, and a response from a model. Please answer yes i"
"f the response contains the correct answer. Otherwise, answer no. If the response is equival"
"ent to the correct answer or contains all the intermediate steps to get the correct answer, "
"you should also answer yes. If the response only contains a subset of the information required"
" by the answer, answer no. \n\nQuestion: {}\n\nCorrect Answer: {}\n\nModel Response: {}\n\nIs"
" the model response correct? Answer yes or no only."
)
prompt = template.format(question, answer, response)
elif task == "temporal-reasoning":
template = (
"I will give you a question, a correct answer, and a response from a model. Please answer yes"
" if the response contains the correct answer. Otherwise, answer no. If the response is equiv"
"alent to the correct answer or contains all the intermediate steps to get the correct answer"
", you should also answer yes. If the response only contains a subset of the information requ"
"ired by the answer, answer no. In addition, do not penalize off-by-one errors for the numbe"
"r of days. If the question asks for the number of days/weeks/months, etc., and the model ma"
"kes off-by-one errors (e.g., predicting 19 days when the answer is 18), the model's respon"
"se is still correct. \n\nQuestion: {}\n\nCorrect Answer: {}\n\nModel Response: {}\n\nIs th"
"e model response correct? Answer yes or no only."
)
prompt = template.format(question, answer, response)
elif task == "knowledge-update":
template = (
"I will give you a question, a correct answer, and a response from a model. Please answer yes "
"if the response contains the correct answer. Otherwise, answer no. If the response contains "
"some previous information along with an updated answer, the response should be considered "
"as correct as long as the updated answer is the required answer.\n\nQuestion: {}\n\nCorrec"
"t Answer: {}\n\nModel Response: {}\n\nIs the model response correct? Answer yes or no only."
)
prompt = template.format(question, answer, response)
elif task == "single-session-preference":
template = (
"I will give you a question, a rubric for desired personalized response, and a response from a"
" model. Please answer yes if the response satisfies the desired response. Otherwise, answer"
" no. The model does not need to reflect all the points in the rubric. The response is corr"
"ect as long as it recalls and utilizes the user's personal information correctly.\n\nQuest"
"ion: {}\n\nRubric: {}\n\nModel Response: {}\n\nIs the model response correct? Answer yes o"
"r no only."
)
prompt = template.format(question, answer, response)
else:
# Default template
template = (
"I will give you a question, a correct answer, and a response from a model. Please answer y"
"es if the response contains the correct answer. Otherwise, answer no. If the response is "
"equivalent to the correct answer or contains all the intermediate steps to get the correc"
"t answer, you should also answer yes. If the response only contains a subset of the infor"
"mation required by the answer, answer no. \n\nQuestion: {}\n\nCorrect Answer: {}\n\nModel"
" Response: {}\n\nIs the model response correct? Answer yes or no only."
)
prompt = template.format(question, answer, response)
else:
template = (
"I will give you an unanswerable question, an explanation, and a response from a model. Please "
"answer yes if the model correctly identifies the question as unanswerable. The model could say "
"that the information is incomplete, or some other information is given but the asked informati"
"on is not.\n\nQuestion: {}\n\nExplanation: {}\n\nModel Response: {}\n\nDoes the model correct"
"ly identify the question as unanswerable? Answer yes or no only."
)
prompt = template.format(question, answer, response)
return prompt
# ==================== Utilities ====================
class DataLoader:
"""Handles loading and parsing of LongMemEval data."""
@staticmethod
def load_json(file_path: str) -> list[dict]:
"""Load all entries from a JSON file."""
with open(file_path, "r", encoding="utf-8") as f:
return json.load(f)
@staticmethod
def filter_by_type(data: list[dict], samples_per_type: int = -1) -> list[tuple[int, dict]]:
"""Filter data by question type with specified number of samples per type.
Args:
data: List of question entries
samples_per_type: Number of samples per type, -1 for all
Returns:
List of tuples (original_index, entry) for selected samples
"""
if samples_per_type == -1:
# Return all with original indices
return list(enumerate(data))
# Group by question type
type_groups: dict[str, list[tuple[int, dict]]] = {}
for i, entry in enumerate(data):
qtype = entry.get("question_type", "unknown")
if qtype not in type_groups:
type_groups[qtype] = []
type_groups[qtype].append((i, entry))
# Select samples from each type
selected = []
for qtype, entries in type_groups.items():
count = min(samples_per_type, len(entries))
selected.extend(entries[:count])
logger.info(f" {qtype}: selected {count}/{len(entries)} samples")
# Sort by original index to maintain order
selected.sort(key=lambda x: x[0])
return selected
class FileManager:
"""Manages file I/O operations."""
def __init__(self, base_dir: str):
self.base_dir = Path(base_dir)
self.base_dir.mkdir(parents=True, exist_ok=True)
def save_question_result(self, idx: int, question_id: str, data: dict):
"""Save result for a single question."""
file_path = self.base_dir / f"question_{idx:04d}_{question_id}.json"
with open(file_path, "w", encoding="utf-8") as f:
json.dump(data, f, indent=4, ensure_ascii=False)
logger.info(f"✅ Saved question result to {file_path}")
def load_question_result(self, idx: int, question_id: str) -> Optional[dict]:
"""Load result for a single question if exists."""
file_path = self.base_dir / f"question_{idx:04d}_{question_id}.json"
if not file_path.exists():
return None
with open(file_path, "r", encoding="utf-8") as f:
return json.load(f)
def save_summary(self, results: list[dict]):
"""Save summary of all results."""
file_path = self.base_dir / "summary.json"
with open(file_path, "w", encoding="utf-8") as f:
json.dump(results, f, indent=4, ensure_ascii=False)
logger.info(f"✅ Saved summary to {file_path}")
# ==================== Evaluation Functions ====================
async def answer_question_with_memories(
reme: ReMe,
question: str,
memories: str,
user_id: str = None,
model_name: str = "qwen3-max",
):
"""
Answer a question using retrieved memories with PROMPT_MEMZERO_JSON template.
Args:
reme: ReMe instance with default_llm and prompt_handler
question: The question to answer
memories: The retrieved memories (formatted as context)
user_id: Optional user ID for context formatting
model_name: Model name to use for LLM request
Returns:
dict with 'reasoning' and 'answer' fields
"""
# Format context with memories
if user_id:
context = reme.prompt_handler.prompt_format(
"TEMPLATE_MEMOS",
user_id=user_id,
memories=memories,
)
else:
context = f"Memories:\n{memories}"
# Use PROMPT_MEMZERO_JSON template for structured JSON response
prompt = reme.prompt_handler.prompt_format(
"PROMPT_MEMZERO_JSON",
context=context,
question=question,
)
result = await reme.get_llm(name=model_name).simple_request_for_json(
prompt=prompt,
model_name=None,
)
return result
# ==================== Memory Operations ====================
class RetrieveProcessor:
"""Handles ReMe memory retrieve operations only."""
def __init__(
self,
reme: ReMe,
reme_model_name: str = "qwen-flash",
eval_model_name: str = "qwen3-max",
algo_version: str = "v1",
enable_thinking_params: bool = False,
):
self.reme = reme
self.reme_model_name = reme_model_name
self.eval_model_name = eval_model_name
self.algo_version = algo_version
self.enable_thinking_params = enable_thinking_params
async def search_memory(
self,
query: str,
user_id: str,
top_k: int = 20,
) -> tuple[dict, list, float]:
"""
Search memory using ReMe and return structured answer with reasoning.
Returns:
tuple: (answer_dict, agent_messages, duration_ms)
answer_dict contains: {"reasoning": str, "answer": str, "memories": str}
"""
start = time.time()
# Retrieve memories from ReMe using new API
result = await self.reme.retrieve_memory(
llm_config_name="qwen3-max",
query=query,
retrieve_top_k=top_k,
user_name=user_id,
version=self.algo_version,
return_dict=True,
enable_time_filter=True,
enable_thinking_params=True,
)
# Extract memories from response
memories = result["answer"]
agent_messages = [x.simple_dump(enable_argument_dict=True) for x in result["messages"]]
retrieved_nodes = [x.model_dump(exclude_none=True) for x in result["retrieved_nodes"]]
# Use LLM to generate structured answer from memories
answer_result = await answer_question_with_memories(
reme=self.reme,
question=query,
memories=memories,
user_id=user_id,
model_name=self.eval_model_name,
)
# Add original memories to the result
answer_result["memories"] = memories
answer_result["retrieved_nodes"] = retrieved_nodes
duration_ms = (time.time() - start) * 1000
return answer_result, agent_messages, duration_ms
# ==================== Answer Judge ====================
class LongMemEvalJudge:
"""LongMemEval answer judge using LLM."""
def __init__(self, reme: ReMe, model: str = "qwen3-max"):
self.reme = reme
self.model = model
async def judge_answer(
self,
question_type: str,
question: str,
answer: str,
response: str,
abstention: bool = False,
) -> dict:
"""
Judge if the model's response is correct.
Returns:
dict with is_correct, llm_response, and judge_prompt
"""
prompt = get_anscheck_prompt(question_type, question, answer, response, abstention)
try:
llm_response = await self.reme.get_llm("default").simple_request(
prompt=prompt,
model_name=self.model,
)
llm_response_lower = llm_response.strip().lower()
is_correct = llm_response_lower.startswith("yes")
return {
"is_correct": is_correct,
"llm_response": llm_response,
"judge_prompt": prompt,
}
except Exception as e:
return {
"is_correct": None,
"error": str(e),
"judge_prompt": prompt,
}
# ==================== Metrics ====================
class MetricsAggregator:
"""Aggregates evaluation metrics for LongMemEval."""
@staticmethod
def compute_metrics(results: list[dict]) -> dict[str, Any]:
"""Compute overall and per-type metrics."""
total = len(results)
correct = sum(1 for r in results if r.get("judgment", {}).get("is_correct") is True)
incorrect = sum(1 for r in results if r.get("judgment", {}).get("is_correct") is False)
error = total - correct - incorrect
metrics = {
"total": total,
"correct": correct,
"incorrect": incorrect,
"error": error,
"accuracy": correct / total if total > 0 else 0,
"accuracy_valid": correct / (correct + incorrect) if (correct + incorrect) > 0 else 0,
}
# Per question type statistics
type_stats = {}
for r in results:
qtype = r.get("question_type", "unknown")
if qtype not in type_stats:
type_stats[qtype] = {"total": 0, "correct": 0, "incorrect": 0}
type_stats[qtype]["total"] += 1
if r.get("judgment", {}).get("is_correct") is True:
type_stats[qtype]["correct"] += 1
elif r.get("judgment", {}).get("is_correct") is False:
type_stats[qtype]["incorrect"] += 1
metrics["by_question_type"] = {
qtype: {
**stats,
"accuracy": stats["correct"] / stats["total"] if stats["total"] > 0 else 0,
"accuracy_valid": (
stats["correct"] / (stats["correct"] + stats["incorrect"])
if (stats["correct"] + stats["incorrect"]) > 0
else 0
),
}
for qtype, stats in type_stats.items()
}
return metrics
@staticmethod
def compute_timing_stats(results: list[dict]) -> dict[str, Any]:
"""Compute timing statistics."""
retrieve_times = []
for r in results:
retrieve_ms = r.get("retrieve_duration_ms", 0)
if retrieve_ms > 0:
retrieve_times.append(retrieve_ms)
def compute_stats(times: list[float]) -> dict:
if not times:
return {"count": 0, "total_ms": 0, "avg_ms": 0, "min_ms": 0, "max_ms": 0}
return {
"count": len(times),
"total_ms": sum(times),
"total_min": sum(times) / 1000 / 60,
"avg_ms": sum(times) / len(times),
"min_ms": min(times),
"max_ms": max(times),
}
return {
"retrieve": compute_stats(retrieve_times),
"total_time_min": sum(retrieve_times) / 1000 / 60,
}
# ==================== Main Pipeline ====================
class LongMemEvalRetrieveEvaluator:
"""Retrieve-only evaluator for LongMemEval benchmark using ReMe."""
def __init__(self, config: RetrieveEvalConfig):
self.config = config
self.reme = ReMe(
default_llm_config={
"model_name": self.config.reme_model_name,
},
llms={
"qwen-plus-think": {
"backend": "openai",
"model_name": "qwen-plus",
"extra_body": {
"enable_thinking": True,
},
},
"qwen3-max-think": {
"backend": "openai",
"model_name": "qwen3-max",
"extra_body": {
"enable_thinking": True,
},
},
"qwen3-max": {
"backend": "openai",
"model_name": "qwen3-max",
"extra_body": {
"enable_thinking": False,
},
},
},
)
# Load evaluation prompts into ReMe's prompt handler
prompts_yaml_path = Path(__file__).parent / "eval_reme.yaml"
self.reme.prompt_handler.load_prompt_by_file(prompts_yaml_path)
self.file_manager = FileManager(config.output_dir)
self.retrieve_processor = RetrieveProcessor(
self.reme,
config.reme_model_name,
config.eval_model_name,
config.algo_version,
config.enable_thinking_params,
)
self.judge = LongMemEvalJudge(self.reme, config.eval_model_name)
self.data_loader = DataLoader()
async def __aenter__(self):
"""Async context manager entry."""
await self.reme.start()
return self
async def __aexit__(self, exc_type, exc_val, exc_tb):
"""Async context manager exit with cleanup."""
await self.reme.close()
return False
async def process_question_entry(self, entry: dict, idx: int) -> dict:
"""Process a single question entry (retrieve + judge only).
Args:
entry: A question entry from LongMemEval dataset
idx: Index of the question
Returns:
Result dictionary
"""
question_id = entry["question_id"]
question = entry["question"]
answer = entry["answer"]
question_type = entry["question_type"]
question_date = entry.get("question_date", "")
haystack_dates = entry["haystack_dates"]
haystack_session_ids = entry["haystack_session_ids"]
haystack_sessions = entry["haystack_sessions"]
# Use question_id as user_id for isolation (same as full eval)
user_id = f"longmemeval_{question_id}"
logger.info(f"\n{'='*60}")
logger.info(f"Question ID: {question_id}")
logger.info(f"Question Type: {question_type}")
logger.info(f"Question: {question}")
logger.info(f"Question_date: {question_date}")
logger.info(f"Answer: {answer}")
logger.info(f"Number of sessions: {len(haystack_sessions)}")
logger.info(f"{'='*60}")
# Skip summary phase - directly search memory and answer question
logger.info(" Retrieving and answering question using ReMe...")
answer_dict, retrieve_messages, retrieve_duration_ms = await self.retrieve_processor.search_memory(
query=f"[Question_date: {question_date} | Question_type: {question_type}] " + question,
user_id=user_id,
top_k=self.config.top_k,
)
# Extract answer and reasoning from the structured response
model_response = answer_dict.get("answer", "")
model_reasoning = answer_dict.get("reasoning", "")
retrieved_memories = answer_dict.get("memories", "")
retrieved_nodes = answer_dict.get("retrieved_nodes", [])
# Judge answer correctness
logger.info(" Judging answer correctness...")
judgment = await self.judge.judge_answer(
question_type=question_type,
question=question,
answer=answer,
response=model_response,
)
is_correct = judgment.get("is_correct")
logger.info(
f" → Answer judgment: {'Correct' if is_correct else 'Incorrect' if is_correct is False else 'Error'}",
)
result = {
"question_id": question_id,
"question_type": question_type,
"question": question,
"answer": answer,
"question_date": question_date,
"haystack_dates": haystack_dates,
"haystack_session_ids": haystack_session_ids,
"num_sessions": len(haystack_sessions),
"model_response": model_response,
"model_reasoning": model_reasoning,
"retrieved_memories": retrieved_memories,
"retrieved_nodes": retrieved_nodes,
"judgment": judgment,
"retrieve_duration_ms": retrieve_duration_ms,
"retrieve_messages": retrieve_messages,
}
# Save individual result
self.file_manager.save_question_result(idx, question_id, result)
logger.info(f" Question {question_id} - Completed")
return result
async def run_evaluation(self):
"""Run the retrieve-only evaluation pipeline with parallel processing."""
start_time = time.time()
# NOTE: Do NOT clear vector store - we assume memories are already there from previous summary run
# Load dataset
logger.info(f"Loading dataset from: {self.config.data_path}")
all_data = self.data_loader.load_json(self.config.data_path)
logger.info(f"Total questions in dataset: {len(all_data)}")
# Filter by question type
logger.info(f"Filtering by type (samples_per_type={self.config.samples_per_type}):")
filtered_data = self.data_loader.filter_by_type(all_data, self.config.samples_per_type)
logger.info(f"Selected {len(filtered_data)} questions after filtering")
# Apply start_index and end_index on filtered data
end_index = self.config.end_index or len(filtered_data)
start_index = self.config.start_index
end_index = min(end_index, len(filtered_data))
# Get the slice we want to process
data_to_process = filtered_data[start_index:end_index]
total_questions = len(data_to_process)
logger.info(f"Processing {total_questions} questions (index {start_index} to {end_index - 1})")
print("\n" + "=" * 80)
print("LONGMEMEVAL EVALUATION - REME (RETRIEVE ONLY)")
print(f"Samples per type: {self.config.samples_per_type} (-1 = all)")
print(f"Questions to process: {total_questions} | Top-K: {self.config.top_k}")
print(f"Max Concurrency: {self.config.max_concurrency}")
print(f"ReMe Model: {self.config.reme_model_name} | Eval Model: {self.config.eval_model_name}")
print(f"Algo Version: {self.config.algo_version}")
print("⚠️ NOTE: Assumes memories are already in vector store from previous summary run")
print("=" * 80 + "\n")
# Use semaphore to control concurrency
semaphore = asyncio.Semaphore(self.config.max_concurrency)
async def process_with_semaphore(idx: int, original_idx: int, entry: dict) -> Optional[dict]:
"""Process a question with semaphore for concurrency control."""
async with semaphore:
question_id = entry["question_id"]
# Check cache first (use original index for cache file naming)
cached_result = self.file_manager.load_question_result(original_idx, question_id)
if cached_result:
print(f"⚡ [{idx}/{total_questions}] Skipping question {original_idx} (cached)")
return cached_result
print(f"\n{'#'*60}")
print(f"### [{idx}/{total_questions}] Processing Question {original_idx} ###")
print(f"{'#'*60}")
try:
result = await self.process_question_entry(entry, original_idx)
print(f"✅ [{idx}/{total_questions}] Completed question {original_idx}")
return result
except Exception as e:
logger.error(f"❌ Error processing question {original_idx}: {e}")
import traceback
traceback.print_exc()
return {
"question_id": question_id,
"error": str(e),
"question_type": entry.get("question_type", "unknown"),
"question": entry.get("question", ""),
"answer": entry.get("answer", ""),
"judgment": {"is_correct": None, "error": str(e)},
}
# Create all tasks from filtered data (each item is a tuple of (original_idx, entry))
tasks = [
process_with_semaphore(idx + 1, original_idx, entry)
for idx, (original_idx, entry) in enumerate(data_to_process)
]
# Execute in parallel with controlled concurrency
all_results = await asyncio.gather(*tasks, return_exceptions=False)
# Filter out None results if any
all_results = [r for r in all_results if r is not None]
# Save summary
self.file_manager.save_summary(all_results)
elapsed = time.time() - start_time
print(f"\n✅ Processing completed in {elapsed:.2f}s")
if total_questions > 0:
print(f" Average time per question: {elapsed / total_questions:.2f}s")
# Compute and report metrics
self._report_metrics(all_results)
return all_results
def _report_metrics(self, results: list[dict]):
"""Report evaluation metrics."""
metrics = MetricsAggregator.compute_metrics(results)
timing_stats = MetricsAggregator.compute_timing_stats(results)
print("\n" + "=" * 80)
print("EVALUATION SUMMARY - LONGMEMEVAL - REME (RETRIEVE ONLY)")
print("=" * 80 + "\n")
print("📊 Overall Results:")
print(f" ✅ Correct: {metrics['correct']}/{metrics['total']} ({100*metrics['accuracy']:.2f}%)")
print(
f" ❌ Incorrect: {metrics['incorrect']}/{metrics['total']}"
f" ({100*metrics['incorrect']/metrics['total'] if metrics['total'] > 0 else 0:.2f}%)",
)
if metrics["error"] > 0:
print(f" ⚠️ Error: {metrics['error']}/{metrics['total']} ({100*metrics['error']/metrics['total']:.2f}%)")
print(f" Accuracy (valid): {100*metrics['accuracy_valid']:.2f}%")
print("\n📊 Accuracy by Question Type:")
print("-" * 60)
print(f"{'Question Type':<30} {'Correct':<10} {'Total':<10} {'Accuracy':<10}")
print("-" * 60)
for qtype in sorted(metrics["by_question_type"].keys()):
stats = metrics["by_question_type"][qtype]
print(f"{qtype:<30} {stats['correct']:<10} {stats['total']:<10} {100*stats['accuracy']:.2f}%")
print("-" * 60)
print("\n⏱️ Timing Statistics (Retrieve Only):")
retrieve = timing_stats["retrieve"]
print(" Memory Retrieval:")
print(f" Total Time: {retrieve['total_min']:.2f} min")
print(f" Avg per Q: {retrieve['avg_ms']:.0f} ms")
print(f" Total Time: {timing_stats['total_time_min']:.2f} min")
# Save metrics
final_results = {
"accuracy": metrics,
"timing": timing_stats,
}
metrics_file = self.file_manager.base_dir / "eval_statistics.json"
with open(metrics_file, "w", encoding="utf-8") as f:
json.dump(final_results, f, indent=4, ensure_ascii=False)
print(f"\n📁 Statistics saved to: {metrics_file}")
print("\n" + "=" * 80)
# ==================== Entry Point ====================
async def main_async(
data_path: str,
top_k: int = 20,
start_index: int = 0,
end_index: Optional[int] = None,
max_concurrency: int = 1,
output_dir: str = "bench_results/longmemeval_reme_retrieve",
reme_model_name: str = "qwen-flash",
eval_model_name: str = "qwen3-max",
algo_version: str = "v1",
samples_per_type: int = -1,
enable_thinking_params: bool = False,
summary_results_dir: Optional[str] = None,
):
"""Main async entry point for LongMemEval retrieve-only evaluation with proper resource cleanup."""
config = RetrieveEvalConfig(
data_path=data_path,
top_k=top_k,
start_index=start_index,
end_index=end_index,
max_concurrency=max_concurrency,
output_dir=output_dir,
reme_model_name=reme_model_name,
eval_model_name=eval_model_name,
algo_version=algo_version,
samples_per_type=samples_per_type,
enable_thinking_params=enable_thinking_params,
summary_results_dir=summary_results_dir,
)
# Use async context manager for automatic cleanup
async with LongMemEvalRetrieveEvaluator(config) as evaluator:
await evaluator.run_evaluation()
def main(
data_path: str,
top_k: int = 20,
start_index: int = 0,
end_index: Optional[int] = None,
max_concurrency: int = 1,
output_dir: str = "bench_results/longmemeval_reme_retrieve",
reme_model_name: str = "qwen-flash",
eval_model_name: str = "qwen3-max",
algo_version: str = "v1",
samples_per_type: int = -1,
enable_thinking_params: bool = False,
summary_results_dir: Optional[str] = None,
):
"""Main entry point for LongMemEval retrieve-only evaluation."""
asyncio.run(
main_async(
data_path=data_path,
top_k=top_k,
start_index=start_index,
end_index=end_index,
max_concurrency=max_concurrency,
output_dir=output_dir,
reme_model_name=reme_model_name,
eval_model_name=eval_model_name,
algo_version=algo_version,
samples_per_type=samples_per_type,
enable_thinking_params=enable_thinking_params,
summary_results_dir=summary_results_dir,
),
)
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(
description="Evaluate ReMe on LongMemEval benchmark (Retrieve Phase Only)",
)
parser.add_argument(
"--data_path",
type=str,
# default="/Users/zhouwk/PycharmProjects/MemAgent/dataset/longmemeval/longmemeval_s_cleaned.json",
default="/Users/zhouwk/PycharmProjects/MemAgent/dataset/longmemeval/longmemeval_oracle.json",
help="Path to LongMemEval JSON file",
)
parser.add_argument(
"--top_k",
type=int,
default=10,
help="Number of memories to retrieve (default: 10)",
)
parser.add_argument(
"--start_index",
type=int,
default=0,
help="Start index for processing questions (default: 0)",
)
parser.add_argument(
"--end_index",
type=int,
default=None,
help="End index for processing questions (default: None, process all)",
)
parser.add_argument(
"--max_concurrency",
type=int,
default=8,
help="Maximum concurrent question processing (default: 1)",
)
parser.add_argument(
"--output_dir",
type=str,
default="bench_results/longmemeval_reme_retrieve_gpt4",
help="Output directory for results",
)
parser.add_argument(
"--reme_model_name",
type=str,
default="gpt-4o-mini-2024-07-18",
help="Model name for ReMe operations (default: gpt-4o-mini-2024-07-18)",
)
parser.add_argument(
"--eval_model_name",
type=str,
default="gpt-4o-mini-2024-07-18",
help="Model name for evaluation/judgment (default: gpt-4o-mini-2024-07-18)",
)
parser.add_argument(
"--algo_version",
type=str,
default="longmemeval",
help="Algorithm version for retrieval (default: longmemeval)",
)
parser.add_argument(
"--samples_per_type",
type=int,
default=4,
help="Number of samples per question type, -1 for all (default: 4)",
)
parser.add_argument(
"--enable_thinking_params",
action="store_true",
default=False,
help="Enable thinking parameters for retrieval (default: False)",
)
parser.add_argument(
"--summary_results_dir",
type=str,
default="/Users/zhouwk/PycharmProjects/ReMe/benchmark/longmemeval/bench_results",
help="Optional: path to previous summary results directory (for reference)",
)
parser.add_argument(
"--no_cache",
action="store_true",
default=False,
help="Ignore cached results and re-run all questions (default: False)",
)
args = parser.parse_args()
print(f"args={args}!")
main(
data_path=args.data_path,
top_k=args.top_k,
start_index=args.start_index,
end_index=args.end_index,
max_concurrency=args.max_concurrency,
output_dir=args.output_dir,
reme_model_name=args.reme_model_name,
eval_model_name=args.eval_model_name,
algo_version=args.algo_version,
samples_per_type=args.samples_per_type,
enable_thinking_params=args.enable_thinking_params,
summary_results_dir=args.summary_results_dir,
)

View file

@ -1,232 +0,0 @@
"""Evaluation tools for ReMe LongMemEval benchmark."""
from pathlib import Path
import yaml
from reme.reme import ReMe
# Load prompts from YAML file
_YAML_PATH = Path(__file__).parent / "eval_reme.yaml"
with open(_YAML_PATH, "r", encoding="utf-8") as f:
_PROMPTS = yaml.safe_load(f)
async def evaluation_for_memory_integrity(
reme: ReMe,
extract_memories: str,
target_memory: str,
model_name: str = "qwen3-max",
) -> dict:
"""
Memory Integrity Evaluation
Args:
reme: ReMe instance
extract_memories: A formatted string concatenating all memory points extracted by the memory system.
target_memory: The target key memory point.
model_name: Model name for evaluation
Returns:
dict with 'reasoning' and 'score' fields
"""
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_MEMORY_INTEGRITY"].format(
memories=extract_memories,
expected_memory_point=target_memory,
)
result = await reme.llm.simple_request_for_json(
prompt=prompt,
model_name=model_name,
)
return result
async def evaluation_for_memory_accuracy(
reme: ReMe,
dialogue: str,
golden_memories: str,
candidate_memory: str,
model_name: str = "qwen3-max",
) -> dict:
"""
Memory Accuracy Evaluation
Args:
reme: ReMe instance
dialogue: The complete human-machine dialogue record.
golden_memories: The core memory points for this dialogue segment in the evaluation set .
candidate_memory: A specific memory point extracted by the memory system being evaluated.
model_name: Model name for evaluation
Returns:
dict with 'accuracy_score', 'is_included_in_golden_memories', and 'reason' fields
"""
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_MEMORY_ACCURACY"].format(
dialogue=dialogue,
golden_memories=golden_memories,
candidate_memory=candidate_memory,
)
result = await reme.llm.simple_request_for_json(
prompt=prompt,
model_name=model_name,
)
return result
async def evaluation_for_update_memory(
reme: ReMe,
extract_memories: str,
target_update_memory: str,
original_memory: str,
model_name: str = "qwen3-max",
) -> dict:
"""
Memory Update Evaluation
Args:
reme: ReMe instance
extract_memories: A formatted string concatenating all memory points extracted by the memory system .
target_update_memory: The target updated memory point.
original_memory: A formatted string concatenating all original memory points corresponding.
model_name: Model name for evaluation
Returns:
dict with 'reason' and 'evaluation_result' fields
"""
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_UPDATE_MEMORY"].format(
memories=extract_memories,
updated_memory=target_update_memory,
original_memory=original_memory,
)
result = await reme.llm.simple_request_for_json(
prompt=prompt,
model_name=model_name,
)
return result
async def evaluation_for_question(
reme: ReMe,
question: str,
reference_answer: str,
key_memory_points: str,
response: str,
model_name: str = "qwen3-max",
) -> dict:
"""
Question-Answering Evaluation
Args:
reme: ReMe instance
question: The question string to be evaluated.
reference_answer: The reference (gold-standard) answer.
key_memory_points: The memory points used to derive the reference answer.
response: The answer produced by the memory system.
model_name: Model name for evaluation
Returns:
dict with 'reasoning' and 'evaluation_result' fields
"""
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION"].format(
question=question,
reference_answer=reference_answer,
key_memory_points=key_memory_points,
response=response,
)
result = await reme.llm.simple_request_for_json(
prompt=prompt,
model_name=model_name,
)
return result
async def evaluation_for_question2(
reme: ReMe,
question: str,
reference_answer: str,
key_memory_points: str,
response: str,
dialogue: str = "",
model_name: str = "qwen3-max",
) -> dict:
"""
Question-Answering Evaluation with Dialogue Context (Version 2)
Args:
reme: ReMe instance
question: The question string to be evaluated.
reference_answer: The reference (gold-standard) answer.
key_memory_points: The memory points used to derive the reference answer.
response: The answer produced by the memory system.
dialogue: The formatted dialogue history (role, content, time_created).
model_name: Model name for evaluation
Returns:
dict with 'reasoning' and 'evaluation_result' fields
"""
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION2"].format(
question=question,
reference_answer=reference_answer,
key_memory_points=key_memory_points,
response=response,
dialogue=dialogue if dialogue else "",
)
result = await reme.llm.simple_request_for_json(
prompt=prompt,
model_name=model_name,
)
return result
async def answer_question_with_memories(
reme: ReMe,
question: str,
memories: str,
user_id: str = None,
model_name: str = "qwen3-max",
) -> dict:
"""
Answer a question using retrieved memories with PROMPT_MEMZERO_JSON template.
Args:
reme: ReMe instance
question: The question to answer
memories: The retrieved memories (formatted as context)
user_id: Optional user ID for context formatting
model_name: Model name for LLM request
Returns:
dict with 'reasoning' and 'answer' fields
"""
# Format context with memories
if user_id:
context = _PROMPTS["TEMPLATE_MEMOS"].format(
user_id=user_id,
memories=memories,
)
else:
context = f"Memories:\n{memories}"
# Use PROMPT_MEMZERO_JSON template for structured JSON response
prompt = _PROMPTS["PROMPT_MEMZERO_JSON"].format(
context=context,
question=question,
)
result = await reme.llm.simple_request_for_json(
prompt=prompt,
model_name=model_name,
)
return result

View file

@ -1,97 +0,0 @@
"""LLM utilities for LongMemEval benchmark evaluation."""
import asyncio
import json
import logging
import re
from tenacity import retry, stop_after_attempt, wait_random_exponential, before_sleep_log
from reme.core.schema import Message
from reme.core.utils import load_env
from reme.reme import ReMe
logger = logging.getLogger(__name__)
load_env()
WAIT_TIME_LOWER = 1
WAIT_TIME_UPPER = 60
RETRY_TIMES = 5
@retry(
wait=wait_random_exponential(min=WAIT_TIME_LOWER, max=WAIT_TIME_UPPER),
stop=stop_after_attempt(3),
reraise=True,
before_sleep=before_sleep_log(logger, logging.WARNING),
)
async def llm_request(reme: ReMe, prompt: str, model_name: str = "qwen3-max", **kwargs) -> str:
"""Make an LLM request using ReMe's LLM with optional model override.
Args:
reme: ReMe instance
prompt: The prompt to send to the LLM
model_name: Optional model name to override the default model (default: "qwen3-max")
**kwargs: Additional arguments to pass to the chat method
Returns:
The assistant's response content
"""
assistant_message = await reme.llm.chat(
messages=[
Message(role="user", content=prompt),
],
model_name=model_name,
**kwargs,
)
return assistant_message.content
@retry(
wait=wait_random_exponential(min=WAIT_TIME_LOWER, max=WAIT_TIME_UPPER),
stop=stop_after_attempt(RETRY_TIMES),
reraise=True,
before_sleep=before_sleep_log(logger, logging.WARNING),
)
async def llm_request_for_json(reme: ReMe, prompt: str, model_name: str = "qwen-flash", **kwargs) -> dict:
"""Make an LLM request expecting JSON response using ReMe's LLM.
Args:
reme: ReMe instance
prompt: The prompt to send to the LLM
model_name: Optional model name to override the default model (default: "qwen-flash")
**kwargs: Additional arguments to pass to the chat method
Returns:
Parsed JSON object from the LLM response
Raises:
ValueError: If no JSON block is found in the model output
"""
content = await llm_request(reme, prompt, model_name=model_name, **kwargs)
match = re.search(r"```json\s*(\{.*?\})\s*```", content, re.DOTALL)
if not match:
raise ValueError(f"No JSON block found in model output: {content}")
json_str = match.group(1).strip()
return json.loads(json_str)
if __name__ == "__main__":
async def test():
"""Simple manual test for JSON LLM request."""
reme = ReMe()
await reme.start()
try:
r = await llm_request_for_json(
reme,
'hello? answer in ```json\n{"answer": "..."}```',
)
print(r)
finally:
await reme.close()
asyncio.run(test())

View file

@ -1,23 +1,26 @@
"""ReMe"""
"""ReMe CLI package."""
__version__ = "0.4.0.0"
from . import config
from . import core
from . import extension
from . import memory
from . import constants
from . import enumeration
from . import schema
from . import steps
from . import utils
from .application import Application
from .components import BaseComponent
from .reme import ReMe
__version__ = "0.3.1.10"
__all__ = [
"config",
"core",
"extension",
"memory",
"Application",
"BaseComponent",
"ReMe",
# submodules
"config",
"constants",
"enumeration",
"schema",
"steps",
"utils",
]
"""
conda create -n fl_test2 python=3.10
conda activate fl_test2
conda env remove -n fl_test2
"""

View file

@ -1,5 +1,8 @@
"""config"""
"""Config"""
from .reme_config_parser import ReMeConfigParser
from .config_parser import parse_args, resolve_app_config
__all__ = ["ReMeConfigParser"]
__all__ = [
"parse_args",
"resolve_app_config",
]

View file

@ -1,7 +0,0 @@
"""Configuration parser for ReMe framework."""
from ..core.utils import PydanticConfigParser
class ReMeConfigParser(PydanticConfigParser):
"""Configuration parser for ReMe framework."""

View file

@ -1,51 +0,0 @@
"""Core"""
from . import as_llm
from . import as_llm_formatter
from . import as_token_counter
from . import embedding
from . import enumeration
from . import file_store
from . import file_watcher
from . import flow
from . import llm
from . import op
from . import schema
from . import service
from . import token_counter
from . import utils
from . import vector_store
from .application import Application
from .base_dict import BaseDict
from .prompt_handler import PromptHandler
from .registry_factory import R, Registry, RegistryFactory
from .runtime_context import RuntimeContext
from .service_context import ServiceContext
__all__ = [
# Submodules
"as_llm",
"as_llm_formatter",
"as_token_counter",
"embedding",
"enumeration",
"file_watcher",
"flow",
"llm",
"file_store",
"op",
"schema",
"service",
"token_counter",
"utils",
"vector_store",
# Classes
"Application",
"BaseDict",
"PromptHandler",
"R",
"Registry",
"RegistryFactory",
"RuntimeContext",
"ServiceContext",
]

View file

@ -1,657 +0,0 @@
"""High-level entry point for configuring and running ReMe services and flows."""
import asyncio
import os
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from .embedding import BaseEmbeddingModel
from .file_store import BaseFileStore
from .file_watcher import BaseFileWatcher
from .flow import BaseFlow
from .llm import BaseLLM
from .prompt_handler import PromptHandler
from .registry_factory import R
from .schema import (
EmbeddingModelConfig,
Response,
ServiceConfig,
LLMConfig,
VectorStoreConfig,
FileStoreConfig,
FileWatcherConfig,
TokenCounterConfig,
)
from .service_context import ServiceContext
from .token_counter import BaseTokenCounter
from .utils import execute_stream_task, PydanticConfigParser, init_logger, MCPClient, print_logo, get_logger, load_env
from .vector_store import BaseVectorStore
logger = get_logger()
class Application:
"""Application wrapper that wires together service context, flows, and runtimes."""
def __init__(
self,
*args,
llm_api_key: str | None = None,
llm_base_url: str | None = None,
embedding_api_key: str | None = None,
embedding_base_url: str | None = None,
working_dir: str | None = None,
config_path: str | None = None,
enable_logo: bool = True,
log_to_console: bool = True,
log_to_file: bool = True,
enable_load_env: bool = True,
parser: type[PydanticConfigParser] | None = None,
default_as_llm_config: dict | None = None,
default_as_llm_formatter_config: dict | None = None,
default_llm_config: dict | None = None,
default_embedding_model_config: dict | None = None,
default_vector_store_config: dict | None = None,
default_file_store_config: dict | None = None,
default_token_counter_config: dict | None = None,
default_file_watcher_config: dict | None = None,
**kwargs,
):
if enable_load_env:
load_env()
self.llm_api_key = llm_api_key or os.getenv("LLM_API_KEY", "")
self.llm_base_url = llm_base_url or os.getenv("LLM_BASE_URL", "")
self.embedding_api_key = embedding_api_key or os.getenv("EMBEDDING_API_KEY", "")
self.embedding_base_url = embedding_base_url or os.getenv("EMBEDDING_BASE_URL", "")
self.service_context = ServiceContext(
*args,
service_config=None,
parser=parser,
working_dir=working_dir,
config_path=config_path,
enable_logo=enable_logo,
log_to_console=log_to_console,
log_to_file=log_to_file,
default_as_llm_config=default_as_llm_config,
default_as_llm_formatter_config=default_as_llm_formatter_config,
default_llm_config=default_llm_config,
default_embedding_model_config=default_embedding_model_config,
default_vector_store_config=default_vector_store_config,
default_file_store_config=default_file_store_config,
default_token_counter_config=default_token_counter_config,
default_file_watcher_config=default_file_watcher_config,
**kwargs,
)
self.prompt_handler = PromptHandler(language=self.service_config.language)
# NOTE: flows are initialized here to start service!
self.init_flows()
self._started: bool = False
@classmethod
async def create(cls, *args, **kwargs) -> "Application":
"""Create and start an Application instance asynchronously."""
instance = cls(*args, **kwargs)
await instance.start()
return instance
def init_flows(self):
"""Initialize flows."""
expression_flow_cls = None
for name, flow_cls in R.flows.items():
if not self._filter_flows(name):
continue
if name == "ExpressionFlow":
expression_flow_cls = flow_cls
else:
flow: "BaseFlow" = flow_cls(name=name, service_context=self.service_context)
self.service_context.flows[flow.name] = flow
if expression_flow_cls is not None:
for name, flow_config in self.service_config.flows.items():
if not self._filter_flows(name):
continue
flow_config.name = name
flow: BaseFlow = expression_flow_cls( # noqa
flow_config=flow_config,
service_context=self.service_context,
)
self.service_context.flows[flow.name] = flow
else:
logger.info("No expression flow found, please check your configuration.")
def _filter_flows(self, name: str) -> bool:
"""Filter flows based on enabled_flows and disabled_flows configuration."""
if self.service_config.enabled_flows:
return name in self.service_config.enabled_flows
elif self.service_config.disabled_flows:
return name not in self.service_config.disabled_flows
else:
return True
@property
def service_config(self) -> ServiceConfig:
"""Get the service configuration."""
return self.service_context.service_config
async def start(self):
"""Start the service context by initializing all configured components."""
if self._started:
logger.warning("Application has already started.")
return self
init_logger(
log_to_console=self.service_config.log_to_console,
log_to_file=self.service_config.log_to_file,
)
logger.info(f"Init ReMe with config: {self.service_config.model_dump_json()}")
working_path = Path(self.service_config.working_dir)
working_path.mkdir(parents=True, exist_ok=True)
if self.service_config.ray_max_workers > 1:
import ray
if not ray.is_initialized():
ray.init(num_cpus=self.service_config.ray_max_workers)
if self.service_config.thread_pool_max_workers > 0 and (
self.service_context.thread_pool is None
or self.service_context.thread_pool._shutdown # pylint: disable=protected-access
):
self.service_context.thread_pool = ThreadPoolExecutor(
max_workers=self.service_config.thread_pool_max_workers,
)
elif self.service_config.thread_pool_max_workers <= 0:
logger.info("Thread pool is disabled (thread_pool_max_workers <= 0)")
if self.service_context.service_config.enable_logo:
print_logo(service_config=self.service_config)
for name, config in self.service_config.as_llms.items():
if config.backend not in R.as_llms:
logger.warning(f"AS LLM backend {config.backend} is not supported.")
else:
try:
config_dict = config.model_dump(exclude={"backend"})
if not config_dict.get("api_key", ""):
config_dict["api_key"] = self.llm_api_key
if "client_kwargs" not in config_dict:
config_dict["client_kwargs"] = {}
if not config_dict["client_kwargs"].get("base_url", ""):
config_dict["client_kwargs"]["base_url"] = self.llm_base_url
self.service_context.as_llms[name] = R.as_llms[config.backend](**config_dict)
except Exception as e:
logger.error(f"Failed to initialize AS LLM '{name}': {e}")
for name, config in self.service_config.as_llm_formatters.items():
if config.backend not in R.as_llm_formatters:
logger.warning(f"AS LLM formatter backend {config.backend} is not supported.")
else:
try:
config_dict = config.model_dump(exclude={"backend"})
self.service_context.as_llm_formatters[name] = R.as_llm_formatters[config.backend](**config_dict)
except Exception as e:
logger.error(f"Failed to initialize AS LLM formatter '{name}': {e}")
for name, config in self.service_config.as_token_counters.items():
if config.backend not in R.as_token_counters:
logger.warning(f"Token counter backend {config.backend} is not supported.")
else:
try:
config_dict = config.model_dump(exclude={"backend"})
self.service_context.as_token_counters[name] = R.as_token_counters[config.backend](**config_dict)
except Exception as e:
logger.error(f"Failed to initialize AS token counter '{name}': {e}")
for name, config in self.service_config.llms.items():
if config.backend not in R.llms:
logger.warning(f"LLM backend {config.backend} is not supported.")
else:
config_dict = config.model_dump(exclude={"backend"})
config_dict.setdefault("api_key", self.llm_api_key)
config_dict.setdefault("base_url", self.llm_base_url)
self.service_context.llms[name] = R.llms[config.backend](**config_dict)
await self.service_context.llms[name].start()
for name, config in self.service_config.embedding_models.items():
if config.backend not in R.embedding_models:
logger.warning(f"Embedding model backend {config.backend} is not supported.")
else:
config_dict = config.model_dump(exclude={"backend"})
config_dict.setdefault("api_key", self.embedding_api_key)
config_dict.setdefault("base_url", self.embedding_base_url)
config_dict.setdefault("cache_dir", working_path / "embedding_cache")
self.service_context.embedding_models[name] = R.embedding_models[config.backend](**config_dict)
await self.service_context.embedding_models[name].start()
for name, config in self.service_config.token_counters.items():
if config.backend not in R.token_counters:
logger.warning(f"Token counter backend {config.backend} is not supported.")
else:
config_dict = config.model_dump(exclude={"backend"})
self.service_context.token_counters[name] = R.token_counters[config.backend](**config_dict)
for name, config in self.service_config.vector_stores.items():
if config.backend not in R.vector_stores:
logger.warning(f"Vector store backend {config.backend} is not supported.")
else:
config_dict = config.model_dump(exclude={"backend", "embedding_model"})
config_dict.update(
{
"embedding_model": self.service_context.embedding_models[config.embedding_model],
"db_path": working_path / "vector_store",
},
)
self.service_context.vector_stores[name] = R.vector_stores[config.backend](**config_dict)
await self.service_context.vector_stores[name].start()
for name, config in self.service_config.file_stores.items():
if config.backend not in R.file_stores:
logger.warning(f"File store backend {config.backend} is not supported.")
else:
config_dict = config.model_dump(exclude={"backend", "embedding_model"})
config_dict.update(
{
"embedding_model": self.service_context.embedding_models[config.embedding_model],
"db_path": working_path / "file_store",
},
)
self.service_context.file_stores[name] = R.file_stores[config.backend](**config_dict)
await self.service_context.file_stores[name].start()
for name, config in self.service_config.file_watchers.items():
if config.backend not in R.file_watchers:
logger.warning(f"File watcher backend {config.backend} is not supported.")
else:
config_dict = config.model_dump(exclude={"backend", "file_store"})
config_dict["file_store"] = self.service_context.file_stores[config.file_store]
self.service_context.file_watchers[name] = R.file_watchers[config.backend](**config_dict)
await self.service_context.file_watchers[name].start()
if self.service_config.mcp_servers:
await self.prepare_mcp_servers()
self._started = True
logger.info("ReMe Application started")
return self
# pylint: disable=too-many-statements
async def restart(self, restart_config: dict):
"""Restart the application with new config."""
working_path = Path(self.service_config.working_dir)
working_path.mkdir(parents=True, exist_ok=True)
# as_llms
if "as_llms" in restart_config:
as_llms_config = restart_config["as_llms"]
assert isinstance(as_llms_config, dict)
for name, config in as_llms_config.items():
if name in self.service_context.as_llms:
del self.service_context.as_llms[name]
if config.get("backend") not in R.as_llms:
logger.warning(f"AS LLM backend {config.get('backend')} is not supported.")
continue
try:
config_dict = {k: v for k, v in config.items() if k != "backend"}
if not config_dict.get("api_key", ""):
config_dict["api_key"] = self.llm_api_key
if "client_kwargs" not in config_dict:
config_dict["client_kwargs"] = {}
if not config_dict["client_kwargs"].get("base_url", ""):
config_dict["client_kwargs"]["base_url"] = self.llm_base_url
self.service_context.as_llms[name] = R.as_llms[config["backend"]](**config_dict)
logger.info(f"Restarted AS LLM: {name}")
except Exception as e:
logger.error(f"Failed to restart AS LLM '{name}': {e}")
# as_llm_formatters
if "as_llm_formatters" in restart_config:
as_llm_formatters_config = restart_config["as_llm_formatters"]
assert isinstance(as_llm_formatters_config, dict)
for name, config in as_llm_formatters_config.items():
if name in self.service_context.as_llm_formatters:
del self.service_context.as_llm_formatters[name]
if config.get("backend") not in R.as_llm_formatters:
logger.warning(f"AS LLM formatter backend {config.get('backend')} is not supported.")
continue
try:
config_dict = {k: v for k, v in config.items() if k != "backend"}
self.service_context.as_llm_formatters[name] = R.as_llm_formatters[config["backend"]](**config_dict)
logger.info(f"Restarted AS LLM formatter: {name}")
except Exception as e:
logger.error(f"Failed to restart AS LLM formatter '{name}': {e}")
# as_token_counters
if "as_token_counters" in restart_config:
as_token_counters_config = restart_config["as_token_counters"]
assert isinstance(as_token_counters_config, dict)
for name, config in as_token_counters_config.items():
if name in self.service_context.as_token_counters:
del self.service_context.as_token_counters[name]
if config.get("backend") not in R.as_token_counters:
logger.warning(f"Token counter backend {config.get('backend')} is not supported.")
continue
try:
config_dict = {k: v for k, v in config.items() if k != "backend"}
self.service_context.as_token_counters[name] = R.as_token_counters[config["backend"]](**config_dict)
logger.info(f"Restarted AS token counter: {name}")
except Exception as e:
logger.error(f"Failed to restart AS token counter '{name}': {e}")
# llms
if "llms" in restart_config:
llms_config = restart_config["llms"]
assert isinstance(llms_config, dict)
for name, config in llms_config.items():
if name in self.service_context.llms:
llm = self.service_context.llms.pop(name)
await llm.close()
if isinstance(config, dict):
config = LLMConfig(**config)
if config.backend not in R.llms:
logger.warning(f"LLM backend {config.backend} is not supported.")
continue
config_dict = config.model_dump(exclude={"backend"})
config_dict.setdefault("api_key", self.llm_api_key)
config_dict.setdefault("base_url", self.llm_base_url)
self.service_context.llms[name] = R.llms[config.backend](**config_dict)
await self.service_context.llms[name].start()
logger.info(f"Restarted LLM: {name}")
# embedding_models
if "embedding_models" in restart_config:
embedding_models_config = restart_config["embedding_models"]
assert isinstance(embedding_models_config, dict)
updated_names = set()
for name, config in embedding_models_config.items():
if name in self.service_context.embedding_models:
embedding_model = self.service_context.embedding_models.pop(name)
await embedding_model.close()
if isinstance(config, dict):
config = EmbeddingModelConfig(**config)
if config.backend not in R.embedding_models:
logger.warning(f"Embedding model backend {config.backend} is not supported.")
continue
config_dict = config.model_dump(exclude={"backend"})
config_dict.setdefault("api_key", self.embedding_api_key)
config_dict.setdefault("base_url", self.embedding_base_url)
config_dict.setdefault("cache_dir", working_path / "embedding_cache")
self.service_context.embedding_models[name] = R.embedding_models[config.backend](**config_dict)
await self.service_context.embedding_models[name].start()
logger.info(f"Restarted embedding model: {name}")
updated_names.add(name)
# update embedding_model attribute for existing vector_stores and file_stores
for name in updated_names:
for vs_name, vs_config in self.service_config.vector_stores.items():
if vs_config.embedding_model == name and vs_name in self.service_context.vector_stores:
self.service_context.vector_stores[vs_name].embedding_model = (
self.service_context.embedding_models[name]
)
logger.info(f"Updated embedding model for vector store: {vs_name}")
for fs_name, fs_config in self.service_config.file_stores.items():
if fs_config.embedding_model == name and fs_name in self.service_context.file_stores:
self.service_context.file_stores[fs_name].embedding_model = (
self.service_context.embedding_models[name]
)
logger.info(f"Updated embedding model for file store: {fs_name}")
# token_counters
if "token_counters" in restart_config:
token_counters_config = restart_config["token_counters"]
assert isinstance(token_counters_config, dict)
for name, config in token_counters_config.items():
if name in self.service_context.token_counters:
del self.service_context.token_counters[name]
if isinstance(config, dict):
config = TokenCounterConfig(**config)
if config.backend not in R.token_counters:
logger.warning(f"Token counter backend {config.backend} is not supported.")
continue
config_dict = config.model_dump(exclude={"backend"})
self.service_context.token_counters[name] = R.token_counters[config.backend](**config_dict)
logger.info(f"Restarted token counter: {name}")
# vector_stores
if "vector_stores" in restart_config:
vector_stores_config = restart_config["vector_stores"]
assert isinstance(vector_stores_config, dict)
for name, config in vector_stores_config.items():
if name in self.service_context.vector_stores:
vector_store = self.service_context.vector_stores.pop(name)
await vector_store.close()
if isinstance(config, dict):
config = VectorStoreConfig(**config)
if config.backend not in R.vector_stores:
logger.warning(f"Vector store backend {config.backend} is not supported.")
continue
config_dict = config.model_dump(exclude={"backend", "embedding_model"})
config_dict.update(
{
"embedding_model": self.service_context.embedding_models[config.embedding_model],
"db_path": working_path / "vector_store",
},
)
self.service_context.vector_stores[name] = R.vector_stores[config.backend](**config_dict)
await self.service_context.vector_stores[name].start()
logger.info(f"Restarted vector store: {name}")
# file_stores
if "file_stores" in restart_config:
file_stores_config = restart_config["file_stores"]
assert isinstance(file_stores_config, dict)
for name, config in file_stores_config.items():
if name in self.service_context.file_stores:
file_store = self.service_context.file_stores.pop(name)
await file_store.close()
if isinstance(config, dict):
config = FileStoreConfig(**config)
if config.backend not in R.file_stores:
logger.warning(f"File store backend {config.backend} is not supported.")
continue
config_dict = config.model_dump(exclude={"backend", "embedding_model"})
config_dict.update(
{
"embedding_model": self.service_context.embedding_models[config.embedding_model],
"db_path": working_path / "file_store",
},
)
self.service_context.file_stores[name] = R.file_stores[config.backend](**config_dict)
await self.service_context.file_stores[name].start()
logger.info(f"Restarted file store: {name}")
# file_watchers
if "file_watchers" in restart_config:
file_watchers_config = restart_config["file_watchers"]
assert isinstance(file_watchers_config, dict)
for name, config in file_watchers_config.items():
if name in self.service_context.file_watchers:
file_watcher = self.service_context.file_watchers.pop(name)
await file_watcher.close()
if isinstance(config, dict):
config = FileWatcherConfig(**config)
if config.backend not in R.file_watchers:
logger.warning(f"File watcher backend {config.backend} is not supported.")
continue
config_dict = config.model_dump(exclude={"backend", "file_store"})
config_dict["file_store"] = self.service_context.file_stores[config.file_store]
self.service_context.file_watchers[name] = R.file_watchers[config.backend](**config_dict)
await self.service_context.file_watchers[name].start()
logger.info(f"Restarted file watcher: {name}")
async def prepare_mcp_servers(self):
"""Prepare and initialize MCP server connections."""
mcp_client = MCPClient(config={"mcpServers": self.service_config.mcp_servers})
for server_name in self.service_config.mcp_servers.keys():
try:
tool_calls = await mcp_client.list_tool_calls(server_name=server_name, return_dict=False)
self.service_context.mcp_server_mapping[server_name] = {
tool_call.name: tool_call for tool_call in tool_calls
}
for tool_call in tool_calls:
logger.info(f"list_tool_calls: {server_name}@{tool_call.name} {tool_call.simple_input_dump()}")
except Exception as e:
logger.exception(f"list_tool_calls: {server_name} error: {e}")
async def close(self) -> bool:
"""Close all service components asynchronously."""
if not self._started:
logger.warning("Application is not started")
return True
for name, file_watcher in self.service_context.file_watchers.items():
logger.info(f"Closing file watcher: {name}")
await file_watcher.close()
for name, file_store in self.service_context.file_stores.items():
logger.info(f"Closing file store: {name}")
await file_store.close()
for name, vector_store in self.service_context.vector_stores.items():
logger.info(f"Closing vector store: {name}")
await vector_store.close()
for name, llm in self.service_context.llms.items():
logger.info(f"Closing LLM: {name}")
await llm.close()
for name, embedding_model in self.service_context.embedding_models.items():
logger.info(f"Closing embedding model: {name}")
await embedding_model.close()
self.shutdown_thread_pool()
self.shutdown_ray()
self._started = False
logger.info("ReMe Application closed")
return False
def shutdown_thread_pool(self, wait: bool = True):
"""Shutdown the thread pool executor."""
if self.service_context.thread_pool is not None:
self.service_context.thread_pool.shutdown(wait=wait)
def shutdown_ray(self, wait: bool = True):
"""Shutdown Ray cluster if it was initialized."""
if self.service_config and self.service_config.ray_max_workers > 1:
import ray
ray.shutdown(_exiting_interpreter=not wait)
async def __aenter__(self):
"""Async context manager entry."""
return await self.start()
async def __aexit__(self, exc_type=None, exc_val=None, exc_tb=None):
"""Async context manager exit."""
return await self.close()
async def execute_flow(self, name: str, **kwargs) -> Response:
"""Execute a flow with the given name and parameters."""
assert name in self.service_context.flows, f"Flow {name} not found"
flow: BaseFlow = self.service_context.flows[name]
return await flow.call(**kwargs)
async def execute_stream_flow(self, name: str, **kwargs):
"""Execute a stream flow with the given name and parameters."""
assert name in self.service_context.flows, f"Flow {name} not found"
flow: BaseFlow = self.service_context.flows[name]
assert flow.stream is True, "non-stream flow is not supported in execute_stream_flow!"
stream_queue = asyncio.Queue()
task = asyncio.create_task(flow.call(stream_queue=stream_queue, **kwargs))
async for chunk in execute_stream_task(
stream_queue=stream_queue,
task=task,
task_name=name,
output_format="str",
):
yield chunk
@property
def default_llm(self) -> BaseLLM:
"""Get the default LLM instance."""
return self.service_context.llms.get("default")
def get_llm(self, name: str):
"""Get an LLM instance by name."""
return self.service_context.llms.get(name)
def update_default_llm_name(self, name: str):
"""Update the default LLM name."""
self.default_llm.model_name = name
@property
def default_embedding_model(self) -> BaseEmbeddingModel:
"""Get the default embedding model instance."""
return self.service_context.embedding_models.get("default")
def get_embedding_model(self, name: str):
"""Get an embedding model instance by name."""
return self.service_context.embedding_models.get(name)
def update_default_embedding_name(self, name: str):
"""Update the default embedding model name."""
self.default_embedding_model.model_name = name
@property
def default_vector_store(self) -> BaseVectorStore:
"""Get the default vector store instance."""
return self.service_context.vector_stores.get("default")
def get_vector_store(self, name: str):
"""Get a vector store instance by name."""
return self.service_context.vector_stores.get(name)
@property
def default_file_store(self) -> BaseFileStore:
"""Get the default file store instance."""
return self.service_context.file_stores.get("default")
def get_file_store(self, name: str):
"""Get a file store instance by name."""
return self.service_context.file_stores.get(name)
@property
def default_file_watcher(self) -> BaseFileWatcher:
"""Get the default file watcher instance."""
return self.service_context.file_watchers.get("default")
def get_file_watcher(self, name: str):
"""Get a file watcher instance by name."""
return self.service_context.file_watchers.get(name)
@property
def default_token_counter(self) -> BaseTokenCounter:
"""Get the default token counter instance."""
return self.service_context.token_counters.get("default")
def get_token_counter(self, name: str):
"""Get a token counter instance by name."""
return self.service_context.token_counters.get(name)
def run_service(self):
"""Run the configured service (HTTP, MCP, or CMD)."""
import warnings
warnings.filterwarnings("ignore", category=DeprecationWarning)
service = R.services[self.service_config.backend](app=self)
service.run()
async def reset_default_collection(self, collection_name: str):
"""Reset the default vector store."""
await self.service_context.vector_stores["default"].reset_collection(collection_name)

View file

@ -1,9 +0,0 @@
"""Module for registering AgentScope LLM models."""
from agentscope.model import DashScopeChatModel
from agentscope.model import OpenAIChatModel
from ..registry_factory import R
R.as_llms.register("openai")(OpenAIChatModel)
R.as_llms.register("dashscope")(DashScopeChatModel)

View file

@ -1,9 +0,0 @@
"""Module for registering AgentScope LLM formatters."""
from agentscope.formatter import DashScopeChatFormatter
from .reme_openai_chat_formatter import ReMeOpenAIChatFormatter
from ..registry_factory import R
R.as_llm_formatters.register("openai")(ReMeOpenAIChatFormatter)
R.as_llm_formatters.register("dashscope")(DashScopeChatFormatter)

View file

@ -1,215 +0,0 @@
"""ReMeOpenAIChatFormatter"""
import json
from typing import Any
from agentscope.formatter import OpenAIChatFormatter
from agentscope.formatter._openai_formatter import (
_format_openai_image_block,
_to_openai_audio_data,
)
from agentscope.message import Msg, TextBlock, ImageBlock, URLSource
from loguru import logger
def _format_openai_video_block(video_block: dict) -> dict[str, Any]:
"""Format a video block for OpenAI API.
Args:
video_block: The video block to format.
Returns:
A dictionary with video content in OpenAI format.
"""
source = video_block["source"]
if source["type"] == "url":
url = source["url"]
elif source["type"] == "base64":
data = source["data"]
media_type = source["media_type"]
url = f"data:{media_type};base64,{data}"
else:
raise ValueError(f"Unsupported video source type: {source['type']}")
return {
"type": "video_url",
"video_url": {
"url": url,
},
}
class ReMeOpenAIChatFormatter(OpenAIChatFormatter):
"""ReMeOpenAIChatFormatter"""
async def _format(
self,
msgs: list[Msg],
) -> list[dict[str, Any]]:
"""Format message objects into OpenAI API required format.
Args:
msgs (`list[Msg]`):
The list of Msg objects to format.
Returns:
`list[dict[str, Any]]`:
A list of dictionaries, where each dictionary has "name",
"role", and "content" keys.
"""
self.assert_list_of_msgs(msgs)
messages: list[dict] = []
i = 0
while i < len(msgs):
msg = msgs[i]
content_blocks = []
tool_calls = []
reasoning_content_blocks = []
for block in msg.get_content_blocks():
typ = block.get("type")
if typ == "text":
content_blocks.append({**block})
elif typ == "thinking":
# Collect thinking blocks for reasoning_content field
# This is compatible with models like DeepSeek that support
# extended thinking via reasoning_content field
reasoning_content_blocks.append({**block})
elif typ == "tool_use":
tool_calls.append(
{
"id": block.get("id"),
"type": "function",
"function": {
"name": block.get("name"),
"arguments": json.dumps(
block.get("input", {}),
ensure_ascii=False,
),
},
},
)
elif typ == "tool_result":
(
textual_output,
multimodal_data,
) = self.convert_tool_result_to_string(block["output"])
messages.append(
{
"role": "tool",
"tool_call_id": block.get("id"),
"content": (textual_output), # type: ignore[arg-type]
"name": block.get("name"),
},
)
# Then, handle the multimodal data if any
promoted_blocks: list = []
for url, multimodal_block in multimodal_data:
if multimodal_block["type"] == "image" and self.promote_tool_result_images:
promoted_blocks.extend(
[
TextBlock(
type="text",
text=f"\n- The image from '{url}': ",
),
ImageBlock(
type="image",
source=URLSource(
type="url",
url=url,
),
),
],
)
if promoted_blocks:
# Insert promoted blocks as new user message(s)
promoted_blocks = [
TextBlock(
type="text",
text="<system-info>The following are "
"the image contents from the tool "
f"result of '{block['name']}':",
),
*promoted_blocks,
TextBlock(
type="text",
text="</system-info>",
),
]
msgs.insert(
i + 1,
Msg(
name="user",
content=promoted_blocks,
role="user",
),
)
elif typ == "image":
content_blocks.append(
_format_openai_image_block(
block, # type: ignore[arg-type]
),
)
elif typ == "audio":
# Filter out audio content when the multimodal model
# outputs both text and audio, to prevent errors in
# subsequent model calls
if msg.role == "assistant":
continue
input_audio = _to_openai_audio_data(block["source"])
content_blocks.append(
{
"type": "input_audio",
"input_audio": input_audio,
},
)
elif typ == "video":
# Filter out video content when the multimodal model
# outputs both text and video, to prevent errors in
# subsequent model calls
if msg.role == "assistant":
continue
content_blocks.append(
_format_openai_video_block(block),
)
else:
logger.warning(
"Unsupported block type %s in the message, skipped.",
typ,
)
msg_openai = {
"role": msg.role,
"name": msg.name,
"content": content_blocks or None,
}
if tool_calls:
msg_openai["tool_calls"] = tool_calls
# Add reasoning_content for thinking blocks (compatible with DeepSeek, etc.)
if reasoning_content_blocks:
reasoning_msg = "\n".join(reasoning.get("thinking", "") for reasoning in reasoning_content_blocks)
if reasoning_msg:
msg_openai["reasoning_content"] = reasoning_msg
# When both content and tool_calls are None, skipped
if msg_openai["content"] or msg_openai.get("tool_calls"):
messages.append(msg_openai)
# Move to next message
i += 1
return messages

View file

@ -1,8 +0,0 @@
"""Module for registering AgentScope token counters."""
from .reme_token_counter import ReMeTokenCounter
from .rule_token_counter import RuleTokenCounter
from ..registry_factory import R
R.as_token_counters.register("hf")(ReMeTokenCounter)
R.as_token_counters.register("rule")(RuleTokenCounter)

View file

@ -1,123 +0,0 @@
"""Token counter for ReMe."""
import os
from typing import Any
from agentscope.token import HuggingFaceTokenCounter
from ..utils import get_logger
logger = get_logger()
class ReMeTokenCounter(HuggingFaceTokenCounter):
"""Token counter for CoPaw with configurable tokenizer support.
This class extends HuggingFaceTokenCounter to provide token counting
functionality with support for both local and remote tokenizers,
as well as HuggingFace mirror for users in China.
Attributes:
pretrained_model_name_or_path: The tokenizer model path or "default" for local tokenizer.
use_mirror: Whether to use HuggingFace mirror.
token_count_estimate_divisor: Divisor for token estimation.
"""
def __init__(
self,
pretrained_model_name_or_path: str,
use_mirror: bool = True,
token_count_estimate_divisor: float = 3.75,
**kwargs,
):
"""Initialize the token counter with the specified configuration.
Args:
pretrained_model_name_or_path: The tokenizer model path.
use_mirror: Whether to use the HuggingFace mirror
(https://hf-mirror.com) for downloading tokenizers. Useful for
users in China.
token_count_estimate_divisor: Divisor for estimating tokens when
tokenizer is unavailable. Defaults to 3.75.
**kwargs: Additional keyword arguments passed to HuggingFaceTokenCounter.
"""
self.pretrained_model_name_or_path = pretrained_model_name_or_path
self.use_mirror = use_mirror
self.token_count_estimate_divisor = token_count_estimate_divisor
# Set HuggingFace endpoint for mirror support
if use_mirror:
mirror = "https://hf-mirror.com"
else:
mirror = "https://huggingface.co"
os.environ["HF_ENDPOINT"] = mirror
# if the huggingface is already imported in other dependencies,
# we need to set the endpoint manually
import huggingface_hub.constants
huggingface_hub.constants.ENDPOINT = mirror
huggingface_hub.constants.HUGGINGFACE_CO_URL_TEMPLATE = mirror + "/{repo_id}/resolve/{revision}/{filename}"
try:
super().__init__(
pretrained_model_name_or_path=self.pretrained_model_name_or_path,
use_mirror=use_mirror,
use_fast=True,
trust_remote_code=True,
**kwargs,
)
self._tokenizer_available = True
except Exception as e:
logger.error(f"Failed to initialize tokenizer {e}")
self._tokenizer_available = False
async def count(
self,
messages: list[dict],
tools: list[dict] | None = None,
text: str | None = None,
**kwargs: Any,
) -> int:
"""Count tokens in messages or text.
If text is provided, counts tokens directly in the text string.
Otherwise, counts tokens in the messages using the parent class method.
Args:
messages: List of message dictionaries in chat format.
tools: Optional list of tool definitions for token counting.
text: Optional text string to count tokens directly.
**kwargs: Additional keyword arguments passed to parent count method.
Returns:
The number of tokens, guaranteed to be at least the estimated minimum.
"""
if text:
if self._tokenizer_available:
try:
token_ids = self.tokenizer.encode(text)
return max(len(token_ids), self.estimate_tokens(text))
except Exception as e:
logger.exception("Failed to encode text with tokenizer: %s", e)
return self.estimate_tokens(text)
else:
return self.estimate_tokens(text)
else:
return await super().count(messages, tools, **kwargs)
def estimate_tokens(self, text: str) -> int:
"""Estimate the number of tokens in a text string.
Provides a fast character-based estimation as a fallback or lower bound.
Uses the configured divisor from instance settings.
Args:
text: The text string to estimate tokens for.
Returns:
The estimated number of tokens in the text string.
"""
return int(len(text.encode("utf-8")) / self.token_count_estimate_divisor + 0.5)

View file

@ -1,78 +0,0 @@
"""Rule-based token counter for fast estimation without loading tokenizer."""
from typing import Any
from agentscope.token import HuggingFaceTokenCounter
class RuleTokenCounter(HuggingFaceTokenCounter):
"""Lightweight token counter using rule-based estimation only.
This class provides fast token estimation without loading any tokenizer,
useful when exact token counts are not critical or for quick approximations.
Attributes:
token_count_estimate_divisor: Divisor for token estimation.
"""
def __init__(
self,
token_count_estimate_divisor: float = 3.75,
**_kwargs,
):
"""Initialize the rule-based token counter.
Args:
token_count_estimate_divisor: Divisor for estimating tokens.
Defaults to 3.75 (approximately 4 characters per token).
**kwargs: Additional keyword arguments (ignored).
"""
self.token_count_estimate_divisor = token_count_estimate_divisor
# Skip tokenizer initialization from parent
self._tokenizer_available = False
async def count(
self,
messages: list[dict],
_tools: list[dict] | None = None,
text: str | None = None,
**_kwargs: Any,
) -> int:
"""Count tokens using rule-based estimation.
Args:
messages: List of message dictionaries in chat format.
_tools: Optional list of tool definitions (ignored).
text: Optional text string to count tokens directly.
**_kwargs: Additional keyword arguments (ignored).
Returns:
The estimated number of tokens.
"""
if text:
return self.estimate_tokens(text)
# Estimate from messages
total_text = ""
for msg in messages:
content = msg.get("content", "")
if isinstance(content, str):
total_text += content
elif isinstance(content, list):
for part in content:
if isinstance(part, dict) and "text" in part:
total_text += part["text"]
return self.estimate_tokens(total_text)
def estimate_tokens(self, text: str) -> int:
"""Estimate the number of tokens in a text string.
Uses character-based estimation with the configured divisor.
Args:
text: The text string to estimate tokens for.
Returns:
The estimated number of tokens in the text string.
"""
return int(len(text.encode("utf-8")) / self.token_count_estimate_divisor + 0.5)

View file

@ -1,41 +0,0 @@
"""Module providing a dictionary subclass with attribute-style access and pickling support."""
from typing import Generic, TypeVar
_KT = TypeVar("_KT")
_VT = TypeVar("_VT")
class BaseDict(dict, Generic[_KT, _VT]):
"""A dictionary subclass that enables accessing and modifying keys as attributes."""
def __getattr__(self, name: str) -> _VT:
"""Retrieve a dictionary item as an attribute."""
try:
return self[name]
except KeyError as e:
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e
def __setattr__(self, name: str, value: _VT) -> None:
"""Assign a value to a dictionary item using attribute syntax."""
self[name] = value
def __delattr__(self, name: str) -> None:
"""Remove a dictionary item using attribute syntax."""
try:
# Delete item from dict via key
del self[name]
except KeyError as e:
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e
def __getstate__(self) -> dict:
"""Return the dictionary representation for pickling."""
return dict(self)
def __setstate__(self, state: dict) -> None:
"""Restore the dictionary state from a pickled object."""
self.update(state)
def __reduce__(self):
"""Define the reconstruction logic for pickling processes."""
return self.__class__, (), self.__getstate__()

View file

@ -1,15 +0,0 @@
"""embedding"""
from .base_embedding_model import BaseEmbeddingModel
from .openai_embedding_model import OpenAIEmbeddingModel
from .openai_embedding_model_sync import OpenAIEmbeddingModelSync
from ..registry_factory import R
__all__ = [
"BaseEmbeddingModel",
"OpenAIEmbeddingModel",
"OpenAIEmbeddingModelSync",
]
R.embedding_models.register("openai")(OpenAIEmbeddingModel)
R.embedding_models.register("openai_sync")(OpenAIEmbeddingModelSync)

View file

@ -1,599 +0,0 @@
"""Base embedding model interface for ReMe.
Defines the abstract base class and standard API for all embedding model implementations.
"""
import asyncio
import hashlib
import json
import time
from abc import ABC
from collections import OrderedDict
from pathlib import Path
from loguru import logger
from ..schema import VectorNode, MemoryChunk
class BaseEmbeddingModel(ABC):
"""Abstract base class for embedding model implementations.
Provides a standard interface for text-to-vector generation with
built-in batching, retry logic, and error handling.
"""
def __init__(
self,
api_key: str | None = None,
base_url: str | None = None,
model_name: str = "",
dimensions: int = 1024,
use_dimensions: bool = False,
max_batch_size: int = 10,
max_retries: int = 3,
raise_exception: bool = True,
max_input_length: int = 8192,
cache_dir: str | Path = ".reme",
max_cache_size: int = 2000,
enable_cache: bool = True,
**kwargs,
):
"""Initialize model configuration and parameters.
Args:
api_key: API key for the embedding service
base_url: Base URL for the embedding service
model_name: Name of the embedding model
dimensions: Vector dimensions of the embeddings
use_dimensions: Whether to pass dimensions parameter to API (some APIs don't support it)
max_batch_size: Maximum batch size for embedding requests
max_retries: Maximum number of retry attempts on failure
raise_exception: Whether to raise exceptions on failure
max_input_length: Maximum input text length
max_cache_size: Maximum number of embeddings to cache in memory (LRU)
enable_cache: Whether to enable embedding cache
**kwargs: Additional model-specific parameters
"""
self.api_key: str | None = api_key
self.base_url: str | None = base_url
self.model_name = model_name
self.dimensions = dimensions
self.use_dimensions = use_dimensions
self.max_batch_size = max_batch_size
self.max_retries = max_retries
self.raise_exception = raise_exception
self.max_input_length = max_input_length
self.cache_dir = cache_dir
self.max_cache_size = max_cache_size
self.enable_cache = enable_cache
self.kwargs = kwargs
# Initialize LRU cache for embeddings
self._embedding_cache: OrderedDict[str, list[float]] = OrderedDict()
self._cache_hits = 0
self._cache_misses = 0
self.cache_path: Path = Path(self.cache_dir)
self.cache_path.mkdir(parents=True, exist_ok=True)
def _truncate_text(self, text: str) -> str:
"""Truncate text to max_input_length if it exceeds the limit."""
if len(text) > self.max_input_length:
logger.warning(f"Text length {len(text)} exceeds {self.max_input_length}, truncating")
return text[: self.max_input_length]
return text
def _truncate_texts(self, texts: list[str]) -> list[str]:
"""Truncate a list of texts to max_input_length."""
return [self._truncate_text(text) for text in texts]
def _validate_and_adjust_embedding(self, embedding: list[float]) -> list[float]:
"""Validate and adjust embedding dimensions to match expected dimensions.
Args:
embedding: The embedding vector to validate
Returns:
Embedding vector adjusted to match self.dimensions
"""
actual_len = len(embedding)
if actual_len == self.dimensions:
return embedding
elif actual_len < self.dimensions:
logger.warning(
f"[ACTUAL_EMB_LENGTH]Embedding dimensions {actual_len} is less than expected {self.dimensions}, "
f"padding with zeros",
)
return embedding + [0.0] * (self.dimensions - actual_len)
else:
logger.warning(
f"[ACTUAL_EMB_LENGTH]Embedding dimensions {actual_len} is greater than expected {self.dimensions}, "
f"truncating to {self.dimensions}",
)
return embedding[: self.dimensions]
def _get_cache_key(self, text: str, dimensions: int) -> str:
"""Generate a cache key by hashing text + model_name + dimensions.
This ensures that the same text produces different cache keys when
using different models or dimensions.
Args:
text: Input text to hash
dimensions: Vector dimensions of the embeddings
Returns:
SHA256 hash combining text, model name, and dimensions
"""
# Combine text, model_name, and dimensions to create unique cache key
cache_string = f"{text}|{self.model_name}|{dimensions}"
return hashlib.sha256(cache_string.encode("utf-8")).hexdigest()
def _get_cache_file_path(self) -> Path:
"""Get the path to the cache file.
Returns:
Path to the embedding cache JSONL file
"""
return self.cache_path / "embedding_cache.jsonl"
def _load_cache(self) -> None:
"""Load embedding cache from disk (JSONL format).
Each line in the JSONL file contains a JSON object with:
- key: the cache key (SHA256 hash)
- embedding: the embedding vector (list of floats)
Loads in reverse order (newest first) to prioritize recent embeddings
when max_cache_size is smaller than the file content.
"""
if not self.enable_cache:
return
cache_file = self._get_cache_file_path()
if not cache_file.exists():
logger.info(f"No cache file found at {cache_file}, starting with empty cache")
return
try:
load_start = time.time()
# Read all lines first (to load in reverse order)
with open(cache_file, "r", encoding="utf-8") as f:
lines = f.readlines()
loaded_count = 0
# Load in reverse order (newest entries first)
for _, line in enumerate(reversed(lines), 1):
line = line.strip()
if not line:
continue
try:
data = json.loads(line)
if not data:
continue
# Each line is {cache_key: embedding}
cache_key, embedding = next(iter(data.items()))
if cache_key and embedding and isinstance(embedding, list):
# Skip if already loaded (keep the newest)
if cache_key in self._embedding_cache:
continue
if len(embedding) != self.dimensions:
logger.warning(
f"Embedding dimensions mismatch for cache key {cache_key}, "
f"expected {self.dimensions}, got {len(embedding)}",
)
continue
# Respect max_cache_size during loading
if len(self._embedding_cache) >= self.max_cache_size:
logger.info(
f"Cache size limit reached ({self.max_cache_size}), "
f"loaded {loaded_count} newest entries",
)
break
self._embedding_cache[cache_key] = embedding
loaded_count += 1
except json.JSONDecodeError as e:
logger.warning(f"Failed to parse line in cache file: {e}")
continue
logger.info(
f"Loaded {loaded_count} embeddings from cache file: {cache_file} in {time.time() - load_start:.2f}s",
)
except Exception as e:
logger.error(f"Failed to load cache from {cache_file}: {e}, deleting cache file")
try:
cache_file.unlink()
logger.info(f"Deleted corrupted cache file: {cache_file}")
except Exception as del_e:
logger.error(f"Failed to delete cache file {cache_file}: {del_e}")
def _save_cache(self) -> None:
"""Save embedding cache to disk (JSONL format).
Each line contains a JSON object with the cache key and embedding vector.
Only saves if cache is non-empty.
"""
if not self.enable_cache:
return
logger.info(f"Attempting to save cache, current size: {len(self._embedding_cache)}")
if not self._embedding_cache:
logger.info("Cache is empty, skipping save")
return
cache_file = self._get_cache_file_path()
try:
with open(cache_file, "w", encoding="utf-8") as f:
for cache_key, embedding in self._embedding_cache.items():
if len(embedding) != self.dimensions:
logger.warning(
f"Embedding dimensions mismatch for cache key {cache_key}, "
f"expected {self.dimensions}, got {len(embedding)}",
)
continue
cache_entry = {cache_key: embedding}
f.write(json.dumps(cache_entry, ensure_ascii=False) + "\n")
logger.info(f"Saved {len(self._embedding_cache)} embeddings to cache file: {cache_file}")
except Exception as e:
logger.error(f"Failed to save cache to {cache_file}: {e}")
def _get_from_cache(self, text: str) -> list[float] | None:
"""Retrieve embedding from cache if it exists.
Args:
text: Input text to look up
Returns:
Cached embedding vector or None if not found
"""
if not self.enable_cache:
return None
cache_key = self._get_cache_key(text, self.dimensions)
if cache_key in self._embedding_cache:
embeddings: list[float] = self._embedding_cache[cache_key]
# Validate embedding dimensions match expected dimensions
if len(embeddings) != self.dimensions:
logger.warning(
f"Cached embedding dimensions mismatch: expected {self.dimensions}, "
f"got {len(embeddings)}. Removing invalid cache entry.",
)
del self._embedding_cache[cache_key]
self._cache_misses += 1
return None
# Move to end (most recently used)
self._embedding_cache.move_to_end(cache_key)
self._cache_hits += 1
text_preview = text[:50] + "..." if len(text) > 50 else text
logger.info(f"Cache hit for text: {text_preview} (hits: {self._cache_hits}, misses: {self._cache_misses})")
return embeddings
self._cache_misses += 1
return None
def _put_to_cache(self, text: str, embedding: list[float]) -> None:
"""Store embedding in cache with LRU eviction.
Args:
text: Input text used as cache key
embedding: Embedding vector to cache
"""
if not self.enable_cache:
return
if self.max_cache_size <= 0:
return
cache_key = self._get_cache_key(text, self.dimensions)
if len(embedding) != self.dimensions:
logger.warning(
f"[PUT_TO_CACHE] Embedding dimensions mismatch for cache key {cache_key}, "
f"expected {self.dimensions}, got real length {len(embedding)}",
)
return
# Remove the oldest entry if cache is full
if len(self._embedding_cache) >= self.max_cache_size and cache_key not in self._embedding_cache:
self._embedding_cache.popitem(last=False)
self._embedding_cache[cache_key] = embedding
self._embedding_cache.move_to_end(cache_key)
def get_cache_stats(self) -> dict[str, int]:
"""Get cache statistics.
Returns:
Dictionary with cache size, hits, misses, and hit rate
"""
total_requests = self._cache_hits + self._cache_misses
hit_rate = self._cache_hits / total_requests if total_requests > 0 else 0.0
return {
"cache_size": len(self._embedding_cache),
"max_cache_size": self.max_cache_size,
"cache_hits": self._cache_hits,
"cache_misses": self._cache_misses,
"hit_rate": hit_rate,
}
def clear_cache(self) -> None:
"""Clear the embedding cache and reset statistics."""
self._embedding_cache.clear()
self._cache_hits = 0
self._cache_misses = 0
async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Internal async implementation for calling the embedding API with batch input."""
def _get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Internal synchronous implementation for calling the embedding API with batch input."""
async def get_embedding(self, input_text: str, **kwargs) -> list[float]:
"""Async get embedding for a single text with exponential backoff retries."""
truncated_text = self._truncate_text(input_text)
# Check cache first
cached_embedding = self._get_from_cache(truncated_text)
if cached_embedding is not None:
return cached_embedding
# Cache miss - compute embedding
for i in range(self.max_retries):
try:
result = await self._get_embeddings([truncated_text], **kwargs)
embedding = self._validate_and_adjust_embedding(result[0])
# Store in cache
self._put_to_cache(truncated_text, embedding)
return embedding
except Exception as e:
logger.error(f"Model {self.model_name} failed: {e}")
if i == self.max_retries - 1:
if self.raise_exception:
raise
return []
await asyncio.sleep(i + 1)
return []
async def get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Async get embeddings with automatic batching and exponential backoff retries."""
# Truncate all input texts first
truncated_texts = self._truncate_texts(input_text)
# Check cache for each text and separate cached vs uncached
results: list[list[float] | None] = [None] * len(truncated_texts)
texts_to_compute: list[tuple[int, str]] = [] # (original_index, text)
for idx, text in enumerate(truncated_texts):
cached = self._get_from_cache(text)
if cached is not None:
results[idx] = cached
else:
texts_to_compute.append((idx, text))
# If all texts were cached, return early
if not texts_to_compute:
return [r for r in results if r is not None]
# Compute embeddings for uncached texts in batches
uncached_texts = [text for _, text in texts_to_compute]
for i in range(0, len(uncached_texts), self.max_batch_size):
batch_texts = uncached_texts[i : i + self.max_batch_size]
batch_indices = [idx for idx, _ in texts_to_compute[i : i + self.max_batch_size]]
# Process each batch with retry logic
for retry in range(self.max_retries):
try:
batch_embeddings = await self._get_embeddings(batch_texts, **kwargs)
if batch_embeddings:
# Store results and cache them
for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings):
adjusted_embedding = self._validate_and_adjust_embedding(embedding)
results[orig_idx] = adjusted_embedding
self._put_to_cache(text, adjusted_embedding)
break
except Exception as e:
logger.error(f"Model {self.model_name} batch failed: {e}")
if retry == self.max_retries - 1:
if self.raise_exception:
raise
else:
await asyncio.sleep(retry + 1)
return [r for r in results if r is not None]
def get_embedding_sync(self, input_text: str, **kwargs) -> list[float]:
"""Synchronous get embedding for a single text with retry logic."""
truncated_text = self._truncate_text(input_text)
# Check cache first
cached_embedding = self._get_from_cache(truncated_text)
if cached_embedding is not None:
return cached_embedding
# Cache miss - compute embedding
for i in range(self.max_retries):
try:
result = self._get_embeddings_sync([truncated_text], **kwargs)
embedding = self._validate_and_adjust_embedding(result[0])
# Store in cache
self._put_to_cache(truncated_text, embedding)
return embedding
except Exception as exc:
logger.error(f"Model {self.model_name} failed: {exc}")
if i == self.max_retries - 1:
if self.raise_exception:
raise
return []
time.sleep(i + 1)
return []
def get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Synchronous get embeddings with automatic batching and retry logic."""
# Truncate all input texts first
truncated_texts = self._truncate_texts(input_text)
# Check cache for each text and separate cached vs uncached
results: list[list[float] | None] = [None] * len(truncated_texts)
texts_to_compute: list[tuple[int, str]] = [] # (original_index, text)
for idx, text in enumerate(truncated_texts):
cached = self._get_from_cache(text)
if cached is not None:
results[idx] = cached
else:
texts_to_compute.append((idx, text))
# If all texts were cached, return early
if not texts_to_compute:
return [r for r in results if r is not None]
# Compute embeddings for uncached texts in batches
uncached_texts = [text for _, text in texts_to_compute]
for i in range(0, len(uncached_texts), self.max_batch_size):
batch_texts = uncached_texts[i : i + self.max_batch_size]
batch_indices = [idx for idx, _ in texts_to_compute[i : i + self.max_batch_size]]
# Process each batch with retry logic
for retry in range(self.max_retries):
try:
batch_embeddings = self._get_embeddings_sync(batch_texts, **kwargs)
if batch_embeddings:
# Store results and cache them
for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings):
adjusted_embedding = self._validate_and_adjust_embedding(embedding)
results[orig_idx] = adjusted_embedding
self._put_to_cache(text, adjusted_embedding)
break
except Exception as exc:
logger.error(f"Model {self.model_name} batch failed: {exc}")
if retry == self.max_retries - 1:
if self.raise_exception:
raise
else:
time.sleep(retry + 1)
return [r for r in results if r is not None]
async def get_node_embedding(self, node: VectorNode, **kwargs) -> VectorNode:
"""Async generate and populate vector field for a single VectorNode object."""
node.vector = await self.get_embedding(node.content, **kwargs)
return node
async def get_node_embeddings(self, nodes: list[VectorNode], **kwargs) -> list[VectorNode]:
"""Async generate and populate vector fields for a batch of VectorNode objects."""
contents = [node.content for node in nodes]
embeddings: list[list[float]] = await self.get_embeddings(contents, **kwargs)
if len(embeddings) == len(nodes):
for node, vec in zip(nodes, embeddings):
node.vector = vec
else:
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(nodes)} nodes")
return nodes
def get_node_embedding_sync(self, node: VectorNode, **kwargs) -> VectorNode:
"""Synchronously generate and populate vector field for a single VectorNode object."""
node.vector = self.get_embedding_sync(node.content, **kwargs)
return node
def get_node_embeddings_sync(self, nodes: list[VectorNode], **kwargs) -> list[VectorNode]:
"""Synchronously generate and populate vector fields for a batch of VectorNode objects."""
contents = [node.content for node in nodes]
embeddings: list[list[float]] = self.get_embeddings_sync(contents, **kwargs)
if len(embeddings) == len(nodes):
for node, vec in zip(nodes, embeddings):
node.vector = vec
else:
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(nodes)} nodes")
return nodes
async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
"""Async generate and populate embedding field for a single MemoryChunk object.
Args:
chunk: MemoryChunk object containing text to embed
**kwargs: Additional arguments passed to the embedding model
Returns:
The same MemoryChunk object with populated embedding field
"""
chunk.embedding = await self.get_embedding(chunk.text, **kwargs)
return chunk
async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
"""Async generate and populate embedding fields for a batch of MemoryChunk objects.
Args:
chunks: List of MemoryChunk objects containing text to embed
**kwargs: Additional arguments passed to the embedding model
Returns:
The same list of MemoryChunk objects with populated embedding fields
"""
texts = [chunk.text for chunk in chunks]
embeddings: list[list[float]] = await self.get_embeddings(texts, **kwargs)
if len(embeddings) == len(chunks):
for chunk, vec in zip(chunks, embeddings):
chunk.embedding = vec
else:
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(chunks)} chunks")
return chunks
def get_chunk_embedding_sync(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
"""Synchronously generate and populate embedding field for a single MemoryChunk object.
Args:
chunk: MemoryChunk object containing text to embed
**kwargs: Additional arguments passed to the embedding model
Returns:
The same MemoryChunk object with populated embedding field
"""
chunk.embedding = self.get_embedding_sync(chunk.text, **kwargs)
return chunk
def get_chunk_embeddings_sync(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
"""Synchronously generate embeddings for a batch of MemoryChunk objects.
Args:
chunks: List of MemoryChunk objects containing text to embed
**kwargs: Additional arguments passed to the embedding model
Returns:
The same list of MemoryChunk objects with populated embedding fields
"""
texts = [chunk.text for chunk in chunks]
embeddings: list[list[float]] = self.get_embeddings_sync(texts, **kwargs)
if len(embeddings) == len(chunks):
for chunk, vec in zip(chunks, embeddings):
chunk.embedding = vec
else:
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(chunks)} chunks")
return chunks
def start_sync(self):
"""Synchronously initialize resources and load cache."""
self._load_cache()
async def start(self):
"""Asynchronously initialize resources and load cache."""
self._load_cache()
def close_sync(self):
"""Synchronously release resources and close connections."""
self._save_cache()
async def close(self):
"""Asynchronously release resources and close connections."""
self._save_cache()

View file

@ -1,62 +0,0 @@
"""Asynchronous OpenAI-compatible embedding model implementation for ReMe."""
from typing import Literal
from openai import AsyncOpenAI
from .base_embedding_model import BaseEmbeddingModel
class OpenAIEmbeddingModel(BaseEmbeddingModel):
"""Asynchronous embedding model implementation compatible with OpenAI-style APIs."""
def __init__(self, encoding_format: Literal["float", "base64"] = "float", **kwargs):
"""Initialize the OpenAI async embedding model with API credentials and configuration."""
super().__init__(**kwargs)
self.encoding_format: Literal["float", "base64"] = encoding_format
# Lazy client initialization
self._client = None
def _create_client(self):
"""Create and return an internal AsyncOpenAI client instance."""
return AsyncOpenAI(api_key=self.api_key, base_url=self.base_url)
@property
def client(self):
"""Lazily create and return the AsyncOpenAI client."""
if self._client is None:
self._client = self._create_client()
return self._client
async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Fetch embeddings from the API for a batch of strings."""
create_kwargs: dict = {
"model": self.model_name,
"input": input_text,
"encoding_format": self.encoding_format,
**self.kwargs,
**kwargs,
}
if self.use_dimensions:
create_kwargs["dimensions"] = self.dimensions
completion = await self.client.embeddings.create(**create_kwargs)
result_emb = [[] for _ in range(len(input_text))]
for emb in completion.data:
# BGE-M3 returns dense_embedding instead of embedding; use as fallback
vec = getattr(emb, "embedding", None) or getattr(emb, "dense_embedding", None)
result_emb[emb.index] = list(vec) if vec is not None else []
return result_emb
async def start(self):
"""Initialize the asynchronous OpenAI embedding model and load cache."""
await super().start()
async def close(self):
"""Close the asynchronous OpenAI client and release network resources."""
if self._client is not None:
await self._client.close()
self._client = None
await super().close()

View file

@ -1,39 +0,0 @@
"""Synchronous OpenAI-compatible embedding model implementation for ReMe."""
from openai import OpenAI
from .openai_embedding_model import OpenAIEmbeddingModel
class OpenAIEmbeddingModelSync(OpenAIEmbeddingModel):
"""Synchronous embedding model implementation that extends the asynchronous OpenAI model."""
def _create_client(self):
"""Create and return an internal synchronous OpenAI client instance."""
return OpenAI(api_key=self.api_key, base_url=self.base_url)
def _get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Fetch embeddings synchronously from the API for a batch of strings."""
create_kwargs: dict = {
"model": self.model_name,
"input": input_text,
"encoding_format": self.encoding_format,
**self.kwargs,
**kwargs,
}
if self.use_dimensions:
create_kwargs["dimensions"] = self.dimensions
completion = self.client.embeddings.create(**create_kwargs)
result_emb = [[] for _ in range(len(input_text))]
for emb in completion.data:
result_emb[emb.index] = emb.embedding
return result_emb
def close_sync(self):
"""Close the synchronous OpenAI client and release network resources."""
if self._client is not None:
self._client.close()
self._client = None
super().close_sync()

View file

@ -1,19 +0,0 @@
"""enumeration"""
from .chunk_enum import ChunkEnum
from .http_enum import HttpEnum
from .json_schema_enum import JsonSchemaEnum
from .memory_source import MemorySource
from .memory_type import MemoryType
from .registry_enum import RegistryEnum
from .role import Role
__all__ = [
"ChunkEnum",
"HttpEnum",
"JsonSchemaEnum",
"MemorySource",
"MemoryType",
"RegistryEnum",
"Role",
]

View file

@ -1,31 +0,0 @@
"""Defines the types of data chunks used in streaming responses."""
from enum import Enum
class ChunkEnum(str, Enum):
"""Enumeration of possible chunk categories for stream processing."""
# Internal reasoning or chain-of-thought process
THINK = "think"
# The final generated response content
ANSWER = "answer"
# Metadata or calls related to external tools
TOOL = "tool"
# Resource consumption and token usage statistics
USAGE = "usage"
# Error messages or exception details
ERROR = "error"
# Signal indicating the start of a new ReAct step
STEP_START = "step_start"
# Tool execution result
TOOL_RESULT = "tool_result"
# Final signal indicating the completion of the stream
DONE = "done"

View file

@ -1,22 +0,0 @@
"""Provides a collection of standard HTTP request methods."""
from enum import Enum
class HttpEnum(str, Enum):
"""Enumeration of supported HTTP methods for network requests."""
# Retrieves data from a specified resource
GET = "get"
# Submits data to be processed to a specified resource
POST = "post"
# Identical to GET but only retrieves the response headers
HEAD = "head"
# Uploads or replaces the representation of a target resource
PUT = "put"
# Deletes the specified resource from the server
DELETE = "delete"

View file

@ -1,38 +0,0 @@
"""Defines the standard data types supported by JSON Schema.
This enum maps common JSON Schema primitive types to their corresponding
Python runtime types, and provides a convenient string representation
compatible with JSON Schema (`"string"`, `"number"`, etc.).
"""
from enum import Enum
class JsonSchemaEnum(Enum):
"""Enumeration of valid JSON Schema data types.
The enum value is the corresponding Python type, while the string
representation (`str(...)`) is the canonical JSON Schema type name.
"""
# Textual data
STRING = str
# Numeric values, including integers and floats
NUMBER = float
# Integer-only numeric values
INTEGER = int
# JSON objects (key-value mappings)
OBJECT = dict
# Ordered JSON lists/arrays
ARRAY = list
# Boolean values: true / false
BOOLEAN = bool
def __str__(self) -> str:
"""Return the lowercase JSON Schema type name for this enum member."""
return self.name.lower()

View file

@ -1,11 +0,0 @@
"""Memory source types."""
from enum import Enum
class MemorySource(str, Enum):
"""Source of memory data."""
MEMORY = "memory"
SESSIONS = "sessions"

View file

@ -1,33 +0,0 @@
"""Defines the high-level categories of memory managed by ReMe.
This enumeration is used across the system to tag, route, and store different
kinds of memories (identity, personal context, procedures, tools, etc.).
"""
from enum import Enum
class MemoryType(str, Enum):
"""Enumeration of memory categories used by the memory subsystem.
These types describe *what* a piece of memory is about, which guides
storage, retrieval, and summarization strategies.
"""
# Long‑term, relatively stable attributes about the user (name, roles, etc.)
IDENTITY = "identity"
# User-specific preferences, habits, and evolving personal context
PERSONAL = "personal"
# How‑to knowledge, workflows, and step‑by‑step instructions
PROCEDURAL = "procedural"
# Information learned about tools, APIs, and their usage patterns
TOOL = "tool"
# Condensed representation of larger memory collections
SUMMARY = "summary"
# Raw chronological interaction history, typically before summarization
HISTORY = "history"

View file

@ -1,31 +0,0 @@
"""Defines the registry categories for core components of the system."""
from enum import Enum
class RegistryEnum(str, Enum):
"""Enumeration of component types registered within the application lifecycle."""
# Large Language Model interfaces
LLM = "llm"
# Models used for generating vector embeddings
EMBEDDING_MODEL = "embedding_model"
# Databases or storage systems for vector search
VECTOR_STORE = "vector_store"
# Databases or storage systems for long-term file storage
FILE_STORE = "file_store"
# Atomic operations or functional units
OP = "op"
# Orchestrated sequences of operations or workflows
FLOW = "flow"
# External APIs or shared internal services
SERVICE = "service"
# Utilities for tracking and limiting token consumption
TOKEN_COUNTER = "token_counter"

View file

@ -1,19 +0,0 @@
"""Defines the participant roles in a chat completion sequence."""
from enum import Enum
class Role(str, Enum):
"""Enumeration of standard personas involved in a conversation flow."""
# High-level instructions to guide the model's behavior
SYSTEM = "system"
# Input or queries provided by the human user
USER = "user"
# Responses or messages generated by the AI model
ASSISTANT = "assistant"
# Output or results returned from external tool executions
TOOL = "tool"

View file

@ -1,34 +0,0 @@
"""File store module for persistent memory management.
This module provides storage backends for memory chunks and file metadata,
including SQLite-based, ChromaDB-based, seekdb-based, and pure-Python local
implementations with vector and full-text search.
"""
from .base_file_store import BaseFileStore
from .chroma_file_store import ChromaFileStore
from .local_file_store import LocalFileStore
from .sqlite_file_store import SqliteFileStore
from .zvec_file_store import ZvecFileStore
from ..registry_factory import R
__all__ = [
"BaseFileStore",
"ChromaFileStore",
"LocalFileStore",
"SqliteFileStore",
"ZvecFileStore",
]
R.file_stores.register("sqlite")(SqliteFileStore)
R.file_stores.register("chroma")(ChromaFileStore)
R.file_stores.register("local")(LocalFileStore)
R.file_stores.register("zvec")(ZvecFileStore)
try:
from .seekdb_file_store import SeekdbFileStore
R.file_stores.register("seekdb")(SeekdbFileStore)
__all__.append("SeekdbFileStore")
except ImportError:
pass

View file

@ -1,227 +0,0 @@
"""Base storage interface for file store."""
import re
from abc import ABC, abstractmethod
from pathlib import Path
from ..embedding import BaseEmbeddingModel
from ..enumeration import MemorySource
from ..schema import FileMetadata, MemoryChunk, MemorySearchResult
from ..utils import get_logger
logger = get_logger()
class BaseFileStore(ABC):
"""Abstract base class for file storage backends."""
def __init__(
self,
store_name: str,
db_path: str | Path,
embedding_model: BaseEmbeddingModel | None = None,
vector_enabled: bool = False,
fts_enabled: bool = True,
**kwargs,
):
"""Initialize"""
# Validate store_name to prevent SQL injection
# Only allow alphanumeric characters and underscores
if not re.match(r"^[a-zA-Z0-9_]+$", store_name):
raise ValueError(f"Invalid '{store_name}'. Only alphanumeric characters and underscores are allowed.")
# Ensure at least one search method is enabled
if not vector_enabled and not fts_enabled:
raise ValueError("At least one of vector_enabled or fts_enabled must be True.")
# Ensure embedding_model is provided when vector search is enabled
if vector_enabled and embedding_model is None:
raise ValueError("embedding_model is required when vector_enabled is True.")
self.store_name: str = store_name
self.db_path: Path = Path(db_path)
self.db_path.mkdir(parents=True, exist_ok=True)
self.embedding_model: BaseEmbeddingModel | None = embedding_model
self.vector_enabled: bool = vector_enabled
self.fts_enabled: bool = fts_enabled
self.kwargs: dict = kwargs
@property
def embedding_dim(self) -> int:
"""Get the embedding model's dimensionality."""
if self.embedding_model is None:
return 1024
return self.embedding_model.dimensions
def _get_mock_embedding(self) -> list[float]:
"""Generate a zero vector based on embedding model dimensions."""
return [0.0] * self.embedding_dim
def _disable_vector_search(self, reason: str = "embedding API error") -> None:
"""Disable vector search and log a warning."""
if self.vector_enabled:
logger.warning(
f"[{self.store_name}] Disabling vector search due to {reason}. "
"Falling back to full-text search only.",
)
self.vector_enabled = False
async def get_embedding(self, query: str, **kwargs) -> list[float]:
"""Get embedding for a single query string."""
if not self.vector_enabled:
return self._get_mock_embedding()
try:
return await self.embedding_model.get_embedding(query, **kwargs)
except Exception as e:
self._disable_vector_search(str(e))
return self._get_mock_embedding()
async def get_embeddings(self, queries: list[str], **kwargs) -> list[list[float]]:
"""Get embeddings for a batch of query strings."""
if not self.vector_enabled:
return [self._get_mock_embedding() for _ in queries]
try:
return await self.embedding_model.get_embeddings(queries, **kwargs)
except Exception as e:
self._disable_vector_search(str(e))
return [self._get_mock_embedding() for _ in queries]
async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
"""Generate and populate embedding field for a single MemoryChunk object."""
if not self.vector_enabled:
chunk.embedding = self._get_mock_embedding()
return chunk
try:
return await self.embedding_model.get_chunk_embedding(chunk, **kwargs)
except Exception as e:
self._disable_vector_search(str(e))
chunk.embedding = self._get_mock_embedding()
return chunk
async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
"""Generate and populate embedding fields for a batch of MemoryChunk objects."""
if not self.vector_enabled:
mock_embedding = self._get_mock_embedding()
for chunk in chunks:
chunk.embedding = mock_embedding.copy()
return chunks
try:
return await self.embedding_model.get_chunk_embeddings(chunks, **kwargs)
except Exception as e:
self._disable_vector_search(str(e))
mock_embedding = self._get_mock_embedding()
for chunk in chunks:
chunk.embedding = mock_embedding.copy()
return chunks
@abstractmethod
async def start(self):
"""Initialize the storage backend."""
@abstractmethod
async def upsert_file(self, file_meta: FileMetadata, source: MemorySource, chunks: list[MemoryChunk]):
"""Insert or update a file and its chunks."""
@abstractmethod
async def delete_file(self, path: str, source: MemorySource):
"""Delete a file and all its chunks."""
@abstractmethod
async def delete_file_chunks(self, path: str, chunk_ids: list[str]):
"""Delete chunks for a file."""
@abstractmethod
async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource):
"""Insert or update specific chunks without affecting other chunks."""
@abstractmethod
async def list_files(self, source: MemorySource) -> list[str]:
"""List all indexed file paths for a source."""
@abstractmethod
async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None:
"""Get full file metadata with statistics."""
@abstractmethod
async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None:
"""Update file metadata without affecting chunks.
This is useful for incremental updates where only metadata needs to be updated
(e.g., after adding/removing chunks in delta file watcher).
Args:
file_meta: Updated file metadata (hash, mtime_ms, size, chunk_count)
source: Memory source
"""
@abstractmethod
async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
"""Get all chunks for a file."""
@abstractmethod
async def vector_search(
self,
query: str,
limit: int,
sources: list[MemorySource] | None = None,
) -> list[MemorySearchResult]:
"""Perform vector similarity search.
Args:
query: Query embedding vector
limit: Maximum number of results
sources: Optional list of sources to filter
Returns:
List of search results sorted by similarity
"""
@abstractmethod
async def keyword_search(
self,
query: str,
limit: int,
sources: list[MemorySource] | None = None,
) -> list[MemorySearchResult]:
"""Perform keyword/full-text search.
Args:
query: Search query text
limit: Maximum number of results
sources: Optional list of sources to filter
Returns:
List of search results sorted by relevance
"""
@abstractmethod
async def hybrid_search(
self,
query: str,
limit: int,
sources: list[MemorySource] | None = None,
vector_weight: float = 0.7,
candidate_multiplier: float = 3.0,
) -> list[MemorySearchResult]:
"""Perform hybrid search combining vector and keyword search.
Args:
query: Search query text
limit: Maximum number of results
sources: Optional list of sources to filter
vector_weight: Weight for vector search results (0.0-1.0).
Keyword weight = 1.0 - vector_weight.
candidate_multiplier: Multiplier for candidate pool size.
candidates = limit * candidate_multiplier
Returns:
List of search results sorted by combined relevance score
"""
@abstractmethod
async def clear_all(self):
"""Clear all indexed data."""
@abstractmethod
async def close(self):
"""Close storage and release resources."""

View file

@ -1,633 +0,0 @@
"""ChromaDB storage backend for file store."""
import json
import random
import time
from pathlib import Path
from .base_file_store import BaseFileStore
from ..enumeration import MemorySource
from ..schema import FileMetadata, MemoryChunk, MemorySearchResult
from ..utils import get_logger
logger = get_logger()
try:
import chromadb
from chromadb.config import Settings
_CHROMADB_IMPORT_ERROR: Exception | None = None
except Exception as e:
_CHROMADB_IMPORT_ERROR = e
chromadb = None
Settings = None
class ChromaFileStore(BaseFileStore):
"""ChromaDB file storage with vector and full-text search.
Inherits embedding methods from BaseFileStore:
- get_chunk_embedding / get_chunk_embeddings (async)
- get_embedding / get_embeddings (async)
Provides ChromaDB-backed persistent storage with:
- Vector similarity search (native ChromaDB)
- Full-text search (via ChromaDB where_document filter)
- Efficient chunk and file metadata management
"""
def __init__(
self,
**kwargs,
):
if _CHROMADB_IMPORT_ERROR is not None:
raise _CHROMADB_IMPORT_ERROR
super().__init__(**kwargs)
self.client: "chromadb.ClientAPI | None" = None
self.chunks_collection: "chromadb.Collection | None" = None
# Initialize metadata file path (db_path and store_name are set by base class)
self._metadata_file: Path = self.db_path.parent / f"{self.store_name}_file_metadata.json"
self._metadata_cache: dict[str, dict[str, FileMetadata]] = {}
@property
def collection_name(self) -> str:
"""Get the name of the ChromaDB collection for this store."""
return f"chunks_{self.store_name}"
async def _load_metadata(self) -> dict[str, dict[str, FileMetadata]]:
"""Load file metadata from disk.
Returns:
Dictionary mapping source -> path -> FileMetadata
"""
if not self._metadata_file.exists():
return {}
try:
data = self._metadata_file.read_text(encoding="utf-8")
metadata_dict = json.loads(data)
# Convert dict to FileMetadata objects
result = {}
for source, files in metadata_dict.items():
result[source] = {}
for path, meta in files.items():
result[source][path] = FileMetadata(**meta)
logger.debug(f"Loaded file metadata from {self._metadata_file}")
return result
except Exception as e:
logger.warning(f"Failed to load file metadata from {self._metadata_file}: {e}")
return {}
async def _save_metadata(self, metadata: dict[str, dict[str, FileMetadata]]) -> None:
"""Save file metadata to disk.
Args:
metadata: Dictionary mapping source -> path -> FileMetadata
"""
try:
# Convert FileMetadata objects to dict for JSON serialization
metadata_dict = {}
for source, files in metadata.items():
metadata_dict[source] = {}
for path, meta in files.items():
metadata_dict[source][path] = {
"path": meta.path,
"hash": meta.hash,
"mtime_ms": meta.mtime_ms,
"size": meta.size,
"chunk_count": meta.chunk_count,
}
data = json.dumps(metadata_dict, indent=2, ensure_ascii=False)
self._metadata_file.write_text(data, encoding="utf-8")
logger.debug(f"Saved file metadata to {self._metadata_file}")
except Exception as e:
logger.error(f"Failed to save file metadata to {self._metadata_file}: {e}")
async def start(self) -> None:
"""Initialize ChromaDB client and collection."""
if self.client is not None:
return
# Initialize persistent ChromaDB client
self.client = chromadb.PersistentClient(
path=str(self.db_path),
settings=Settings(
anonymized_telemetry=False,
allow_reset=True,
),
)
# Get or create the chunks collection
# ChromaDB uses cosine distance by default for similarity
self.chunks_collection = self.client.get_or_create_collection(
name=self.collection_name,
metadata={"hnsw:space": "cosine"},
)
# Load metadata into cache
self._metadata_cache = await self._load_metadata()
logger.info(f"ChromaDB initialized with collection: {self.collection_name}")
logger.info(f"File metadata will be persisted to: {self._metadata_file}")
async def upsert_file(
self,
file_meta: FileMetadata,
source: MemorySource,
chunks: list[MemoryChunk],
) -> None:
"""Insert or update file and its chunks."""
if not chunks:
return
# Delete existing chunks for this file first
await self.delete_file(file_meta.path, source)
# Batch generate embeddings for all chunks
# (base class returns mock embeddings when vector_enabled=False)
chunks = await self.get_chunk_embeddings(chunks)
# Prepare data for ChromaDB batch upsert
ids = []
documents = []
embeddings = []
metadatas = []
now = int(time.time() * 1000)
for chunk in chunks:
ids.append(chunk.id)
documents.append(chunk.text)
embeddings.append(chunk.embedding)
metadatas.append(
{
"path": file_meta.path,
"source": source.value,
"start_line": chunk.start_line,
"end_line": chunk.end_line,
"hash": chunk.hash,
"updated_at": now,
},
)
# Batch upsert to ChromaDB (always pass embeddings to prevent default embedding function)
self.chunks_collection.upsert(
ids=ids,
documents=documents,
embeddings=embeddings,
metadatas=metadatas,
)
# Update file metadata in cache
if source.value not in self._metadata_cache:
self._metadata_cache[source.value] = {}
self._metadata_cache[source.value][file_meta.path] = FileMetadata(
hash=file_meta.hash,
mtime_ms=file_meta.mtime_ms,
size=file_meta.size,
path=file_meta.path,
chunk_count=len(chunks),
)
async def delete_file(self, path: str, source: MemorySource) -> None:
"""Delete file and all its chunks."""
# Query for all chunks with this path and source
results = self.chunks_collection.get(
where={"$and": [{"path": path}, {"source": source.value}]},
include=[],
)
if results["ids"]:
self.chunks_collection.delete(
ids=results["ids"],
)
# Remove from file metadata cache
if source.value in self._metadata_cache:
self._metadata_cache[source.value].pop(path, None)
async def delete_file_chunks(self, path: str, chunk_ids: list[str]) -> None:
"""Delete specific chunks for a file."""
if not chunk_ids:
return
self.chunks_collection.delete(
ids=chunk_ids,
)
# Update chunk count in file metadata cache
for source_meta in self._metadata_cache.values():
if path in source_meta:
# Recalculate chunk count
results = self.chunks_collection.get(
where={"path": path},
include=[],
)
source_meta[path].chunk_count = len(results["ids"])
break
async def upsert_chunks(
self,
chunks: list[MemoryChunk],
source: MemorySource,
) -> None:
"""Insert or update specific chunks without affecting other chunks."""
if not chunks:
return
# Batch generate embeddings for all chunks
# (base class returns mock embeddings when vector_enabled=False)
chunks = await self.get_chunk_embeddings(chunks)
ids = []
documents = []
embeddings = []
metadatas = []
now = int(time.time() * 1000)
for chunk in chunks:
ids.append(chunk.id)
documents.append(chunk.text)
embeddings.append(chunk.embedding)
metadatas.append(
{
"path": chunk.path,
"source": source.value,
"start_line": chunk.start_line,
"end_line": chunk.end_line,
"hash": chunk.hash,
"updated_at": now,
},
)
# Always pass embeddings to prevent default embedding function
self.chunks_collection.upsert(
ids=ids,
documents=documents,
embeddings=embeddings,
metadatas=metadatas,
)
async def list_files(self, source: MemorySource) -> list[str]:
"""List all indexed files for a source."""
if source.value not in self._metadata_cache:
return []
return list(self._metadata_cache[source.value].keys())
async def get_file_metadata(
self,
path: str,
source: MemorySource,
) -> FileMetadata | None:
"""Get file metadata with chunk count."""
if source.value not in self._metadata_cache:
return None
return self._metadata_cache[source.value].get(path)
async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None:
"""Update file metadata without affecting chunks."""
if source.value not in self._metadata_cache:
self._metadata_cache[source.value] = {}
self._metadata_cache[source.value][file_meta.path] = FileMetadata(
hash=file_meta.hash,
mtime_ms=file_meta.mtime_ms,
size=file_meta.size,
path=file_meta.path,
chunk_count=file_meta.chunk_count,
)
async def get_file_chunks(
self,
path: str,
source: MemorySource,
) -> list[MemoryChunk]:
"""Get all chunks for a file."""
results = self.chunks_collection.get(
where={"$and": [{"path": path}, {"source": source.value}]},
include=["documents", "embeddings", "metadatas"],
)
chunks = []
for i, chunk_id in enumerate(results["ids"]):
metadata = results["metadatas"][i]
chunks.append(
MemoryChunk(
id=chunk_id,
path=metadata["path"],
source=MemorySource(metadata["source"]),
start_line=metadata["start_line"],
end_line=metadata["end_line"],
text=results["documents"][i],
hash=metadata["hash"],
embedding=results["embeddings"][i] if results["embeddings"] is not None else None,
),
)
# Sort by start_line
chunks.sort(key=lambda c: c.start_line)
return chunks
async def vector_search(
self,
query: str,
limit: int,
sources: list[MemorySource] | None = None,
) -> list[MemorySearchResult]:
"""Perform vector similarity search."""
if not self.vector_enabled or not query:
return []
# Get query embedding
query_embedding = await self.get_embedding(query)
if not query_embedding:
return []
# Build where filter for sources
where_filter = None
if sources:
if len(sources) == 1:
where_filter = {"source": sources[0].value}
else:
where_filter = {"source": {"$in": [s.value for s in sources]}}
# Perform vector search
try:
results = self.chunks_collection.query(
query_embeddings=[query_embedding],
n_results=limit,
where=where_filter,
include=["documents", "metadatas", "distances"],
)
except Exception as e:
logger.error(f"Vector search failed: {e}, falling back to random results")
# Fallback: get some documents without vector search and assign random scores
try:
fallback_results = self.chunks_collection.get(
where=where_filter,
limit=limit,
include=["documents", "metadatas"],
)
search_results = []
if fallback_results["ids"]:
for i, _ in enumerate(fallback_results["ids"]):
metadata = fallback_results["metadatas"][i]
search_results.append(
MemorySearchResult(
path=metadata["path"],
start_line=metadata["start_line"],
end_line=metadata["end_line"],
score=random.uniform(0.3, 0.7), # Random score in middle range
snippet=fallback_results["documents"][i],
source=MemorySource(metadata["source"]),
raw_metric=None,
),
)
return search_results
except Exception as fallback_e:
logger.error(f"Fallback search also failed: {fallback_e}")
return []
search_results = []
if results["ids"] and results["ids"][0]:
for i, _ in enumerate(results["ids"][0]):
metadata = results["metadatas"][0][i]
distance = results["distances"][0][i]
# Convert cosine distance to similarity score
# Cosine distance range is [0, 2], convert to [1, 0] score
score = max(0.0, 1.0 - distance / 2.0)
search_results.append(
MemorySearchResult(
path=metadata["path"],
start_line=metadata["start_line"],
end_line=metadata["end_line"],
score=score,
snippet=results["documents"][0][i],
source=MemorySource(metadata["source"]),
raw_metric=distance,
),
)
# Sort by score descending
search_results.sort(key=lambda r: r.score, reverse=True)
return search_results
async def keyword_search(
self,
query: str,
limit: int,
sources: list[MemorySource] | None = None,
) -> list[MemorySearchResult]:
"""Perform keyword/full-text search.
ChromaDB supports where_document filter for text matching.
Note: ChromaDB's $contains is case-sensitive, so we generate multiple
case variants (original, lowercase, capitalized) for each word to
improve recall while maintaining case-insensitive scoring.
"""
if not self.fts_enabled or not query:
return []
# Normalize whitespace and split into words
words = query.split()
if not words:
return []
# Generate case variants for each word to handle case-sensitive $contains
# Include: original, lowercase, and capitalized forms
word_variants = set()
for word in words:
word_variants.add(word) # original
word_variants.add(word.lower()) # lowercase
word_variants.add(word.capitalize()) # Capitalized
word_variants.add(word.upper()) # UPPERCASE
word_variants_list = list(word_variants)
# Build where filter for sources
where_filter = None
if sources:
if len(sources) == 1:
where_filter = {"source": sources[0].value}
else:
where_filter = {"source": {"$in": [s.value for s in sources]}}
# ChromaDB where_document uses $contains for substring matching (case-sensitive)
# Use multiple case variants to improve recall
if len(word_variants_list) == 1:
where_document: dict = {"$contains": word_variants_list[0]}
else:
where_document = {"$or": [{"$contains": w} for w in word_variants_list]}
# Get all matching documents
results = self.chunks_collection.get(
where=where_filter,
where_document=where_document,
include=["documents", "metadatas"],
)
search_results = []
query_lower = query.lower()
words_lower = [w.lower() for w in words] # lowercase words for scoring
n_words = len(words)
for i, _ in enumerate(results["ids"]):
metadata = results["metadatas"][i]
text = results["documents"][i]
text_lower = text.lower()
# Calculate relevance score based on word matches
match_count = sum(1 for w in words_lower if w in text_lower)
base_score = match_count / n_words
# Bonus for full phrase match (only applies to multi-word queries)
phrase_bonus = 0.2 if n_words > 1 and query_lower in text_lower else 0.0
# Scale base_score and add phrase bonus, max score is 1.0
score = min(1.0, base_score + phrase_bonus)
search_results.append(
MemorySearchResult(
path=metadata["path"],
start_line=metadata["start_line"],
end_line=metadata["end_line"],
score=score,
snippet=text,
source=MemorySource(metadata["source"]),
),
)
# Sort by score descending and limit results
search_results.sort(key=lambda r: r.score, reverse=True)
return search_results[:limit]
async def hybrid_search(
self,
query: str,
limit: int,
sources: list[MemorySource] | None = None,
vector_weight: float = 0.7,
candidate_multiplier: float = 3.0,
) -> list[MemorySearchResult]:
"""Perform hybrid search combining vector and keyword search.
Args:
query: Search query text
limit: Maximum number of results
sources: Optional list of sources to filter
vector_weight: Weight for vector search results (0.0-1.0).
Keyword weight = 1.0 - vector_weight.
candidate_multiplier: Multiplier for candidate pool size.
Returns:
List of search results sorted by combined relevance score
"""
assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be between 0 and 1, got {vector_weight}"
candidates = min(200, max(1, int(limit * candidate_multiplier)))
text_weight = 1.0 - vector_weight
# Perform search based on enabled backends
if self.vector_enabled and self.fts_enabled:
keyword_results = await self.keyword_search(query, candidates, sources)
vector_results = await self.vector_search(query, candidates, sources)
# Log original vector results
logger.info("\n=== Vector Search Results ===")
for i, r in enumerate(vector_results[:10], 1):
snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet
logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}")
# Log original keyword results
logger.info("\n=== Keyword Search Results ===")
for i, r in enumerate(keyword_results[:10], 1):
snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet
logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}")
if not keyword_results:
return vector_results[:limit]
elif not vector_results:
return keyword_results[:limit]
else:
merged = self._merge_hybrid_results(
vector=vector_results,
keyword=keyword_results,
vector_weight=vector_weight,
text_weight=text_weight,
)
# Log merged results
logger.info("\n=== Merged Hybrid Results ===")
for i, r in enumerate(merged[:10], 1):
snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet
logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}")
return merged[:limit]
elif self.vector_enabled:
vector_results = await self.vector_search(query, limit, sources)
return vector_results
elif self.fts_enabled:
keyword_results = await self.keyword_search(query, limit, sources)
return keyword_results
else:
return []
@staticmethod
def _merge_hybrid_results(
vector: list[MemorySearchResult],
keyword: list[MemorySearchResult],
vector_weight: float,
text_weight: float,
) -> list[MemorySearchResult]:
"""Merge vector and keyword search results with weighted scoring."""
merged: dict[str, MemorySearchResult] = {}
# Process vector results
for result in vector:
result.score = result.score * vector_weight
merged[result.merge_key] = result
# Process keyword results
for result in keyword:
key = result.merge_key
if key in merged:
merged[key].score += result.score * text_weight
else:
result.score = result.score * text_weight
merged[key] = result
# Sort by score and return
results = list(merged.values())
results.sort(key=lambda r: r.score, reverse=True)
return results
async def clear_all(self) -> None:
"""Clear all indexed data."""
# Delete and recreate the collection
self.client.delete_collection(
name=self.collection_name,
)
self.chunks_collection = self.client.get_or_create_collection(
name=self.collection_name,
metadata={"hnsw:space": "cosine"},
)
# Clear file metadata cache and disk
self._metadata_cache = {}
await self._save_metadata({})
logger.info(f"Cleared all data from ChromaDB collection: {self.collection_name}")
async def close(self) -> None:
"""Close ChromaDB client and release resources."""
# Persist metadata cache to disk before closing
if self._metadata_cache:
await self._save_metadata(self._metadata_cache)
# ChromaDB PersistentClient handles persistence automatically
self.client = None
self.chunks_collection = None
await super().close()

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