mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
refactor(core): restructure application initialization and component management
This commit is contained in:
parent
3408e8dcf1
commit
e9757bccf6
26 changed files with 406 additions and 696 deletions
|
|
@ -20,8 +20,7 @@ __all__ = [
|
|||
"ReMeFs",
|
||||
]
|
||||
|
||||
__version__ = "0.3.0.0b3"
|
||||
|
||||
__version__ = "0.3.0.0b4"
|
||||
|
||||
"""
|
||||
conda create -n fl_test2 python=3.10
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 doesn’t exist, an error message will be returned).
|
||||
2. Intelligently merge new information with existing content (skip merging if the file doesn’t 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 there’s 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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)})"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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=[
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue