refactor(core): restructure module organization and enhance error handling

This commit is contained in:
jinli.yl 2026-02-26 15:34:36 +08:00
parent 64b497e330
commit ed133ffc30
82 changed files with 180 additions and 125 deletions

View file

@ -1,9 +0,0 @@
"""A simple chatbot."""
from . import chat
from . import memory
__all__ = [
"chat",
"memory",
]

View file

@ -1,16 +0,0 @@
"""chat agent"""
from .fs_cli import FsCli
from .simple_chat import SimpleChat
from .stream_chat import StreamChat
from ...core import R
__all__ = [
"FsCli",
"StreamChat",
"SimpleChat",
]
R.ops.register(FsCli)
R.ops.register(SimpleChat)
R.ops.register(StreamChat)

View file

@ -16,6 +16,7 @@ from ..llm import BaseLLM
from ..prompt_handler import PromptHandler
from ..runtime_context import RuntimeContext
from ..schema import Response, ServiceConfig
from ..schema.service_config import OpConfig
from ..service_context import ServiceContext
from ..token_counter import BaseTokenCounter
from ..utils import camel_to_snake, CacheHandler, timer
@ -174,12 +175,42 @@ class BaseOp(metaclass=ABCMeta):
return self.context.response
def before_execute_sync(self):
"""Prepare context and validate before sync execution."""
"""Prepare context and validate before sync execution.
This method performs the following steps:
1. Apply input mapping to transform context variables
2. Load operator-specific configuration from service config if available
3. Override operator parameters and prompts based on config
"""
self.context.apply_mapping(self.input_mapping)
if self.context.service_context is None:
return
service_config = self.service_context.service_config
if self.name not in service_config.ops:
return
op_config: OpConfig = service_config.ops[self.name]
# Override operator parameters from config
if op_config.params:
for k, v in op_config.params.items():
if hasattr(self, k):
setattr(self, k, v)
logger.info(f"[{self.__class__.__name__}] Set attribute '{k}' = {v}")
else:
self.op_params[k] = v
logger.info(f"[{self.__class__.__name__}] Set op_param '{k}' = {v}")
# Load custom prompt templates from config
if op_config.prompt_dict:
self.prompt.load_prompt_dict(op_config.prompt_dict)
logger.info(f"[{self.__class__.__name__}] Loaded prompt keys={list(op_config.prompt_dict.keys())}")
async def before_execute(self):
"""Prepare context and validate before async execution."""
self.context.apply_mapping(self.input_mapping)
self.before_execute_sync()
def execute_sync(self):
"""Define core sync logic in subclasses."""

View file

@ -34,6 +34,7 @@ class BaseFileTool(BaseTool):
response = await self.execute()
response = await self.after_execute(response)
return response
except Exception as e:
# Return error message to LLM instead of raising
error_msg = f"{self.__class__.__name__} failed: {str(e)}"

View file

@ -0,0 +1,18 @@
"""Extension operations and tools."""
from .simple_chat import SimpleChat
from .stream_chat import StreamChat
from .test_op import TestOp
from .translate_ts import TranslateTs
from ..core.registry_factory import R
__all__ = [
"SimpleChat",
"StreamChat",
"TestOp",
"TranslateTs",
]
for name in __all__:
op_class = globals()[name]
R.ops.register(op_class)

View file

@ -2,9 +2,9 @@
from loguru import logger
from ...core.enumeration import Role
from ...core.op import BaseTool
from ...core.schema import Message, ToolCall
from ..core.enumeration import Role
from ..core.op import BaseTool
from ..core.schema import Message, ToolCall
class SimpleChat(BaseTool):

View file

@ -2,9 +2,9 @@
from loguru import logger
from ...core.enumeration import Role, ChunkEnum
from ...core.op import BaseTool
from ...core.schema import Message, ToolCall
from ..core.enumeration import Role, ChunkEnum
from ..core.op import BaseTool
from ..core.schema import Message, ToolCall
class StreamChat(BaseTool):

View file

@ -2,7 +2,7 @@
from loguru import logger
from ...core.op import BaseOp
from ..core.op import BaseOp
class TestOp(BaseOp):

View file

@ -4,9 +4,9 @@ from pathlib import Path
from loguru import logger
from ...core.enumeration import Role
from ...core.op import BaseOp
from ...core.schema import Message
from ..core.enumeration import Role
from ..core.op import BaseOp
from ..core.schema import Message
class TranslateTs(BaseOp):

View file

@ -1,11 +1,18 @@
"""File system agents for memory management."""
"""File-based memory operations."""
from .fs_cli import FsCli
from .fs_compactor import FsCompactor
from .fs_context_checker import FsContextChecker
from .fs_summarizer import FsSummarizer
from ...core.registry_factory import R
__all__ = [
"FsSummarizer",
"FsCli",
"FsCompactor",
"FsContextChecker",
"FsSummarizer",
]
for name in __all__:
op_class = globals()[name]
R.ops.register(op_class)

View file

@ -9,20 +9,20 @@ from loguru import logger
from ...core.enumeration import Role, ChunkEnum
from ...core.op import BaseReactStream
from ...core.schema import Message, StreamChunk
from ...core.tools import BashTool, LsTool, ReadTool, WriteTool, EditTool
from ...core.utils import format_messages
from ...tool.fs import BashTool, LsTool, ReadTool, WriteTool, EditTool
class FsCli(BaseReactStream):
"""FsCli agent with system prompt."""
def __init__(
self,
working_dir: str,
context_window_tokens: int = 128000,
reserve_tokens: int = 36000,
keep_recent_tokens: int = 20000,
**kwargs,
self,
working_dir: str,
context_window_tokens: int = 128000,
reserve_tokens: int = 36000,
keep_recent_tokens: int = 20000,
**kwargs,
):
super().__init__(**kwargs)
self.working_dir: str = working_dir
@ -50,7 +50,7 @@ class FsCli(BaseReactStream):
remaining_tasks.append(task)
self.summary_tasks = remaining_tasks
from ..fs import FsSummarizer
from .fs_summarizer import FsSummarizer
# Summarize current conversation and save to memory files
current_date = datetime.now().strftime("%Y-%m-%d")
@ -94,7 +94,7 @@ class FsCli(BaseReactStream):
async def context_check(self) -> dict:
"""Check if messages exceed token limits."""
# Import required modules
from ..fs import FsContextChecker
from .fs_context_checker import FsContextChecker
# Step 1: Check and find cut point
checker = FsContextChecker(
@ -120,7 +120,7 @@ class FsCli(BaseReactStream):
return "No history to compact."
# Import required modules
from ..fs import FsCompactor
from .fs_compactor import FsCompactor
# Step 1: Check and find cut point
cut_result = await self.context_check()

View file

@ -2,50 +2,60 @@
from .base_memory_tool import BaseMemoryTool
from .delegate_task import DelegateTask
# chunk tools
from .chunk.memory_get import MemoryGet
from .chunk.memory_search import MemorySearch
# history tools
from .history.add_history import AddHistory
from .history.read_history import ReadHistory
from .history.read_history_v2 import ReadHistoryV2
# profiles tools
from .profiles.add_draft_and_read_all_profiles import AddDraftAndReadAllProfiles
from .profiles.add_profile import AddProfile
from .profiles.delete_profile import DeleteProfile
from .profiles.profile_handler import ProfileHandler
from .profiles.read_all_profiles import ReadAllProfiles
from .profiles.update_profile import UpdateProfile
from .profiles.update_profiles_v1 import UpdateProfilesV1
from .vector.add_and_retrieve_similar_memory import AddAndRetrieveSimilarMemory
from .vector.add_draft_and_retrieve_similar_memory import AddDraftAndRetrieveSimilarMemory
from .vector.add_memory import AddMemory
from .vector.delete_memory import DeleteMemory
from .vector.memory_handler import MemoryHandler
from .vector.retrieve_memory import RetrieveMemory
from .vector.retrieve_recent_memory import RetrieveRecentMemory
from .vector.update_memory import UpdateMemory
from .vector.update_memory_v1 import UpdateMemoryV1
from .vector.update_memory_v2 import UpdateMemoryV2
# record tools
from .record.add_and_retrieve_similar_memory import AddAndRetrieveSimilarMemory
from .record.add_draft_and_retrieve_similar_memory import AddDraftAndRetrieveSimilarMemory
from .record.add_memory import AddMemory
from .record.delete_memory import DeleteMemory
from .record.retrieve_memory import RetrieveMemory
from .record.retrieve_recent_memory import RetrieveRecentMemory
from .record.update_memory import UpdateMemory
from .record.update_memory_v1 import UpdateMemoryV1
from .record.update_memory_v2 import UpdateMemoryV2
from ...core import R
__all__ = [
# Base
# base
"BaseMemoryTool",
"DelegateTask",
# History
# chunk tools
"MemoryGet",
"MemorySearch",
# history tools
"AddHistory",
"ReadHistory",
"ReadHistoryV2",
# Profiles
# profiles tools
"AddDraftAndReadAllProfiles",
"AddProfile",
"ProfileHandler",
"DeleteProfile",
"ReadAllProfiles",
"UpdateProfile",
"DeleteProfile",
"UpdateProfilesV1",
# Vector
# record tools
"AddAndRetrieveSimilarMemory",
"AddDraftAndRetrieveSimilarMemory",
"AddMemory",
"DeleteMemory",
"MemoryHandler",
"RetrieveMemory",
"RetrieveRecentMemory",
"UpdateMemory",
@ -55,5 +65,4 @@ __all__ = [
for name in __all__:
tool_class = globals()[name]
if isinstance(tool_class, type) and issubclass(tool_class, BaseMemoryTool) and tool_class is not BaseMemoryTool:
R.ops.register(tool_class)
R.ops.register(tool_class)

View file

@ -3,16 +3,20 @@
import os
from pathlib import Path
from reme.core.schema import ToolCall
from .base_fs_tool import BaseFsTool
from loguru import logger
from ....core import RuntimeContext
from ....core.op import BaseTool
from ....core.schema import ToolCall
class FsMemoryGet(BaseFsTool):
class MemoryGet(BaseTool):
"""Read specific snippets from memory files."""
def __init__(self, cwd: str | None = None, **kwargs):
"""Initialize memory get tool."""
kwargs.setdefault("name", "memory_get")
kwargs.setdefault("max_retries", 1)
kwargs.setdefault("raise_exception", False)
super().__init__(**kwargs)
self.cwd = cwd or os.getcwd()
@ -59,7 +63,7 @@ class FsMemoryGet(BaseFsTool):
# Check file exists, is not a symlink, and is a regular file
file_path = Path(abs_path)
assert (
file_path.exists() and not file_path.is_symlink() and file_path.is_file()
file_path.exists() and not file_path.is_symlink() and file_path.is_file()
), f"File not found or not a regular file: {abs_path}"
with open(abs_path, "r", encoding="utf-8") as f:
@ -85,5 +89,25 @@ class FsMemoryGet(BaseFsTool):
count = total_lines - start + 1
# Extract slice (1-indexed to 0-indexed conversion)
selected = lines[start - 1 : start - 1 + count]
selected = lines[start - 1: start - 1 + count]
return "\n".join(selected)
async def call(self, context: RuntimeContext = None, **kwargs):
"""Execute the tool with unified error handling.
This method catches all exceptions and returns error messages
to the LLM instead of raising them.
"""
self.context = RuntimeContext.from_context(context, **kwargs)
try:
await self.before_execute()
response = await self.execute()
response = await self.after_execute(response)
return response
except Exception as e:
# Return error message to LLM instead of raising
error_msg = f"{self.__class__.__name__} failed: {str(e)}"
logger.error(error_msg)
return await self.after_execute(error_msg)

View file

@ -2,26 +2,30 @@
import json
from reme.core.enumeration import MemorySource
from reme.core.schema import ToolCall
from .base_fs_tool import BaseFsTool
from loguru import logger
from ....core.enumeration import MemorySource
from ....core.op import BaseTool
from ....core.runtime_context import RuntimeContext
from ....core.schema import ToolCall
class FsMemorySearch(BaseFsTool):
class MemorySearch(BaseTool):
"""Semantically search MEMORY.md and memory files."""
def __init__(
self,
sources: list[MemorySource] | None = None,
min_score: float = 0.1,
max_results: int = 5,
vector_weight: float = 0.7,
candidate_multiplier: float = 3.0,
**kwargs,
self,
sources: list[MemorySource] | None = None,
min_score: float = 0.1,
max_results: int = 5,
vector_weight: float = 0.7,
candidate_multiplier: float = 3.0,
**kwargs,
):
"""Initialize memory search tool."""
assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be between 0 and 1, got {vector_weight}"
kwargs.setdefault("name", "memory_search")
kwargs.setdefault("max_retries", 1)
kwargs.setdefault("raise_exception", False)
super().__init__(**kwargs)
self.sources = sources or [MemorySource.MEMORY]
self.min_score = min_score
@ -67,10 +71,10 @@ class FsMemorySearch(BaseFsTool):
assert query, "Query cannot be empty"
assert (
isinstance(min_score, float) and 0.0 <= min_score <= 1.0
isinstance(min_score, float) and 0.0 <= min_score <= 1.0
), f"min_score must be between 0 and 1, got {min_score}"
assert (
isinstance(max_results, int) and max_results > 0
isinstance(max_results, int) and max_results > 0
), f"max_results must be a positive integer, got {max_results}"
# Use hybrid_search from file_store
@ -86,3 +90,23 @@ class FsMemorySearch(BaseFsTool):
results = [r for r in results if r.score >= min_score]
return json.dumps([result.model_dump(exclude_none=True) for result in results], indent=2, ensure_ascii=False)
async def call(self, context: RuntimeContext = None, **kwargs):
"""Execute the tool with unified error handling.
This method catches all exceptions and returns error messages
to the LLM instead of raising them.
"""
self.context = RuntimeContext.from_context(context, **kwargs)
try:
await self.before_execute()
response = await self.execute()
response = await self.after_execute(response)
return response
except Exception as e:
# Return error message to LLM instead of raising
error_msg = f"{self.__class__.__name__} failed: {str(e)}"
logger.error(error_msg)
return await self.after_execute(error_msg)

View file

@ -3,7 +3,7 @@
from loguru import logger
from .base_memory_tool import BaseMemoryTool
from ...agent.memory import BaseMemoryAgent
from ..vector_based import BaseMemoryAgent
from ...core.enumeration import MemoryType
from ...core.schema import ToolCall

View file

@ -3,7 +3,7 @@
import numpy as np
from loguru import logger
from ....core.runtime_dict import ServiceContext
from ....core import ServiceContext
from ....core.enumeration import MemoryType
from ....core.schema import MemoryNode
from ....core.utils.common_utils import batch_cosine_similarity

View file

@ -1,11 +0,0 @@
"""Tool"""
from . import gallery
from . import memory
from . import search
__all__ = [
"gallery",
"memory",
"search",
]

View file

@ -1,9 +0,0 @@
"""workflow"""
from . import gallery
from . import procedural_memory
__all__ = [
"gallery",
"procedural_memory",
]

View file

@ -1,14 +0,0 @@
"""test"""
from .test_op import TestOp
from .translate_ts import TranslateTs
from ...core import R
__all__ = [
"TestOp",
"TranslateTs",
]
for name in __all__:
agent_class = globals()[name]
R.ops.register(agent_class)