refactor(core): restructure application initialization and component management

This commit is contained in:
jinli.yl 2026-02-24 19:45:47 +08:00
parent 3408e8dcf1
commit e9757bccf6
26 changed files with 406 additions and 696 deletions

View file

@ -20,8 +20,7 @@ __all__ = [
"ReMeFs",
]
__version__ = "0.3.0.0b3"
__version__ = "0.3.0.0b4"
"""
conda create -n fl_test2 python=3.10

View file

@ -11,6 +11,10 @@ from ...core.utils import format_messages
class FsCompactor(BaseOp):
"""Generate summaries for conversation history compaction."""
def __init__(self, return_prompt: bool = False, **kwargs):
super().__init__(**kwargs)
self.return_prompt = return_prompt
@staticmethod
def _normalize_messages(messages: list[Message | dict]) -> list[Message]:
"""Convert dict messages to Message objects."""
@ -27,7 +31,7 @@ class FsCompactor(BaseOp):
return format_messages(
messages=messages,
add_index=False,
add_time=False,
add_time=True,
use_name=True,
add_reasoning=False,
add_tools=True,
@ -58,33 +62,57 @@ class FsCompactor(BaseOp):
system_prompt = self.get_prompt("system_prompt")
conversation_text = self._serialize_conversation(turn_prefix_messages)
turn_prefix_prompt = self.prompt_format("turn_prefix_summarization", conversation_text=conversation_text)
turn_prefix_prompt = self.prompt_format(
"turn_prefix_summarization",
conversation_text=conversation_text,
)
return [
Message(role=Role.SYSTEM, content=system_prompt),
Message(role=Role.USER, content=turn_prefix_prompt),
]
async def execute(self) -> str:
async def execute(self) -> str | dict:
"""Generate summary for conversation history."""
messages_to_summarize = self.context.get("messages_to_summarize", [])
turn_prefix_messages = self.context.get("turn_prefix_messages", [])
previous_summary = self.context.get("previous_summary", "")
messages_to_summarize = self._normalize_messages(messages_to_summarize)
if messages_to_summarize:
history_prompt_messages = self._build_history_prompt(messages_to_summarize, previous_summary)
history_summary = "**History Summary**:\n\n" + await self._generate_summary(history_prompt_messages)
else:
history_summary = ""
turn_prefix_messages = self._normalize_messages(turn_prefix_messages)
if turn_prefix_messages:
turn_prefix_prompt_messages = self._build_turn_prefix_prompt(turn_prefix_messages)
turn_prefix_summary = "**Turn Context**:\n\n" + await self._generate_summary(turn_prefix_prompt_messages)
else:
turn_prefix_summary = ""
summary = "\n\n---".join([history_summary, turn_prefix_summary])
logger.info(f"Generated summary: {summary}")
return summary
if self.return_prompt:
result = {
"system": self.get_prompt("system_prompt"),
}
if messages_to_summarize:
history_prompt_messages = self._build_history_prompt(messages_to_summarize, previous_summary)
if len(history_prompt_messages) == 2:
result["history_user"] = history_prompt_messages[-1].content
if turn_prefix_messages:
turn_prefix_prompt_messages = self._build_turn_prefix_prompt(turn_prefix_messages)
if len(turn_prefix_prompt_messages) == 2:
result["turn_prefix_user"] = turn_prefix_prompt_messages[-1].content
return result
else:
if messages_to_summarize:
history_prompt_messages = self._build_history_prompt(messages_to_summarize, previous_summary)
history_summary = "**History Summary**:\n\n" + await self._generate_summary(history_prompt_messages)
else:
history_summary = ""
if turn_prefix_messages:
turn_prefix_prompt_messages = self._build_turn_prefix_prompt(turn_prefix_messages)
turn_prefix_summary = "**Turn Context**:\n\n" + await self._generate_summary(
turn_prefix_prompt_messages,
)
else:
turn_prefix_summary = ""
summary = "\n\n---".join([history_summary, turn_prefix_summary])
logger.info(f"Generated summary: {summary}")
return summary

View file

@ -13,30 +13,25 @@ from ...core.utils import format_messages
class FsSummarizer(BaseReact):
"""Retrieve personal memories through vector search and history reading."""
def __init__(self, working_dir: str, memory_dir: str = "memory", version: str = "v1", **kwargs):
def __init__(
self,
working_dir: str,
memory_dir: str = "memory",
version: str = "default",
return_prompt: bool = False,
**kwargs,
):
super().__init__(**kwargs)
self.working_dir: str = working_dir
self.memory_dir: str = memory_dir
self.version: str = version
self.return_prompt = return_prompt
async def build_messages(self) -> list[Message]:
messages: list[Message] = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages]
date_str: str = self.context.get("date", datetime.datetime.now().strftime("%Y-%m-%d"))
if self.version == "default":
messages.append(
Message(
role=Role.USER,
content=self.prompt_format(
"user_message_default",
working_dir=self.working_dir,
date=date_str,
memory_dir=self.memory_dir,
),
),
)
elif self.version == "v1":
conversation = format_messages(messages, add_index=False)
messages = [
Message(
@ -52,14 +47,27 @@ class FsSummarizer(BaseReact):
),
]
elif self.version == "v1":
messages.append(
Message(
role=Role.USER,
content=self.prompt_format(
"user_message_default",
working_dir=self.working_dir,
date=date_str,
memory_dir=self.memory_dir,
),
),
)
else:
messages.extend(
[
Message(role=Role.SYSTEM, content=self.get_prompt("system_prompt")),
Message(role=Role.SYSTEM, content=self.get_prompt("system_prompt_deprecated")),
Message(
role=Role.USER,
content=self.prompt_format(
"user_message",
"user_message_deprecated",
date=date_str,
memory_dir=self.memory_dir,
),
@ -69,7 +77,13 @@ class FsSummarizer(BaseReact):
return messages
async def execute(self):
result = await super().execute()
answer = result["answer"]
logger.info(f"[{self.__class__.__name__}] answer={answer}")
if self.return_prompt:
result = {}
messages: list[Message] = await self.build_messages()
result["prompt"] = messages[-1].content
return result
else:
result = await super().execute()
answer = str(result["answer"])
logger.info(f"[{self.__class__.__name__}] answer={answer}")
return result

View file

@ -1,39 +1,38 @@
system_prompt: |
system_prompt_deprecated: |
Pre-compaction memory flush turn.
The session is near auto-compaction; capture durable memories to disk.
You may reply, but usually [SILENT] is correct.
user_message: |
user_message_deprecated: |
Pre-compaction memory flush.
Current Date: {date}
Store durable memories now (use {memory_dir}/YYYY-MM-DD.md; create {memory_dir}/ if needed).
If nothing to store, reply with [SILENT].
user_message_default: |
Pre-compaction memory flush turn.
The session is near auto-compaction; capture durable memories to disk.
Memory Pre-compression Flush Cycle Initiated
The current session is about to enter the automatic compression phase. Please capture persistent memory and write it to disk.
Current Date: {date}
Working Dir: {working_dir}
Current date: {date}
Working directory: {working_dir}
Store durable memories now (use {memory_dir}/YYYY-MM-DD.md).
Immediately store persistent memory to: {memory_dir}/YYYY-MM-DD.md
Workflow:
1. Use read_tool to read {memory_dir}/YYYY-MM-DD.md (if file doesn't exist, read_tool tool will return an error)
2. Intelligently merge new information with existing content (skip if file doesn't exist):
- Avoid duplicating information that's already recorded
- Enrich existing entries with new details when relevant
- Maintain chronological order when applicable
1. First, `read` {memory_dir}/YYYY-MM-DD.md (if the file doesnt exist, an error message will be returned).
2. Intelligently merge new information with existing content (skip merging if the file doesnt exist):
- Avoid duplicating already recorded information
- Enrich existing entries with new details where relevant
- Maintain chronological order wherever applicable
3. Write the updated content:
- Use edit_tool to update specific sections when possible
- Use write_tool to overwrite the entire file if major restructuring is needed
4. Create {memory_dir}/ if it doesn't exist
- Prefer using `edit` to update specific sections when possible
- Use `write` to overwrite the entire file only if substantial restructuring is required
Principles:
- Always preserve timestamps, dates, and time-related context
- Only add truly new or enriching information
- Keep entries concise but complete
- If nothing meaningful to store, reply with [SILENT]
- Always preserve timestamps and any date/time-related context
- Add only genuinely new or meaningfully enriching information
- Keep entries concise yet complete
- If theres nothing to store, respond with [SILENT]
user_message_default_zh: |
@ -46,83 +45,17 @@ user_message_default_zh: |
立即存储持久化记忆(使用路径 {memory_dir}/YYYY-MM-DD.md
工作流程:
1. 使用 read_tool 读取 {memory_dir}/YYYY-MM-DD.md如文件不存在read_tool 会返回错误提示)
1. 先 `read` {memory_dir}/YYYY-MM-DD.md如文件不存在会返回错误提示
2. 智能合并新信息与现有内容(若文件不存在则跳过合并):
- 避免重复已记录的信息
- 在相关时丰富现有条目的新细节
- 在适用时保持时间顺序
3. 写入更新后的内容:
- 尽可能使用 edit_tool 更新特定部分
- 如需大幅重构则使用 write_tool 覆盖整个文件
4. 如 {memory_dir}/ 不存在则创建
- 尽可能使用 `edit` 更新特定部分
- 如需大幅重构则使用 `write` 覆盖整个文件
原则:
- 始终保留时间戳、日期和时间相关上下文
- 仅添加真正新的或有丰富价值的信息
- 保持条目简洁但完整
- 若无有意义的内容可存储,请回复 [SILENT]
user_message_v1: |
<conversation>
{conversation}
</conversation>
The conversation is about to be compacted. Please extract persistent memories to disk.
Current Date: {date}
Working Dir: {working_dir}
Execution Flow:
1. Determine if the conversation contains information worth storing
- If no: Reply with reason + [SILENT]
- If yes: Continue to step 2
2. Use Read tool to read {memory_dir}/YYYY-MM-DD.md (use actual date)
- If file doesn't exist (Read returns error): Write new memories directly
- If file exists:
a) Compare and identify new/updated information from the read content
c) Prefer `edit_tool` for precise additions (preserves existing content); `write_tool` overwrites entire file
d) If no new information: Reply with explanation + [SILENT]
Update Principles:
- Only add unrecorded information, preserve all existing content
- Intelligently merge duplicate information, retain key details like timestamps
Example (information merging):
Existing: "Alice joined Company A as Software Engineer on 2023-01-01"
New: "Alice joined Company B as Senior Engineer on 2024-01-01"
Result: "Alice joined Company A as Software Engineer on 2023-01-01, moved to Company B as Senior Engineer on 2024-01-01"
Please store persistent memories, keeping entries concise and well-structured.
user_message_v1_zh: |
<conversation>
{conversation}
</conversation>
conversation即将压缩请提取持久性记忆存储至磁盘。
当前日期:{date}
工作目录: {working_dir}
执行流程:
1. 判断对话是否包含值得存储的信息
- 若无:回复原因 + [SILENT]
- 若有:继续步骤 2
2. 使用 Read 工具读取 {memory_dir}/YYYY-MM-DD.md使用实际日期
- 文件不存在Read 返回错误):直接使用 `write_tool` 写入新记忆
- 文件已存在:
a) 从读取的内容中对比识别新增/更新信息
c) 优先使用 `edit_tool` 精准添加新信息(保留已有内容)
d) 若无新信息:回复说明 + [SILENT]
更新原则:
- 仅添加未记录的信息,保留所有已有内容
- 智能合并重复信息,保留时间等关键细节
示例(信息合并):
已有:"Alice 于 2023-01-01 加入 A 公司任软件工程师"
新增:"Alice 于 2024-01-01 加入 B 公司任高级工程师"
结果:"Alice 于 2023-01-01 加入 A 公司任软件工程师2024-01-01 转入 B 公司任高级工程师"
请存储持久性记忆,保持条目简洁、结构清晰。
- 若无内容可存储,请回复 [SILENT]

View file

@ -1,12 +1,12 @@
"""memory agent"""
from .base_memory_agent import BaseMemoryAgent
from .personal.personal_halumem_retriever import PersonalHalumemRetriever
from .personal.personal_halumem_summarizer import PersonalHalumemSummarizer
from .personal.personal_retriever import PersonalRetriever
from .personal.personal_summarizer import PersonalSummarizer
from .personal.personal_v1_retriever import PersonalV1Retriever
from .personal.personal_v1_summarizer import PersonalV1Summarizer
from .personal.personal_halumem_retriever import PersonalHalumemRetriever
from .personal.personal_halumem_summarizer import PersonalHalumemSummarizer
from .procedural.procedural_retriever import ProceduralRetriever
from .procedural.procedural_summarizer import ProceduralSummarizer
from .reme_retriever import ReMeRetriever

View file

@ -23,7 +23,7 @@ embedding_models:
memory_stores:
default:
backend: chroma
# backend: local
# backend: local
db_name: reme.db
store_name: reme
embedding_model: default
@ -34,8 +34,8 @@ file_watchers:
default:
backend: full
memory_store: default
watch_paths: [".reme", ".reme/memory"]
suffix_filters: [".md"]
watch_paths: [ ".reme", ".reme/memory" ]
suffix_filters: [ ".md" ]
recursive: false
scan_on_start: true

View file

@ -21,9 +21,9 @@ llms:
default:
backend: openai
model_name: qwen3-30b-a3b-instruct-2507
# model_name: qwen3-30b-a3b-thinking-2507
# model_name: qwen3-30b-a3b-thinking-2507
request_interval: 1
# temperature: 0.0001
# temperature: 0.0001
qwen3_max_instruct:
backend: openai
@ -46,7 +46,7 @@ embedding_models:
vector_stores:
default:
backend: chroma
# backend: local
# backend: local
embedding_model: default
collection_name: reme

View file

@ -3,8 +3,8 @@ backend: cmd
llms:
default:
backend: openai
# model_name: qwen3-30b-a3b-instruct-2507
# model_name: qwen3-30b-a3b-thinking-2507
# model_name: qwen3-30b-a3b-instruct-2507
# model_name: qwen3-30b-a3b-thinking-2507
model_name: qwen3-235b-a22b-thinking-2507
request_interval: 1
# temperature: 0.0001
@ -17,9 +17,9 @@ embedding_models:
memory_stores:
default:
# backend: sqlite
# backend: sqlite
backend: chroma
# backend: local
# backend: local
db_name: reme.db
store_name: reme
embedding_model: default
@ -30,8 +30,8 @@ file_watchers:
default:
backend: full
memory_store: default
watch_paths: [".reme", ".reme/memory"]
suffix_filters: [".md"]
watch_paths: [ ".reme", ".reme/memory" ]
suffix_filters: [ ".md" ]
recursive: false
scan_on_start: true

View file

@ -1,8 +1,12 @@
"""High-level entry point for configuring and running ReMe services and flows."""
import asyncio
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from .context import PromptHandler, ServiceContext
from loguru import logger
from .context import PromptHandler, ServiceContext, R
from .embedding import BaseEmbeddingModel
from .file_watcher import BaseFileWatcher
from .flow import BaseFlow
@ -10,7 +14,7 @@ from .llm import BaseLLM
from .memory_store import BaseMemoryStore
from .schema import Response, ServiceConfig
from .token_counter import BaseTokenCounter
from .utils import execute_stream_task, PydanticConfigParser
from .utils import execute_stream_task, PydanticConfigParser, init_logger, print_logo, MCPClient
from .vector_store import BaseVectorStore
@ -57,80 +61,207 @@ class Application:
default_file_watcher_config=default_file_watcher_config,
**kwargs,
)
self.prompt_handler = PromptHandler(language=self.service_context.language)
self.prompt_handler = PromptHandler(language=self.service_config.language)
self._started: bool = False
def update_api_envs(
self,
llm_api_key: str | None = None,
llm_base_url: str | None = None,
embedding_api_key: str | None = None,
embedding_base_url: str | None = None,
):
"""Update the API environment variables."""
self.service_context.update_api_envs(
llm_api_key=llm_api_key,
llm_base_url=llm_base_url,
embedding_api_key=embedding_api_key,
embedding_base_url=embedding_base_url,
)
@classmethod
async def create(
cls,
*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,
enable_logo: bool = True,
parser: type[PydanticConfigParser] | None = None,
llm: dict | None = None,
embedding_model: dict | None = None,
vector_store: dict | None = None,
memory_store: dict | None = None,
token_counter: dict | None = None,
file_watcher: dict | None = None,
**kwargs,
) -> "Application":
async def create(cls, *args, **kwargs) -> "Application":
"""Create and start an Application instance asynchronously."""
instance = cls(
*args,
llm_api_key=llm_api_key,
llm_base_url=llm_base_url,
embedding_api_key=embedding_api_key,
embedding_base_url=embedding_base_url,
enable_logo=enable_logo,
parser=parser,
default_llm_config=llm,
default_embedding_model_config=embedding_model,
default_vector_store_config=vector_store,
default_memory_store_config=memory_store,
default_token_counter_config=token_counter,
default_file_watcher_config=file_watcher,
**kwargs,
)
instance = cls(*args, **kwargs)
await instance.start()
return instance
@property
def service_config(self) -> ServiceConfig:
"""Get the service configuration."""
return self.service_context.service_config
async def start(self):
"""Start the application."""
"""Start the service context by initializing all configured components."""
if self._started:
return self
else:
await self.service_context.start()
self._started = True
logger.warning("Application has already started.")
return self
async def close(self):
"""Close the application."""
if self._started:
await self.service_context.close()
self._started = False
init_logger(log_to_console=self.service_config.log_to_console)
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.enable_logo:
print_logo(service_config=self.service_config)
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_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,
)
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:
raise RuntimeError("Application is not started")
logger.info("No expression flow found, please check your configuration.")
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"})
self.service_context.llms[name] = R.llms[config.backend](**config_dict)
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["cache_dir"] = working_path / "embedding_cache"
self.service_context.embedding_models[name] = R.embedding_models[config.backend](**config_dict)
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],
"thread_pool": self.service_context.thread_pool,
},
)
self.service_context.vector_stores[name] = R.vector_stores[config.backend](**config_dict)
await self.service_context.vector_stores[name].create_collection(config.collection_name)
for name, config in self.service_config.memory_stores.items():
if config.backend not in R.memory_stores:
logger.warning(f"Memory 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],
"thread_pool": self.service_context.thread_pool,
"db_path": working_path / "memory_store",
},
)
self.service_context.memory_stores[name] = R.memory_stores[config.backend](**config_dict)
await self.service_context.memory_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", "memory_store"})
config_dict["memory_store"] = self.service_context.memory_stores[config.memory_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
return self
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
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, vector_store in self.service_context.vector_stores.items():
logger.info(f"Closing vector store: {name}")
await vector_store.close()
for name, memory_store in self.service_context.memory_stores.items():
logger.info(f"Closing memory store: {name}")
await memory_store.close()
for name, file_watcher in self.service_context.file_watchers.items():
logger.info(f"Closing file watcher: {name}")
await file_watcher.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
return False
def shutdown_thread_pool(self, wait: bool = True):
"""Shutdown the thread pool executor."""
if self.service_context.thread_pool:
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()
@ -218,11 +349,6 @@ class Application:
"""Get the default token counter instance."""
return self.service_context.token_counters.get("default")
@property
def service_config(self) -> ServiceConfig:
"""Get the service configuration."""
return self.service_context.service_config
def get_token_counter(self, name: str):
"""Get a token counter instance by name."""
return self.service_context.token_counters.get(name)
@ -232,4 +358,9 @@ class Application:
import warnings
warnings.filterwarnings("ignore", category=DeprecationWarning)
self.service_context.service.run()
service = R.services[self.service_config.backend](service_context=self.service_context)
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,12 +1,4 @@
"""Module for managing and formatting prompt templates from files or dictionaries.
This module provides a PromptHandler class that:
- Loads prompts from YAML/JSON files or dictionaries
- Supports multi-language prompts with automatic suffix handling
- Provides conditional line filtering using boolean flags
- Formats prompts with template variable substitution
- Validates format strings and provides helpful error messages
"""
"""Module for managing and formatting prompt templates from files or dictionaries."""
import json
from pathlib import Path
@ -19,54 +11,10 @@ from loguru import logger
from .base_context import BaseContext
class PromptNotFoundError(KeyError):
"""Exception raised when a requested prompt template is not found."""
def __init__(self, prompt_name: str, available_prompts: list[str]):
self.prompt_name = prompt_name
self.available_prompts = available_prompts
super().__init__(
f"Prompt '{prompt_name}' not found. "
f"Available prompts: {', '.join(available_prompts[:10])}"
f"{'...' if len(available_prompts) > 10 else ''}",
)
class PromptFormattingError(ValueError):
"""Exception raised when prompt formatting fails."""
class PromptHandler(BaseContext):
"""A context-aware handler for loading, retrieving, and formatting prompt templates.
This handler supports:
- Loading prompts from YAML/JSON files or dictionaries
- Multi-language prompt support with automatic language suffix
- Conditional line filtering using boolean flags (e.g., [debug], [verbose])
- Template variable substitution with validation
- Method chaining for fluent API
Examples:
>>> handler = PromptHandler(language="en")
>>> handler.load_prompt_dict({
... "greeting_en": "Hello, {name}!",
... "farewell_en": "[debug]Debug mode\\nGoodbye, {name}!"
... })
>>> handler.prompt_format("greeting", name="Alice")
'Hello, Alice!'
>>> handler.prompt_format("farewell", name="Bob", debug=False)
'Goodbye, Bob!'
"""
"""A context-aware handler for loading, retrieving, and formatting prompt templates."""
def __init__(self, language: str = "", **kwargs):
"""Initialize the PromptHandler with optional language configuration.
Args:
language: Language code to append as suffix (e.g., "en", "zh", "ja").
If provided, get_prompt will automatically try to find
prompts with this suffix (e.g., "greeting" -> "greeting_en").
**kwargs: Additional key-value pairs to initialize the context.
"""
super().__init__(**kwargs)
self.language: str = language.strip()
@ -75,25 +23,7 @@ class PromptHandler(BaseContext):
prompt_file_path: Optional[Union[Path, str]] = None,
overwrite: bool = True,
) -> "PromptHandler":
"""Load prompt configurations from a YAML or JSON file into the context.
Supports both YAML (.yaml, .yml) and JSON (.json) file formats.
Non-existent files are silently skipped.
Args:
prompt_file_path: Path to the prompt configuration file.
If None, returns self without changes.
overwrite: If True, allows overwriting existing prompts with warnings.
If False, skips existing prompts without overwriting.
Returns:
Self for method chaining.
Raises:
ValueError: If file format is not supported.
yaml.YAMLError: If YAML parsing fails.
json.JSONDecodeError: If JSON parsing fails.
"""
"""Load prompt configurations from a YAML or JSON file."""
if prompt_file_path is None:
return self
@ -105,23 +35,15 @@ class PromptHandler(BaseContext):
suffix = prompt_file_path.suffix.lower()
try:
with prompt_file_path.open(encoding="utf-8") as f:
if suffix in [".yaml", ".yml"]:
prompt_dict = yaml.safe_load(f)
elif suffix == ".json":
prompt_dict = json.load(f)
else:
raise ValueError(
f"Unsupported file format: {suffix}. " f"Supported formats: .yaml, .yml, .json",
)
self.load_prompt_dict(prompt_dict, overwrite=overwrite)
except (yaml.YAMLError, json.JSONDecodeError) as e:
logger.error(f"Failed to parse prompt file {prompt_file_path}: {e}")
raise
with prompt_file_path.open(encoding="utf-8") as f:
if suffix in [".yaml", ".yml"]:
prompt_dict = yaml.safe_load(f)
elif suffix == ".json":
prompt_dict = json.load(f)
else:
raise ValueError(f"Unsupported file format: {suffix}")
self.load_prompt_dict(prompt_dict, overwrite=overwrite)
return self
def load_prompt_dict(
@ -129,233 +51,95 @@ class PromptHandler(BaseContext):
prompt_dict: Optional[Dict[str, Any]] = None,
overwrite: bool = True,
) -> "PromptHandler":
"""Merge a dictionary of prompt strings into the current context.
Only string values are stored as prompts. Non-string values are skipped.
Args:
prompt_dict: Dictionary mapping prompt names to prompt template strings.
overwrite: If True, allows overwriting existing prompts with warnings.
If False, skips existing prompts without overwriting.
Returns:
Self for method chaining.
"""
"""Merge a dictionary of prompt strings into the current context."""
if not prompt_dict:
return self
for key, value in prompt_dict.items():
if not isinstance(value, str):
logger.debug(f"Skipping non-string prompt: key={key}, type={type(value)}")
continue
if key in self:
if overwrite:
logger.warning(
f"Overwriting prompt '{key}': " f"old length={len(self[key])}, new length={len(value)}",
)
logger.warning(f"Overwriting prompt '{key}'")
self[key] = value
else:
logger.debug(f"Skipping existing prompt: key={key}")
else:
logger.debug(f"Adding new prompt: key={key}, length={len(value)}")
self[key] = value
return self
def get_prompt(self, prompt_name: str, fallback_to_base: bool = True) -> str:
"""Retrieve a prompt by name with automatic language suffix handling.
If a language is configured, this method will:
1. First try to find the prompt with language suffix (e.g., "greeting_en")
2. If not found and fallback_to_base is True, try the base name (e.g., "greeting")
3. Otherwise, raise PromptNotFoundError
Args:
prompt_name: Name of the prompt to retrieve.
fallback_to_base: If True and language-specific prompt not found,
fallback to prompt without language suffix.
Returns:
The prompt template string, stripped of leading/trailing whitespace.
Raises:
PromptNotFoundError: If the prompt is not found.
"""
# Try with language suffix first
"""Retrieve a prompt by name with automatic language suffix handling."""
if self.language and not prompt_name.endswith(f"_{self.language}"):
key_with_lang = f"{prompt_name}_{self.language}"
if key_with_lang in self:
return self[key_with_lang].strip()
# Try base name
if prompt_name in self:
return self[prompt_name].strip()
# Try fallback if enabled
if fallback_to_base and self.language:
# Check if prompt_name already has language suffix, try without it
if prompt_name.endswith(f"_{self.language}"):
base_name = prompt_name[: -(len(self.language) + 1)]
if base_name in self:
return self[base_name].strip()
if fallback_to_base and self.language and prompt_name.endswith(f"_{self.language}"):
base_name = prompt_name[: -(len(self.language) + 1)]
if base_name in self:
return self[base_name].strip()
# Not found, raise error with helpful message
available = list(self.keys())
raise PromptNotFoundError(prompt_name, available)
raise KeyError(f"Prompt '{prompt_name}' not found. Available: {list(self.keys())[:10]}")
def has_prompt(self, prompt_name: str) -> bool:
"""Check if a prompt exists (with or without language suffix).
Args:
prompt_name: Name of the prompt to check.
Returns:
True if the prompt exists, False otherwise.
"""
"""Check if a prompt exists."""
try:
self.get_prompt(prompt_name)
return True
except PromptNotFoundError:
except KeyError:
return False
def list_prompts(self, language_filter: Optional[str] = None) -> list[str]:
"""List all available prompt names.
Args:
language_filter: If provided, only return prompts for this language.
If None, return all prompts.
Returns:
List of prompt names.
"""
"""List all available prompt names."""
if language_filter is None:
return list(self.keys())
suffix = f"_{language_filter.strip()}"
return [key for key in self.keys() if key.endswith(suffix)]
@staticmethod
def _extract_format_fields(template: str) -> set[str]:
"""Extract all format field names from a template string.
Args:
template: Template string with {variable} placeholders.
Returns:
Set of field names used in the template.
"""
"""Extract all format field names from a template string."""
return {field_name for _, field_name, _, _ in Formatter().parse(template) if field_name is not None}
@staticmethod
def _filter_conditional_lines(prompt: str, flags: Dict[str, bool]) -> str:
"""Filter lines based on boolean flags.
Lines starting with [flag_name] are conditionally included based on
the value of flags[flag_name]. If True, the line is included (without
the flag marker). If False, the line is excluded.
Args:
prompt: The prompt text with conditional markers.
flags: Dictionary of flag names to boolean values.
Returns:
Filtered prompt text.
"""
"""Filter lines based on boolean flags."""
filtered_lines = []
for line in prompt.split("\n"):
# Check each flag
matched_flag = None
for flag_name in flags:
marker = f"[{flag_name}]"
if line.startswith(marker):
if line.startswith(f"[{flag_name}]"):
matched_flag = flag_name
break
if matched_flag is None:
# No flag marker, always include
filtered_lines.append(line)
elif flags[matched_flag]:
# Flag is True, include without marker
marker = f"[{matched_flag}]"
filtered_lines.append(line[len(marker) :])
# else: Flag is False, skip this line
filtered_lines.append(line[len(f"[{matched_flag}]") :])
return "\n".join(filtered_lines)
def prompt_format(
self,
prompt_name: str,
validate: bool = True,
**kwargs,
) -> str:
"""Format a prompt with conditional line filtering and variable substitution.
This method performs two-stage formatting:
1. Conditional line filtering: Lines marked with [flag] are included only
if the corresponding boolean kwarg is True.
2. Variable substitution: Template variables {var} are replaced with
provided values.
Args:
prompt_name: Name of the prompt to format.
validate: If True, check that all required template variables are provided.
**kwargs: Keyword arguments for formatting. Boolean values are treated as
conditional flags, other values are used for template substitution.
Returns:
Formatted prompt string.
Raises:
PromptNotFoundError: If the prompt is not found.
PromptFormattingError: If validation fails or formatting errors occur.
Examples:
>>> handler = PromptHandler()
>>> handler["test"] = "[debug]Debug: {info}\\nResult: {value}"
>>> handler.prompt_format("test", debug=False, info="test", value=42)
'Result: 42'
>>> handler.prompt_format("test", debug=True, info="test", value=42)
'Debug: test\\nResult: 42'
"""
# Get the prompt template
def prompt_format(self, prompt_name: str, validate: bool = True, **kwargs) -> str:
"""Format a prompt with conditional line filtering and variable substitution."""
prompt = self.get_prompt(prompt_name)
# Separate boolean flags from format variables
flag_kwargs = {k: v for k, v in kwargs.items() if isinstance(v, bool)}
format_kwargs = {k: v for k, v in kwargs.items() if not isinstance(v, bool)}
# Step 1: Filter conditional lines
if flag_kwargs:
prompt = self._filter_conditional_lines(prompt, flag_kwargs)
# Step 2: Validate required fields if requested
if validate:
required_fields = self._extract_format_fields(prompt)
missing_fields = required_fields - set(format_kwargs.keys())
if missing_fields:
raise PromptFormattingError(
f"Missing required format variables for prompt '{prompt_name}': "
f"{', '.join(sorted(missing_fields))}",
)
raise ValueError(f"Missing format variables for '{prompt_name}': {sorted(missing_fields)}")
# Step 3: Format with variables
try:
if format_kwargs:
prompt = prompt.format(**format_kwargs)
except KeyError as e:
raise PromptFormattingError(
f"Format error in prompt '{prompt_name}': missing variable {e}",
) from e
except (ValueError, IndexError) as e:
raise PromptFormattingError(
f"Format error in prompt '{prompt_name}': {e}",
) from e
if format_kwargs:
prompt = prompt.format(**format_kwargs)
return prompt.strip()
def __repr__(self) -> str:
"""Return a string representation of the PromptHandler."""
return f"PromptHandler(language='{self.language}', " f"num_prompts={len(self)})"
return f"PromptHandler(language='{self.language}', num_prompts={len(self)})"

View file

@ -70,7 +70,11 @@ class RuntimeContext(BaseContext):
self[target] = self[source]
return self
def validate_required_keys(self, required_keys: dict[str, bool], context_name: str = "context") -> "RuntimeContext":
def validate_required_keys(
self,
required_keys: dict[str, bool],
context_name: str = "context",
) -> "RuntimeContext":
"""Ensure all required keys are present in the context.
Args:

View file

@ -2,15 +2,13 @@
import os
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from typing import TYPE_CHECKING
from loguru import logger
from .base_context import BaseContext
from .registry_factory import R
from ..schema import ServiceConfig
from ..utils import load_env, MCPClient, print_logo, PydanticConfigParser, init_logger
from ..utils import load_env, PydanticConfigParser
if TYPE_CHECKING:
from ..llm import BaseLLM
@ -19,7 +17,6 @@ if TYPE_CHECKING:
from ..memory_store import BaseMemoryStore
from ..token_counter import BaseTokenCounter
from ..flow import BaseFlow
from ..service import BaseService
from ..file_watcher import BaseFileWatcher
@ -49,8 +46,14 @@ class ServiceContext(BaseContext):
):
super().__init__()
# Load environment variables
load_env()
self.update_api_envs(llm_api_key, llm_base_url, embedding_api_key, embedding_base_url)
# Update common environment variables for LLM and embedding services.
self.update_env("REME_LLM_API_KEY", llm_api_key)
self.update_env("REME_LLM_BASE_URL", llm_base_url)
self.update_env("REME_EMBEDDING_API_KEY", embedding_api_key)
self.update_env("REME_EMBEDDING_BASE_URL", embedding_base_url)
if service_config is None:
parser_class = parser if parser is not None else PydanticConfigParser
@ -85,37 +88,16 @@ class ServiceContext(BaseContext):
service_config = parser.parse_args(*input_args, **kwargs)
self.service_config: ServiceConfig = service_config
init_logger(log_to_console=self.service_config.log_to_console)
logger.info(f"ReMe Config: {service_config.model_dump_json()}")
if self.service_config.working_dir:
self.working_path = Path(self.service_config.working_dir)
self.working_path.mkdir(parents=True, exist_ok=True)
if self.service_config.enable_logo:
print_logo(service_config=self.service_config)
self.language: str = self.service_config.language
self.thread_pool: ThreadPoolExecutor = ThreadPoolExecutor(
max_workers=self.service_config.thread_pool_max_workers,
)
if self.service_config.ray_max_workers > 1:
import ray
ray.init(num_cpus=self.service_config.ray_max_workers)
self.thread_pool: ThreadPoolExecutor | None = None
self.llms: dict[str, "BaseLLM"] = {}
self.embedding_models: dict[str, "BaseEmbeddingModel"] = {}
self.token_counters: dict[str, "BaseTokenCounter"] = {}
self.vector_stores: dict[str, "BaseVectorStore"] = {}
self.memory_stores: dict[str, "BaseMemoryStore"] = {}
self.file_watchers: dict[str, "BaseFileWatcher"] = {}
self.flows: dict[str, "BaseFlow"] = {}
self.mcp_server_mapping: dict[str, dict] = {}
self.service: "BaseService" = R.services[self.service_config.backend](service_context=self)
self._build_flows()
@staticmethod
def update_env(key: str, value: str | None):
@ -123,19 +105,6 @@ class ServiceContext(BaseContext):
if value:
os.environ[key] = value
def update_api_envs(
self,
llm_api_key: str | None = None,
llm_base_url: str | None = None,
embedding_api_key: str | None = None,
embedding_base_url: str | None = None,
):
"""Update common environment variables for LLM and embedding services."""
self.update_env("REME_LLM_API_KEY", llm_api_key)
self.update_env("REME_LLM_BASE_URL", llm_base_url)
self.update_env("REME_EMBEDDING_API_KEY", embedding_api_key)
self.update_env("REME_EMBEDDING_BASE_URL", embedding_base_url)
@staticmethod
def _update_section_config(config: dict, section_name: str, **kwargs):
"""Update a specific section of the service config with new values."""
@ -144,170 +113,3 @@ class ServiceContext(BaseContext):
if "default" not in config[section_name]:
config[section_name]["default"] = {}
config[section_name]["default"].update(kwargs)
def _build_flows(self):
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)
self.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(flow_config=flow_config, service_context=self) # noqa
self.flows[flow.name] = flow
else:
logger.info("No expression flow found, please check your configuration.")
async def start(self):
"""Start the service context by initializing all configured components."""
# Recreate thread pool if it was shut down
if self.thread_pool is None or self.thread_pool._shutdown: # pylint: disable=protected-access
self.thread_pool = ThreadPoolExecutor(
max_workers=self.service_config.thread_pool_max_workers,
)
# Re-initialize Ray if it was shut down
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)
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:
self.llms[name] = R.llms[config.backend](model_name=config.model_name, **config.model_extra)
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:
self.embedding_models[name] = R.embedding_models[config.backend](
cache_dir=self.working_path / "embedding_cache",
model_name=config.model_name,
**config.model_extra,
)
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:
self.token_counters[name] = R.token_counters[config.backend](
model_name=config.model_name,
**config.model_extra,
)
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:
# Extract config dict and replace special fields with actual instances
config_dict = config.model_dump(exclude={"backend", "embedding_model"})
config_dict.update(
{
"embedding_model": self.embedding_models[config.embedding_model],
"thread_pool": self.thread_pool,
},
)
self.vector_stores[name] = R.vector_stores[config.backend](**config_dict)
await self.vector_stores[name].create_collection(config.collection_name)
for name, config in self.service_config.memory_stores.items():
if config.backend not in R.memory_stores:
logger.warning(f"Memory store backend {config.backend} is not supported.")
else:
# Extract config dict and replace embedding_model string with actual instance
config_dict = config.model_dump(exclude={"backend", "embedding_model"})
config_dict.update(
{
"embedding_model": self.embedding_models[config.embedding_model],
"thread_pool": self.thread_pool,
"db_path": self.working_path / config.db_name,
},
)
self.memory_stores[name] = R.memory_stores[config.backend](**config_dict)
await self.memory_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", "memory_store"})
config_dict["memory_store"] = self.memory_stores[config.memory_store]
self.file_watchers[name] = R.file_watchers[config.backend](**config_dict)
await self.file_watchers[name].start()
if self.service_config.mcp_servers:
await self.prepare_mcp_servers()
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
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.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 reset_default_collection(self, collection_name: str):
"""Reset the default vector store."""
await self.vector_stores["default"].reset_collection(collection_name)
async def close(self):
"""Close all service components asynchronously."""
for name, vector_store in self.vector_stores.items():
logger.info(f"Closing vector store: {name}")
await vector_store.close()
for name, memory_store in self.memory_stores.items():
logger.info(f"Closing memory store: {name}")
await memory_store.close()
for name, file_watcher in self.file_watchers.items():
logger.info(f"Closing file watcher: {name}")
await file_watcher.close()
for name, llm in self.llms.items():
logger.info(f"Closing LLM: {name}")
await llm.close()
for name, embedding_model in self.embedding_models.items():
logger.info(f"Closing embedding model: {name}")
await embedding_model.close()
self.shutdown_thread_pool()
self.shutdown_ray()
def shutdown_thread_pool(self, wait: bool = True):
"""Shutdown the thread pool executor."""
if self.thread_pool:
self.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)

View file

@ -37,6 +37,7 @@ class BaseEmbeddingModel(ABC):
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.
@ -51,6 +52,7 @@ class BaseEmbeddingModel(ABC):
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 = api_key
@ -63,6 +65,7 @@ class BaseEmbeddingModel(ABC):
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
@ -89,9 +92,7 @@ class BaseEmbeddingModel(ABC):
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 max_input_length {self.max_input_length}, truncating",
)
logger.warning(f"Text length {len(text)} exceeds {self.max_input_length}, truncating")
return text[: self.max_input_length]
return text
@ -133,6 +134,9 @@ class BaseEmbeddingModel(ABC):
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")
@ -151,8 +155,10 @@ class BaseEmbeddingModel(ABC):
continue
try:
data = json.loads(line)
cache_key = data.get("key")
embedding = data.get("embedding")
if not data:
continue
# Each line is {cache_key: embedding}
cache_key, embedding = next(iter(data.items()))
if cache_key and embedding:
# Skip if already loaded (keep the newest)
@ -174,7 +180,12 @@ class BaseEmbeddingModel(ABC):
logger.info(f"Loaded {loaded_count} embeddings from cache file: {cache_file}")
except Exception as e:
logger.error(f"Failed to load cache from {cache_file}: {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).
@ -182,6 +193,9 @@ class BaseEmbeddingModel(ABC):
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")
@ -191,7 +205,7 @@ class BaseEmbeddingModel(ABC):
try:
with open(cache_file, "w", encoding="utf-8") as f:
for cache_key, embedding in self._embedding_cache.items():
cache_entry = {"key": cache_key, "embedding": embedding}
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}")
@ -207,6 +221,9 @@ class BaseEmbeddingModel(ABC):
Returns:
Cached embedding vector or None if not found
"""
if not self.enable_cache:
return None
cache_key = self._get_cache_key(text)
if cache_key in self._embedding_cache:
# Move to end (most recently used)
@ -227,6 +244,9 @@ class BaseEmbeddingModel(ABC):
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

View file

@ -134,10 +134,10 @@ class BaseFileWatcher:
existing_files.add((Change.added, str(file_path)))
if existing_files:
logger.info(f"Found {len(existing_files)} existing files to process")
logger.info(f"[SCAN_ON_START] Found {len(existing_files)} existing files matching watch criteria")
await self.on_changes(existing_files)
else:
logger.info("No existing files found matching watch criteria")
logger.info("[SCAN_ON_START] No existing files found matching watch criteria")
files: list[str] = await self.memory_store.list_files(MemorySource.MEMORY)
for file_path in files:

View file

@ -38,6 +38,7 @@ class BaseMemoryStore(ABC):
self.store_name: str = store_name
self.db_path: Path = Path(db_path)
self.db_path.mkdir(parents=True, exist_ok=True)
self.thread_pool: ThreadPoolExecutor = thread_pool
self.embedding_model: BaseEmbeddingModel = embedding_model
self.vector_enabled: bool = vector_enabled

View file

@ -112,8 +112,6 @@ class ChromaMemoryStore(BaseMemoryStore):
if self.client is not None:
return
self.db_path.mkdir(parents=True, exist_ok=True)
# Initialize persistent ChromaDB client
self.client = chromadb.PersistentClient(
path=str(self.db_path),

View file

@ -51,8 +51,8 @@ class LocalMemoryStore(BaseMemoryStore):
self._chunks: dict[str, _ChunkRecord] = {}
self._files: dict[str, dict[str, FileMetadata]] = {} # source -> path -> meta
# Persistence paths (mirror ChromaMemoryStore convention)
self._chunks_file: Path = self.db_path.parent / f"{self.store_name}_chunks.jsonl"
self._metadata_file: Path = self.db_path.parent / f"{self.store_name}_file_metadata.json"
self._chunks_file: Path = self.db_path / f"{self.store_name}_chunks.jsonl"
self._metadata_file: Path = self.db_path / f"{self.store_name}_file_metadata.json"
# ------------------------------------------------------------------
# Persistence helpers
@ -142,7 +142,6 @@ class LocalMemoryStore(BaseMemoryStore):
if self._started:
return
self._started = True
self.db_path.mkdir(parents=True, exist_ok=True)
await self._load_metadata()
await self._load_chunks()
logger.info(

View file

@ -4,7 +4,6 @@ import json
import sqlite3
import struct
import time
from pathlib import Path
from loguru import logger
@ -63,9 +62,7 @@ class SqliteMemoryStore(BaseMemoryStore):
if self.conn is not None:
return
Path(self.db_path).parent.mkdir(parents=True, exist_ok=True)
self.conn = sqlite3.connect(self.db_path, check_same_thread=False)
self.conn = sqlite3.connect(self.db_path / "reme.db", check_same_thread=False)
# Only load sqlite-vec extension if vector search is enabled
if self.vector_enabled:

View file

@ -83,7 +83,6 @@ class MemoryStoreConfig(BaseModel):
model_config = ConfigDict(extra="allow")
backend: str = Field(default="sqlite")
db_name: str = Field(default="reme.db")
store_name: str = Field(default="reme")
embedding_model: str = Field(default="default")
@ -114,7 +113,7 @@ class ServiceConfig(BaseModel):
backend: str = Field(default="")
app_name: str = Field(default=os.getenv("APP_NAME", "ReMe"))
working_dir: str | None = Field(default=None)
working_dir: str = Field(default=".reme")
enable_logo: bool = Field(default=True)
language: str = Field(default="")
thread_pool_max_workers: int = Field(default=16)
@ -122,8 +121,8 @@ class ServiceConfig(BaseModel):
log_to_console: bool = Field(default=True)
disabled_flows: list[str] = Field(default_factory=list)
enabled_flows: list[str] = Field(default_factory=list)
mcp_servers: dict[str, dict] = Field(default_factory=dict)
mcp_servers: dict[str, dict] = Field(default_factory=dict)
mcp: MCPConfig = Field(default_factory=MCPConfig)
http: HttpConfig = Field(default_factory=HttpConfig)
cmd: CmdConfig = Field(default_factory=CmdConfig)

View file

@ -7,6 +7,7 @@ from .chunking_utils import chunk_markdown
from .common_utils import run_coro_safely, execute_stream_task, hash_text, cosine_similarity, batch_cosine_similarity
from .env_utils import load_env
from .execute_utils import exec_code, run_shell_command, async_exec_code
from .horse import play_horse_easter_egg
from .http_client import HttpClient
from .llm_utils import extract_content, format_messages, deduplicate_memories
from .logger_utils import init_logger
@ -32,6 +33,7 @@ __all__ = [
"exec_code",
"async_exec_code",
"run_shell_command",
"play_horse_easter_egg",
"HttpClient",
"extract_content",
"format_messages",

View file

@ -5,8 +5,8 @@ with support for async execution and output capture.
"""
import asyncio
import contextlib
import concurrent.futures
import contextlib
from io import StringIO

View file

@ -20,7 +20,7 @@ def _mirror_frame(frame: str) -> str:
return "\n".join(mirrored)
def _play_horse_easter_egg() -> None:
def play_horse_easter_egg() -> None:
"""Play the /horse Easter egg: fireworks, galloping horse, and a blessing."""
cols = shutil.get_terminal_size((80, 24)).columns
rows = shutil.get_terminal_size((80, 24)).lines

View file

@ -11,7 +11,7 @@ from reme.core.op import BaseTool
from .agent.chat import FsCli
from .core.enumeration import ChunkEnum
from .core.schema import StreamChunk
from .core.utils import execute_stream_task
from .core.utils import execute_stream_task, play_horse_easter_egg
from .reme_fs import ReMeFs
from .tool.fs import (
BashTool,
@ -23,7 +23,6 @@ from .tool.fs import (
)
from .tool.gallery import ExecuteCode
from .tool.search import DashscopeSearch, TavilySearch
from .horse import _play_horse_easter_egg
class ReMeCli(ReMeFs):
@ -142,7 +141,7 @@ class ReMeCli(ReMeFs):
continue
if user_input == "/horse":
_play_horse_easter_egg()
play_horse_easter_egg()
continue
# Stream processing state

View file

@ -101,7 +101,7 @@ class ReMeFs(Application):
previous_summary: str = "",
language: str = "zh",
**kwargs,
) -> str:
) -> str | dict:
"""Compact messages into a summary."""
compactor = FsCompactor(language=language, **kwargs)
return await compactor.call(
@ -118,7 +118,7 @@ class ReMeFs(Application):
version: str = "default",
language: str = "zh",
**kwargs,
):
) -> str | dict:
"""Generate a summary of the given messages."""
summarizer = FsSummarizer(
tools=[

View file

@ -4,9 +4,9 @@ This package exposes and registers summarization-related operators such as
`TrajectoryPreprocess` and `SuccessExtraction` to the global operator registry.
"""
from ....core import R
from .trajectory_preprocess import TrajectoryPreprocess
from .success_extraction import SuccessExtraction
from .trajectory_preprocess import TrajectoryPreprocess
from ....core import R
__all__ = ["TrajectoryPreprocess", "SuccessExtraction"]

View file

@ -3,7 +3,7 @@
import os
import time
from reme.horse import _mirror_frame
from reme.core.utils.horse import _mirror_frame
def clear_screen():