Merge remote-tracking branch 'origin/main'

This commit is contained in:
方应 2026-05-19 17:41:50 +08:00
commit d27ca21386
110 changed files with 12528 additions and 43 deletions

43
.github/workflows/unittest.yml vendored Normal file
View file

@ -0,0 +1,43 @@
name: Tests ReMe
on:
push:
branches: [main, master, dev, develop]
pull_request:
branches: [main, master, dev, develop]
workflow_dispatch:
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
unit-tests:
name: Unit Tests - py${{ matrix.python-version }}
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
python-version: ["3.10", "3.13"]
steps:
- uses: actions/checkout@v4
- name: Set up Python ${{ matrix.python-version }}
uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}
cache: 'pip'
- name: Install dependencies
run: |
python -m pip install --upgrade pip setuptools wheel
pip install -e ".[dev,core]"
- name: Run tests4 unit tests
run: |
pytest tests4/unittest \
-v \
--tb=long \
-s \
--log-cli-level=WARNING

2
.gitignore vendored
View file

@ -42,4 +42,4 @@ meta_memory/*
**/data/*.json
*.db
memories/*
.reme/*
.reme/*

183
docs4/reme_design.md Normal file
View file

@ -0,0 +1,183 @@
# 快速测试
```bash
# 终端 A:启动服务
reme4 start
# 终端 B:调用 version 验证服务可用
reme4 version
# 预期输出:✅ ReMe v{__version__}
```
# 基础Job
@jinli
入口:`reme4/reme.py::main()` → `parse_args(*sys.argv[1:])` 解析首个位置参数为 `action`,后续 `key=value` 解析为 kwargs(支持
`service.port=8080` 的 dot notation;自动剥离 `--` / `-` 前缀;值会做 bool / int / float / JSON 转换)。
调用模式:
- `start`:本地启动 `ReMe(Application)` 服务(不经过 client)
- `find_reme`:本地探测正在运行的 reme,不调用服务
- `list`:在 client 端拦截,不转发到服务端,直接返回 action 目录
- 其他 action:通过 `call_server(action, **kwargs)` → `R.get(ComponentEnum.CLIENT, backend)` 实例化客户端并流式打印(任意未列出的
step register name 都按本规则透传)
通用可选参数 `backend:str=http`(取值 `http` / `mcp`,对应 `reme4/components/client/{http_client,mcp_client}.py` 中
`@R.register` 注册名);服务端默认 host/port 见 `reme4/constants.py`,可由 `start` 端通过 `service.host=` / `service.port=`
覆盖。
说明:📥 输入参数 | 📤 输出 | ⭐ 必填 | 🎚️ 默认值 | 🛠️ 内部行为 | 📊 metadata
| 分类 | 指令 (register name) | 入口 | 参数 & 行为 |
|------------|--------------------------------------------------|-------------------------------------------------------------|---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
| 🚀 本地 | 🟢 `start` | `reme.py:30` → `ReMe(**kwargs).run_app()` | 📥 可选 `config=<name\|path>`(默认加载 `reme4/config/default.yaml`,`.yaml/.yml/.json` 都支持,含 `${ENV:-default}` 占位符)| 可选 `service.host=` / `service.port=` 等任意 dot-notation 覆盖 | 🛠️ 流程:`load_env()` → `resolve_app_config(**kwargs)` deep merge → `precheck_start(svc)`(`utils/service_utils.py:72`:目标 host:port 已有 reme → 打印 `reme already running ...` 直接返回;端口被其他进程占用 → stderr 提示 `port {port} occupied. Start on another port: reme4 start service.port=<other_port>` 并 `sys.exit(1)`)→ 启动服务 |
| 🚀 本地 | 🧭 `find_reme` | `reme.py:36` → `utils/service_utils.py:89` | 📥 无 | 📤 发现服务则 stdout 打印 `HOST={host} PORT={port} PID={pid or 'unknown'}`;未发现则 stderr 提示 `reme not started. Try: reme start` 并 `sys.exit(1)` | 🛠️ 流程:先探 `REME_DEFAULT_HOST:REME_DEFAULT_PORT`(`health_check` 命中算 `reme`),再 `pgrep -af "reme.* start"` 扫描其他端口 |
| 🛰️ 客户端 | 📜 `list` | `components/client/base_client.py:36` | 📥 无 | 📤 服务端可用 action 目录(JSON,`indent=2 ensure_ascii=False`)| 🛠️ 在 `BaseClient.__call__` 中拦截,不进入 `_execute`,直接调用 `list_actions()`(HTTP/MCP backend 各自实现) |
| 🌐 通用 step | 🆘 `help` (`help_step`) | `call_server("help")` | 📥 无 | 📤 `answer` 一行一个 job:`🛠️ \`{name}\` — {description} 📥 {params}`,参数渲染为 `name:type*`(必填) / `name:type={default}` / `name:type` | 📊 `metadata.job_count` | 🛠️ 自动跳过名为 `help` 的 job |
| 🌐 通用 step | 🩺 `health_check` (`health_check_step`) | `call_server("health_check")` | 📥 无 | 📤 `answer = "✅/❌ ReMe v{version} - healthy/unhealthy"` | 📊 `metadata.health = {version, healthy, components}` | 🧩 覆盖组件:`embedding_model`(🟢 is_started/is_healthy/model_name/dimensions/cache_size/memory) · `file_graph`(🕸️ n_nodes/n_edges/n_virtual\|n_pending/memory) · `file_store`(📦 n_chunks/n_chunks_with_embedding/memory) · `file_watcher`(👀 background_running/watch_paths) · `keyword_index`(🔤 n_docs/vocab_size/memory) | 🛠️ deep sizeof(含 numpy.nbytes),未启动 / 后台未跑 / embedding 不健康 → ❌ |
| 🌐 通用 step | 🏷️ `version` (`version_step`) | `call_server("version")` | 📥 无 | 📤 `answer = reme4.__version__` | 📊 `metadata.version` |
| 🌐 通用 step | 🔄 `reindex` (`reindex_step`) | `call_server("reindex")` | 📥 无 | 📤 `answer = "🔄 Reindexed {added} file(s)"` | 📊 `metadata.counts = {added, ...}` | 🛠️ 流程:`file_watcher.close()` → `file_store.clear()` → `file_watcher.update_store()` → `file_watcher.start()`(finally 保证重启) |
| 🔎 search | 🔍 `search` (`search_step`) | `call_server("search", query=…, …)` | 📥 `query:str` ⭐ | 🎚️ `limit:int=5`(>0) | 🎚️ `min_score:float=0.0` | ⚖️ `vector_weight:float=0.7` ∈[0,1](keyword 权 = 1-vw)| 🔀 `candidate_multiplier:float=3.0`(candidates = min(200, limit×mult))| 🔗 `expand_links:bool=True` | 🔢 `max_links_per_direction:int=10` | 🎚️ `search_filter:dict={}` | 📤 `answer` 每命中一行 `path:start-end [score=… vector=… keyword=…] text` + 缩进的 `→ outlinks (n)` / `← inlinks (n)` + `via predicate=… anchor=#…` | 📊 `metadata.results` / `metadata.link_expansion` / `metadata.counts={vector,keyword,returned,hybrid}` | 🛠️ 并行 `vector_search` + `keyword_search` → RRF 融合(K=60,按 chunk.id 合并)→ `min_score` 过滤 → `limit` 截断 → 邻居 meta 注入 |
| 🧪 demo | 🪄 `demo_echo` (`demo_echo_step1` + `step2`) | `call_server("demo_echo", query=…, min_score=…)` | 📥 `query:str=""` | 🎚️ `min_score:float=0.5` | 🛠️ step1:`processed_query = query.strip().lower()`,`adjusted_min_score = min_score * 0.9`,写回 context | 📤 step2:`answer = "echo: {processed_query} (min_score={adjusted_min_score})"` | 📊 `metadata = {step, query, min_score, processed_query, adjusted_min_score}` |
| 🌊 demo | 🌊 `stream_demo` (`stream_demo_step1` + `step2`) | `call_server("stream_demo", query=…, repeat=…, interval=…)` | 📥 `query:str=""` | 🎚️ `repeat:int=10` | 🎚️ `interval:float=0.1`(秒/字符)| 🛠️ step1:`stream_text = query * repeat` 写回 context | 📤 step2:按字符 `add_stream_string(ch, ChunkEnum.CONTENT)` 流式输出,`asyncio.sleep(interval)` 节流 |
| 📂 crud | 📖 `read` (`read_step`) | `call_server("read", path=…, …)` | 📥 `path:str` ⭐(**完整相对路径**,相对于 `working_dir`;绝对路径会被拒绝;非 `.md` 后缀拒绝)| 🎚️ `start_line:int=null`(1-based, 含端点)| 🎚️ `end_line:int=null`(1-based, 含端点)| 🎚️ `max_bytes:int=51200`(截断阈值)| 📤 `answer = 选中的行内容`,超过 `max_bytes` 时附加 `--- TRUNCATED ---` 续读指引(`start_line=…`)| 📊 `metadata.path` / `metadata.total_lines`(出错路径才会附带)| 🛠️ 流程:`BaseStep.resolve_path(raw, require_md=True)` → `aiofiles.os.stat` → `read_file_safe`(utf-8-sig BOM 容忍、UnicodeDecodeError fallback `errors=ignore`)→ `split("\n")` 切片 `[s-1:e]` → `truncate_text_output` 按字节截断保行 |
使用示例:
```bash
# 启动(默认 default.yaml)
reme4 start
# 指定 config 与服务端口
reme4 start config=paw.yaml service.port=8181
# 查找在跑的 reme
reme4 find_reme
# HOST=127.0.0.1 PORT=8000 PID=12345
# 列出所有可用 action(client 端处理,不转服务端)
reme4 list
# 转发到服务端的 step:所有 key=value 透传为 step kwargs
reme4 help
reme4 health_check
reme4 version
reme4 reindex
reme4 search query="latency 问题" limit=10 min_score=0.2 vector_weight=0.6
# 读取 working_dir 下的 markdown(完整相对路径;无后缀自动补 .md;可按行切片或限制字节)
reme4 read path=Templates/Recipe.md
reme4 read path=Notes start_line=1 end_line=20
reme4 read path=Big.md max_bytes=4096
# 通过 MCP backend 调用
reme4 search query="..." backend=mcp
```
@sen
| tags | stat | 返回特定tag信息 |
| tags | list | 返回所有tag列表 |
| crud | upload/download | 其他文件 |
| file | stat | path |
| file | list | path |
| property | property:read | |
| property | property:update | path="My Note" status=done xx=xxx |
| property | property:delete | keys="[xxxx, xxxx]" |
| graph | traverse | path="My Note" directtion=forward/backward depth=1 predicat=xxx |
@wangce
| crud | create | path="New Note" content="# Hello" title="xxx" tags="[]" status="" |
| crud | read | path="Templates/Recipe.md" |
| crud | edit | path="Templates/Recipe.md" old="xxx" new="xxx" |
| crud | append | path="My Note" content="New line" |
| crud | prepend | path="My Note" content="New line" |
| crud | delete | path="My Note
| daily:crud | daily:xxx | 与 crud 参数保持一致 |
# 日记类型
| 类型 | 路径 | 说明 |
|-----------|-----------------------------------------------|-----------------------------|
| daily | {daily}/xxxx-mm-dd.md + xxxx-mm-dd/{event}.md | 按日期归档的原始信息记录 |
| topic | topic/{topic:-personal(agent)}/{xxxx}.md | 按主题聚类的二次加工内容 |
| proactive | todo | 基于 daily / topic 思考后主动推送的消息 |
# 生成Job
| 任务 | 输入 | 输出 | 触发时机 | 说明 |
|-------------------------|---------------|-----------------------------------------------|-----------------------------|------------------------------------------------------|
| 日记summary @sen @wangce | msg | {daily}/xxxx-mm-dd.md + xxxx-mm-dd/{event}.md | freq (every_n_turn、compact) | 把 msg 的信息写入 daily 目录 |
| 主题dream + 生成链接 @sen | daily/xxx | knowledge/xxx | /dream | 把 daily 目录的内容按主题聚类合并到 topic 目录, 主动在文档中建立 [[link]] 关联 |
| 主动proactive @wangce | daily / topic | proactive_query | pre_query | 思考 daily / topic 信息,主动决定推送给用户的消息 |
2. file_parser
a. 抽象基类 parse: @jinli
ⅰ. 输入是path:相对路径
ⅱ. 输出是FileMetadata & list[FileChunks] & list[FileEdge]
b. default parser 兼容老方案 @jinli
ⅰ. 带overlap的chunking策略 ,不输出FileEdge
c. markdown parser @sen
ⅰ. 根据markdown ast做chunk,不需要overlap
ⅱ. 增加一个索引的chunk chunk_type @锦鲤 file_chunk_type content/index
ⅲ. 增加link的正则解析:predicate:: [[path#anchor]]
3. file_store @sen
a. 抽象存储:
ⅰ. filenode = file + path + st_mtime + metadata + list[FileEdge]
ⅱ. graph=dict[str, filenode] 内存+json
ⅲ. list[FileChunk] 存db
b. 抽象基类
ⅰ. graph:fellow dict的操作 update/get/set
ⅱ. chunks dict[str, list[chunk]]
1. delete_chunks_by_path
2. update_chunks_by_path
3. list_chunks_by_path
4. vector_search/keyword_search
ⅲ. 手写一个bm25检索
ⅳ. 【核心】检索机制 vector bm25 graph 如何进行融合
4. file_watcher @jinli
a. 抽象基类
ⅰ. on_start:
1. file_store 的start 在前,加载graph,file_watcher在后,递归扫描目录
a. 通过ms_time对比graph,on_change 进行改动
ⅱ. on_change:
1. 更新/增加:
a. delete_chunks_by_path 更新数据库
b. upate_chunks_by_path 更新数据库
c. 更新graph
2. 删除
a. delete_chunks_by_path 更新数据库
MemorySchema
1. markdown文件结构 @sen
a. formatter:
ⅰ. title
ⅱ. desc
ⅲ. tags
ⅳ.
2. memory文件结构目录
a. MEMORY.md
b. msg/files -> daily/YYYYMMDD/YYYYMMDD.md + xxxx.md
ⅰ. YYYYMMDD.md
1. xxx -> xxxx.md
2. xxx -> xxxd.md
ⅱ.
c. daily -> topic/topic_l1/topic_l1.md + xxx.md + topic_l2
d. proactive
steps:
1. 治理(算法+LLM):
a. 节点关联P0:现有的链接做补充,挖掘新的LLM的link
ⅰ. /Users/yuli/workspace/ReMe/reme2/component/edge_extractor/llm_edge_extractor.py
ⅱ. 移动到steps
b. 节点整合/节点拆分/节点归档
c. 健康度检查
2. retrieve 调用store的检索
3. 原子steps:reme edit
4. 组合steps:总结:
a. - freq (every_n_turn、compact) -> daily_summarizer
b. topic (/dream ) -> topic_summarizer(daily_xx -> topic_xx)
c. proactive -> proactive_summarizer(personal_xxx -> proactive_query - pre_query

7
docs4/todo.md Normal file
View file

@ -0,0 +1,7 @@
1. 完善mcp_servers config
2. 完善mcp/http的服务测试
3. [PosixPath('.reme')]
4. error
5. meta信息存在一个地方
6. 测试一个完整的Service client的框架,测试各种命令
7. config 默认改成default

View file

@ -33,9 +33,7 @@ classifiers = [
keywords = ["llm", "memory", "experience", "memoryscope", "ai", "mcp", "http", "reme", "personal"]
dependencies = [
"sqlite-vec>=0.1.6",
"prompt_toolkit>=3.0.52",
"rich>=14.2.0",
"aiofiles>=24.1.0",
"asyncpg>=0.31.0",
"chromadb>=1.3.5",
"dashscope>=1.25.1",
@ -43,16 +41,22 @@ dependencies = [
"fastapi>=0.121.3",
"fastmcp>=2.14.1",
"httpx>=0.28.1",
"jieba>=0.42.1",
"loguru>=0.7.3",
"mcp>=1.25.0",
"networkx>=3.4",
"numpy>=2.2.6",
"openai>=2.8.1",
"pandas>=2.3.3",
"prompt_toolkit>=3.0.52",
"pydantic>=2.12.4",
"pyobvector>=0.1.20",
"pyyaml>=6.0.3",
"qdrant-client>=1.16.0",
"rich>=14.2.0",
"sqlite-vec>=0.1.6",
# pyobvector imports Expression from sqlglot; removed from sqlglot 30+ top-level API
"sqlglot>=25,<30",
"qdrant-client>=1.16.0",
"tavily-python>=0.7.13",
"tiktoken>=0.12.0",
"tqdm>=4.67.1",
@ -60,6 +64,8 @@ dependencies = [
"uvicorn>=0.40.0",
"watchfiles>=1.1.1",
"pyyaml>=6.0.3",
"mistletoe",
"neo4j",
]
[project.optional-dependencies]
@ -75,6 +81,8 @@ dev = [
"furo",
"sphinxcontrib-mermaid",
"pre-commit",
"pytest>=8.0",
"pytest-asyncio>=0.23",
]
full = [
@ -86,13 +94,17 @@ litellm = [
]
light = [
"agentscope==1.0.18",
"agentscope==1.0.19",
"flowllm[reme]>=0.2.0.10",
]
core = [
"agentscope==1.0.19",
]
[tool.setuptools.packages.find]
where = ["."]
include = ["reme_ai*", "reme*"]
include = ["reme_ai*", "reme*", "reme4*"]
exclude = ["test*", "cookbook*", "doc*", "library*", "dist*"]
[tool.setuptools.package-data]
@ -108,6 +120,12 @@ reme = [
"**/*.json",
]
reme4 = [
"**/*.yaml",
"**/*.py",
"**/*.json",
]
[tool.setuptools.dynamic]
version = { attr = "reme.__version__" }
@ -120,6 +138,7 @@ Repository = "https://github.com/agentscope-ai/ReMe"
reme = "reme_ai.main:main"
reme2 = "reme.reme:main"
remecli = "reme.reme_cli:main"
reme4 = "reme4.reme:main"
[tool.pytest.ini_options]
asyncio_default_fixture_loop_scope = "function"

View file

@ -6,7 +6,7 @@ from . import extension
from . import memory
from .reme import ReMe
__version__ = "0.3.1.8"
__version__ = "0.3.1.9"
__all__ = [
"config",

View file

@ -76,12 +76,14 @@ class BaseFileWatcher:
if self._running:
return
self._stop_event = asyncio.Event()
self._running = True
async def _initialize_and_watch():
if self.rebuild_index_on_start:
await self.file_store.clear_all()
logger.info("Cleared all indexed data on start")
if self.file_store is not None:
await self.file_store.clear_all()
logger.info("Cleared all indexed data on start")
await self._scan_existing_files()
await self._watch_loop()
@ -183,6 +185,7 @@ class BaseFileWatcher:
logger.info(f"Starting watch on valid paths: {valid_paths}")
async for changes in awatch(
*valid_paths,
force_polling=True,
watch_filter=self.watch_filter,
recursive=self.recursive,
debounce=self.debounce,

View file

@ -119,7 +119,7 @@ class ReMeLight(Application):
The following directory structure will be created:
- {working_dir}/ - Root working directory
- {working_dir}/memory/ - Memory storage files
- {working_dir}/tool_result/ - Compacted tool result files
- {working_dir}/tool_results/ - Compacted tool result files
- {working_dir}/dialog/ - Raw conversation records
"""
# Initialize working directory structure
@ -127,7 +127,7 @@ class ReMeLight(Application):
self.working_path.mkdir(parents=True, exist_ok=True)
self.memory_path = self.working_path / "memory"
self.memory_path.mkdir(parents=True, exist_ok=True)
self.tool_result_path = self.working_path / "tool_result"
self.tool_result_path = self.working_path / "tool_results"
self.tool_result_path.mkdir(parents=True, exist_ok=True)
self.dialog_path = self.working_path / "dialog"
self.dialog_path.mkdir(parents=True, exist_ok=True)
@ -135,12 +135,14 @@ class ReMeLight(Application):
self.vector_weight: float = vector_weight
self.candidate_multiplier: float = candidate_multiplier
# Build the file watcher config: use provided watch_paths if given, otherwise use defaults
_default_watch_paths = [
str(self.working_path / "MEMORY.md"),
str(self.working_path / "memory.md"),
str(self.memory_path),
]
# Pick the existing memory markdown file. On case-insensitive filesystems
# (Windows NTFS, macOS APFS/HFS+) ``MEMORY.md`` and ``memory.md`` are the
# same file, so this also avoids watching it twice. Default to ``MEMORY.md``
# when neither exists yet.
_memory_md = self.working_path / "MEMORY.md"
if not _memory_md.exists() and (self.working_path / "memory.md").exists():
_memory_md = self.working_path / "memory.md"
_default_watch_paths = [str(_memory_md), str(self.memory_path)]
if default_file_watcher_config and default_file_watcher_config.get("watch_paths"):
_merged_file_watcher_config = default_file_watcher_config
else:

26
reme4/__init__.py Normal file
View file

@ -0,0 +1,26 @@
"""ReMe CLI package."""
__version__ = "0.4.0.0"
from . import config
from . import constants
from . import enumeration
from . import schema
from . import steps
from . import utils
from .application import Application
from .components import BaseComponent
from .reme import ReMe
__all__ = [
"Application",
"BaseComponent",
"ReMe",
# submodules
"config",
"constants",
"enumeration",
"schema",
"steps",
"utils",
]

174
reme4/application.py Normal file
View file

@ -0,0 +1,174 @@
"""Main application entry point."""
import asyncio
import heapq
from pathlib import Path
from typing import AsyncGenerator
from .components import BaseComponent, ApplicationContext
from .enumeration import ComponentEnum
from .schema import Response, StreamChunk
from .utils import execute_stream_task, print_logo, get_logger
class Application(BaseComponent):
"""Main application: initializes components, resolves dependencies, runs jobs."""
def __init__(self, **kwargs) -> None:
self.context = ApplicationContext(**kwargs)
working_path = Path(self.config.working_dir).absolute()
working_path.mkdir(parents=True, exist_ok=True)
(working_path / self.config.metadata_dir).mkdir(parents=True, exist_ok=True)
(working_path / self.config.daily_dir).mkdir(parents=True, exist_ok=True)
(working_path / self.config.knowledge_dir).mkdir(parents=True, exist_ok=True)
if self.config.enable_logo:
print_logo(self.config)
logger = get_logger(
log_to_console=self.config.log_to_console,
log_to_file=self.config.log_to_file,
force_init=True,
)
logger.info(f"Initializing {self.config.app_name} Application")
super().__init__()
from .components import R
# Service
service_config = self.config.service
if not service_config.backend:
raise ValueError("Service configuration is missing the required 'backend' field")
service_cls = R.get(ComponentEnum.SERVICE, service_config.backend)
if not service_cls:
raise ValueError(f"Unregistered service backend '{service_config.backend}'")
params = service_config.model_dump()
params["app_context"] = self.context
self.context.service = service_cls(**params)
# Components
for component_type, component_configs in self.config.components.items():
self.context.components[component_type] = {}
for name, config in component_configs.items():
if not config.backend:
raise ValueError(f"Component '{name}' is missing the required 'backend' field")
backend_cls = R.get(component_type, config.backend)
if not backend_cls:
raise ValueError(f"Unregistered backend '{config.backend}' for component '{name}'")
params = config.model_dump()
params.setdefault("name", name)
params["app_context"] = self.context
self.context.components[component_type][name] = backend_cls(**params)
# Jobs
for job_config in self.config.jobs:
if not job_config.backend:
raise ValueError(f"Job '{job_config.name}' is missing the required 'backend' field")
job_cls = R.get(ComponentEnum.JOB, job_config.backend)
if not job_cls:
raise ValueError(f"Unregistered backend '{job_config.backend}' for job '{job_config.name}'")
params = job_config.model_dump()
params["app_context"] = self.context
self.context.jobs[job_config.name] = job_cls(**params)
@property
def config(self):
"""Application configuration."""
return self.context.app_config
def _topological_order(self) -> list[BaseComponent]:
"""Kahn's algorithm. Raises on missing required dep or cycle."""
nodes: dict[tuple[ComponentEnum, str], BaseComponent] = {
(ctype, name): comp for ctype, group in self.context.components.items() for name, comp in group.items()
}
in_degree: dict[tuple[ComponentEnum, str], int] = dict.fromkeys(nodes, 0)
dependents: dict[tuple[ComponentEnum, str], list[tuple[ComponentEnum, str]]] = {k: [] for k in nodes}
for key, comp in nodes.items():
for dep in comp.dependencies:
dep_key = (dep.ctype, dep.name)
if dep_key in nodes:
dependents[dep_key].append(key)
in_degree[key] += 1
elif not dep.optional:
raise ValueError(
f"Component {key[0].value}:{key[1]} depends on {dep.ctype.value}:{dep.name}, not registered",
)
ready = [k for k, d in in_degree.items() if d == 0]
heapq.heapify(ready)
ordered: list[BaseComponent] = []
while ready:
key = heapq.heappop(ready)
ordered.append(nodes[key])
for downstream in dependents[key]:
in_degree[downstream] -= 1
if in_degree[downstream] == 0:
heapq.heappush(ready, downstream)
if len(ordered) != len(nodes):
unresolved = [f"{k[0].value}:{k[1]}" for k, d in in_degree.items() if d > 0]
raise ValueError(f"Circular dependency detected among: {unresolved}")
return ordered
async def _start(self) -> None:
"""Start components in topological order, then jobs."""
start_order = self._topological_order()
order_str = " -> ".join(f"{c.component_type.value}:{c.name}" for c in start_order)
self.logger.info(f"Component start order: {order_str}")
for component in start_order:
try:
await component.start()
except Exception as e:
self.logger.exception(f"Failed to start {component.component_type.value}:{component.name}: {e}")
for name, job in self.context.jobs.items():
try:
await job.start()
except Exception as e:
self.logger.exception(f"Failed to start job '{name}': {e}")
async def _close(self) -> None:
"""Close all jobs, then components in reverse."""
for name, job in self.context.jobs.items():
try:
await job.close()
except Exception as e:
self.logger.exception(f"Failed to close job '{name}': {e}")
for components in self.context.components.values():
for component in components.values():
try:
await component.close()
except Exception as e:
self.logger.exception(f"Failed to close {component.component_type.value}:{component.name}: {e}")
async def run_job(self, name: str, /, **kwargs) -> Response:
"""Execute a registered job by name."""
if name not in self.context.jobs:
raise KeyError(f"Job '{name}' not found")
return await self.context.jobs[name](**kwargs)
async def run_stream_job(self, name: str, /, **kwargs) -> AsyncGenerator[StreamChunk, None]:
"""Execute a streaming job and yield chunks."""
if name not in self.context.jobs:
raise KeyError(f"Job '{name}' not found")
job = self.context.jobs[name]
stream_queue = asyncio.Queue()
task = asyncio.create_task(job(stream_queue=stream_queue, **kwargs))
async for chunk in execute_stream_task(
stream_queue=stream_queue,
task=task,
task_name=name,
output_format="chunk",
):
assert isinstance(chunk, StreamChunk)
yield chunk
def run_app(self):
"""Start the service and serve the application."""
if self.context.service is None:
raise RuntimeError("Service not configured")
self.context.service.run_app(app=self)

View file

@ -0,0 +1,43 @@
"""Components"""
from . import as_llm
from . import as_llm_formatter
from . import as_token_counter
from . import client
from . import embedding
from . import file_graph
from . import file_parser
from . import file_store
from . import file_watcher
from . import job
from . import keyword_index
from . import service
from . import tokenizer
from .application_context import ApplicationContext
from .base_component import BaseComponent
from .component_registry import ComponentRegistry, R
from .prompt_handler import PromptHandler
from .runtime_context import RuntimeContext
__all__ = [
"ApplicationContext",
"BaseComponent",
"ComponentRegistry",
"R",
"PromptHandler",
"RuntimeContext",
# base components
"as_llm",
"as_llm_formatter",
"as_token_counter",
"client",
"embedding",
"file_graph",
"file_parser",
"file_store",
"file_watcher",
"job",
"keyword_index",
"service",
"tokenizer",
]

View file

@ -0,0 +1,28 @@
"""Application context: shared state container for components, jobs, and service."""
from ..enumeration import ComponentEnum
from ..schema import ApplicationConfig
class ApplicationContext:
"""Holds the parsed config and instantiated components, jobs, and service.
Acts as a passive state container. The actual wiring (resolving backends from
the registry and instantiating each component) is performed by Application.
"""
def __init__(self, **kwargs):
# Parse and validate raw config kwargs into a typed ApplicationConfig.
self.app_config: ApplicationConfig = ApplicationConfig(**kwargs)
# Local imports to avoid circular dependencies during module init.
from .base_component import BaseComponent
from .job import BaseJob
from .service import BaseService
# Service endpoint (e.g. HTTP/MCP). Populated by Application.__init__.
self.service: BaseService | None = None
# Components keyed by type then by user-defined name.
self.components: dict[ComponentEnum, dict[str, BaseComponent]] = {}
# Jobs keyed by user-defined name.
self.jobs: dict[str, BaseJob] = {}

View file

@ -0,0 +1,53 @@
"""AgentScope LLM model wrappers."""
from agentscope.model import AnthropicChatModel, ChatModelBase, OpenAIChatModel
from ..base_component import BaseComponent
from ..component_registry import R
from ...enumeration import ComponentEnum
class BaseAsLLM(BaseComponent):
"""Base wrapper for AgentScope chat models. Builds ``self.model`` in ``_start``."""
component_type = ComponentEnum.AS_LLM
def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.model: ChatModelBase | None = None
async def _close(self) -> None:
self.model = None
@R.register("openai")
class OpenAIAsLLM(BaseAsLLM):
"""OpenAI chat model wrapper."""
async def _start(self) -> None:
self.model = OpenAIChatModel(**self.kwargs)
async def _close(self) -> None:
if self.model is not None:
assert isinstance(self.model, OpenAIChatModel)
await self.model.client.close()
@R.register("anthropic")
class AnthropicAsLLM(BaseAsLLM):
"""Anthropic chat model wrapper."""
async def _start(self) -> None:
self.model = AnthropicChatModel(**self.kwargs)
async def _close(self) -> None:
if self.model is not None:
assert isinstance(self.model, AnthropicChatModel)
await self.model.client.close()
__all__ = [
"BaseAsLLM",
"OpenAIAsLLM",
"AnthropicAsLLM",
]

View file

@ -0,0 +1,44 @@
"""AgentScope LLM formatter wrappers."""
from agentscope.formatter import AnthropicChatFormatter, FormatterBase
from .reme_openai_chat_formatter import ReMeOpenAIChatFormatter
from ..base_component import BaseComponent
from ..component_registry import R
from ...enumeration import ComponentEnum
class BaseAsLLMFormatter(BaseComponent):
"""Base wrapper for AgentScope formatters. Builds ``self.formatter`` in ``_start``."""
component_type = ComponentEnum.AS_LLM_FORMATTER
def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.formatter: FormatterBase | None = None
async def _close(self) -> None:
self.formatter = None
@R.register("openai")
class AsOpenAIChatFormatter(BaseAsLLMFormatter):
"""OpenAI chat formatter wrapper (uses ReMe extensions)."""
async def _start(self) -> None:
self.formatter = ReMeOpenAIChatFormatter(**self.kwargs)
@R.register("anthropic")
class AsAnthropicChatFormatter(BaseAsLLMFormatter):
"""Anthropic chat formatter wrapper."""
async def _start(self) -> None:
self.formatter = AnthropicChatFormatter(**self.kwargs)
__all__ = [
"BaseAsLLMFormatter",
"AsOpenAIChatFormatter",
"AsAnthropicChatFormatter",
]

View file

@ -0,0 +1,141 @@
"""OpenAI chat formatter with ReMe extensions: image promotion and reasoning_content."""
import json
from typing import Any
from agentscope.formatter import OpenAIChatFormatter
# noinspection PyProtectedMember
from agentscope.formatter._openai_formatter import (
_format_openai_image_block,
_to_openai_audio_data,
)
from agentscope.message import Msg, TextBlock, ImageBlock, URLSource
def _format_openai_video_block(video_block: dict) -> dict[str, Any]:
"""Convert a video block to OpenAI ``video_url`` content."""
source = video_block["source"]
if source["type"] == "url":
url = source["url"]
elif source["type"] == "base64":
url = f"data:{source['media_type']};base64,{source['data']}"
else:
raise ValueError(f"Unsupported video source type: {source['type']}")
return {"type": "video_url", "video_url": {"url": url}}
class ReMeOpenAIChatFormatter(OpenAIChatFormatter):
"""OpenAIChatFormatter + tool-result image promotion + reasoning_content passthrough."""
async def _format(self, msgs: list[Msg]) -> list[dict[str, Any]]:
"""Format ``Msg`` list into OpenAI chat-completion message dicts."""
self.assert_list_of_msgs(msgs)
messages: list[dict] = []
i = 0
while i < len(msgs):
msg = msgs[i]
content_blocks = []
tool_calls = []
reasoning_content_blocks = []
for block in msg.get_content_blocks():
typ = block.get("type")
if typ == "text":
content_blocks.append({**block})
elif typ == "thinking":
reasoning_content_blocks.append({**block})
elif typ == "tool_use":
tool_calls.append(
{
"id": block.get("id"),
"type": "function",
"function": {
"name": block.get("name"),
"arguments": json.dumps(block.get("input", {}), ensure_ascii=False),
},
},
)
elif typ == "tool_result":
textual_output, multimodal_data = self.convert_tool_result_to_string(block["output"])
messages.append(
{
"role": "tool",
"tool_call_id": block.get("id"),
"content": textual_output,
"name": block.get("name"),
},
)
# OpenAI tool messages can't carry images; promote to a follow-up user message.
promoted_blocks = []
for url, multimodal_block in multimodal_data:
if multimodal_block["type"] == "image" and self.promote_tool_result_images:
promoted_blocks.extend(
[
TextBlock(type="text", text=f"\n- The image from '{url}': "),
ImageBlock(type="image", source=URLSource(type="url", url=url)),
],
)
if promoted_blocks:
promoted_blocks = [
TextBlock(
type="text",
text="<system-info>The following are the image contents from the tool "
f"result of '{block['name']}':",
),
*promoted_blocks,
TextBlock(type="text", text="</system-info>"),
]
msgs.insert(
i + 1,
Msg(name="user", content=promoted_blocks, role="user"),
)
elif typ == "image":
content_blocks.append(_format_openai_image_block(block))
elif typ == "audio":
# Skip assistant audio — not a valid input modality.
if msg.role == "assistant":
continue
content_blocks.append(
{
"type": "input_audio",
"input_audio": _to_openai_audio_data(block["source"]),
},
)
elif typ == "video":
# Skip assistant video — not a valid input modality.
if msg.role == "assistant":
continue
content_blocks.append(_format_openai_video_block(block))
msg_openai = {
"role": msg.role,
"name": msg.name,
"content": content_blocks or None,
}
if tool_calls:
msg_openai["tool_calls"] = tool_calls
# Merge thinking blocks into reasoning_content for compatible models.
if reasoning_content_blocks:
reasoning_msg = "\n".join(r.get("thinking", "") for r in reasoning_content_blocks)
if reasoning_msg:
msg_openai["reasoning_content"] = reasoning_msg
if msg_openai["content"] or msg_openai.get("tool_calls"):
messages.append(msg_openai)
i += 1
return messages

View file

@ -0,0 +1,35 @@
"""AgentScope token counter wrappers."""
from agentscope.token import TokenCounterBase
from .estimate_token_counter import EstimatedTokenCounter
from ..base_component import BaseComponent
from ..component_registry import R
from ...enumeration import ComponentEnum
class BaseAsTokenCounter(BaseComponent):
"""Base wrapper for AgentScope token counters. Builds ``self.token_counter`` in ``_start``."""
component_type = ComponentEnum.AS_TOKEN_COUNTER
def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.token_counter: TokenCounterBase | None = None
async def _close(self) -> None:
self.token_counter = None
@R.register("estimated")
class EstimatedAsTokenCounter(BaseAsTokenCounter):
"""Character-based estimated token counter — fast but approximate."""
async def _start(self) -> None:
self.token_counter = EstimatedTokenCounter(**self.kwargs)
__all__ = [
"BaseAsTokenCounter",
"EstimatedAsTokenCounter",
]

View file

@ -0,0 +1,21 @@
"""Character-based token-count estimator."""
from agentscope.token import TokenCounterBase
class EstimatedTokenCounter(TokenCounterBase):
"""Approximate token count as ``encoded_byte_len / divisor``.
Cheap proxy when exact counts aren't needed; use the model's real
tokenizer for accuracy.
"""
def __init__(self, estimate_divisor: float = 4, encoding: str = "utf-8"):
if estimate_divisor <= 0:
raise ValueError("estimate_divisor must be positive")
self.estimate_divisor: float = estimate_divisor
self.encoding: str = encoding
async def count(self, text: str, **_kwargs) -> int:
"""Estimated token count for ``text``."""
return int(len(text.encode(self.encoding)) / self.estimate_divisor + 0.5)

View file

@ -0,0 +1,185 @@
"""Base class for components."""
import asyncio
from abc import ABC
from pathlib import Path
from typing import TYPE_CHECKING, Any, Callable, TypeVar, cast
from ..enumeration import ComponentEnum
from ..utils import get_logger
if TYPE_CHECKING:
from .application_context import ApplicationContext
T = TypeVar("T", bound="BaseComponent")
class Dependency:
"""Declared dependency: bind() return value, instance attribute placeholder, and topological-sort edge."""
__slots__ = ("ctype", "name", "default_factory", "optional")
def __init__(
self,
ctype: ComponentEnum,
name: str,
default_factory: Callable[[], Any] | None = None,
optional: bool = True,
) -> None:
self.ctype = ctype
self.name = name
self.default_factory = default_factory
self.optional = optional
def __repr__(self) -> str:
suffix = "?" if self.optional else ""
return f"<unresolved {self.ctype.value}:{self.name}{suffix}>"
def __getattr__(self, item: str) -> Any:
# Guard against using the dependency before start() resolves it.
raise RuntimeError(
f"Dependency {self.ctype.value}:{self.name} accessed before start() (attribute '{item}')",
)
class BaseComponent(ABC):
"""Async lifecycle base class with bind-based dependency injection."""
component_type = ComponentEnum.BASE
def __init__(
self,
name: str | None = None,
backend: str = "",
app_context: "ApplicationContext | None" = None,
**kwargs,
) -> None:
self.name: str = name or self.__class__.__name__
self.backend: str = backend
self.app_context: "ApplicationContext | None" = app_context
self.kwargs: dict = dict(kwargs)
self.logger = get_logger()
if hasattr(self.logger, "bind"):
self.logger = self.logger.bind(component=self.name)
self._is_started: bool = False
self._lock: asyncio.Lock = asyncio.Lock()
# Components created from bind() default_factory in standalone mode (auto-managed lifecycle).
self._owned: list["BaseComponent"] = []
@property
def is_started(self) -> bool:
"""Whether the component has been started."""
return self._is_started
# ----- Dependency declaration ----------------------------------------
@staticmethod
def bind(
name: str | None,
base_cls: type[T],
*,
default_factory: Callable[[], T] | None = None,
optional: bool = True,
) -> T | None:
"""Declare a dependency on another component; resolved at start(). Empty name → None."""
if not name:
return None
ctype = getattr(base_cls, "component_type", None)
if not isinstance(ctype, ComponentEnum) or ctype is ComponentEnum.BASE:
raise TypeError(f"{base_cls.__name__} must declare a non-BASE ComponentEnum 'component_type'")
return cast(T, Dependency(ctype, name, default_factory, optional))
@property
def dependencies(self) -> list[Dependency]:
"""All unresolved bindings declared on this instance."""
return [v for v in self.__dict__.values() if isinstance(v, Dependency)]
async def _resolve_bindings(self) -> None:
"""Replace Dependency placeholders with real components (or default_factory / None for optional)."""
for attr, value in list(self.__dict__.items()):
if not isinstance(value, Dependency):
continue
if self.app_context is None:
# Standalone mode: factory or (optional → None) or keep placeholder.
if value.default_factory is not None:
instance = value.default_factory()
setattr(self, attr, instance)
if isinstance(instance, BaseComponent):
self._owned.append(instance)
elif value.optional:
setattr(self, attr, None)
else:
target = self.app_context.components.get(value.ctype, {}).get(value.name)
if target is not None:
setattr(self, attr, target)
elif value.optional:
setattr(self, attr, None)
else:
raise ValueError(f"{value.ctype.value} '{value.name}' not found.")
# ----- Lookup --------------------------------------------------------
@property
def working_path(self) -> Path:
"""Resolved working directory from app context or cwd."""
if self.app_context is None:
return Path.cwd()
return Path(self.app_context.app_config.working_dir)
@property
def working_metadata_path(self) -> Path:
"""Resolved metadata directory: working_path / metadata_dir, or absolute metadata_dir."""
if self.app_context is None:
return Path.cwd() / "metadata"
return self.working_path / self.app_context.app_config.metadata_dir
# ----- Lifecycle -----------------------------------------------------
async def _start(self) -> None:
"""Subclass hook: start logic."""
async def _close(self) -> None:
"""Subclass hook: close logic."""
async def dump(self) -> None:
"""Persist in-memory state to disk. Override in subclasses that need persistence."""
async def load(self) -> None:
"""Restore in-memory state from disk. Override in subclasses that need persistence."""
async def start(self) -> None:
"""Resolve bindings → start owned fallbacks → _start(). No-op if already started."""
async with self._lock:
if self._is_started:
return
await self._resolve_bindings()
for owned in self._owned:
await owned.start()
await self._start()
self._is_started = True
async def close(self) -> None:
"""_close() → close owned fallbacks in reverse. No-op if not started."""
async with self._lock:
if not self._is_started:
return
await self._close()
for owned in reversed(self._owned):
await owned.close()
self._is_started = False
async def restart(self) -> None:
"""Close then start."""
await self.close()
await self.start()
async def __call__(self, **kwargs):
raise NotImplementedError
async def __aenter__(self) -> "BaseComponent":
await self.start()
return self
async def __aexit__(self, exc_type, exc_val, exc_tb) -> None:
await self.close()

View file

@ -0,0 +1,7 @@
"""Client components."""
from .base_client import BaseClient
from .http_client import HttpClient
from .mcp_client import MCPClient
__all__ = ["BaseClient", "HttpClient", "MCPClient"]

View file

@ -0,0 +1,41 @@
"""Base client abstraction."""
import json
from abc import abstractmethod
from collections.abc import AsyncGenerator
from ..base_component import BaseComponent
from ...enumeration import ComponentEnum
class BaseClient(BaseComponent):
"""Abstract base for clients that communicate with ReMe services."""
component_type = ComponentEnum.CLIENT
def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.client = None
async def _start(self) -> None:
"""Initialize the client."""
async def _close(self) -> None:
"""Close the client and release resources."""
@abstractmethod
def _execute(self) -> AsyncGenerator[str, None]:
"""Backend-specific execution; yield text chunks (single yield for non-streaming backends)."""
@abstractmethod
async def list_actions(self) -> list[dict]:
"""Discover available actions on the server; each dict is the raw backend descriptor."""
async def __call__(self) -> AsyncGenerator[str, None]:
"""Dispatch: action='list' returns the action catalog; otherwise delegate to _execute()."""
if getattr(self, "action", None) == "list":
actions = await self.list_actions()
yield json.dumps(actions, indent=2, ensure_ascii=False)
return
async for chunk in self._execute():
yield chunk

View file

@ -0,0 +1,143 @@
"""HTTP client for ReMe services."""
import json
import os
from collections.abc import AsyncGenerator
import httpx
from .base_client import BaseClient
from ..component_registry import R
from ...constants import REME_SERVICE_INFO, REME_DEFAULT_HOST, REME_DEFAULT_PORT
from ...enumeration import ChunkEnum
from ...schema import StreamChunk
@R.register("http")
class HttpClient(BaseClient):
"""HTTP client that auto-adapts to JSON or SSE endpoints via Content-Type."""
def __init__(
self,
action: str,
host: str | None = None,
port: int | None = None,
timeout: float = 30.0,
**kwargs,
):
super().__init__(**kwargs)
# Resolve host/port: explicit args > env var > defaults
if not (host and port):
if service_info := os.environ.get(REME_SERVICE_INFO):
try:
data = json.loads(service_info)
host = data["host"]
port = data["port"]
except Exception:
self.logger.warning(f"Invalid service info: {service_info}")
host, port = REME_DEFAULT_HOST, REME_DEFAULT_PORT
else:
host, port = REME_DEFAULT_HOST, REME_DEFAULT_PORT
self.action = action
self.base_url = f"http://{host}:{port}"
self.timeout = timeout
async def _start(self) -> None:
"""Initialize the HTTP client."""
if self.client is None:
self.client = httpx.AsyncClient(base_url=self.base_url, timeout=self.timeout)
async def _iter_stream_chunks(self) -> AsyncGenerator[StreamChunk, None]:
"""Send request and yield raw StreamChunks; auto-detects JSON vs SSE via Content-Type.
For JSON responses: yields a single CONTENT chunk with the raw response body.
For SSE responses: yields each streaming chunk as it arrives.
"""
if self.client is None:
raise RuntimeError("Client not initialized. Call _start() first.")
async with self.client.stream("POST", f"/{self.action}", json=self.kwargs) as resp:
resp.raise_for_status()
ctype = resp.headers.get("content-type", "")
if ctype.startswith("text/event-stream"):
async for line in resp.aiter_lines():
if not line.startswith("data:"):
continue
payload = line[len("data:") :]
if payload.strip() == "[DONE]":
return
try:
data = json.loads(payload)
except json.JSONDecodeError:
continue
chunk = StreamChunk(**data)
if chunk.chunk_type == ChunkEnum.ERROR:
# Surface server-side errors as exceptions so callers don't
# mistake error chunks for valid content.
raise RuntimeError(str(chunk.chunk))
if chunk.done:
return
yield chunk
else:
body = await resp.aread()
yield StreamChunk(chunk_type=ChunkEnum.CONTENT, chunk=body.decode())
async def stream_chunks(self) -> AsyncGenerator[StreamChunk, None]:
"""HTTP-specific richer access: yield raw StreamChunk objects (no display formatting)."""
async for chunk in self._iter_stream_chunks():
yield chunk
async def list_actions(self) -> list[dict]:
"""Return raw OpenAPI operations; each dict gets an `action` key (path without leading '/')."""
if self.client is None:
raise RuntimeError("Client not initialized. Call _start() first.")
resp = await self.client.get("/openapi.json")
resp.raise_for_status()
spec = resp.json()
actions: list[dict] = []
for path, methods in spec.get("paths", {}).items():
for method, op in methods.items():
actions.append({"action": path.lstrip("/"), "method": method.upper(), **op})
return actions
@staticmethod
def _format_for_display(text: str) -> str:
"""Render a JSON response as human-friendly CLI text; pass through unrecognized payloads."""
try:
data = json.loads(text)
except (ValueError, json.JSONDecodeError):
return text
if not (isinstance(data, dict) and isinstance(data.get("answer"), str)):
return json.dumps(data, indent=2, ensure_ascii=False) if isinstance(data, (dict, list)) else text
d = dict(data)
answer = d.pop("answer")
success = d.pop("success", None)
metadata = d.pop("metadata", None)
parts = [answer]
status_pieces = []
if success is not None:
status_pieces.append("✅" if success else "❌")
if metadata:
status_pieces.append(json.dumps(metadata, ensure_ascii=False))
if status_pieces:
parts.append(" ".join(status_pieces))
if d:
parts.append(json.dumps(d, indent=2, ensure_ascii=False))
return "\n".join(parts)
# pylint: disable=invalid-overridden-method
async def _execute(self) -> AsyncGenerator[str, None]:
"""Yield text chunks for CLI display; JSON responses are pretty-formatted."""
async for chunk in self._iter_stream_chunks():
payload = chunk.chunk
text = payload if isinstance(payload, str) else json.dumps(payload, ensure_ascii=False)
yield self._format_for_display(text)
async def _close(self) -> None:
"""Close the HTTP client."""
if self.client is not None:
await self.client.aclose()
self.client = None

View file

@ -0,0 +1,125 @@
"""MCP client for ReMe services."""
import json
import os
from collections.abc import AsyncGenerator
from typing import Any
from fastmcp import Client
from fastmcp.client import SSETransport, StdioTransport, StreamableHttpTransport
from fastmcp.client.client import CallToolResult
from .base_client import BaseClient
from ..component_registry import R
from ...constants import REME_SERVICE_INFO, REME_DEFAULT_HOST, REME_DEFAULT_PORT
_TRANSPORT_MAP = {
"sse": SSETransport,
"stdio": StdioTransport,
"streamable-http": StreamableHttpTransport,
}
@R.register("mcp")
class MCPClient(BaseClient):
"""MCP client that communicates with ReMe MCP service via fastmcp.Client.
Usage:
# SSE (default)
client = MCPClient(action="my_tool", host="localhost", port=8000, query="hello")
async with client:
async for text in client():
print(text)
# Streamable HTTP
client = MCPClient(action="my_tool", transport="streamable-http", host="localhost", port=8000)
# Stdio
client = MCPClient(action="my_tool", transport="stdio", command="python", args=["server.py"])
# Custom transport object
from fastmcp.client import SSETransport
client = MCPClient(action="my_tool", transport=SSETransport(url="http://host:port/sse"))
"""
def __init__(
self,
action: str,
transport: str | Any = "sse",
host: str | None = None,
port: int | None = None,
timeout: float = 30.0,
**kwargs,
):
super().__init__(**kwargs)
if isinstance(transport, str) and transport not in _TRANSPORT_MAP:
raise ValueError(f"Unknown transport: {transport!r}, expected one of {list(_TRANSPORT_MAP)}")
if isinstance(transport, str) and transport != "stdio":
if not (host and port):
if service_info := os.environ.get(REME_SERVICE_INFO):
try:
data = json.loads(service_info)
host = data["host"]
port = data["port"]
except Exception:
self.logger.warning(f"Invalid service info: {service_info}")
host, port = REME_DEFAULT_HOST, REME_DEFAULT_PORT
else:
host, port = REME_DEFAULT_HOST, REME_DEFAULT_PORT
self.host = host
self.port = port
self.action = action
self.transport = transport
self.timeout = timeout
def _build_transport(self):
if not isinstance(self.transport, str):
return self.transport
cls = _TRANSPORT_MAP[self.transport]
if self.transport == "stdio":
command = self.kwargs.pop("command", "")
args = self.kwargs.pop("args", [])
return cls(command=command, args=args)
path = "/sse" if self.transport == "sse" else "/mcp"
url = f"http://{self.host}:{self.port}{path}"
return cls(url=url)
# pylint: disable=unnecessary-dunder-call
async def _start(self) -> None:
if self.client is None:
self.client = Client(self._build_transport(), timeout=self.timeout)
await self.client.__aenter__()
# pylint: disable=invalid-overridden-method
async def _execute(self) -> AsyncGenerator[str, None]:
if self.client is None:
raise RuntimeError("Client not initialized. Call _start() first.")
result: CallToolResult = await self.client.call_tool(self.action, self.kwargs)
yield self._extract_text(result)
async def list_actions(self) -> list[dict]:
"""Return raw MCP Tool dumps; each dict gets an `action` key (the tool name)."""
if self.client is None:
raise RuntimeError("Client not initialized. Call _start() first.")
tools = await self.client.list_tools()
return [tool.model_dump() for tool in tools]
# pylint: disable=unnecessary-dunder-call
async def _close(self) -> None:
if self.client is not None:
await self.client.__aexit__(None, None, None)
self.client = None
@staticmethod
def _extract_text(result: CallToolResult) -> str:
for block in result.content:
if hasattr(block, "text"):
return block.text
return str(result.content)

View file

@ -0,0 +1,77 @@
"""Global registry mapping (ComponentEnum, name) -> component class."""
from typing import Callable, TypeVar, cast
from .base_component import BaseComponent
from ..enumeration import ComponentEnum
from ..utils import get_logger
T = TypeVar("T", bound=BaseComponent)
class ComponentRegistry:
"""Two-level registry: component_type -> name -> class.
Supports both direct calls — ``R.register(MyClass, "name")`` — and
decorator usage — ``@R.register("name")``.
"""
def __init__(self) -> None:
self._registry: dict[ComponentEnum, dict[str, type[BaseComponent]]] = {}
self.logger = get_logger()
def _do_register(self, cls: type[T], name: str) -> type[T]:
"""Insert `cls` under its `component_type` group; warn on overwrite."""
component_type = getattr(cls, "component_type", None)
if not isinstance(component_type, ComponentEnum):
raise TypeError(f"{cls.__name__} must have a ComponentEnum 'component_type' attribute")
if not name:
raise ValueError("Component name cannot be empty")
group = self._registry.setdefault(component_type, {})
if name in group:
self.logger.warning(f"Component '{name}' already registered for {component_type}, overwriting")
group[name] = cls
return cls
def register(
self,
cls_or_name: type[T] | str,
name: str | None = None,
) -> Callable[[type[T]], type[T]] | type[T]:
"""Register a component class directly, or return a decorator that does so."""
# Direct mode: first arg is the class itself.
if isinstance(cls_or_name, type):
return self._do_register(cast(type[T], cls_or_name), name if name is not None else cls_or_name.__name__)
# Decorator mode: first arg is the registration name.
if not isinstance(cls_or_name, str):
raise TypeError(f"Expected a class or string, got {type(cls_or_name).__name__}")
def decorator(decorated_cls: type[T]) -> type[T]:
return self._do_register(decorated_cls, cls_or_name)
return decorator
def get(self, component_type: ComponentEnum, name: str) -> type[BaseComponent] | None:
"""Look up a registered class; return None if not found."""
return self._registry.get(component_type, {}).get(name)
def get_all(self, component_type: ComponentEnum) -> dict[str, type[BaseComponent]]:
"""Return a shallow copy of all classes registered under `component_type`."""
return dict(self._registry.get(component_type, {}))
def unregister(self, component_type: ComponentEnum, name: str) -> bool:
"""Remove an entry; return True if it existed, False otherwise."""
if (group := self._registry.get(component_type)) and name in group:
del group[name]
return True
return False
def clear(self) -> None:
"""Drop every registered entry."""
self._registry.clear()
# Process-wide singleton used throughout the codebase.
R = ComponentRegistry()

View file

@ -0,0 +1,6 @@
"""Embedding model implementations."""
from .base_embedding_model import BaseEmbeddingModel
from .openai_embedding_model import OpenAIEmbeddingModel
__all__ = ["BaseEmbeddingModel", "OpenAIEmbeddingModel"]

View file

@ -0,0 +1,214 @@
"""Base embedding model with LRU cache and disk persistence."""
import asyncio
import hashlib
import os
from abc import abstractmethod
from collections import OrderedDict
from pathlib import Path
import numpy as np
from ..base_component import BaseComponent
from ...enumeration import ComponentEnum
from ...schema import EmbNode
class BaseEmbeddingModel(BaseComponent):
"""Embedding model with LRU cache and disk persistence."""
component_type = ComponentEnum.EMBEDDING_MODEL
def __init__(
self,
api_key: str | None = None,
base_url: str | None = None,
model_name: str = "",
dimensions: int = 1024,
pass_dimensions: bool = False,
max_batch_size: int = 10,
max_input_length: int = 8192,
max_cache_size: int = 10000,
enable_cache: bool = True,
cache_version: str = "v1",
max_retries: int = 3,
**kwargs,
):
super().__init__(**kwargs)
self.api_key = api_key or os.environ.get("EMBEDDING_API_KEY", "")
self.base_url = base_url or os.environ.get("EMBEDDING_BASE_URL", "")
self.model_name = model_name
self.dimensions = dimensions
self.pass_dimensions = pass_dimensions
self.max_batch_size = max_batch_size
self.max_input_length = max_input_length
self.max_cache_size = max_cache_size
self.enable_cache = enable_cache
self.cache_version = cache_version
self.max_retries = max_retries
self._embedding_cache: OrderedDict[str, np.ndarray] = OrderedDict()
self.is_healthy: bool = True
@property
def cache_path(self) -> Path:
"""Disk path for the embedding cache file."""
return self.working_metadata_path / "embedding_cache" / f"{self.name}_{self.cache_version}.npz"
async def _start(self) -> None:
"""Load cache from disk on startup."""
await self.load()
async def health_check(self, timeout: float = 2.0) -> bool:
"""Probe the provider; sets and returns is_healthy."""
tag = f"[EMBEDDING HEALTH CHECK] name={self.name} model={self.model_name}"
try:
result = await asyncio.wait_for(self._get_embeddings(["ping"]), timeout=timeout)
if not result or result[0] is None:
raise RuntimeError("empty embedding")
self.is_healthy = True
self.logger.info(f"{tag} -> OK")
except asyncio.TimeoutError:
self.is_healthy = False
self.logger.error(f"{tag} -> FAIL timeout({timeout}s)")
except Exception as e:
self.is_healthy = False
self.logger.error(f"{tag} -> FAIL {type(e).__name__}: {e}")
return self.is_healthy
async def _close(self) -> None:
"""Persist cache to disk on shutdown."""
await self.dump()
# -- Public API --
async def get_embedding(self, input_text: str, **kwargs) -> np.ndarray | None:
"""Get embedding for a single text."""
results = await self.get_embeddings([input_text], **kwargs)
return results[0] if results else None
async def get_embeddings(self, input_text: list[str], **kwargs) -> list[np.ndarray | None]:
"""Get embeddings for a list of texts, with caching and batching."""
truncated = [t[: self.max_input_length] for t in input_text]
results: list[np.ndarray | None] = [None] * len(truncated)
to_compute: list[tuple[int, str]] = []
# Split into cache hits and misses
for idx, text in enumerate(truncated):
cached = self._get_from_cache(text)
if cached is not None:
results[idx] = cached
else:
to_compute.append((idx, text))
# Batch-compute misses with retry
if to_compute:
for i in range(0, len(to_compute), self.max_batch_size):
batch = to_compute[i : i + self.max_batch_size]
indices = [idx for idx, _ in batch]
texts = [text for _, text in batch]
embeddings = None
for attempt in range(self.max_retries):
try:
embeddings = await self._get_embeddings(texts, **kwargs)
if embeddings and len(embeddings) == len(texts):
break
except (TimeoutError, ConnectionError, OSError):
if attempt < self.max_retries - 1:
await asyncio.sleep(2**attempt)
except Exception:
self.logger.exception("Embedding request failed")
break
if not embeddings or len(embeddings) != len(texts):
continue
# Normalize dimensions and cache
for orig_idx, text, emb in zip(indices, texts, embeddings):
if emb is None:
continue
emb_array = np.asarray(emb, dtype=np.float16)
if len(emb_array) != self.dimensions:
if len(emb_array) < self.dimensions:
emb_array = np.pad(emb_array, (0, self.dimensions - len(emb_array)))
else:
emb_array = emb_array[: self.dimensions]
results[orig_idx] = emb_array
self._put_to_cache(text, emb_array)
return results
async def get_node_embeddings(self, nodes: list[EmbNode], **kwargs) -> list[EmbNode]:
"""Compute and assign embeddings for EmbNode objects."""
embeddings = await self.get_embeddings([n.text for n in nodes], **kwargs)
if len(embeddings) == len(nodes):
for node, vec in zip(nodes, embeddings):
if vec is not None:
node.embedding = vec
return nodes
@abstractmethod
async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float] | None]:
"""Get raw embeddings from the underlying provider."""
# -- Cache Operations --
def _get_from_cache(self, text: str) -> np.ndarray | None:
"""Lookup text in LRU cache, promoting on hit."""
if not self.enable_cache:
return None
key = self._get_cache_key(text)
if key not in self._embedding_cache:
return None
self._embedding_cache.move_to_end(key)
return self._embedding_cache[key]
def _put_to_cache(self, text: str, embedding: np.ndarray) -> None:
"""Insert into LRU cache, evicting oldest if full."""
if not self.enable_cache or self.max_cache_size <= 0 or len(embedding) != self.dimensions:
return
key = self._get_cache_key(text)
if len(self._embedding_cache) >= self.max_cache_size and key not in self._embedding_cache:
self._embedding_cache.popitem(last=False)
self._embedding_cache[key] = embedding
self._embedding_cache.move_to_end(key)
def _get_cache_key(self, text: str) -> str:
"""Generate cache key from text, model name, and dimensions."""
return hashlib.sha256(f"{text}|{self.model_name}|{self.dimensions}".encode()).hexdigest()
# -- Cache Persistence --
async def load(self) -> None:
"""Load cached embeddings from disk (npz format); replaces in-memory cache."""
self._embedding_cache.clear()
if not self.enable_cache or not self.cache_path.exists():
return
try:
data = np.load(self.cache_path)
except Exception:
self.logger.exception("Failed to load embedding cache, removing")
self.cache_path.unlink(missing_ok=True)
return
for key, emb in zip(data["keys"], data["embeddings"]):
if len(emb) != self.dimensions:
continue
if len(self._embedding_cache) >= self.max_cache_size:
break
self._embedding_cache[str(key)] = emb.astype(np.float16)
self.logger.info(f"Loaded {len(self._embedding_cache)} embeddings from {self.cache_path}")
async def dump(self) -> None:
"""Persist in-memory cache to disk (npz format)."""
if not self.enable_cache or not self._embedding_cache:
return
self.cache_path.parent.mkdir(parents=True, exist_ok=True)
keys = list(self._embedding_cache.keys())
embeddings = np.stack(list(self._embedding_cache.values()))
try:
np.savez(self.cache_path, keys=np.array(keys, dtype=str), embeddings=embeddings)
self.logger.info(f"Saved {len(self._embedding_cache)} embeddings to {self.cache_path}")
except Exception:
self.logger.exception("Failed to save embedding cache")

View file

@ -0,0 +1,52 @@
"""OpenAI-compatible async embedding model."""
from openai import AsyncOpenAI
from .base_embedding_model import BaseEmbeddingModel
from ..component_registry import R
@R.register("openai")
class OpenAIEmbeddingModel(BaseEmbeddingModel):
"""Embedding model backed by any OpenAI-compatible API."""
def __init__(self, **kwargs):
super().__init__(**kwargs)
self._client: AsyncOpenAI | None = None
async def _start(self) -> None:
"""Initialize async OpenAI client."""
self._client = AsyncOpenAI(api_key=self.api_key, base_url=self.base_url, **self.kwargs)
await super()._start()
async def _close(self) -> None:
"""Close the async OpenAI client."""
if self._client:
await self._client.close()
self._client = None
await super()._close()
async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float] | None]:
"""Call the embeddings API and return results aligned to input order."""
if self._client is None:
raise RuntimeError("Client not initialized. Call _start() first.")
create_kwargs: dict = {"model": self.model_name, "input": input_text, **kwargs}
if self.pass_dimensions:
create_kwargs["dimensions"] = self.dimensions
completion = await self._client.embeddings.create(**create_kwargs)
# Map API results back to input order
result: list[list[float] | None] = [None] * len(input_text)
for emb in completion.data:
if 0 <= emb.index < len(input_text):
vec = emb.embedding or getattr(emb, "dense_embedding", None)
if vec is not None:
result[emb.index] = list(vec)
else:
self.logger.warning(f"Empty embedding at index {emb.index}")
else:
self.logger.warning(f"Index {emb.index} out of range for input length {len(input_text)}")
return result

View file

@ -0,0 +1,8 @@
"""File graph module."""
from .base_file_graph import BaseFileGraph
from .local_file_graph import LocalFileGraph
from .neo4j_file_graph import Neo4jFileGraph
from .nx_file_graph import NxFileGraph
__all__ = ["BaseFileGraph", "LocalFileGraph", "Neo4jFileGraph", "NxFileGraph"]

View file

@ -0,0 +1,53 @@
"""Abstract base for file-graph backends."""
from abc import abstractmethod
from pathlib import Path
from ..base_component import BaseComponent
from ...enumeration import ComponentEnum
from ...schema import FileLink, FileNode
class BaseFileGraph(BaseComponent):
"""Abstract base for file-graph backends."""
component_type = ComponentEnum.FILE_GRAPH
def __init__(self, graph_name: str = "default", graph_version: str = "v1", **kwargs):
super().__init__(**kwargs)
self.graph_name: str = graph_name or self.name
self.graph_version: str = graph_version
self.graph_path: Path = self.working_metadata_path / self.component_type.value
self.graph_path.mkdir(parents=True, exist_ok=True)
# -- Node CRUD ---------------------------------------------------------
@abstractmethod
async def upsert_nodes(self, nodes: list[FileNode]) -> None:
"""Insert or update nodes in the graph."""
@abstractmethod
async def delete_nodes(self, paths: list[str]) -> None:
"""Delete nodes by path."""
@abstractmethod
async def get_nodes(self, paths: list[str] | None = None) -> list[FileNode]:
"""Return nodes by paths; None = all real nodes; [] = []."""
@abstractmethod
async def rebuild_links(self) -> None:
"""Rebuild all edges from each node's link payload."""
@abstractmethod
async def clear(self):
"""Remove all nodes and edges."""
# -- Link access -------------------------------------------------------
@abstractmethod
async def get_outlinks(self, path: str) -> list[FileLink]:
"""Return outgoing links for *path*."""
@abstractmethod
async def get_inlinks(self, path: str) -> list[FileLink]:
"""Return incoming links for *path*."""

View file

@ -0,0 +1,138 @@
"""Pure-Python file-graph backend (no external deps)."""
from pathlib import Path
from .base_file_graph import BaseFileGraph
from ..component_registry import R
from ...schema import FileLink, FileNode
@R.register("local")
class LocalFileGraph(BaseFileGraph):
"""Dict-backed file graph; uses FileLink.target_path for adjacency."""
def __init__(self, **kwargs):
super().__init__(**kwargs)
self._nodes: dict[str, FileNode] = {}
self._inverse: dict[str, set[str]] = {} # target → {sources}
self._pending: dict[str, set[str]] = {} # virtual target → {sources}
self._graph_file: Path = self.graph_path / f"{self.graph_name}_{self.graph_version}.jsonl"
# -- Lifecycle ---------------------------------------------------------
async def _start(self) -> None:
await super()._start()
await self.load()
await self.rebuild_links()
async def _close(self) -> None:
await self.dump()
await super()._close()
async def load(self) -> None:
"""Load nodes from JSONL file into memory; keep current state on failure."""
if not self._graph_file.exists():
return
try:
with open(self._graph_file, "r", encoding="utf-8") as f:
self._nodes.update(
(n.path, n) for line in f if line.strip() for n in [FileNode.model_validate_json(line)]
)
self.logger.info(f"Loaded {len(self._nodes)} nodes from {self._graph_file}")
except Exception as e:
self.logger.exception(f"Failed to load {self._graph_file}: {e}")
async def dump(self) -> None:
"""Persist all nodes to JSONL via atomic rename."""
try:
tmp = self._graph_file.with_suffix(".tmp")
with open(tmp, "w", encoding="utf-8") as f:
f.writelines(f"{n.model_dump_json()}\n" for n in self._nodes.values())
tmp.replace(self._graph_file)
self.logger.info(f"Saved {len(self._nodes)} nodes to {self._graph_file}")
except Exception as e:
self.logger.exception(f"Failed to write {self._graph_file}: {e}")
# -- Edge bookkeeping --------------------------------------------------
def _add_edge(self, src: str, target: str) -> None:
"""Register src→target; route to pending if target is virtual."""
bucket = self._inverse if target in self._nodes else self._pending
bucket.setdefault(target, set()).add(src)
def _remove_edge(self, src: str, target: str) -> None:
"""Remove src→target from both inverse and pending buckets."""
for bucket in (self._inverse, self._pending):
srcs = bucket.get(target)
if srcs is None or src not in srcs:
continue
srcs.discard(src)
if not srcs:
del bucket[target]
# -- Node CRUD ---------------------------------------------------------
async def upsert_nodes(self, nodes: list[FileNode]) -> None:
for node in nodes:
path = node.path
old = self._nodes.get(path)
if old is not None:
for link in old.links:
if link.target_path:
self._remove_edge(path, link.target_path)
self._nodes[path] = node
for link in node.links:
if link.target_path:
self._add_edge(path, link.target_path)
# Promote pending edges that now target a real node.
promoted = self._pending.pop(path, None)
if promoted:
self._inverse.setdefault(path, set()).update(promoted)
async def delete_nodes(self, paths: list[str]) -> None:
for path in paths:
node = self._nodes.pop(path, None)
if node is None:
continue
for link in node.links:
if link.target_path:
self._remove_edge(path, link.target_path)
# Demote inbound edges to pending (sources still reference this path).
demoted = self._inverse.pop(path, None)
if demoted:
self._pending.setdefault(path, set()).update(demoted)
async def get_nodes(self, paths: list[str] | None = None) -> list[FileNode]:
if paths is None:
return list(self._nodes.values())
return [self._nodes[p] for p in paths if p in self._nodes]
async def rebuild_links(self) -> None:
"""Rebuild inverse/pending indexes from all node link payloads."""
self._inverse.clear()
self._pending.clear()
for src, node in self._nodes.items():
for link in node.links:
if link.target_path:
self._add_edge(src, link.target_path)
async def clear(self):
self._nodes.clear()
self._inverse.clear()
self._pending.clear()
self._graph_file.unlink(missing_ok=True)
# -- Link access -------------------------------------------------------
async def get_outlinks(self, path: str) -> list[FileLink]:
node = self._nodes.get(path)
if node is None:
return []
return [lnk for lnk in node.links if lnk.target_path and lnk.target_path in self._nodes]
async def get_inlinks(self, path: str) -> list[FileLink]:
if path not in self._nodes:
return []
return [
link for src in self._inverse.get(path, ()) for link in self._nodes[src].links if link.target_path == path
]

View file

@ -0,0 +1,450 @@
"""Neo4j-backed file graph.
Property-graph mapping:
Real node: (:File {path, st_mtime, title, description, tags,
chunk_ids, links_json, extra_json})
Virtual node: (:File {path}) — placeholder created when something
links to a path that hasn't been upserted yet.
Edge: (:File)-[:LINKS {idx, anchor, predicate}]->(:File)
The ``links_json`` property doubles as the "is real" marker — its
presence means the node was upserted with a payload; its absence
means the node exists only because some edge points at it. This
mirrors ``NxFileGraph`` exactly: ``upsert_nodes`` promotes virtuals
in place, ``delete_nodes`` demotes back to virtual (or fully removes
if nothing points here), and ``get_outlinks`` excludes edges into
virtuals so the agent never sees dangling pointers.
``path`` is the unique key (constraint enforced on ``_start``).
Frontmatter goes into flat properties; arbitrary extras land in
``extra_json``. The full ``FileLink[]`` payload is also stored as
``links_json`` so ``rebuild_links`` can rebuild the relationship
graph from per-node payloads after backend repair / migration.
Adjacency policy: trusts ``FileLink.path`` directly — no internal
wikilink resolution. The parser pipeline (with the external
resolver) produces safe links where ``link.path`` is already a
vault-relative target.
Conditional dependency: the ``neo4j`` driver loads lazily; the
import error fires at ``_start`` (boot), not at first call.
"""
from __future__ import annotations
import json
from typing import Any
from .base_file_graph import BaseFileGraph
from ..component_registry import R
from ...schema import FileLink, FileNode
from ...schema.file_node import FileFrontMatter
_TYPED_FRONTMATTER_FIELDS = {"title", "description", "tags"}
_LINK_FIELDS = {"source_path", "target_path", "target_anchor", "predicate"}
# Properties that distinguish a "real" node from a virtual placeholder.
# Listed for the demote query (delete_nodes) so we can REMOVE them all.
_REAL_PROPS = (
"st_mtime",
"title",
"description",
"tags",
"chunk_ids",
"links_json",
"extra_json",
)
@R.register("neo4j")
class Neo4jFileGraph(BaseFileGraph):
"""Neo4j-backed file graph; trusts ``FileLink.path`` for adjacency.
Connection params (constructor kwargs):
uri: bolt URL, e.g. ``bolt://localhost:7687``
user: auth user (default ``neo4j``)
password: auth password
database: target db name (default ``neo4j``)
"""
def __init__(
self,
uri: str = "bolt://localhost:7687",
user: str = "neo4j",
password: str = "neo4j",
database: str = "neo4j",
**kwargs,
):
super().__init__(**kwargs)
self._uri: str = uri
self._user: str = user
self._password: str = password
self._database: str = database
self._driver = None
# -- Lifecycle ---------------------------------------------------------
async def _start(self) -> None:
await super()._start()
try:
from neo4j import AsyncGraphDatabase
except ImportError as e:
raise ImportError(
"Neo4jFileGraph requires the neo4j driver. Install with `pip install neo4j`.",
) from e
self._driver = AsyncGraphDatabase.driver(
self._uri,
auth=(self._user, self._password),
)
async with self._session() as session:
await session.run(
"CREATE CONSTRAINT file_path_unique IF NOT EXISTS FOR (f:File) REQUIRE f.path IS UNIQUE",
)
real, virtual, edges = await self._counts(session)
self.logger.info(
f"Neo4jFileGraph '{self.graph_name}' connected at "
f"{self._uri}/{self._database}: "
f"{real} nodes, {edges} edges, {virtual} virtual",
)
async def _close(self) -> None:
if self._driver is not None:
await self._driver.close()
self._driver = None
await super()._close()
def _session(self):
assert self._driver is not None, "Neo4jFileGraph not started"
return self._driver.session(database=self._database)
@staticmethod
async def _counts(session) -> tuple[int, int, int]:
rec = await session.run(
"""
MATCH (f:File)
WITH count(CASE WHEN f.links_json IS NOT NULL THEN 1 END) AS real,
count(CASE WHEN f.links_json IS NULL THEN 1 END) AS virtual
OPTIONAL MATCH ()-[r:LINKS]->()
RETURN real, virtual, count(r) AS edges
""",
)
row = await rec.single()
if row is None:
return 0, 0, 0
return int(row["real"] or 0), int(row["virtual"] or 0), int(row["edges"] or 0)
# -- Node CRUD ---------------------------------------------------------
async def upsert_nodes(self, nodes: list[FileNode]) -> None:
"""Upsert in one tx: SET props (promotes virtual to real), drop
existing outgoing edges, re-emit edges (auto-creating virtual
nodes for unindexed targets)."""
if not nodes:
return
payload = [
{
"path": node.path,
"props": self._node_props(node),
"links": [
{
"idx": i,
"anchor": link.target_anchor,
"predicate": link.predicate,
"target": link.target_path,
}
for i, link in enumerate(node.links)
if link.target_path
],
}
for node in nodes
]
async with self._session() as session:
await session.execute_write(self._upsert_nodes_tx, payload)
@staticmethod
async def _upsert_nodes_tx(tx, payload):
# 1. Upsert node props (promotes virtual → real where necessary).
await tx.run(
"""
UNWIND $items AS n
MERGE (f:File {path: n.path})
SET f += n.props
""",
items=payload,
)
# 2. Drop existing outgoing edges from these sources.
await tx.run(
"""
UNWIND $paths AS p
MATCH (f:File {path: p})-[r:LINKS]->()
DELETE r
""",
paths=[item["path"] for item in payload],
)
# 3. Re-emit edges; MERGE on target auto-creates virtual nodes
# for unindexed targets.
await tx.run(
"""
UNWIND $items AS n
MATCH (s:File {path: n.path})
UNWIND n.links AS link
MERGE (t:File {path: link.target})
MERGE (s)-[r:LINKS {idx: link.idx}]->(t)
SET r.anchor = link.anchor, r.predicate = link.predicate
""",
items=payload,
)
async def delete_nodes(self, paths: list[str]) -> None:
"""Demote real → virtual to preserve inbound visibility; fully
remove the (now-virtual) node only if no edge points at it."""
if not paths:
return
async with self._session() as session:
await session.execute_write(self._delete_nodes_tx, list(paths))
@staticmethod
async def _delete_nodes_tx(tx, paths):
# 1. Drop outgoing edges, then strip "real" properties (demote).
# Building the REMOVE clause from _REAL_PROPS keeps the list of
# properties in one place (top of module).
remove_clause = ", ".join(f"f.{name}" for name in _REAL_PROPS)
await tx.run(
f"""
UNWIND $paths AS p
MATCH (f:File {{path: p}})
OPTIONAL MATCH (f)-[r:LINKS]->()
DELETE r
WITH DISTINCT f
REMOVE {remove_clause}
""",
paths=paths,
)
# 2. Garbage-collect: drop the virtual node entirely if nothing
# points at it anymore.
await tx.run(
"""
UNWIND $paths AS p
MATCH (f:File {path: p})
WHERE f.links_json IS NULL AND NOT (f)<-[:LINKS]-()
DELETE f
""",
paths=paths,
)
async def get_nodes(self, paths: list[str] | None = None) -> list[FileNode]:
"""Return real nodes (virtual placeholders filtered).
``paths=None`` streams every real node ordered by path. An
explicit ``[]`` returns ``[]`` without hitting the database.
"""
if paths is not None and not paths:
return []
async with self._session() as session:
if paths is None:
rec = await session.run(
"""
MATCH (f:File)
WHERE f.links_json IS NOT NULL
RETURN f
ORDER BY f.path ASC
""",
)
else:
rec = await session.run(
"""
UNWIND $paths AS p
MATCH (f:File {path: p})
WHERE f.links_json IS NOT NULL
RETURN f
""",
paths=list(paths),
)
rows = [row["f"] async for row in rec]
return [self._row_to_node(row) for row in rows]
async def rebuild_links(self) -> None:
"""Defensive full rebuild from each real node's ``links_json``.
Three steps in one tx: drop all LINKS edges; drop all virtual
nodes; re-emit edges from per-node link payloads (re-creating
virtual targets as needed). Useful after manual repair or
schema migration.
"""
async with self._session() as session:
rec = await session.run(
"""
MATCH (f:File)
WHERE f.links_json IS NOT NULL
RETURN f.path AS p, f.links_json AS l
""",
)
rows = [dict(r) async for r in rec]
payload: list[dict] = []
for row in rows:
try:
links = json.loads(row.get("l") or "[]")
except json.JSONDecodeError:
continue
items = [
{
"idx": i,
"anchor": link.get("target_anchor"),
"predicate": link.get("predicate"),
"target": link.get("target_path"),
}
for i, link in enumerate(links)
if isinstance(link, dict) and link.get("target_path")
]
payload.append({"path": row["p"], "links": items})
async with self._session() as session:
await session.execute_write(self._rebuild_links_tx, payload)
@staticmethod
async def _rebuild_links_tx(tx, payload):
# 1. Wipe all edges and all virtual nodes.
await tx.run("MATCH ()-[r:LINKS]->() DELETE r")
await tx.run("MATCH (f:File) WHERE f.links_json IS NULL DELETE f")
if not payload:
return
# 2. Re-emit edges; virtual targets reappear via MERGE.
await tx.run(
"""
UNWIND $items AS n
MATCH (s:File {path: n.path})
UNWIND n.links AS link
MERGE (t:File {path: link.target})
MERGE (s)-[r:LINKS {idx: link.idx}]->(t)
SET r.anchor = link.anchor, r.predicate = link.predicate
""",
items=payload,
)
async def clear(self):
"""Remove every node and edge in the configured database."""
async with self._session() as session:
await session.run("MATCH (f:File) DETACH DELETE f")
# -- Link access -------------------------------------------------------
async def get_outlinks(self, path: str) -> list[FileLink]:
"""Outgoing links from ``path``. Source must be real; targets
into virtual nodes are excluded so dangling refs are invisible."""
async with self._session() as session:
rec = await session.run(
"""
MATCH (s:File {path: $path})
WHERE s.links_json IS NOT NULL
MATCH (s)-[r:LINKS]->(t:File)
WHERE t.links_json IS NOT NULL
RETURN t.path AS target, r.anchor AS anchor,
r.predicate AS predicate, r.idx AS idx
ORDER BY r.idx ASC
""",
path=path,
)
rows = [dict(row) async for row in rec]
return [
FileLink(
source_path=path,
target_path=row["target"],
target_anchor=row.get("anchor"),
predicate=row.get("predicate"),
)
for row in rows
]
async def get_inlinks(self, path: str) -> list[FileLink]:
"""Incoming links to ``path`` (must be real). Sources are always
real because virtual nodes never have outgoing edges."""
async with self._session() as session:
rec = await session.run(
"""
MATCH (t:File {path: $path})
WHERE t.links_json IS NOT NULL
MATCH (s:File)-[r:LINKS]->(t)
RETURN r.anchor AS anchor, r.predicate AS predicate,
r.idx AS idx, s.path AS source
ORDER BY s.path ASC, r.idx ASC
""",
path=path,
)
rows = [dict(row) async for row in rec]
return [
FileLink(
source_path=row["source"],
target_path=path,
target_anchor=row.get("anchor"),
predicate=row.get("predicate"),
)
for row in rows
]
# -- Internal: row ↔ schema marshaling ---------------------------------
@staticmethod
def _node_props(node: FileNode) -> dict[str, Any]:
fm = node.front_matter
extras = dict(fm.__pydantic_extra__ or {})
return {
"path": node.path,
"st_mtime": float(node.st_mtime),
"title": fm.title or "",
"description": fm.description or "",
"tags": list(fm.tags or []),
"chunk_ids": list(node.chunk_ids or []),
"links_json": json.dumps(
[link.model_dump(exclude_none=True) for link in node.links],
ensure_ascii=False,
),
"extra_json": json.dumps(extras, ensure_ascii=False, sort_keys=True),
}
@staticmethod
def _row_to_node(row) -> FileNode:
d = dict(row)
try:
extras = json.loads(d.get("extra_json") or "{}")
except json.JSONDecodeError:
extras = {}
try:
links_raw = json.loads(d.get("links_json") or "[]")
except json.JSONDecodeError:
links_raw = []
links: list[FileLink] = []
for link in links_raw:
if not isinstance(link, dict):
continue
# Defensive: strip any keys the schema doesn't recognise
# (e.g. legacy fields from prior schema versions).
clean = {k: v for k, v in link.items() if k in _LINK_FIELDS}
# Ensure source_path is populated — older payloads (or
# links written before the schema split) only carry the
# target side; default to the owning node's path.
clean.setdefault("source_path", d["path"])
if not clean.get("target_path"):
continue
try:
links.append(FileLink(**clean))
except Exception:
continue
fm_kwargs: dict[str, Any] = {
"title": d.get("title", "") or "",
"description": d.get("description", "") or "",
"tags": d.get("tags") or None,
}
fm_kwargs.update(
{k: v for k, v in extras.items() if k not in _TYPED_FRONTMATTER_FIELDS},
)
return FileNode(
path=d["path"],
st_mtime=float(d.get("st_mtime", 0.0)),
links=links,
chunk_ids=[str(c) for c in (d.get("chunk_ids") or [])],
front_matter=FileFrontMatter(**fm_kwargs),
)

View file

@ -0,0 +1,122 @@
"""Networkx file-graph backend."""
import pickle
from pathlib import Path
try:
import networkx as nx
except ImportError:
nx = None
from .base_file_graph import BaseFileGraph
from ..component_registry import R
from ...schema import FileLink, FileNode
@R.register("nx")
class NxFileGraph(BaseFileGraph):
"""Networkx-backed file graph; uses FileLink.target_path for adjacency."""
def __init__(self, **kwargs):
super().__init__(**kwargs)
if nx is None:
raise ImportError("NxFileGraph requires networkx — pip install networkx")
self._graph: nx.MultiDiGraph = nx.MultiDiGraph()
self._graph_file: Path = self.graph_path / f"{self.graph_name}_{self.graph_version}.pkl"
# -- Lifecycle ---------------------------------------------------------
async def _start(self) -> None:
await super()._start()
await self.load()
async def _close(self) -> None:
await self.dump()
await super()._close()
async def load(self) -> None:
"""Load graph from pickle file; keep current graph on failure."""
if not self._graph_file.exists():
return
try:
with open(self._graph_file, "rb") as f:
self._graph = pickle.load(f)
n_real = sum(1 for _, d in self._graph.nodes(data=True) if "node" in d)
self.logger.info(f"Loaded {n_real} nodes from {self._graph_file}")
except Exception as e:
self.logger.exception(f"Failed to load {self._graph_file}: {e}")
async def dump(self) -> None:
"""Persist graph to pickle via atomic rename."""
try:
tmp = self._graph_file.with_suffix(".tmp")
with open(tmp, "wb") as f:
pickle.dump(self._graph, f, protocol=pickle.HIGHEST_PROTOCOL)
tmp.replace(self._graph_file)
n_real = sum(1 for _, d in self._graph.nodes(data=True) if "node" in d)
self.logger.info(f"Saved {n_real} nodes to {self._graph_file}")
except Exception as e:
self.logger.exception(f"Failed to write {self._graph_file}: {e}")
# -- Node CRUD ---------------------------------------------------------
async def upsert_nodes(self, nodes: list[FileNode]) -> None:
for node in nodes:
path = node.path
if self._graph.has_node(path):
# Drop outgoing edges; inbound stay intact.
self._graph.remove_edges_from(list(self._graph.out_edges(path, keys=True)))
self._graph.add_node(path, node=node) # promotes virtual node if present
# Missing targets become attr-less virtual nodes.
self._graph.add_edges_from((path, lnk.target_path, {"link": lnk}) for lnk in node.links if lnk.target_path)
async def delete_nodes(self, paths: list[str]) -> None:
for path in paths:
if not self._graph.has_node(path):
continue
self._graph.remove_edges_from(list(self._graph.out_edges(path, keys=True)))
# Demote to virtual: keep inbound edges, drop node payload.
self._graph.nodes[path].pop("node", None)
if self._graph.in_degree(path) == 0:
self._graph.remove_node(path) # remove orphan virtual node
async def get_nodes(self, paths: list[str] | None = None) -> list[FileNode]:
nodes_view = self._graph.nodes
if paths is None:
return [d["node"] for _, d in nodes_view(data=True) if "node" in d]
return [nodes_view[path]["node"] for path in paths if path in nodes_view and "node" in nodes_view[path]]
async def rebuild_links(self) -> None:
"""Rebuild all edges from real node payloads; drop virtual nodes."""
self._graph.remove_edges_from(list(self._graph.edges(keys=True)))
virtual = [n for n, d in self._graph.nodes(data=True) if "node" not in d]
self._graph.remove_nodes_from(virtual)
self._graph.add_edges_from(
(path, lnk.target_path, {"link": lnk})
for path, data in self._graph.nodes(data=True)
for lnk in data["node"].links
if lnk.target_path
)
async def clear(self):
"""Remove all nodes and edges, and remove persisted file."""
self._graph.clear()
self._graph_file.unlink(missing_ok=True)
# -- Link access -------------------------------------------------------
async def get_outlinks(self, path: str) -> list[FileLink]:
nodes_view = self._graph.nodes
if path not in nodes_view or "node" not in nodes_view[path]:
return []
return [
d["link"]
for _, target, d in self._graph.out_edges(path, data=True)
if "link" in d and "node" in nodes_view[target]
]
async def get_inlinks(self, path: str) -> list[FileLink]:
nodes_view = self._graph.nodes
if path not in nodes_view or "node" not in nodes_view[path]:
return []
return [d["link"] for _, _, d in self._graph.in_edges(path, data=True) if "link" in d]

View file

@ -0,0 +1,8 @@
"""File parser components."""
from .bare_file_parser import BareFileParser
from .base_file_parser import BaseFileParser
from .default_file_parser import DefaultFileParser
from .linked_file_parser import LinkedFileParser
__all__ = ["BareFileParser", "BaseFileParser", "DefaultFileParser", "LinkedFileParser"]

View file

@ -0,0 +1,22 @@
"""Stat-only parser for attachment/binary files."""
from pathlib import Path
from .base_file_parser import BaseFileParser
from ..component_registry import R
from ...schema import FileChunk, FileNode
@R.register("bare")
class BareFileParser(BaseFileParser):
"""Stat-only parser for attachment/binary files.
No content read, no chunking, no link extraction. The resulting FileNode
has empty links and chunk_ids; front_matter carries mime and size so
retrieval can filter by file type without reopening the file.
"""
async def parse(self, path: str | Path) -> tuple[FileNode, list[FileChunk]]:
file_path = Path(path)
stat = file_path.stat()
return FileNode(path=self._get_relative_path(path), st_mtime=stat.st_mtime, links=[], chunk_ids=[]), []

View file

@ -0,0 +1,30 @@
"""Abstract base for file parsers."""
from abc import abstractmethod
from pathlib import Path
from ..base_component import BaseComponent
from ...enumeration import ComponentEnum
from ...schema import FileChunk, FileNode
class BaseFileParser(BaseComponent):
"""Abstract base for file parsers. Subclasses implement `parse`."""
component_type = ComponentEnum.FILE_PARSER
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.working_dir = self.app_context.app_config.working_dir if self.app_context else ""
def _get_relative_path(self, path: str | Path) -> str:
"""Return path relative to working_dir, or absolute path if outside."""
file_path = Path(path).absolute()
try:
return str(file_path.relative_to(Path(self.working_dir).absolute()))
except ValueError:
return str(file_path)
@abstractmethod
async def parse(self, path: str | Path) -> tuple[FileNode, list[FileChunk]]:
"""Parse a file into (node, chunks)."""

View file

@ -0,0 +1,123 @@
"""Default file parser with byte-based overlapping chunking."""
import re
from bisect import bisect_right
from pathlib import Path
import aiofiles
import yaml
from .base_file_parser import BaseFileParser
from ..component_registry import R
from ...schema import FileChunk, FileFrontMatter, FileLink, FileNode
# Single-pass wikilink + optional Dataview predicate.
# Covers: [[X]] / [[X#h]] / [[X|alias]] / pred:: [[X]] / [pred:: [[X]]]
# - predicate group: optional leading '[' (Dataview inline-bracket form), an identifier,
# then '::' — the whole prefix is non-capturing-optional so bare wikilinks still match.
# - target / anchor: target stops before '#', '|', '[', ']'; anchor stops before '|', '[', ']'.
# - alias '|...': consumed but not captured (we don't need display text).
_LINK_RE = re.compile(
r"(?:\[?\s*(?P<predicate>[A-Za-z][\w-]*)\s*::\s*)?"
r"\[\[\s*(?P<target>[^\[\]|#]+?)"
r"(?:#(?P<anchor>[^\[\]|]+?))?"
r"\s*(?:\|[^\[\]]*?)?\s*\]\]",
)
@R.register("default")
class DefaultFileParser(BaseFileParser):
"""Parser that splits files into byte-based overlapping chunks."""
def __init__(self, encoding: str = "utf-8", chunk_byte_size: int = 10000, overlap_byte_size: int = 100, **kwargs):
super().__init__(**kwargs)
self.encoding = encoding
self.chunk_byte_size = max(100, chunk_byte_size)
self.overlap_byte_size = max(4, overlap_byte_size)
@staticmethod
def parse_links(content: str, source_path: str) -> list[FileLink]:
"""Extract wikilinks with optional Dataview predicate as outgoing FileLinks."""
links: list[FileLink] = []
for m in _LINK_RE.finditer(content):
target = m["target"].strip()
if not target:
continue
anchor = m["anchor"]
links.append(
FileLink(
source_path=source_path,
target_path=target,
target_anchor=anchor.strip() if anchor else None,
predicate=m["predicate"],
),
)
return links
@staticmethod
def _parse_front_matter(text: str) -> tuple[FileFrontMatter, str]:
"""Parse YAML front matter delimited by ---, return (front_matter, remaining)."""
if not text.startswith("---"):
return FileFrontMatter(), text
end_idx = text.find("\n---", 3)
if end_idx == -1:
return FileFrontMatter(), text
try:
data = yaml.safe_load(text[3:end_idx].strip()) or {}
front_matter = FileFrontMatter(**(data if isinstance(data, dict) else {}))
except yaml.YAMLError:
front_matter = FileFrontMatter()
return front_matter, text[end_idx + 4 :].lstrip("\n")
async def parse(self, path: str | Path) -> tuple[FileNode, list[FileChunk]]:
file_path = Path(path)
stat = file_path.stat()
rel_path = self._get_relative_path(path)
async with aiofiles.open(file_path, encoding=self.encoding) as f:
text = await f.read()
if not text:
return FileNode(path=rel_path, st_mtime=stat.st_mtime), []
front_matter, content = self._parse_front_matter(text)
if not content:
return FileNode(path=rel_path, st_mtime=stat.st_mtime, front_matter=front_matter), []
links = self.parse_links(content, rel_path)
chunks = self._chunk_content(content, rel_path)
chunk_ids = [c.id for c in chunks]
return (
FileNode(
path=rel_path,
st_mtime=stat.st_mtime,
front_matter=front_matter,
links=links,
chunk_ids=chunk_ids,
),
chunks,
)
def _chunk_content(self, content: str, rel_path: str) -> list[FileChunk]:
"""Split content into overlapping byte-range chunks with line numbers."""
content_bytes = content.encode(self.encoding)
newline_positions = [i for i, b in enumerate(content_bytes) if b == ord("\n")]
chunks: list[FileChunk] = []
step = self.chunk_byte_size - self.overlap_byte_size
start = 0
while start < len(content_bytes):
end = min(start + self.chunk_byte_size, len(content_bytes))
chunk_text = content_bytes[start:end].decode(self.encoding, errors="ignore")
start_line = bisect_right(newline_positions, start - 1) + 1
end_line = bisect_right(newline_positions, end - 1) + 1
if content_bytes[end - 1] == ord("\n"):
end_line -= 1
chunks.append(
FileChunk(path=rel_path, start_line=start_line, end_line=end_line, text=chunk_text).set_hash_id(),
)
if end >= len(content_bytes):
break
start += step
return chunks

View file

@ -0,0 +1,729 @@
"""Markdown file parser — frontmatter + wikilink graph + AST tree chunks.
Each chunk carries the **complete heading skeleton** of the document
with its content inlined under the section that owns it; other sections
appear as bare headings so the reader always sees a full document map.
Pipeline: build mistletoe AST → ``MdNode`` tree (sections nest by
heading level) → recursive chunk (try whole subtree; on overflow walk
children — body siblings pack as a run, subsections recurse). Leaf
blocks (table / code / list / paragraph) split on internal boundaries
and each piece is annotated ``[Part X/N]``. Wikilinks in the body are
extracted as graph edges, with optional Dataview-style typed predicates
(line-level ``predicate:: [[X]]`` or inline-bracketed ``[predicate:: [[X]]]``).
"""
from __future__ import annotations
import re
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
import frontmatter
from .base_file_parser import BaseFileParser
from ..component_registry import R
from ..file_graph import BaseFileGraph
from ...enumeration import ComponentEnum
from ...schema import (
FileChunk,
FileLink,
FileFrontMatter,
FileNode,
)
# -- Wikilink resolution --------------------------------------------------
#
# Wikilinks are a markdown user-facing convention: ``[[Alice]]`` should
# resolve to ``topics/Alice/Alice.md`` (or wherever the file lives).
# This short-form / implicit-``.md`` / folder-note resolution lives here
# at the markdown boundary rather than as a generic utility — file-IO
# steps require full vault-relative paths and never use these helpers.
def _complete_md(target: str) -> str:
"""Apply implicit ``.md`` rule for wikilink targets."""
if not target:
return target
last = target.rsplit("/", 1)[-1]
return target if "." in last else target + ".md"
def _filter_folder_note(target: str, paths: list[str]) -> list[str]:
"""Apply folder-note rule: when both ``X.md`` and ``X/X.md`` exist,
prefer ``X/X.md``. Sorted for determinism.
"""
if not paths:
return []
stem = Path(target).stem
folder_hits = sorted(p for p in paths if Path(p).parent.name == stem)
return folder_hits or sorted(paths)
async def _resolve_wikilink(graph: BaseFileGraph, target: str) -> list[str]:
"""Resolve a wikilink target to vault-relative path(s).
Returns:
``[path]`` for an unambiguous match,
``[path, path, ...]`` for short-form ambiguity (caller may
fan out one FileLink per candidate), or
``[]`` when nothing matches (dangling — caller drops the link).
"""
if not target:
return []
target = _complete_md(target)
if "/" in target:
nodes = await graph.get_nodes([target])
return [target] if nodes else []
matches = [n.path for n in await graph.get_nodes() if Path(n.path).name == target]
return _filter_folder_note(target, matches)
# -- Wikilink extraction --------------------------------------------------
_WIKILINK_RE = re.compile(
r"""
(?:!)?
\[\[
(?P<target>[^\]\|\#\n]+?)
(?:\#(?P<anchor>[^\]\|\n]+))?
(?:\|[^\]\n]+)?
\]\]
""",
re.VERBOSE,
)
_DATAVIEW_LINE_RE = re.compile(
r"^[ \t]*(?:[-*+][ \t]+)?(?P<predicate>[A-Za-z][A-Za-z0-9_]*)\s*::\s*(?P<value>.+?)\s*$",
re.MULTILINE,
)
_INLINE_FIELD_OPEN_RE = re.compile(r"\[(?P<predicate>[A-Za-z][A-Za-z0-9_]*)\s*::\s*")
def _iter_inline_fields(text: str) -> list[tuple[int, int, str]]:
"""Find inline-bracketed ``[predicate:: …]`` field spans by depth scan."""
out: list[tuple[int, int, str]] = []
for m in _INLINE_FIELD_OPEN_RE.finditer(text):
depth = 1
i = m.end()
n = len(text)
while i < n:
c = text[i]
if c == "\n":
break
if c == "[":
depth += 1
elif c == "]":
depth -= 1
if depth == 0:
out.append((m.start(), i + 1, m.group("predicate")))
break
i += 1
return out
def _predicate_for(
text: str,
pos: int,
inline_spans: list[tuple[int, int, str]],
) -> str | None:
"""Resolve the predicate governing a wikilink at offset ``pos``.
Precedence: inline-bracketed > line-level Dataview > none.
"""
for field_start, field_end, predicate in inline_spans:
if field_start <= pos < field_end:
return predicate
line_start = text.rfind("\n", 0, pos) + 1
line_end = text.find("\n", pos)
if line_end == -1:
line_end = len(text)
m = _DATAVIEW_LINE_RE.match(text[line_start:line_end])
if m and line_start + m.start("value") <= pos:
return m.group("predicate")
return None
async def _extract_links(
graph: BaseFileGraph,
text: str,
source_path: str,
) -> list[FileLink]:
"""Find every wikilink in ``text``, resolve targets, emit FileLinks.
Short-path ambiguity **expands** into one FileLink per candidate so
the body's wikilink is recorded against every plausible target.
Dangling targets are dropped. Results are deduped by
``(target_path, predicate, target_anchor)`` preserving order.
"""
if not text:
return []
inline_spans = _iter_inline_fields(text)
out: list[FileLink] = []
seen: set[tuple] = set()
for wm in _WIKILINK_RE.finditer(text):
target = wm.group("target").strip()
if not target:
continue
anchor_raw = wm.group("anchor")
anchor = anchor_raw.strip() if anchor_raw else None
predicate = _predicate_for(text, wm.start(), inline_spans)
resolved_paths = await _resolve_wikilink(graph, target)
if not resolved_paths:
continue
for resolved in resolved_paths:
key = (resolved, predicate, anchor)
if key in seen:
continue
seen.add(key)
out.append(
FileLink(
source_path=source_path,
target_path=resolved,
target_anchor=anchor,
predicate=predicate,
),
)
return out
# -- AST node + helpers ---------------------------------------------------
@dataclass
class MdNode:
"""``root`` / ``section`` (heading + children until equal-or-shallower
heading) / ``body`` (one mistletoe block; ``block`` keeps the original).
``text`` is the rendered subtree (own heading excluded for sections).
``desc_toc`` caches the section-only DFS outline of descendants
(own heading excluded), used as the TOC suffix when emitting chunks
inside a section. Line ranges span the full subtree.
"""
kind: str # "root" | "section" | "body"
heading: str | None = None
level: int = 0
children: list["MdNode"] = field(default_factory=list)
block: Any = None
text: str = ""
start_line: int = 0
end_line: int = 0
desc_toc: str = ""
def _heading_text(node: Any, renderer) -> str:
"""Heading text without `#` markers (for outline)."""
rendered = renderer.render(node).rstrip("\n")
if rendered.startswith("#"):
return rendered.lstrip("#").strip()
return rendered.split("\n", 1)[0].strip()
def _finalize(n: MdNode) -> None:
"""Bottom-up pass: propagate line ranges, populate ``n.text`` (rendered
subtree, own heading excluded for sections) and ``n.desc_toc`` (DFS
section outline of descendants)."""
parts: list[str] = []
desc_lines: list[str] = []
for c in n.children:
_finalize(c)
if c.kind == "section":
heading = f"{'#' * c.level} {c.heading or ''}"
parts.append(f"{heading}\n\n{c.text}" if c.text else heading)
desc_lines.append(f"{heading}\n\n{c.desc_toc}" if c.desc_toc else heading)
elif c.text:
parts.append(c.text)
if n.children:
first = n.children[0].start_line
n.start_line = min(n.start_line, first) if n.start_line else first
n.end_line = max(c.end_line for c in n.children)
elif n.end_line < n.start_line:
n.end_line = n.start_line
if n.kind != "body":
n.text = "\n\n".join(parts)
n.desc_toc = "\n\n".join(desc_lines)
def _toc_join(*parts: str) -> str:
"""Concatenate TOC fragments with ``\\n\\n``, skipping empty ones."""
return "\n\n".join(p for p in parts if p)
def _subtree_toc(n: MdNode) -> str:
"""Section's heading + descendants TOC — its contribution to a parent's
``desc_toc``. For root (no own heading) this is just ``desc_toc``."""
if n.kind != "section" or n.heading is None:
return n.desc_toc
heading = f"{'#' * n.level} {n.heading}"
return f"{heading}\n\n{n.desc_toc}" if n.desc_toc else heading
# -- Parser ---------------------------------------------------------------
@R.register("md")
class LinkedFileParser(BaseFileParser):
"""Markdown parser: frontmatter + wikilink edges + full-skeleton chunks."""
def __init__(
self,
encoding: str = "utf-8",
chunk_chars: int = 2000,
embed_toc: bool = True,
file_graph: str = "default",
**kwargs,
):
super().__init__(**kwargs)
self.encoding = encoding
self.chunk_chars = max(100, chunk_chars)
self.embed_toc = embed_toc
self._file_graph_name: str = file_graph
def _resolve_file_graph(self) -> BaseFileGraph | None:
"""Lazily fetch the configured file_graph from app_context.
Lazy (rather than ``_start``) so the parser doesn't impose a
component start-order constraint, and so tests can construct
the parser without a graph wired up.
"""
if self.app_context is None:
return None
graphs = self.app_context.components.get(ComponentEnum.FILE_GRAPH, {})
graph = graphs.get(self._file_graph_name)
if graph is None:
return None
if not isinstance(graph, BaseFileGraph):
raise TypeError(
f"Expected BaseFileGraph, got {type(graph).__name__}",
)
return graph
async def parse(self, path: str | Path) -> tuple[FileNode, list[FileChunk]]:
from mistletoe.markdown_renderer import MarkdownRenderer
from mistletoe.block_token import Document
file_path = Path(path)
rel_path = self._get_relative_path(path)
post = frontmatter.loads(file_path.read_text(encoding=self.encoding))
chunks: list[FileChunk] = []
if post.content and post.content.strip():
with MarkdownRenderer() as renderer:
tree = self._build_tree(Document(post.content), renderer)
chunks = self._chunk_node(tree, "", "", rel_path, renderer)
links: list[FileLink] = []
graph = self._resolve_file_graph()
if graph is not None:
links = await _extract_links(graph, post.content, rel_path)
node = FileNode(
path=rel_path,
st_mtime=file_path.stat().st_mtime,
chunk_ids=[chunk.id for chunk in chunks],
links=links,
front_matter=FileFrontMatter(**dict(post.metadata)),
)
return node, chunks
def _build_tree(self, doc: Any, renderer) -> MdNode:
"""Heading-level stack folds mistletoe's flat children into nested
sections; non-headings attach as ``body`` to the current section
(or root before the first heading)."""
from mistletoe.markdown_renderer import BlankLine
from mistletoe.block_token import (
Heading,
SetextHeading,
)
root = MdNode(kind="root", start_line=1, end_line=1)
stack: list[MdNode] = [root]
for child in doc.children or []:
if isinstance(child, BlankLine):
continue
line = getattr(child, "line_number", None) or stack[-1].start_line
if isinstance(child, (Heading, SetextHeading)):
level = max(1, getattr(child, "level", 1))
while len(stack) > 1 and stack[-1].level >= level:
stack.pop()
sec = MdNode(
kind="section",
heading=_heading_text(child, renderer),
level=level,
start_line=line,
)
stack[-1].children.append(sec)
stack.append(sec)
continue
rendered = renderer.render(child).rstrip("\n")
if not rendered:
continue
stack[-1].children.append(
MdNode(
kind="body",
block=child,
text=rendered,
start_line=line,
end_line=line + rendered.count("\n"),
),
)
_finalize(root)
return root
# -- Recursive chunker ------------------------------------------------
def _chunk_node(
self,
node: MdNode,
before: str,
after: str,
path: str,
renderer,
) -> list[FileChunk]:
"""Try the whole subtree; on overflow split (leaf) or descend.
``before``/``after`` are TOC fragments that bracket each emitted
chunk's content (chunk text = ``before + content + after``).
As we descend, the prefix grows with section headings already
passed and the suffix shrinks correspondingly.
"""
if not node.text:
return []
if node.kind == "section":
heading_line = f"{'#' * node.level} {node.heading or ''}"
before_self = _toc_join(before, heading_line)
else:
before_self = before
if len(node.text) <= self.chunk_chars:
return [
self._make_chunk(
before_self,
node.text,
after,
node.start_line,
node.end_line,
path,
),
]
if node.kind == "body":
return self._split_leaf(node, before, after, path, renderer)
after_inside = _toc_join(node.desc_toc, after)
sub_tocs = [_subtree_toc(c) for c in node.children if c.kind == "section"]
chunks: list[FileChunk] = []
accumulated = before_self
sec_idx = 0
run: list[MdNode] = []
for c in node.children:
if c.kind == "section":
if run:
chunks.extend(
self._chunk_body_run(
run,
before_self,
after_inside,
path,
renderer,
),
)
run = []
remaining = "\n\n".join(sub_tocs[sec_idx + 1 :])
chunks.extend(
self._chunk_node(
c,
accumulated,
_toc_join(remaining, after),
path,
renderer,
),
)
accumulated = _toc_join(accumulated, sub_tocs[sec_idx])
sec_idx += 1
else:
run.append(c)
if run:
chunks.extend(
self._chunk_body_run(
run,
before_self,
after_inside,
path,
renderer,
),
)
return chunks
def _chunk_body_run(
self,
run: list[MdNode],
before: str,
after: str,
path: str,
renderer,
) -> list[FileChunk]:
"""Greedy-pack consecutive body siblings under the same TOC slot.
No ``[Part X/N]`` markers — distinct blocks, not a leaf split.
Oversized single body recurses to ``_split_leaf``."""
composite_size = sum(len(b.text) for b in run) + 2 * max(0, len(run) - 1)
if composite_size <= self.chunk_chars:
return [
self._make_chunk(
before,
"\n\n".join(b.text for b in run),
after,
run[0].start_line,
run[-1].end_line,
path,
),
]
chunks: list[FileChunk] = []
bucket: list[MdNode] = []
bucket_chars = 0
def flush() -> None:
nonlocal bucket, bucket_chars
if not bucket:
return
chunks.append(
self._make_chunk(
before,
"\n\n".join(b.text for b in bucket),
after,
bucket[0].start_line,
bucket[-1].end_line,
path,
),
)
bucket = []
bucket_chars = 0
for body in run:
if len(body.text) > self.chunk_chars:
flush()
chunks.extend(self._split_leaf(body, before, after, path, renderer))
continue
sep = 2 if bucket else 0
if bucket_chars + sep + len(body.text) > self.chunk_chars:
flush()
sep = 0
bucket.append(body)
bucket_chars += sep + len(body.text)
flush()
return chunks
# -- Leaf splitters: build (text, start, end) units, hand off to packer
def _split_leaf(
self,
body: MdNode,
before: str,
after: str,
path: str,
renderer,
) -> list[FileChunk]:
from mistletoe.block_token import (
CodeFence,
List,
Table,
)
block = body.block
if isinstance(block, Table):
return self._split_table(body, before, after, path)
if isinstance(block, CodeFence):
return self._split_code(body, before, after, path)
if isinstance(block, List):
return self._split_list(body, before, after, path, renderer)
return self._split_lines(body, before, after, path)
def _split_table(
self,
body: MdNode,
before: str,
after: str,
path: str,
) -> list[FileChunk]:
"""Repeat header + separator on every chunk."""
from mistletoe.block_token import TableRow
lines = body.text.split("\n")
header, data = "\n".join(lines[:2]), lines[2:]
rows = [r for r in (body.block.children or []) if isinstance(r, TableRow)]
base = body.start_line + 2
def line_of(i: int) -> int:
return rows[i].line_number if i < len(rows) and rows[i].line_number else base + i
units = [(text, line_of(i), line_of(i)) for i, text in enumerate(data)]
return self._emit_packed(
units,
before,
after,
path,
joiner="\n",
wrap=f"{header}\n{{inner}}",
)
def _split_code(
self,
body: MdNode,
before: str,
after: str,
path: str,
) -> list[FileChunk]:
"""Repeat fence opener + closer on every chunk."""
code = body.block
indent = " " * (code.indentation or 0)
fence = f"{indent}{code.delimiter}"
opener = f"{fence}{code.info_string or ''}"
raw = (code.children[0].content if code.children else "").rstrip("\n")
if not raw:
return []
start = body.start_line + 1
units = [(indent + ln, start + i, start + i) for i, ln in enumerate(raw.split("\n"))]
return self._emit_packed(
units,
before,
after,
path,
joiner="\n",
wrap=f"{opener}\n{{inner}}\n{fence}",
allow_empty=True,
)
def _split_list(
self,
body: MdNode,
before: str,
after: str,
path: str,
renderer,
) -> list[FileChunk]:
"""Pack list items; oversized items emit alone (overflow accepted)."""
from mistletoe.block_token import ListItem
items = [c for c in (body.block.children or []) if isinstance(c, ListItem)]
if not items:
return self._split_lines(body, before, after, path)
units: list[tuple[str, int, int]] = []
for it in items:
text = renderer.render(it).rstrip("\n")
if not text:
continue
line = it.line_number or body.start_line
units.append((text, line, line + text.count("\n")))
return self._emit_packed(
units,
before,
after,
path,
joiner="\n",
wrap="{inner}",
)
def _split_lines(
self,
body: MdNode,
before: str,
after: str,
path: str,
) -> list[FileChunk]:
"""Last-resort line-greedy split for paragraphs / quotes / html."""
start = body.start_line
units = [(line, start + i, start + i) for i, line in enumerate(body.text.split("\n"))]
return self._emit_packed(
units,
before,
after,
path,
joiner="\n",
wrap="{inner}",
)
def _emit_packed(
self,
units: list[tuple[str, int, int]],
before: str,
after: str,
path: str,
joiner: str,
wrap: str,
allow_empty: bool = False,
) -> list[FileChunk]:
"""Greedy-pack units into ``wrap`` envelopes; emit each piece.
Envelope (table header, code fence) counts against ``chunk_chars``;
TOC (when on) is additive prefix/suffix downstream. Oversized
units overflow rather than truncate. Multi-piece outputs get
``[Part X/N]`` markers; single pieces don't.
"""
envelope = len(wrap.replace("{inner}", ""))
budget = max(64, self.chunk_chars - envelope)
sep_len = len(joiner)
parts: list[tuple[str, int, int]] = []
bucket: list[tuple[str, int, int]] = []
bucket_chars = 0
def flush() -> None:
nonlocal bucket, bucket_chars
if not bucket:
return
inner = joiner.join(t for t, _, _ in bucket)
parts.append((inner, bucket[0][1], bucket[-1][2]))
bucket = []
bucket_chars = 0
for text, s, e in units:
if not text and not allow_empty:
continue
sep = sep_len if bucket else 0
if bucket_chars + sep + len(text) > budget:
flush()
sep = 0
bucket.append((text, s, e))
bucket_chars += sep + len(text)
flush()
total = len(parts)
return [
self._make_chunk(
before,
(
f"[Part {idx}/{total}]\n\n{wrap.replace('{inner}', inner)}"
if total > 1
else wrap.replace("{inner}", inner)
),
after,
s,
e,
path,
)
for idx, (inner, s, e) in enumerate(parts, 1)
]
# -- Emit -------------------------------------------------------------
def _make_chunk(
self,
before: str,
content: str,
after: str,
start_line: int,
end_line: int,
path: str,
) -> FileChunk:
"""Build one ``FileChunk`` — text is ``before + content + after``
when ``embed_toc``, otherwise just ``content``."""
text = _toc_join(before, content, after) if self.embed_toc else content
return FileChunk(
path=path,
start_line=start_line,
end_line=end_line,
text=text,
).set_hash_id()

View file

@ -0,0 +1,14 @@
"""File store module.
In-memory + JSONL backend for the (file → chunks) graph. Subclass
`BaseFileStore` to add other backends; only `LocalFileStore` is
shipped today.
"""
from .base_file_store import BaseFileStore
from .local_file_store import LocalFileStore
__all__ = [
"BaseFileStore",
"LocalFileStore",
]

View file

@ -0,0 +1,100 @@
"""Abstract base for file store backends."""
from abc import abstractmethod
from ..base_component import BaseComponent
from ..embedding import BaseEmbeddingModel
from ..file_graph import BaseFileGraph
from ..keyword_index import BaseKeywordIndex
from ...enumeration import ComponentEnum
from ...schema import FileChunk, FileNode, FileLink
class BaseFileStore(BaseComponent):
"""Abstract base for file store backends."""
component_type = ComponentEnum.FILE_STORE
def __init__(
self,
store_name: str,
embedding_model: str = "default",
keyword_index: str = "default",
file_graph: str = "default",
store_version: str = "v1",
**kwargs,
):
super().__init__(**kwargs)
from ..embedding import OpenAIEmbeddingModel
from ..file_graph import LocalFileGraph
from ..keyword_index import BM25Index
self.store_name = store_name or self.name
self.store_version = store_version
if not embedding_model and not keyword_index:
raise ValueError("At least one of embedding_model or keyword_index must be set.")
self.embedding_model = self.bind(embedding_model, BaseEmbeddingModel, default_factory=OpenAIEmbeddingModel)
self.keyword_index = self.bind(keyword_index, BaseKeywordIndex, default_factory=BM25Index)
self.file_graph = self.bind(file_graph, BaseFileGraph, default_factory=LocalFileGraph)
self.store_path = self.working_metadata_path / self.component_type.value / store_name
self.store_path.mkdir(parents=True, exist_ok=True)
async def _start(self) -> None:
"""Probe embedding model; disable vector capability if it fails."""
if self.embedding_model is None:
return
if not await self.embedding_model.health_check():
self.logger.warning(f"{self.store_name}: embedding unhealthy, vector disabled")
self.embedding_model = None
def _disable_embedding(self, reason: str) -> None:
"""Drop embedding after a runtime failure; keyword search still works."""
if self.embedding_model is None:
return
self.logger.error(f"{self.store_name}: embedding disabled, {reason}")
self.embedding_model = None
async def upsert_file(
self,
file: tuple[FileNode, list[FileChunk]] | list[tuple[FileNode, list[FileChunk]]],
) -> None:
"""Upsert a file and its chunks into the store."""
async def delete_by_path(self, path: str | list[str]) -> None:
"""Delete files by their paths from the store."""
async def clear(self):
"""Clear the store of all files and chunks."""
@abstractmethod
async def vector_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]:
"""Perform vector similarity search."""
@abstractmethod
async def keyword_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]:
"""Perform full-text keyword search."""
async def rebuild_links(self) -> None:
"""Rebuild all edges from each node's link payload."""
if not self.file_graph:
raise RuntimeError("file_graph is required for delete_by_path")
return await self.file_graph.rebuild_links()
async def get_nodes(self, paths: list[str]) -> list[FileNode]:
"""Return file nodes for the given paths (missing paths are skipped)."""
if not self.file_graph:
raise RuntimeError("file_graph is required for get_nodes")
return await self.file_graph.get_nodes(paths)
async def get_outlinks(self, path: str) -> list[FileLink]:
"""Return outgoing links for *path*."""
if not self.file_graph:
raise RuntimeError("file_graph is required for delete_by_path")
return await self.file_graph.get_outlinks(path)
async def get_inlinks(self, path: str) -> list[FileLink]:
"""Return incoming links for *path*."""
if not self.file_graph:
raise RuntimeError("file_graph is required for delete_by_path")
return await self.file_graph.get_inlinks(path)

View file

@ -0,0 +1,177 @@
"""In-memory file store with JSONL persistence on close."""
import aiofiles
import numpy as np
from .base_file_store import BaseFileStore
from ..component_registry import R
from ...schema import FileChunk, FileNode
from ...utils import batch_cosine_similarity
@R.register("local")
class LocalFileStore(BaseFileStore):
"""In-memory file store with deferred JSONL persistence."""
def __init__(self, encoding: str = "utf-8", **kwargs):
super().__init__(**kwargs)
self.encoding = encoding
self.file_chunks: dict[str, FileChunk] = {}
self.chunks_path = self.store_path / f"file_chunks_{self.store_version}.jsonl"
# Lifecycle
async def _start(self) -> None:
await super()._start()
await self.load()
async def _close(self) -> None:
await self.dump()
self.file_chunks.clear()
await super()._close()
async def load(self) -> None:
"""Load chunks from JSONL file into memory."""
if not self.chunks_path.exists():
return
try:
async with aiofiles.open(self.chunks_path, encoding=self.encoding) as f:
async for line in f:
line = line.strip()
if line:
chunk = FileChunk.model_validate_json(line)
self.file_chunks[chunk.id] = chunk
self.logger.info(f"Loaded {len(self.file_chunks)} chunks from {self.chunks_path}")
except Exception as e:
self.logger.exception(f"Failed to load {self.chunks_path}: {e}")
async def dump(self) -> None:
"""Persist chunks to JSONL via atomic rename, then cascade to keyword_index and file_graph."""
try:
tmp = self.chunks_path.with_suffix(".tmp")
async with aiofiles.open(tmp, "w", encoding=self.encoding) as f:
await f.write("\n".join(c.model_dump_json() for c in self.file_chunks.values()))
tmp.replace(self.chunks_path)
self.logger.info(f"Saved {len(self.file_chunks)} chunks to {self.chunks_path}")
except Exception as e:
self.logger.exception(f"Failed to write {self.chunks_path}: {e}")
if self.keyword_index:
await self.keyword_index.dump()
if self.file_graph:
await self.file_graph.dump()
# Base class interface
async def upsert_file(
self,
file: tuple[FileNode, list[FileChunk]] | list[tuple[FileNode, list[FileChunk]]],
) -> None:
if not self.file_graph:
raise RuntimeError("file_graph is required for upsert_file")
if isinstance(file, tuple):
file = [file]
old_map = {n.path: n for n in await self.file_graph.get_nodes([node.path for node, _ in file])}
new_nodes: list[FileNode] = []
needs_embed: list[FileChunk] = []
keyword_docs: dict[str, str] = {}
for node, chunks in file:
old_node: FileNode | None = old_map.get(node.path)
cached = {}
if old_node and self.embedding_model:
for cid in old_node.chunk_ids:
old = self.file_chunks.pop(cid, None)
if old and old.embedding is not None:
cached[cid] = old.embedding
node.chunk_ids = []
for c in chunks:
if self.embedding_model and c.embedding is None:
if c.id in cached:
c.embedding = cached[c.id]
elif c.text:
needs_embed.append(c)
node.chunk_ids.append(c.id)
self.file_chunks[c.id] = c
if c.text:
keyword_docs[c.id] = c.text
new_nodes.append(node)
await self.file_graph.upsert_nodes(new_nodes)
if needs_embed and self.embedding_model:
try:
await self.embedding_model.get_node_embeddings(needs_embed)
except Exception as e:
self._disable_embedding(f"upsert: {type(e).__name__}: {e}")
if self.keyword_index and keyword_docs:
await self.keyword_index.add_docs(keyword_docs)
async def delete_by_path(self, path: str | list[str]) -> None:
if not self.file_graph:
raise RuntimeError("file_graph is required for delete_by_path")
if isinstance(path, str):
path = [path]
nodes = await self.file_graph.get_nodes(path)
if not nodes:
return
deleted_chunk_ids = [cid for n in nodes for cid in n.chunk_ids]
for cid in deleted_chunk_ids:
self.file_chunks.pop(cid, None)
await self.file_graph.delete_nodes([n.path for n in nodes])
if self.keyword_index and deleted_chunk_ids:
await self.keyword_index.delete_docs(deleted_chunk_ids)
async def clear(self) -> None:
if not self.file_graph:
raise RuntimeError("file_graph is required for clear")
self.file_chunks.clear()
self.chunks_path.unlink(missing_ok=True)
if self.keyword_index:
await self.keyword_index.clear()
await self.file_graph.clear()
# Search
async def vector_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]:
if self.embedding_model is None or not query:
return []
try:
query_embedding = await self.embedding_model.get_embedding(query)
except Exception as e:
self._disable_embedding(f"search: {type(e).__name__}: {e}")
return []
if query_embedding is None:
return []
candidates = [c for c in self.file_chunks.values() if c.embedding is not None]
if not candidates:
return []
candidate_embeddings = np.stack([c.embedding for c in candidates])
similarities = batch_cosine_similarity(query_embedding.reshape(1, -1), candidate_embeddings)[0]
results = [
c.model_copy(update={"scores": {"vector": float(s), "score": float(s)}})
for c, s in zip(candidates, similarities)
]
results.sort(key=lambda r: r.score, reverse=True)
return results[:limit]
async def keyword_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]:
if not self.keyword_index:
return []
query = query.strip()
if not query:
return []
doc_id_score_dict = await self.keyword_index.retrieve(query, limit=limit)
results = []
for doc_id, score in doc_id_score_dict.items():
chunk = self.file_chunks.get(doc_id)
if chunk:
results.append(chunk.model_copy(update={"scores": {"keyword": score, "score": score}}))
return results

View file

@ -0,0 +1,9 @@
"""File watcher implementations for monitoring file system changes."""
from .base_file_watcher import BaseFileWatcher
from .lite_file_watcher import LiteFileWatcher
__all__ = [
"BaseFileWatcher",
"LiteFileWatcher",
]

View file

@ -0,0 +1,118 @@
"""Abstract base for file watchers."""
import asyncio
from abc import abstractmethod
from pathlib import Path
from watchfiles import Change
from ..base_component import BaseComponent
from ..file_parser import BaseFileParser
from ..file_store import BaseFileStore
from ...enumeration import ComponentEnum
class BaseFileWatcher(BaseComponent):
"""Abstract base for file watchers. Subclasses implement watch_loop and event handlers."""
component_type = ComponentEnum.FILE_WATCHER
def __init__(
self,
watch_paths: list[str] | str,
suffix_filters: list[str] | None = None,
recursive: bool = True,
force_polling: bool = True,
debounce: int = 2000,
poll_delay_ms: int = 2000,
file_store: str = "default",
file_parser: str = "default",
**kwargs,
):
super().__init__(**kwargs)
from ..file_parser import DefaultFileParser
from ..file_store import LocalFileStore
watch_paths = [watch_paths] if isinstance(watch_paths, str) else watch_paths
base = self.working_path
self.watch_paths: list[Path] = [base / x for x in watch_paths if (base / x).exists()]
self.suffix_filters: list[str] = suffix_filters or ["md"]
self.recursive: bool = recursive
self.force_polling: bool = force_polling
self.debounce: int = debounce
self.poll_delay_ms: int = poll_delay_ms
self.file_store = self.bind(file_store, BaseFileStore, default_factory=LocalFileStore)
self.file_parser = self.bind(file_parser, BaseFileParser, default_factory=DefaultFileParser)
self._stop_event: asyncio.Event = asyncio.Event()
self._background_task: asyncio.Task | None = None
self._retry_interval: float = 10
async def _start(self):
self._stop_event = asyncio.Event()
self._background_task = asyncio.create_task(self._background_run())
self.logger.info(f"Started watching: {[str(p) for p in self.watch_paths]}")
async def _background_run(self):
"""Sync store then enter watch loop."""
await self.update_store()
await self.watch_loop()
async def _close(self):
self._stop_event.set()
if self._background_task:
await self._background_task
self.logger.info("Stopped watching")
def watch_filter(self, _change: Change, path: str) -> bool:
"""Return True if the file suffix matches the filter list."""
if not self.suffix_filters:
return True
return any(path.endswith("." + s.strip(".")) for s in self.suffix_filters)
def _get_relative_path(self, path: str | Path) -> str:
"""Return path relative to working_dir, or absolute path if outside."""
file_path = Path(path).absolute()
try:
return str(file_path.relative_to(self.working_path.absolute()))
except ValueError:
return str(file_path)
def _get_absolute_path(self, path: str | Path) -> Path:
"""Return absolute path; relative paths are resolved against working_dir."""
p = Path(path)
return p if p.is_absolute() else self.working_path / p
async def scan_existing_files(self) -> dict[str, Path]:
"""Collect watchable files under watch_paths as {relative_path: absolute_path}."""
files: dict[str, Path] = {}
for path in self.watch_paths:
if not path.exists():
continue
candidates = [path] if path.is_file() else (path.rglob("*") if self.recursive else path.iterdir())
for p in candidates:
if p.is_file() and self.watch_filter(Change.added, str(p)):
files[self._get_relative_path(p)] = p.absolute()
return files
@abstractmethod
async def watch_loop(self):
"""Watch for file changes and dispatch events."""
@abstractmethod
async def update_store(self, dump: bool = True) -> dict[str, int]:
"""Sync the store with watch_paths; dump store if any changes and dump=True.
Returns counts {"added": int, "modified": int, "deleted": int}.
"""
@abstractmethod
async def on_added(self, path: str | list[str]):
"""Handle file added event (relative paths)."""
@abstractmethod
async def on_modified(self, path: str | list[str]):
"""Handle file modified event (relative paths)."""
@abstractmethod
async def on_deleted(self, path: str | list[str]):
"""Handle file deleted event (relative paths)."""

View file

@ -0,0 +1,129 @@
"""Polling-based file watcher using watchfiles."""
import asyncio
from watchfiles import Change, awatch
from .base_file_watcher import BaseFileWatcher
from ..component_registry import R
from ...schema import FileChunk, FileNode
@R.register("lite")
class LiteFileWatcher(BaseFileWatcher):
"""Polling-based file watcher using watchfiles awatch."""
async def _interruptible_sleep(self):
"""Sleep until stop or timeout, whichever comes first."""
try:
await asyncio.wait_for(self._stop_event.wait(), timeout=self._retry_interval)
except asyncio.TimeoutError:
pass
async def watch_loop(self):
if not self.watch_paths:
self.logger.warning("No watch paths specified")
return
while not self._stop_event.is_set():
valid_paths = [p for p in self.watch_paths if p.exists()]
if not valid_paths:
self.logger.warning(f"No valid paths, retrying in {self._retry_interval}s...")
await self._interruptible_sleep()
continue
invalid = set(self.watch_paths) - set(valid_paths)
if invalid:
self.logger.warning(f"Skipping invalid paths: {[str(p) for p in invalid]}")
try:
self.logger.info(f"Watching: {[str(p) for p in valid_paths]}")
async for changes in awatch(
*valid_paths,
watch_filter=self.watch_filter,
recursive=self.recursive,
force_polling=self.force_polling,
debounce=self.debounce,
poll_delay_ms=self.poll_delay_ms,
stop_event=self._stop_event,
):
if self._stop_event.is_set():
break
await self._dispatch_changes(changes)
except Exception:
self.logger.exception(f"Watch error, retrying in {self._retry_interval}s...")
if not self._stop_event.is_set():
await self._interruptible_sleep()
async def _dispatch_changes(self, changes: set[tuple[Change, str]]):
"""Classify raw changes and dispatch to event handlers."""
buckets: dict[Change, list[str]] = {Change.added: [], Change.modified: [], Change.deleted: []}
for c, p in changes:
if c in buckets:
buckets[c].append(self._get_relative_path(p))
for change, handler, label in (
(Change.added, self.on_added, "added"),
(Change.modified, self.on_modified, "modified"),
(Change.deleted, self.on_deleted, "deleted"),
):
if buckets[change]:
self.logger.info(f"Detected {len(buckets[change])} {label} file(s)")
await handler(buckets[change])
async def update_store(self, dump: bool = True) -> dict[str, int]:
if self.file_store is None:
raise ValueError("file_store is not initialized!")
existing: dict[str, float] = {
rel: abs_p.stat().st_mtime for rel, abs_p in (await self.scan_existing_files()).items()
}
indexed: dict[str, float] = {n.path: n.st_mtime for n in await self.file_store.file_graph.get_nodes()}
to_delete = list(indexed.keys() - existing.keys())
to_add = list(existing.keys() - indexed.keys())
to_modify = [p for p in existing.keys() & indexed.keys() if existing[p] != indexed[p]]
if to_modify:
self.logger.info(f"Updating {len(to_modify)} modified file(s)")
await self.on_modified(to_modify)
if to_delete:
self.logger.info(f"Removing {len(to_delete)} deleted file(s)")
await self.on_deleted(to_delete)
if to_add:
self.logger.info(f"Indexing {len(to_add)} new file(s)")
await self.on_added(to_add)
changed = bool(to_add or to_modify or to_delete)
if not changed:
self.logger.info("Store is up to date")
if dump and changed:
await self.file_store.dump()
return {"added": len(to_add), "modified": len(to_modify), "deleted": len(to_delete)}
async def _parse_and_upsert(self, paths: list[str], action: str):
"""Parse files and upsert into store. Shared by on_added / on_modified."""
if self.file_parser is None or self.file_store is None:
raise RuntimeError("file_parser or file_store is not initialized!")
parsed: list[tuple[FileNode, list[FileChunk]]] = []
for rel in paths:
abs_path = self._get_absolute_path(rel)
if abs_path.is_file():
self.logger.info(f"{action} file: {rel}")
parsed.append(await self.file_parser.parse(abs_path))
if parsed:
await self.file_store.delete_by_path([n.path for n, _ in parsed])
await self.file_store.upsert_file(parsed)
async def on_added(self, path: str | list[str]):
await self._parse_and_upsert([path] if isinstance(path, str) else path, "Adding")
async def on_modified(self, path: str | list[str]):
await self._parse_and_upsert([path] if isinstance(path, str) else path, "Updating")
async def on_deleted(self, path: str | list[str]):
if self.file_store is None:
raise RuntimeError("file_store is not initialized!")
paths = [path] if isinstance(path, str) else path
self.logger.info(f"Deleting {len(paths)} file(s)")
await self.file_store.delete_by_path(paths)

View file

@ -0,0 +1,6 @@
"""Job components for executing workflows."""
from .base_job import BaseJob
from .stream_job import StreamJob
__all__ = ["BaseJob", "StreamJob"]

View file

@ -0,0 +1,54 @@
"""Base job component for sequential step execution."""
from ..base_component import BaseComponent
from ..component_registry import R
from ..runtime_context import RuntimeContext
from ...enumeration import ComponentEnum
from ...schema import ComponentConfig, Response
@R.register("base")
class BaseJob(BaseComponent):
"""Job that executes steps sequentially and returns a Response."""
component_type = ComponentEnum.JOB
def __init__(self, description: str, parameters: dict, steps: list[ComponentConfig | dict], **kwargs):
super().__init__(**kwargs)
self.description = description
self.parameters = parameters or {}
self.step_configs = steps or []
from ...steps import BaseStep
self.step_components: list[BaseStep] = []
async def _start(self) -> None:
"""Resolve step configs into instantiated step components."""
assert self.app_context is not None, "app_context must be provided"
for raw in self.step_configs:
config = raw if isinstance(raw, ComponentConfig) else ComponentConfig(**raw)
if not config.backend:
raise ValueError("Step is missing the required 'backend' field")
step_cls = R.get(ComponentEnum.STEP, config.backend)
if not step_cls:
raise ValueError(f"Unregistered backend '{config.backend}' of type '{ComponentEnum.STEP}'")
params = config.model_dump()
params["app_context"] = self.app_context
self.step_components.append(step_cls(**params))
async def _close(self) -> None:
"""Release all step components."""
self.step_components.clear()
async def __call__(self, **kwargs) -> Response:
"""Execute all steps in order and return the final response."""
context = RuntimeContext(**kwargs)
try:
for step in self.step_components:
await step(context)
except Exception as e:
self.logger.exception(f"Failed to execute job: {e}")
context.response.success = False
context.response.answer = str(e)
return context.response

View file

@ -0,0 +1,21 @@
"""Streaming job for real-time output delivery."""
from .base_job import BaseJob
from ..component_registry import R
from ..runtime_context import RuntimeContext
from ...enumeration import ChunkEnum
@R.register("stream")
class StreamJob(BaseJob):
"""Job that streams chunks to a queue instead of returning a Response."""
async def __call__(self, **kwargs) -> None:
"""Execute steps and stream output; errors are sent as ERROR chunks."""
context = RuntimeContext(**kwargs)
try:
for step in self.step_components:
await step(context)
except Exception as e:
await context.add_stream_string(str(e), ChunkEnum.ERROR)
await context.add_stream_done()

View file

@ -0,0 +1,6 @@
"""Keyword index components."""
from .base_keyword_index import BaseKeywordIndex
from .bm25_index import BM25Index
__all__ = ["BaseKeywordIndex", "BM25Index"]

View file

@ -0,0 +1,70 @@
"""Abstract base class for keyword index implementations."""
from abc import abstractmethod
from pathlib import Path
from ..base_component import BaseComponent
from ..tokenizer import BaseTokenizer
from ...enumeration import ComponentEnum
class BaseKeywordIndex(BaseComponent):
"""Abstract base class for keyword index implementations."""
component_type = ComponentEnum.KEYWORD_INDEX
def __init__(self, tokenizer: str = "default", index_version: str = "v1", **kwargs):
super().__init__(**kwargs)
from ..tokenizer import RegexTokenizer
self.tokenizer = self.bind(tokenizer, BaseTokenizer, default_factory=RegexTokenizer)
self.index_version = index_version
self.index_path = self.working_metadata_path / self.component_type.value
self.index_path.mkdir(parents=True, exist_ok=True)
async def _start(self) -> None:
"""Load existing index from disk if available."""
await self.load()
async def _close(self) -> None:
"""Save index to disk on shutdown."""
await self.dump()
@property
def index_file(self) -> Path:
"""Return the pickle file path derived from tokenizer name."""
if self.tokenizer is None:
raise RuntimeError("Tokenizer not initialized. Call start() first.")
name = type(self.tokenizer).__name__.replace("Tokenizer", "").lower()
return self.index_path / f"bm25_{name}_{self.index_version}.pkl"
def _tokenize(self, text: str) -> list[str]:
"""Tokenize a text string into tokens."""
if self.tokenizer is None:
raise RuntimeError("Tokenizer not initialized. Call start() first.")
return self.tokenizer.tokenize([text])[0]
@abstractmethod
async def add_docs(self, docs_dict: dict[str, str]) -> None:
"""Index or update documents. Mapping of doc_id to content."""
@abstractmethod
async def delete_docs(self, doc_ids: list[str]) -> None:
"""Remove documents by their IDs."""
@abstractmethod
async def retrieve(self, query: str, limit: int = 3) -> dict[str, float]:
"""Search documents. Returns {doc_id: score} sorted descending."""
@abstractmethod
async def clear(self) -> None:
"""Reset index to empty state."""
async def reset_index(self, docs_dict: dict[str, str]) -> None:
"""Clear index, re-add all documents, and persist."""
await self.clear()
await self.add_docs(docs_dict)
await self.dump()
async def optimize_index(self) -> None:
"""Optimize index for performance. Override in subclass if needed."""

View file

@ -0,0 +1,206 @@
"""BM25 search engine with persistent index support.
Implements Okapi BM25 ranking with an inverted index for efficient
document lookup, incremental updates, and pickle-based persistence.
"""
import math
import pickle
from collections import Counter
from typing import TypedDict
from .base_keyword_index import BaseKeywordIndex
from ..component_registry import R
class DocMeta(TypedDict):
"""Per-document metadata: token count and unique token ID set."""
len: int
token_ids: set[int]
@R.register("bm25")
class BM25Index(BaseKeywordIndex):
"""BM25 search engine with file-based persistence.
Args:
k1: Term frequency saturation parameter (default 1.5).
b: Document length normalization parameter (default 0.75).
"""
def __init__(self, k1: float = 1.5, b: float = 0.75, **kwargs):
super().__init__(**kwargs)
self.k1 = k1
self.b = b
self.vocab: dict[str, int] = {} # token -> token_id
self.inverted_index: dict[int, dict[str, int]] = {} # token_id -> {doc_id: tf}
self.doc_meta: dict[str, DocMeta] = {} # doc_id -> metadata
self.total_len: int = 0
self._idf_cache: dict[int, float] = {}
# -- Properties -----------------------------------------------------------
@property
def n_docs(self) -> int:
"""Number of indexed documents."""
return len(self.doc_meta)
@property
def avg_len(self) -> float:
"""Average document length in tokens."""
return self.total_len / self.n_docs if self.n_docs > 0 else 0.0
# -- Internal helpers -----------------------------------------------------
def _tokens_to_ids(self, tokens: list[str]) -> list[int]:
"""Map tokens to integer IDs, assigning new IDs on first encounter."""
ids = []
for token in tokens:
token = token.strip()
if token:
ids.append(self.vocab.setdefault(token, len(self.vocab)))
return ids
def _remove_doc(self, doc_id: str) -> None:
"""Remove a single document from all internal structures."""
if doc_id not in self.doc_meta:
return
meta = self.doc_meta[doc_id]
self.total_len -= meta["len"]
for tid in meta["token_ids"]:
if tid in self.inverted_index:
self.inverted_index[tid].pop(doc_id, None)
if not self.inverted_index[tid]:
del self.inverted_index[tid]
del self.doc_meta[doc_id]
def _get_idf(self, token_id: int) -> float:
"""Compute and cache IDF for a token ID."""
if token_id in self._idf_cache:
return self._idf_cache[token_id]
df = len(self.inverted_index.get(token_id, {}))
self._idf_cache[token_id] = math.log(1 + (self.n_docs - df + 0.5) / (df + 0.5)) if df else 0.0
return self._idf_cache[token_id]
# -- Public API -----------------------------------------------------------
async def add_docs(self, docs_dict: dict[str, str]) -> None:
"""Index or update multiple documents. Mapping of doc_id to content."""
for doc_id, content in docs_dict.items():
if doc_id in self.doc_meta:
self._remove_doc(doc_id)
tokens = self._tokenize(content)
if not tokens:
continue
token_ids = self._tokens_to_ids(tokens)
token_counts = Counter(token_ids)
for tid, tf in token_counts.items():
self.inverted_index.setdefault(tid, {})[doc_id] = tf
self.doc_meta[doc_id] = {"len": len(token_ids), "token_ids": set(token_counts)}
self.total_len += len(token_ids)
self._idf_cache = {}
async def delete_docs(self, doc_ids: list[str]) -> None:
"""Remove documents by their IDs."""
for doc_id in doc_ids:
self._remove_doc(doc_id)
self._idf_cache = {}
async def retrieve(self, query: str, limit: int = 3) -> dict[str, float]:
"""Search documents. Returns {doc_id: score} sorted descending."""
query_ids = [self.vocab[t] for t in self._tokenize(query) if t in self.vocab]
if not query_ids or self.n_docs == 0:
return {}
scores: dict[str, float] = {}
avg_len = self.avg_len
for tid in query_ids:
if tid not in self.inverted_index:
continue
idf = self._get_idf(tid)
for doc_id, tf in self.inverted_index[tid].items():
doc_len = self.doc_meta[doc_id]["len"]
tf_score = tf * (self.k1 + 1) / (tf + self.k1 * (1 - self.b + self.b * doc_len / avg_len))
scores[doc_id] = scores.get(doc_id, 0.0) + idf * tf_score
return dict(sorted(scores.items(), key=lambda x: x[1], reverse=True)[:limit]) if scores else {}
async def dump(self) -> None:
"""Persist index to disk via pickle (atomic rename)."""
try:
tmp = self.index_file.with_suffix(".tmp")
with open(tmp, "wb") as f:
pickle.dump(
{
"vocab": self.vocab,
"inverted_index": self.inverted_index,
"doc_meta": self.doc_meta,
"total_len": self.total_len,
"k1": self.k1,
"b": self.b,
},
f,
)
tmp.replace(self.index_file)
self.logger.info(f"Saved {self.n_docs} docs to {self.index_file}")
except Exception as e:
self.logger.exception(f"Failed to write {self.index_file}: {e}")
async def load(self) -> None:
"""Load index from disk. No-op if file missing; clears index on corruption."""
if not self.index_file.exists():
return
try:
with open(self.index_file, "rb") as f:
data = pickle.load(f)
self.vocab = data["vocab"]
self.inverted_index = data["inverted_index"]
self.doc_meta = data["doc_meta"]
self.total_len = data.get("total_len", 0)
self.k1 = data.get("k1", 1.5)
self.b = data.get("b", 0.75)
self._idf_cache = {}
self.logger.info(f"Loaded {self.n_docs} docs from {self.index_file}")
except Exception as e:
self.logger.exception(f"Failed to load index: {e}")
self.index_file.unlink(missing_ok=True)
await self.clear()
async def clear(self) -> None:
"""Reset index to empty state and remove persisted file."""
self.vocab = {}
self.inverted_index = {}
self.doc_meta = {}
self.total_len = 0
self._idf_cache = {}
self.index_file.unlink(missing_ok=True)
async def optimize_index(self) -> None:
"""Rebuild vocab to remove unused tokens and compact token IDs."""
used_token_ids: set[int] = set()
for tid in self.inverted_index:
used_token_ids.add(tid)
if not used_token_ids:
await self.clear()
return
# Build compact ID mapping
old_to_new: dict[int, int] = {}
new_vocab: dict[str, int] = {}
for token, old_tid in self.vocab.items():
if old_tid in used_token_ids:
new_tid = len(new_vocab)
new_vocab[token] = new_tid
old_to_new[old_tid] = new_tid
# Rebuild inverted index and doc_meta with new IDs
new_inverted_index: dict[int, dict[str, int]] = {}
for old_tid, postings in self.inverted_index.items():
new_inverted_index[old_to_new[old_tid]] = postings
for meta in self.doc_meta.values():
meta["token_ids"] = {old_to_new[t] for t in meta["token_ids"] if t in old_to_new}
self.vocab = new_vocab
self.inverted_index = new_inverted_index
self._idf_cache = {}

View file

@ -0,0 +1,125 @@
"""Prompt template loader and formatter with conditional-line and i18n support."""
import inspect
import json
import re
from pathlib import Path
from string import Formatter
import yaml
# Matches a leading flag tag like "[verbose] some text".
_FLAG_PATTERN = re.compile(r"^\[(\w+)]")
class PromptHandler:
"""Loads prompts from YAML/JSON or class-adjacent files and formats them.
Templates may carry a language suffix (``key_en``, ``key_zh``); ``get_prompt``
falls back to the bare key when no localized variant exists. ``prompt_format``
additionally supports per-line flags such as ``[verbose] extra text`` that
are kept only when the matching flag kwarg is truthy.
"""
_SUPPORTED_EXTENSIONS = {".yaml", ".yml", ".json"}
def __init__(self, language: str = "", **kwargs):
# Only string entries are treated as prompts; other kwargs are ignored.
self.data: dict[str, str] = {k: v for k, v in kwargs.items() if isinstance(v, str)}
self.language: str = language.strip()
def load_prompt_by_file(
self,
prompt_file_path: str | Path | None = None,
overwrite: bool = True,
) -> "PromptHandler":
"""Load prompts from a YAML or JSON file; silently skip on any error."""
if prompt_file_path is None:
return self
path = Path(prompt_file_path)
if not path.exists() or path.suffix.lower() not in self._SUPPORTED_EXTENSIONS:
return self
try:
with path.open(encoding="utf-8") as f:
prompt_dict = yaml.safe_load(f) if path.suffix.lower() in (".yaml", ".yml") else json.load(f)
except (json.JSONDecodeError, yaml.YAMLError, OSError):
return self
return self.load_prompt_dict(prompt_dict, overwrite)
def load_prompt_by_class(self, cls: type, overwrite: bool = True) -> "PromptHandler":
"""Load prompts from ``<class_module>.yaml`` (or ``.yml``) next to `cls`."""
try:
base_path = Path(inspect.getfile(cls)).with_suffix("")
except (TypeError, OSError):
return self
for ext in (".yaml", ".yml"):
if (prompt_path := base_path.with_suffix(ext)).exists():
return self.load_prompt_by_file(prompt_path, overwrite)
return self
def load_prompt_dict(self, prompt_dict: dict | None = None, overwrite: bool = True) -> "PromptHandler":
"""Merge string entries from `prompt_dict` into the in-memory store."""
if not isinstance(prompt_dict, dict):
return self
for key, value in prompt_dict.items():
if isinstance(value, str) and (overwrite or key not in self.data):
self.data[key] = value
return self
def get_prompt(self, prompt_name: str) -> str:
"""Return the template, preferring the language-suffixed variant when set."""
for key in (f"{prompt_name}_{self.language}", prompt_name) if self.language else (prompt_name,):
if key in self.data:
return self.data[key].strip()
raise KeyError(f"Prompt '{prompt_name}' not found. Available: {list(self.data.keys())[:10]}")
def has_prompt(self, prompt_name: str) -> bool:
"""True if either the localized or bare prompt is registered."""
keys = (f"{prompt_name}_{self.language}", prompt_name) if self.language else (prompt_name,)
return any(k in self.data for k in keys)
def list_prompts(self, language_filter: str | None = None) -> list[str]:
"""List all keys, optionally filtered to those ending with ``_<language>``."""
if language_filter is None:
return list(self.data.keys())
suffix = f"_{language_filter.strip()}"
return [k for k in self.data if k.endswith(suffix)]
def prompt_format(self, prompt_name: str, validate: bool = True, **kwargs) -> str:
"""Render a prompt: strip inactive flag-lines, then ``str.format`` it.
Boolean kwargs are treated as flags controlling ``[flag]`` line filtering.
Remaining kwargs become positional substitutions for ``{var}`` placeholders.
With `validate=True`, missing substitutions raise ``ValueError``.
"""
prompt = self.get_prompt(prompt_name)
flags = {k: v for k, v in kwargs.items() if isinstance(v, bool)}
formats = {k: v for k, v in kwargs.items() if not isinstance(v, bool)}
# Keep lines without flags; otherwise keep when at least one flag is enabled.
if flags:
lines = []
for line in prompt.split("\n"):
active_flags = _FLAG_PATTERN.findall(line)
cleaned = _FLAG_PATTERN.sub("", line).lstrip()
if not active_flags or any(flags.get(f, False) for f in active_flags):
lines.append(cleaned)
prompt = "\n".join(lines)
if validate:
required = {f for _, f, _, _ in Formatter().parse(prompt) if f is not None}
if missing := required - set(formats.keys()):
raise ValueError(f"Missing format variables for '{prompt_name}': {sorted(missing)}")
return prompt.format(**formats).strip() if formats else prompt
def __repr__(self) -> str:
return f"PromptHandler(language='{self.language}', num_prompts={len(self.data)})"

View file

@ -0,0 +1,87 @@
"""Per-request runtime context shared across steps and jobs."""
import asyncio
from ..enumeration import ChunkEnum
from ..schema import Response, StreamChunk
class RuntimeContext:
"""Scratch space for a single execution.
Holds the response object, an optional stream queue, and a free-form
data dict accessed via mapping-style operators.
"""
def __init__(
self,
response: Response | None = None,
stream_queue: asyncio.Queue | None = None,
**kwargs,
):
self.response: Response = response or Response()
self.stream_queue: asyncio.Queue | None = stream_queue
self.data: dict = kwargs
def get(self, key: str, default=None):
"""Get a value from the data dict."""
return self.data.get(key, default)
def update(self, data: dict) -> "RuntimeContext":
"""Merge data into the context."""
self.data.update(data)
return self
def __getitem__(self, key: str):
return self.data[key]
def __setitem__(self, key: str, value):
self.data[key] = value
def __delitem__(self, key: str):
del self.data[key]
def __contains__(self, key: str) -> bool:
return key in self.data
@property
def stream(self) -> bool:
"""Whether streaming is enabled."""
return self.stream_queue is not None
@classmethod
def from_context(cls, context: "RuntimeContext | None" = None, **kwargs) -> "RuntimeContext":
"""Reuse or create a RuntimeContext."""
# Reuse the existing context (merging kwargs) or create a new one.
if context is None:
return cls(**kwargs)
context.update(kwargs)
return context
async def _enqueue(self, chunk: StreamChunk) -> None:
"""Put a chunk on the stream queue."""
if self.stream_queue is None:
raise RuntimeError("Stream queue not initialized")
await self.stream_queue.put(chunk)
async def add_stream_string(self, chunk: str, chunk_type: ChunkEnum) -> "RuntimeContext":
"""Emit a text chunk to the stream queue."""
# Emit a text chunk to the stream queue.
await self._enqueue(StreamChunk(chunk_type=chunk_type, chunk=chunk))
return self
async def add_stream_done(self) -> "RuntimeContext":
"""Emit the terminal DONE marker to close the stream."""
# Emit the terminal DONE marker to close the stream.
await self._enqueue(StreamChunk(chunk_type=ChunkEnum.DONE, chunk="", done=True))
return self
def apply_mapping(self, mapping: dict[str, str]) -> "RuntimeContext":
"""Copy data[source] into data[target] for each mapping pair."""
# Copy data[source] into data[target] for each {source: target} pair.
if not mapping:
return self
for source, target in mapping.items():
if source in self.data:
self.data[target] = self.data[source]
return self

View file

@ -0,0 +1,11 @@
"""Service components for exposing jobs via different protocols."""
from .base_service import BaseService
from .http_service import HttpService
from .mcp_service import MCPService
__all__ = [
"BaseService",
"HttpService",
"MCPService",
]

View file

@ -0,0 +1,48 @@
"""Base service class for exposing jobs via HTTP, MCP, etc."""
from abc import abstractmethod
from typing import TYPE_CHECKING
from ..base_component import BaseComponent
from ..job.base_job import BaseJob
from ...enumeration import ComponentEnum
if TYPE_CHECKING:
from ...application import Application
class BaseService(BaseComponent):
"""Base class for services that expose jobs via HTTP, MCP, etc."""
component_type = ComponentEnum.SERVICE
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.service = None
@abstractmethod
def build_service(self, app: "Application") -> None:
"""Initialize the underlying service framework."""
@abstractmethod
def add_job(self, job: BaseJob) -> None:
"""Register a single job with the service."""
@abstractmethod
def start_service(self, app: "Application") -> None:
"""Start serving requests."""
def add_jobs(self, app: "Application") -> None:
"""Register all jobs from the application context."""
for name, job in app.context.jobs.items():
try:
self.add_job(job)
self.logger.info(f"Added job: {name}")
except Exception as e:
self.logger.error(f"Failed to add job {name}: {e}")
def run_app(self, app: "Application") -> None:
"""Build, populate, and start the service."""
self.build_service(app)
self.add_jobs(app)
self.start_service(app)

View file

@ -0,0 +1,99 @@
"""HTTP service implementation for ReMe."""
import asyncio
import json
import os
import warnings
from collections.abc import AsyncGenerator
from contextlib import asynccontextmanager
from typing import TYPE_CHECKING
import uvicorn
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import StreamingResponse
from .base_service import BaseService
from ..component_registry import R
from ..job import BaseJob, StreamJob
from ...constants import REME_DEFAULT_HOST, REME_DEFAULT_PORT, REME_SERVICE_INFO
from ...schema import Request, Response
from ...utils import execute_stream_task
if TYPE_CHECKING:
from ...application import Application
@R.register("http")
class HttpService(BaseService):
"""HTTP service: normal jobs -> JSON endpoints, stream jobs -> SSE endpoints."""
def __init__(self, host: str = REME_DEFAULT_HOST, port: int = REME_DEFAULT_PORT, **kwargs):
super().__init__(**kwargs)
self.host: str = host
self.port: int = port
def _add_job(self, job: BaseJob) -> None:
async def execute_endpoint(request: Request) -> Response:
return await job(**request.model_dump(exclude_none=True))
self.service.post(path=f"/{job.name}", response_model=Response, description=job.description)(execute_endpoint)
def _add_stream_job(self, job: StreamJob) -> None:
async def execute_stream_endpoint(request: Request) -> StreamingResponse:
stream_queue = asyncio.Queue()
task = asyncio.create_task(job(stream_queue=stream_queue, **request.model_dump(exclude_none=True)))
async def generate_stream() -> AsyncGenerator[bytes, None]:
async for chunk in execute_stream_task(
stream_queue=stream_queue,
task=task,
task_name=job.name,
output_format="bytes",
):
assert isinstance(chunk, bytes)
yield chunk
return StreamingResponse(generate_stream(), media_type="text/event-stream")
self.service.post(f"/{job.name}")(execute_stream_endpoint)
def add_job(self, job: BaseJob) -> None:
if isinstance(job, StreamJob):
self._add_stream_job(job)
else:
self._add_job(job)
def build_service(self, app: "Application") -> None:
@asynccontextmanager
async def lifespan(_: FastAPI):
await app.start()
service_info = json.dumps({"host": self.host, "port": self.port})
os.environ[REME_SERVICE_INFO] = service_info
self.logger.info(f"ReMe Service started: {REME_SERVICE_INFO}={service_info}")
yield
await app.close()
self.service = FastAPI(title=app.config.app_name, lifespan=lifespan)
self.service.add_middleware(
CORSMiddleware, # type: ignore[arg-type]
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
def start_service(self, app: "Application") -> None:
# uvicorn 0.41 still imports websockets.legacy / WebSocketServerProtocol
# on startup; silence those specific lines since we don't use WebSocket.
warnings.filterwarnings(
"ignore",
category=DeprecationWarning,
message=r".*websockets\.legacy is deprecated.*",
)
warnings.filterwarnings(
"ignore",
category=DeprecationWarning,
message=r".*WebSocketServerProtocol is deprecated.*",
)
uvicorn.run(self.service, host=self.host, port=self.port, **self.kwargs)

View file

@ -0,0 +1,71 @@
"""MCP (Model Context Protocol) service implementation."""
import json
import os
from contextlib import asynccontextmanager
from typing import TYPE_CHECKING
from fastmcp import FastMCP
from fastmcp.server.server import Transport
from fastmcp.tools import FunctionTool
from .base_service import BaseService
from ..component_registry import R
from ..job import StreamJob, BaseJob
from ...constants import REME_DEFAULT_HOST, REME_DEFAULT_PORT, REME_SERVICE_INFO
if TYPE_CHECKING:
from ...application import Application
@R.register("mcp")
class MCPService(BaseService):
"""Expose jobs as MCP (Model Context Protocol) tools."""
def __init__(
self,
transport: Transport = "sse",
host: str = REME_DEFAULT_HOST,
port: int = REME_DEFAULT_PORT,
**kwargs,
):
super().__init__(**kwargs)
self.transport: Transport = transport
self.host: str = host
self.port: int = port
def build_service(self, app: "Application") -> None:
@asynccontextmanager
async def lifespan(_: FastMCP):
await app.start()
service_info = json.dumps({"host": self.host, "port": self.port})
os.environ[REME_SERVICE_INFO] = service_info
self.logger.info(f"ReMe MCP Service started: {REME_SERVICE_INFO}={service_info}")
yield
await app.close()
self.service = FastMCP(name=app.config.app_name, lifespan=lifespan)
def add_job(self, job: "BaseJob") -> None:
if isinstance(job, StreamJob):
return
async def execute_tool(**kwargs):
response = await job(**kwargs)
return response.answer
self.service.add_tool(
FunctionTool(
name=job.name,
description=job.description,
fn=execute_tool,
parameters=job.parameters or None,
),
)
def start_service(self, app: "Application") -> None:
transport_kwargs = {}
if self.transport != "stdio":
transport_kwargs["host"] = self.host
transport_kwargs["port"] = self.port
self.service.run(transport=self.transport, show_banner=False, **transport_kwargs)

View file

@ -0,0 +1,11 @@
"""Tokenizer component module."""
from .base_tokenizer import BaseTokenizer
from .jieba_tokenizer import JiebaTokenizer
from .regex_tokenizer import RegexTokenizer
__all__ = [
"BaseTokenizer",
"JiebaTokenizer",
"RegexTokenizer",
]

View file

@ -0,0 +1,44 @@
"""Abstract base class for tokenizers."""
from abc import abstractmethod
from pathlib import Path
import aiofiles
from ..base_component import BaseComponent
from ...enumeration import ComponentEnum
class BaseTokenizer(BaseComponent):
"""Base tokenizer. Subclasses must implement `tokenize`. Loads stopwords on start."""
component_type = ComponentEnum.TOKENIZER
DEFAULT_STOPWORDS_PATH = Path(__file__).parent / "stopwords"
def __init__(self, stopwords_path: str | Path | None = None, **kwargs):
super().__init__(**kwargs)
self.stopwords_path = Path(stopwords_path) if stopwords_path else self.DEFAULT_STOPWORDS_PATH
self._stopwords: set[str] = set()
async def _start(self) -> None:
"""Load stopwords from file."""
if not self.stopwords_path.exists():
self.logger.warning(f"Stopwords file not found: {self.stopwords_path}")
return
async with aiofiles.open(self.stopwords_path, encoding="utf-8") as f:
content = await f.read()
self._stopwords = {line.strip().lower() for line in content.splitlines() if line.strip()}
self.logger.info(f"Loaded {len(self._stopwords)} stopwords from {self.stopwords_path}")
async def _close(self) -> None:
"""Clear stopwords."""
self._stopwords.clear()
@property
def stopwords(self) -> set[str]:
"""Get the loaded stopwords."""
return self._stopwords
@abstractmethod
def tokenize(self, texts: list[str], **kwargs) -> list[list[str]]:
"""Tokenize a list of texts."""

View file

@ -0,0 +1,27 @@
"""Jieba tokenizer for Chinese text segmentation."""
from .base_tokenizer import BaseTokenizer
from ..component_registry import R
@R.register("jieba")
class JiebaTokenizer(BaseTokenizer):
"""Tokenizer using jieba for Chinese text segmentation."""
def __init__(self, filter_stopwords: bool = True, **kwargs):
super().__init__(**kwargs)
self.filter_stopwords = filter_stopwords
def tokenize(self, texts: list[str], lower: bool = True, **kwargs) -> list[list[str]]:
"""Tokenize texts using jieba."""
import jieba
result = []
for text in texts:
tokens = jieba.cut(text)
if lower:
tokens = [x.lower() for x in tokens]
if self.filter_stopwords and self._stopwords:
tokens = [t for t in tokens if t not in self._stopwords]
result.append(tokens)
return result

View file

@ -0,0 +1,31 @@
"""Regex tokenizer with Chinese character splitting."""
import re
from .base_tokenizer import BaseTokenizer
from ..component_registry import R
@R.register("regex")
class RegexTokenizer(BaseTokenizer):
"""Tokenizer using regex: splits Chinese chars individually, extracts non-Chinese words."""
WORD_PATTERN = re.compile(r"(?u)\b\w\w+\b") # 2+ char words
CHINESE_PATTERN = re.compile(r"[一-鿿]") # single Chinese char
def __init__(self, filter_stopwords: bool = True, **kwargs):
super().__init__(**kwargs)
self.filter_stopwords = filter_stopwords
def tokenize(self, texts: list[str], lower: bool = True, **kwargs) -> list[list[str]]:
"""Tokenize texts. Extracts Chinese chars, then non-Chinese words from remaining text."""
result = []
for text in texts:
# Extract Chinese chars individually, then non-Chinese words
tokens = self.CHINESE_PATTERN.findall(text)
tokens.extend(self.WORD_PATTERN.findall(self.CHINESE_PATTERN.sub(" ", text)))
if lower:
tokens = [t.lower() for t in tokens]
if self.filter_stopwords and self._stopwords:
tokens = [t for t in tokens if t not in self._stopwords]
result.append(tokens)
return result

File diff suppressed because it is too large Load diff

8
reme4/config/__init__.py Normal file
View file

@ -0,0 +1,8 @@
"""Config"""
from .config_parser import parse_args, resolve_app_config
__all__ = [
"parse_args",
"resolve_app_config",
]

View file

@ -0,0 +1,219 @@
"""Parser for YAML config with CLI argument overrides."""
import json
import os
import re
from pathlib import Path
from typing import Any
import yaml
# Config files are looked up relative to this module's directory
_CONFIG_DIR = Path(__file__).parent
# Extensions in priority order: yaml > yml > json when stems collide
_SUPPORTED_EXTS = (".yaml", ".yml", ".json")
_ENV_VAR_RE = re.compile(r"\$\{([A-Za-z_][A-Za-z0-9_]*)(?::-([^}]*))?}")
# Strings like "007" / "00501" must stay as strings, not be coerced to numbers
_LEADING_ZERO_RE = re.compile(r"^-?0\d")
def _repl(m: re.Match) -> str:
name: str = m.group(1)
# group(2) is None when the placeholder has no `:-default` part
default: str | None = m.group(2)
v = os.environ.get(name)
if v is None:
if default is not None:
return default
raise ValueError(f"Config references undefined env var: {name}")
return v
def _expand_env_vars(value: Any) -> Any:
"""Recursively expand `${VAR}` / `${VAR:-default}` placeholders in strings."""
if isinstance(value, str):
return _ENV_VAR_RE.sub(_repl, value)
if isinstance(value, dict):
return {k: _expand_env_vars(v) for k, v in value.items()}
if isinstance(value, list):
return [_expand_env_vars(v) for v in value]
return value
def _discover_configs() -> dict[str, Path]:
"""Pre-scan config directory: maps file stem (name without ext) -> Path."""
discovered: dict[str, Path] = {}
if _CONFIG_DIR.is_dir():
# Sort by ext priority so registration order is deterministic across filesystems
files = sorted(
(p for p in _CONFIG_DIR.iterdir() if p.is_file() and p.suffix in _SUPPORTED_EXTS),
key=lambda p: (_SUPPORTED_EXTS.index(p.suffix), p.name),
)
for p in files:
discovered.setdefault(p.stem, p)
return discovered
_CONFIG_REGISTRY = _discover_configs()
def parse_dot_notation(dot_list: list[str]) -> dict:
"""Parse "key.subkey=value" strings into nested dict."""
result: dict = {}
for item in dot_list:
if "=" not in item:
raise ValueError(f"Invalid dot notation format (missing '='): {item}")
key_path, value_str = item.split("=", 1)
keys = key_path.split(".")
current = result
for key in keys[:-1]:
if key in current and not isinstance(current[key], dict):
raise ValueError(f"Cannot set nested key '{key_path}': '{key}' is already a value")
current = current.setdefault(key, {})
# Symmetric to the prefix check above: refuse scalar-over-dict overwrite
last_key = keys[-1]
if last_key in current and isinstance(current[last_key], dict):
raise ValueError(f"Cannot overwrite nested dict at '{key_path}' with scalar value")
current[last_key] = _convert_value(value_str)
return result
def _convert_value(value_str: str) -> Any:
"""Convert string to appropriate Python type.
Only converts "true"/"false" (case-insensitive) to boolean.
Use JSON format (e.g., '"yes"', '"no"') to preserve these as strings.
Leading-zero strings (e.g., "007", "00501") are kept as strings.
"""
s = value_str.strip()
lower = s.lower()
# Handle special values (null, bool)
if lower in ("none", "null"):
return None
if lower == "true":
return True
if lower == "false":
return False
# Skip int/float for leading-zero strings to keep zip codes / ids intact
if not _LEADING_ZERO_RE.match(s):
for converter in (int, float):
try:
return converter(s)
except ValueError:
continue
# JSON handles lists, dicts, and explicitly-quoted strings
try:
return json.loads(s)
except (ValueError, json.JSONDecodeError):
pass
# Fallback to original string
return s
def _load_config(name_or_path: str, encoding: str = "utf-8") -> dict:
"""Load a YAML or JSON config file.
First check if name_or_path matches a pre-discovered config (key in _CONFIG_REGISTRY).
If not, treat as a file path and load directly.
"""
# 1. Try pre-discovered configs first
if name_or_path in _CONFIG_REGISTRY:
return _read_config_file(_CONFIG_REGISTRY[name_or_path], encoding)
# 2. Treat as file path
p = Path(name_or_path)
if p.suffix in _SUPPORTED_EXTS:
if not p.exists():
raise FileNotFoundError(f"Config file not found: {p}")
return _read_config_file(p, encoding)
known = ", ".join(sorted(_CONFIG_REGISTRY)) if _CONFIG_REGISTRY else "none"
raise FileNotFoundError(f"Config file not found: {name_or_path}. Available: {known}")
def _read_config_file(path: Path, encoding: str = "utf-8") -> dict:
"""Read YAML or JSON file based on extension. Expands ${ENV_VAR}."""
with path.open(encoding=encoding) as f:
if path.suffix == ".json":
result = json.load(f)
else:
result = yaml.safe_load(f)
if result is None:
return {}
return _expand_env_vars(result)
def _deep_merge(base: dict, update: dict) -> dict:
"""Recursively merge dicts."""
result = base.copy()
for k, v in update.items():
if k in result and isinstance(result[k], dict) and isinstance(v, dict):
result[k] = _deep_merge(result[k], v)
else:
result[k] = v
return result
def _strip_arg_dashes(arg: str) -> str:
"""Strip a single leading `--` or `-` prefix (not all leading dashes)."""
if arg.startswith("--"):
return arg[2:]
if arg.startswith("-"):
return arg[1:]
return arg
def parse_args(*args) -> tuple[str, dict]:
"""Parse CLI args: first arg is action, rest are key=value pairs.
Usage: reme app config=paw.yaml service.name=test
Returns: (action, parsed_kv_dict)
"""
if not args:
raise ValueError("No arguments provided")
first = _strip_arg_dashes(args[0])
if "=" in first:
raise ValueError(f"First argument must be action, got: {args[0]}")
kvs: list[str] = []
for raw in args[1:]:
arg = _strip_arg_dashes(raw)
if "=" in arg:
kvs.append(arg)
parsed = parse_dot_notation(kvs) if kvs else {}
return first, parsed
def resolve_app_config(**kwargs) -> dict:
"""Resolve full app-start config: load `config=path` file, fall back to
`default`, then deep-merge with the remaining kwargs as overrides.
"""
from ..utils import get_logger
logger = get_logger()
configs: list[dict] = []
# `config=path` arrives as a string here; `config.foo=bar` arrives as a
# nested dict and is left in `kwargs` to be merged as a normal override.
config_value = kwargs.get("config")
if isinstance(config_value, str):
kwargs.pop("config")
logger.info(f"Loading config: {config_value}")
configs.append(_load_config(config_value))
elif "default" in _CONFIG_REGISTRY:
logger.info("No config specified, loading 'default'")
configs.append(_load_config("default"))
configs.append(kwargs)
merged: dict = {}
for cfg in configs:
merged = _deep_merge(merged, cfg)
return merged

189
reme4/config/default.yaml Normal file
View file

@ -0,0 +1,189 @@
service:
backend: http
# backend: mcp
jobs:
- backend: base
name: demo
description: "demo job description"
parameters:
type: object
properties:
query:
type: string
description: "query"
min_score:
type: number
description: "min score"
default: 0.5
required:
- query
steps:
- backend: demo_echo_step1
- backend: demo_echo_step2
- backend: base
name: version
description: "return reme4 package version"
parameters:
type: object
properties: {}
steps:
- backend: version_step
- backend: base
name: health_check
description: "return a concise health-check snapshot of reme4 components"
parameters:
type: object
properties: {}
steps:
- backend: health_check_step
- backend: base
name: help
description: "list all registered jobs with their metadata"
parameters:
type: object
properties: {}
steps:
- backend: help_step
- backend: base
name: reindex
description: "wipe the file store and rebuild it from the watcher's tracked files"
parameters:
type: object
properties: {}
steps:
- backend: reindex_step
- backend: base
name: search
description: "hybrid search over file_store: vector + keyword fused via RRF"
parameters:
type: object
properties:
query:
type: string
description: "search query"
limit:
type: integer
description: "max results to return"
default: 5
min_score:
type: number
description: "minimum fused score threshold (RRF scores are small; default 0 disables filter)"
default: 0.0
vector_weight:
type: number
description: "weight for vector results in [0, 1]; keyword weight = 1 - vector_weight"
default: 0.7
candidate_multiplier:
type: number
description: "candidate pool multiplier per branch (capped at 200)"
default: 3.0
expand_links:
type: boolean
description: "attach outlinks/inlinks (with neighbor meta) to each result"
default: true
max_links_per_direction:
type: integer
description: "max neighbors shown per direction per result"
default: 10
required:
- query
steps:
- backend: search_step
- backend: base
name: read
description: "read a markdown file (relative path under working_dir)"
parameters:
type: object
properties:
path:
type: string
description: "relative path under the working_dir (no absolute paths); markdown only"
start_line:
type: integer
description: "Optional, first line to read (1-based, inclusive)"
end_line:
type: integer
description: "Optional, last line to read (1-based, inclusive)"
required:
- path
steps:
- backend: read_step
- backend: stream
name: stream_demo
description: "stream demo job: repeat query 10x and stream char-by-char"
parameters:
type: object
properties:
query:
type: string
description: "query to echo"
repeat:
type: integer
description: "number of times to repeat the query"
default: 10
interval:
type: number
description: "seconds between chunks"
default: 0.1
required:
- query
steps:
- backend: stream_demo_step1
- backend: stream_demo_step2
components:
# 1. tokenizer — no dependencies
tokenizer:
default:
backend: regex
# 2. embedding_model — no dependencies
embedding_model:
default:
backend: openai
model_name: text-embedding-v4
dimensions: 1024
# 3. file_graph — no dependencies
file_graph:
default:
backend: local
# 4. file_parser — no dependencies
file_parser:
default:
backend: default
# 5. keyword_index — depends on tokenizer
keyword_index:
default:
backend: bm25
tokenizer: default
# 6. file_store — depends on embedding_model / keyword_index / file_graph
file_store:
default:
backend: local
store_name: default
# embedding_model: default
embedding_model: ""
keyword_index: default
file_graph: default
# 7. file_watcher — depends on file_store / file_parser
file_watcher:
default:
backend: lite
watch_paths:
- MEMORY.md
- memory
file_store: default
file_parser: default

12
reme4/constants.py Normal file
View file

@ -0,0 +1,12 @@
"""Constants"""
REME_SERVICE_INFO = "REME_SERVICE_INFO"
REME_DEFAULT_HOST = "127.0.0.1"
REME_DEFAULT_PORT = 2333
# CRUD steps: file IO limits and truncation marker (shared across CRUD steps).
DEFAULT_MAX_BYTES = 50 * 1024
MAX_FILE_READ_BYTES = 200 * 1024 * 1024
TRUNCATION_NOTICE_MARKER = "<<TRUNCATION_NOTICE>>"

View file

@ -0,0 +1,9 @@
"""Enumeration"""
from .chunk_enum import ChunkEnum
from .component_enum import ComponentEnum
__all__ = [
"ChunkEnum",
"ComponentEnum",
]

View file

@ -0,0 +1,21 @@
"""Chunk enumeration module."""
from enum import Enum
class ChunkEnum(str, Enum):
"""Enumeration of possible chunk categories for stream processing."""
THINK = "think"
CONTENT = "content"
TOOL_CALL = "tool_call"
TOOL_RESULT = "tool_result"
USAGE = "usage"
ERROR = "error"
DONE = "done"

View file

@ -0,0 +1,37 @@
"""Component enumeration module."""
from enum import Enum
class ComponentEnum(str, Enum):
"""Enumeration of component types for dependency injection and registration."""
BASE = "base"
AS_LLM = "as_llm"
AS_LLM_FORMATTER = "as_llm_formatter"
AS_TOKEN_COUNTER = "as_token_counter"
EMBEDDING_MODEL = "embedding_model"
FILE_PARSER = "file_parser"
FILE_STORE = "file_store"
FILE_GRAPH = "file_graph"
FILE_WATCHER = "file_watcher"
KEYWORD_INDEX = "keyword_index"
SERVICE = "service"
CLIENT = "client"
STEP = "step"
JOB = "job"
TOKENIZER = "tokenizer"

43
reme4/reme.py Normal file
View file

@ -0,0 +1,43 @@
"""ReMe memory management application entry point."""
import asyncio
import sys
from .application import Application
from .components import R
from .config import parse_args, resolve_app_config
from .enumeration import ComponentEnum
from .utils import cli_find_reme, load_env, precheck_start
class ReMe(Application):
"""ReMe memory management application."""
async def call_server(action: str, **kwargs):
"""Call the appropriate server component."""
backend: str = kwargs.pop("backend", "http")
client_cls = R.get(ComponentEnum.CLIENT, backend)
async with client_cls(action=action, **kwargs) as client:
async for chunk in client():
print(chunk, end="", flush=True)
print()
def main():
"""Parse CLI arguments and launch the appropriate mode."""
action, kwargs = parse_args(*sys.argv[1:])
if action == "start":
load_env()
kwargs = resolve_app_config(**kwargs)
if not precheck_start(kwargs.get("service")):
return
ReMe(**kwargs).run_app()
elif action == "find_reme":
cli_find_reme()
else:
asyncio.run(call_server(action, **kwargs))
if __name__ == "__main__":
main()

25
reme4/schema/__init__.py Normal file
View file

@ -0,0 +1,25 @@
"""Schema"""
from .application_config import ApplicationConfig, ComponentConfig, JobConfig
from .emb_node import EmbNode
from .file_chunk import FileChunk
from .file_front_matter import FileFrontMatter
from .file_link import FileLink
from .file_node import FileNode
from .request import Request
from .response import Response
from .stream_chunk import StreamChunk
__all__ = [
"ApplicationConfig",
"ComponentConfig",
"JobConfig",
"EmbNode",
"FileChunk",
"FileFrontMatter",
"FileLink",
"FileNode",
"Request",
"Response",
"StreamChunk",
]

View file

@ -0,0 +1,45 @@
"""Application configuration models."""
import os
from pydantic import BaseModel, ConfigDict, Field
from ..enumeration import ComponentEnum
class ComponentConfig(BaseModel):
"""Base config for a component; extra fields allowed for backend-specific options."""
model_config = ConfigDict(extra="allow")
backend: str = Field(default="", description="Backend implementation class name")
class JobConfig(ComponentConfig):
"""Config for a job — an ordered sequence of step components."""
name: str = Field(default="", description="Unique job identifier")
description: str = Field(default="", description="Human-readable description")
parameters: dict = Field(default_factory=dict, description="Job-level parameters")
steps: list[ComponentConfig] = Field(default_factory=list, description="Ordered step configs")
class ApplicationConfig(BaseModel):
"""Root config for the ReMe application."""
app_name: str = Field(default=os.getenv("APP_NAME", "ReMe"), description="Application display name")
working_dir: str = Field(default=".reme", description="Working directory for runtime files")
metadata_dir: str = Field(default="reme_metadata", description="Subdirectory for ReMe persistent state")
daily_dir: str = Field(default="daily", description="Subdirectory for daily memory")
knowledge_dir: str = Field(default="knowledge", description="Subdirectory for knowledge")
enable_logo: bool = Field(default=True, description="Show ASCII logo on startup")
language: str = Field(default="", description="Default language for LLM interactions")
log_to_console: bool = Field(default=True, description="Log to console")
log_to_file: bool = Field(default=True, description="Log to file")
mcp_servers: dict[str, dict] = Field(default_factory=dict, description="MCP server configs by name")
service: ComponentConfig = Field(default_factory=ComponentConfig, description="Service endpoint config")
jobs: list[JobConfig] = Field(default_factory=list, description="Job definitions")
components: dict[ComponentEnum, dict[str, ComponentConfig]] = Field(
default_factory=dict,
description="Component registry keyed by type then name",
)

34
reme4/schema/emb_node.py Normal file
View file

@ -0,0 +1,34 @@
"""Embedding node — base record carrying text and its vector."""
from uuid import uuid4
import numpy as np
from pydantic import BaseModel, ConfigDict, Field, field_serializer, field_validator
class EmbNode(BaseModel):
"""A text record with an optional embedding vector and metadata."""
model_config = ConfigDict(arbitrary_types_allowed=True)
id: str = Field(default_factory=lambda: uuid4().hex, description="Unique node id")
text: str = Field(default="", description="Text content")
embedding: np.ndarray | None = Field(default=None, description="Embedding vector (float16)")
metadata: dict = Field(default_factory=dict, description="Arbitrary metadata")
@field_validator("embedding", mode="before")
@classmethod
def validate_embedding(cls, v):
"""Coerce list/tuple to float16 ndarray."""
# Coerce list/tuple inputs into a float16 ndarray for compact storage.
if v is None:
return v
return np.array(v, dtype=np.float16)
@field_serializer("embedding")
def serialize_embedding(self, v: np.ndarray | None, _info):
"""Serialize ndarray to a JSON-friendly list."""
# ndarray is not JSON-serializable; emit a plain list.
if v is None:
return None
return v.tolist()

View file

@ -0,0 +1,26 @@
"""File chunk — an embedding node tied to a line range in a file."""
from pydantic import Field
from .emb_node import EmbNode
class FileChunk(EmbNode):
"""A chunk of a file with positional info and per-stage retrieval scores."""
path: str = Field(default="", description="Vault-relative file path")
start_line: int = Field(default=0, description="Inclusive start line (0-based)")
end_line: int = Field(default=0, description="Exclusive end line")
scores: dict[str, float] = Field(default_factory=dict, description="Retrieval scores keyed by stage")
@property
def score(self) -> float:
"""Final aggregated score; 0.0 if not yet computed."""
return self.scores.get("score", 0.0)
def set_hash_id(self):
"""Replace ``id`` with a deterministic hash of (path, range, text)."""
from ..utils import hash_text
self.id = hash_text(" ".join([self.path, str(self.start_line), str(self.end_line), self.text]))
return self

View file

@ -0,0 +1,24 @@
"""FileFrontMatter — parsed Markdown front matter."""
from typing import Any
from pydantic import BaseModel, ConfigDict, Field
class FileFrontMatter(BaseModel):
"""Markdown front matter; unknown keys are preserved as extras."""
model_config = ConfigDict(extra="allow")
title: str = Field(default="", description="Document title")
description: str = Field(default="", description="Document description")
tags: list[str] | None = Field(default=None, description="Tags; None if absent")
@property
def model_extra(self) -> dict[str, Any] | None:
"""Get extra fields set during validation.
Returns:
A dictionary of extra fields, or `None` if `config.extra` is not set to `"allow"`.
"""
return self.__pydantic_extra__

18
reme4/schema/file_link.py Normal file
View file

@ -0,0 +1,18 @@
"""FileLink"""
from pydantic import BaseModel, ConfigDict, Field
class FileLink(BaseModel):
"""file link
[[target_path]]
[[target_path#target_anchor]]
predicate:: [[target_*]]
[predicate:: [[target_*]]]
"""
model_config = ConfigDict(extra="forbid")
source_path: str = Field(default=..., description="source file path relative to working dir")
target_path: str = Field(default=..., description="target file path relative to working dir")
target_anchor: str | None = Field(default=None, description="Heading or block anchor (text after '#')")
predicate: str | None = Field(default=None, description="Dataview-style typed-link predicate")

16
reme4/schema/file_node.py Normal file
View file

@ -0,0 +1,16 @@
"""File node — a file's metadata, links, and chunk references in the graph."""
from pydantic import BaseModel, Field
from .file_front_matter import FileFrontMatter
from .file_link import FileLink
class FileNode(BaseModel):
"""A vault file as a graph node."""
path: str = Field(default=..., description="Vault-relative file path")
st_mtime: float = Field(default=..., description="Filesystem mtime (seconds)")
links: list[FileLink] = Field(default_factory=list, description="Outgoing wikilinks")
chunk_ids: list[str] = Field(default_factory=list, description="Owned FileChunk ids")
front_matter: FileFrontMatter = Field(default_factory=FileFrontMatter, description="Parsed front matter")

11
reme4/schema/request.py Normal file
View file

@ -0,0 +1,11 @@
"""Request schema for service endpoints."""
from pydantic import BaseModel, ConfigDict, Field
class Request(BaseModel):
"""Incoming service request; extra fields are allowed for endpoint-specific payloads."""
model_config = ConfigDict(extra="allow")
metadata: dict = Field(default_factory=dict, description="Request metadata for context")

15
reme4/schema/response.py Normal file
View file

@ -0,0 +1,15 @@
"""Response schema for service endpoints and LLM calls."""
from typing import Any
from pydantic import BaseModel, ConfigDict, Field
class Response(BaseModel):
"""Standard response envelope; extra fields allowed for endpoint-specific output."""
model_config = ConfigDict(extra="allow")
answer: str | Any = Field(default="", description="Response content or result data")
success: bool = Field(default=True, description="Whether the operation succeeded")
metadata: dict = Field(default_factory=dict, description="Additional context and diagnostics")

View file

@ -0,0 +1,14 @@
"""Stream chunk schema for incremental responses (e.g. LLM streaming)."""
from pydantic import BaseModel, Field
from ..enumeration import ChunkEnum
class StreamChunk(BaseModel):
"""A single chunk in a streaming response sequence."""
chunk_type: ChunkEnum = Field(default=ChunkEnum.CONTENT, description="Type of chunk content")
chunk: str | dict | list = Field(default="", description="Chunk payload")
done: bool = Field(default=False, description="Whether this is the final chunk")
metadata: dict = Field(default_factory=dict, description="Chunk metadata")

11
reme4/steps/__init__.py Normal file
View file

@ -0,0 +1,11 @@
"""steps"""
from . import common
from . import crud
from .base_step import BaseStep
__all__ = [
"common",
"crud",
"BaseStep",
]

188
reme4/steps/base_step.py Normal file
View file

@ -0,0 +1,188 @@
"""Base step class for LLM workflow execution."""
import copy
from abc import abstractmethod, ABC
from pathlib import Path
from typing import TypeVar, TYPE_CHECKING
from agentscope.formatter import FormatterBase
from agentscope.message import TextBlock
from agentscope.model import ChatModelBase
from agentscope.token import TokenCounterBase
from agentscope.tool import Toolkit, ToolResponse
from ..components.embedding import BaseEmbeddingModel
from ..components.file_parser import BaseFileParser
from ..components.file_store import BaseFileStore
from ..components.file_watcher import BaseFileWatcher
from ..components.prompt_handler import PromptHandler
from ..components.runtime_context import RuntimeContext
from ..enumeration import ComponentEnum
from ..schema import Response
from ..utils import get_logger
if TYPE_CHECKING:
from ..components import ApplicationContext
from ..components.job import BaseJob
T = TypeVar("T")
class BaseStep(ABC):
"""Composable unit of an LLM workflow."""
component_type = ComponentEnum.STEP
def __new__(cls, *args, **kwargs):
# Snapshot init args so copy() can rebuild an equivalent instance later.
instance = object.__new__(cls)
instance._init_args = copy.copy(args)
instance._init_kwargs = copy.copy(kwargs)
return instance
def __init__(
self,
name: str | None = None,
backend: str = "",
app_context: "ApplicationContext | None" = None,
language: str = "",
prompt_dict: dict[str, str] | None = None,
input_mapping: dict[str, str] | None = None,
output_mapping: dict[str, str] | None = None,
**kwargs,
):
super().__init__()
self.name: str = name or self.__class__.__name__
self.backend: str = backend
self.app_context: "ApplicationContext | None" = app_context
self.language: str = language
self.input_mapping = input_mapping
self.output_mapping = output_mapping
self.kwargs: dict = kwargs
self.context: RuntimeContext | None = None
self.logger = get_logger()
if hasattr(self.logger, "bind"):
self.logger = self.logger.bind(component=self.name)
# Load class-level prompts first, then overlay caller-provided overrides.
self.prompt = PromptHandler(language=self.language)
self.prompt.load_prompt_by_class(self.__class__).load_prompt_dict(prompt_dict)
@abstractmethod
async def execute(self):
"""Run the step's logic against ``self.context``."""
async def __call__(self, context: RuntimeContext | None = None, **kwargs):
# Build runtime context, then apply key remapping around execute().
self.context = RuntimeContext.from_context(context, **kwargs)
assert self.context is not None
if self.input_mapping:
self.context.apply_mapping(self.input_mapping)
result = await self.execute()
if self.output_mapping:
self.context.apply_mapping(self.output_mapping)
return result
@property
def working_path(self) -> Path:
"""Resolved working directory from app context or cwd."""
if self.app_context is None:
return Path.cwd()
return Path(self.app_context.app_config.working_dir)
def _resolve(
self,
key: str,
base_cls: type[T],
comp_enum: ComponentEnum,
attr: str | None = None,
) -> T:
"""Return a kwargs-supplied instance, or look one up by name in the app registry."""
# 1. Step init kwargs, 2. Runtime context (run_job kwargs), 3. App registry by name.
for source in (self.kwargs, self.context or {}):
value = source.get(key)
if isinstance(value, base_cls):
return value
name = self.kwargs.get(key, "default")
assert self.app_context is not None
comp = self.app_context.components[comp_enum][name]
return getattr(comp, attr) if attr else comp
@property
def as_llm(self) -> ChatModelBase:
"""Return the chat model component."""
return self._resolve("as_llm", ChatModelBase, ComponentEnum.AS_LLM, "model")
@property
def as_llm_formatter(self) -> FormatterBase:
"""Return the LLM formatter component."""
return self._resolve("as_llm_formatter", FormatterBase, ComponentEnum.AS_LLM_FORMATTER, "formatter")
@property
def as_token_counter(self) -> TokenCounterBase:
"""Return the token counter component."""
return self._resolve("as_token_counter", TokenCounterBase, ComponentEnum.AS_TOKEN_COUNTER, "token_counter")
@property
def file_parser(self) -> BaseFileParser:
"""Return the file parser component."""
return self._resolve("file_parser", BaseFileParser, ComponentEnum.FILE_PARSER)
@property
def file_store(self) -> BaseFileStore:
"""Return the file store component."""
return self._resolve("file_store", BaseFileStore, ComponentEnum.FILE_STORE)
@property
def embedding(self) -> BaseEmbeddingModel:
"""Return the embedding model component."""
return self._resolve("embedding", BaseEmbeddingModel, ComponentEnum.EMBEDDING_MODEL)
@property
def file_watcher(self) -> BaseFileWatcher:
"""Return the file watcher component."""
return self._resolve("file_watcher", BaseFileWatcher, ComponentEnum.FILE_WATCHER)
def prompt_format(self, prompt_name: str, **kwargs) -> str:
"""Format a named prompt template with the given kwargs."""
return self.prompt.prompt_format(prompt_name=prompt_name, **kwargs)
def get_prompt(self, prompt_name: str) -> str:
"""Return a named prompt template as-is."""
return self.prompt.get_prompt(prompt_name=prompt_name)
def copy(self, **kwargs) -> "BaseStep":
"""Construct a new instance from the original init args, applying overrides."""
return self.__class__(*self._init_args, **{**self._init_kwargs, **kwargs})
def get_job(self, name: str) -> "BaseJob | None":
"""Return a job by name."""
if self.app_context is None:
raise RuntimeError("Cannot get job without an app context")
return self.app_context.jobs.get(name)
async def run_job(self, name: str, **kwargs) -> Response:
"""Execute a job by name and kwargs, return the final response."""
job: "BaseJob | None" = self.get_job(name)
if job is None:
raise RuntimeError(f"Job {name} not found")
return await job(**kwargs)
def add_as_tool(self, toolkit: Toolkit, job_name: str) -> None:
"""Add the step as a tool to the toolkit."""
job: "BaseJob | None" = self.get_job(job_name)
if job is None:
raise RuntimeError(f"Job {job_name} not found")
async def run_job(**kwargs) -> ToolResponse:
response = await job(**kwargs)
return ToolResponse(content=[TextBlock(type="text", text=response.answer)])
toolkit.register_tool_function(
tool_func=run_job,
func_name=job_name,
func_description=job.description,
json_schema=job.parameters,
)

View file

@ -0,0 +1,21 @@
"""Common steps."""
from .demo import DemoEchoStep1, DemoEchoStep2
from .health_check import HealthCheckStep
from .help import HelpStep
from .reindex import ReindexStep
from .search import SearchStep
from .stream_demo import StreamDemoStep1, StreamDemoStep2
from .version import VersionStep
__all__ = [
"DemoEchoStep1",
"DemoEchoStep2",
"HealthCheckStep",
"HelpStep",
"ReindexStep",
"SearchStep",
"StreamDemoStep1",
"StreamDemoStep2",
"VersionStep",
]

View file

@ -0,0 +1,53 @@
"""Demo steps for smoke-testing the application stack."""
from ..base_step import BaseStep
from ...components import R
@R.register("demo_echo_step1")
class DemoEchoStep1(BaseStep):
"""Read query/min_score from context, normalize, and write back for Step2."""
async def execute(self):
assert self.context is not None
query = self.context.get("query", "")
min_score = self.context.get("min_score", 0.5)
self.logger.info(f"[{self.name}] query={query!r}, min_score={min_score}")
processed_query = query.strip().lower()
adjusted_min_score = float(min_score) * 0.9
self.context["processed_query"] = processed_query
self.context["adjusted_min_score"] = adjusted_min_score
return self.context.response
@R.register("demo_echo_step2")
class DemoEchoStep2(BaseStep):
"""Consume Step1's outputs from context and emit the final response."""
async def execute(self):
assert self.context is not None
query = self.context.get("query", "")
min_score = self.context.get("min_score", 0.5)
processed_query = self.context.get("processed_query", "")
adjusted_min_score = self.context.get("adjusted_min_score", min_score)
self.logger.info(
f"[{self.name}] query={query!r}, min_score={min_score}, "
f"processed_query={processed_query!r}, adjusted_min_score={adjusted_min_score}",
)
self.context.response.answer = f"echo: {processed_query} (min_score={adjusted_min_score})"
self.context.response.metadata.update(
{
"step": self.name,
"query": query,
"min_score": min_score,
"processed_query": processed_query,
"adjusted_min_score": adjusted_min_score,
},
)
return self.context.response

View file

@ -0,0 +1,164 @@
"""Return a concise health check snapshot of ReMe runtime components."""
import sys
from collections.abc import Mapping
import numpy as np
from ..base_step import BaseStep
from ... import __version__
from ...components import R
from ...enumeration import ComponentEnum
def _deep_size(obj, _seen: set | None = None) -> int:
"""Recursive sizeof; uses ndarray.nbytes for numpy and walks containers / __dict__."""
if _seen is None:
_seen = set()
obj_id = id(obj)
if obj_id in _seen:
return 0
_seen.add(obj_id)
if isinstance(obj, np.ndarray):
return int(obj.nbytes) + sys.getsizeof(obj)
size = sys.getsizeof(obj)
if isinstance(obj, (str, bytes, bytearray, int, float, bool, type(None))):
return size
if isinstance(obj, Mapping):
size += sum(_deep_size(k, _seen) + _deep_size(v, _seen) for k, v in obj.items())
elif isinstance(obj, (list, tuple, set, frozenset)):
size += sum(_deep_size(item, _seen) for item in obj)
elif hasattr(obj, "__dict__"):
size += _deep_size(vars(obj), _seen)
elif hasattr(obj, "__slots__"):
for slot in obj.__slots__:
if hasattr(obj, slot):
size += _deep_size(getattr(obj, slot), _seen)
return size
def _mb_str(*objs) -> str:
"""Return summed deep size of objs formatted as 'X.XX MB'."""
seen: set = set()
total = sum(_deep_size(o, seen) for o in objs)
return f"{total / (1024 * 1024):.2f} MB"
def _embedding_status(comp) -> dict:
return {
"is_started": comp.is_started,
"is_healthy": getattr(comp, "is_healthy", None),
"model_name": getattr(comp, "model_name", None),
"dimensions": getattr(comp, "dimensions", None),
"cache_size": len(getattr(comp, "_embedding_cache", {}) or {}),
"memory": _mb_str(getattr(comp, "_embedding_cache", {}) or {}),
}
def _file_graph_status(comp) -> dict:
# Nx backend: single _graph attribute holds nodes/edges, virtuals are nodes without "node" payload.
g = getattr(comp, "_graph", None)
if g is not None:
n_real = sum(1 for _, d in g.nodes(data=True) if "node" in d)
return {
"is_started": comp.is_started,
"n_nodes": n_real,
"n_edges": g.number_of_edges(),
"n_virtual": g.number_of_nodes() - n_real,
"memory": _mb_str(g),
}
# Local backend: separate dicts for nodes, resolved inverse edges, and pending edges.
nodes = getattr(comp, "_nodes", {}) or {}
inverse = getattr(comp, "_inverse", {}) or {}
pending = getattr(comp, "_pending", {}) or {}
return {
"is_started": comp.is_started,
"n_nodes": len(nodes),
"n_edges": sum(len(s) for s in inverse.values()),
"n_pending": sum(len(s) for s in pending.values()),
"memory": _mb_str(nodes, inverse, pending),
}
def _file_store_status(comp) -> dict:
chunks = getattr(comp, "file_chunks", {}) or {}
return {
"is_started": comp.is_started,
"n_chunks": len(chunks),
"n_chunks_with_embedding": sum(1 for c in chunks.values() if getattr(c, "embedding", None) is not None),
"memory": _mb_str(chunks),
}
def _file_watcher_status(comp) -> dict:
bg = getattr(comp, "_background_task", None)
return {
"is_started": comp.is_started,
"background_running": bool(bg and not bg.done()),
"watch_paths": [str(p) for p in (getattr(comp, "watch_paths", []) or [])],
}
def _keyword_index_status(comp) -> dict:
return {
"is_started": comp.is_started,
"n_docs": getattr(comp, "n_docs", None),
"vocab_size": len(getattr(comp, "vocab", {}) or {}),
"memory": _mb_str(
getattr(comp, "vocab", {}) or {},
getattr(comp, "inverted_index", {}) or {},
getattr(comp, "doc_meta", {}) or {},
getattr(comp, "_idf_cache", {}) or {},
),
}
_HANDLERS = {
ComponentEnum.EMBEDDING_MODEL: _embedding_status,
ComponentEnum.FILE_GRAPH: _file_graph_status,
ComponentEnum.FILE_STORE: _file_store_status,
ComponentEnum.FILE_WATCHER: _file_watcher_status,
ComponentEnum.KEYWORD_INDEX: _keyword_index_status,
}
def _is_status_healthy(ctype: ComponentEnum, status: dict) -> bool:
"""Per-component health rule. Unstarted = unhealthy; type-specific extras checked."""
if not status.get("is_started"):
return False
if ctype is ComponentEnum.EMBEDDING_MODEL and status.get("is_healthy") is False:
return False
if ctype is ComponentEnum.FILE_WATCHER and not status.get("background_running"):
return False
return True
@R.register("health_check_step")
class HealthCheckStep(BaseStep):
"""Collect a concise health check snapshot of the relevant components."""
async def execute(self):
assert self.context is not None
components: dict = {}
healthy = True
if self.app_context is not None:
for ctype, handler in _HANDLERS.items():
comp_map = self.app_context.components.get(ctype, {})
bucket = {}
for name, comp in comp_map.items():
s = handler(comp)
bucket[name] = s
if not _is_status_healthy(ctype, s):
healthy = False
components[ctype.value] = bucket
health = {"version": __version__, "healthy": healthy, "components": components}
self.logger.info(f"[{self.name}] health collected: {health}")
status_emoji = "✅" if healthy else "❌"
self.context.response.answer = f"{status_emoji} ReMe v{__version__} - {'healthy' if healthy else 'unhealthy'}"
self.context.response.metadata["health"] = health
return self.context.response

View file

@ -0,0 +1,42 @@
"""Return a one-line summary of every registered job for LLM consumption."""
from ..base_step import BaseStep
from ...components import R
@R.register("help_step")
class HelpStep(BaseStep):
"""List all registered jobs (excluding self) as compact one-liners for an LLM."""
@staticmethod
def _format_params(parameters: dict) -> str:
props = (parameters or {}).get("properties") or {}
if not props:
return "no args"
required = set((parameters or {}).get("required") or [])
parts = []
for pname, pschema in props.items():
ptype = pschema.get("type", "any")
if pname in required:
parts.append(f"{pname}:{ptype}*")
elif "default" in pschema:
parts.append(f"{pname}:{ptype}={pschema['default']}")
else:
parts.append(f"{pname}:{ptype}")
return ", ".join(parts)
async def execute(self):
assert self.context is not None
lines = []
if self.app_context is not None:
for name, job in self.app_context.jobs.items():
if name == "help":
continue
lines.append(f"🛠️ `{name}` — {job.description} 📥 {self._format_params(job.parameters)}")
self.logger.info(f"[{self.name}] returning {len(lines)} jobs")
self.context.response.answer = "\n".join(lines)
self.context.response.metadata["job_count"] = len(lines)
return self.context.response

View file

@ -0,0 +1,24 @@
"""Wipe the file store and rebuild it from the watcher's tracked files."""
from ..base_step import BaseStep
from ...components import R
@R.register("reindex_step")
class ReindexStep(BaseStep):
"""Full re-index: stop watcher, clear store, sync from disk, then restart."""
async def execute(self):
assert self.context is not None
await self.file_watcher.close()
try:
await self.file_store.clear()
counts = await self.file_watcher.update_store()
finally:
await self.file_watcher.start()
self.logger.info(f"[{self.name}] reindexed {counts}")
self.context.response.answer = f"🔄 Reindexed {counts['added']} file(s)"
self.context.response.metadata["counts"] = counts
return self.context.response

View file

@ -0,0 +1,227 @@
"""Hybrid search over file_store using RRF fusion of vector + keyword results."""
import asyncio
from ..base_step import BaseStep
from ...components import R
from ...schema import FileChunk, FileLink, FileNode
_RRF_K = 60
_MAX_CANDIDATES = 200
@R.register("search_step")
class SearchStep(BaseStep):
"""Hybrid search: run vector + keyword in parallel, fuse via RRF, filter, truncate."""
@staticmethod
def _rrf_merge(
vector: list[FileChunk],
keyword: list[FileChunk],
vector_weight: float,
) -> list[FileChunk]:
"""Fuse two ranked lists with Reciprocal Rank Fusion, keyed by chunk.id."""
text_weight = 1.0 - vector_weight
merged: dict[str, FileChunk] = {}
for rank, chunk in enumerate(vector, start=1):
contrib = vector_weight / (_RRF_K + rank)
c = chunk.model_copy(deep=False)
c.scores = {**chunk.scores, "vector": chunk.scores.get("vector", chunk.score), "score": contrib}
merged[c.id] = c
for rank, chunk in enumerate(keyword, start=1):
contrib = text_weight / (_RRF_K + rank)
existing = merged.get(chunk.id)
if existing is not None:
existing.scores = {
**existing.scores,
"keyword": chunk.scores.get("keyword", chunk.score),
"score": existing.scores["score"] + contrib,
}
else:
c = chunk.model_copy(deep=False)
c.scores = {**chunk.scores, "keyword": chunk.scores.get("keyword", chunk.score), "score": contrib}
merged[c.id] = c
results = list(merged.values())
results.sort(key=lambda r: r.score, reverse=True)
return results
@staticmethod
def _format_scores(scores: dict[str, float], hybrid: bool) -> str:
"""Format scores for the answer line: always show fused; show per-branch when hybrid."""
parts = [f"score={scores.get('score', 0.0):.4f}"]
if hybrid:
for k in ("vector", "keyword"):
v = scores.get(k)
parts.append(f"{k}={v:.4f}" if v is not None else f"{k}=-")
return " ".join(parts)
@staticmethod
def _group_by_neighbor(links: list[FileLink], key_attr: str) -> dict[str, list[dict]]:
"""Group edges by neighbor path (insertion-ordered), each value a list of {predicate, anchor}."""
out: dict[str, list[dict]] = {}
for lnk in links:
neighbor = getattr(lnk, key_attr)
if not neighbor:
continue
out.setdefault(neighbor, []).append(
{"predicate": lnk.predicate, "anchor": lnk.target_anchor},
)
return out
@staticmethod
def _node_meta(node: FileNode | None) -> dict:
"""Extract a compact meta dict (title/description/tags) from a FileNode."""
if node is None:
return {}
fm = node.front_matter
meta: dict = {}
if fm.title:
meta["title"] = fm.title
if fm.description:
meta["description"] = fm.description
if fm.tags:
meta["tags"] = list(fm.tags)
return meta
@staticmethod
def _format_meta_inline(meta: dict) -> str:
"""One-line render of node meta for the answer; '(no meta)' when empty."""
parts = []
if "title" in meta:
parts.append(f'title="{meta["title"]}"')
if "tags" in meta:
parts.append(f"tags={meta['tags']}")
return " ".join(parts) if parts else "(no meta)"
@staticmethod
def _format_via(edge: dict) -> str:
"""Render a single (predicate, anchor) edge as a 'via ...' descriptor."""
bits = []
if edge.get("predicate"):
bits.append(f"predicate={edge['predicate']}")
if edge.get("anchor"):
bits.append(f"anchor=#{edge['anchor']}")
return ", ".join(bits) if bits else "plain"
async def _expand_links(
self,
chunk_paths: list[str],
max_per_direction: int,
) -> dict[str, dict]:
"""Fetch out/in links for each chunk path; attach neighbor meta. Returns per-path expansion."""
if not chunk_paths:
return {}
out_lists, in_lists = await asyncio.gather(
asyncio.gather(*(self.file_store.get_outlinks(p) for p in chunk_paths)),
asyncio.gather(*(self.file_store.get_inlinks(p) for p in chunk_paths)),
)
# Pre-group + cap per direction so we only fetch meta for displayed neighbors.
out_grouped = [
dict(list(self._group_by_neighbor(outs, "target_path").items())[:max_per_direction]) for outs in out_lists
]
in_grouped = [
dict(list(self._group_by_neighbor(ins, "source_path").items())[:max_per_direction]) for ins in in_lists
]
neighbor_paths = sorted({n for g in out_grouped for n in g} | {n for g in in_grouped for n in g})
nodes = await self.file_store.get_nodes(neighbor_paths) if neighbor_paths else []
meta_by_path = {n.path: self._node_meta(n) for n in nodes}
def _attach(grouped: dict[str, list[dict]]) -> list[dict]:
return [
{"path": npath, "meta": meta_by_path.get(npath, {}), "edges": edges} for npath, edges in grouped.items()
]
return {
cp: {"outlinks": _attach(og), "inlinks": _attach(ig)}
for cp, og, ig in zip(chunk_paths, out_grouped, in_grouped)
}
@classmethod
def _render_expansion_lines(cls, expansion: dict) -> list[str]:
"""Render outlinks/inlinks blocks for one chunk path; return zero or more indented lines."""
lines: list[str] = []
for direction, arrow, items in (
("outlinks", "→", expansion.get("outlinks") or []),
("inlinks", "←", expansion.get("inlinks") or []),
):
if not items:
continue
lines.append(f" {direction} ({len(items)}):")
for item in items:
lines.append(f" {arrow} {item['path']} {cls._format_meta_inline(item['meta'])}")
for edge in item["edges"]:
lines.append(f" via {cls._format_via(edge)}")
return lines
async def execute(self):
assert self.context is not None
query: str = (self.context.get("query", "") or "").strip()
limit: int = int(self.context.get("limit", 5))
min_score: float = float(self.context.get("min_score", 0.0))
vector_weight: float = float(self.context.get("vector_weight", 0.7))
candidate_multiplier: float = float(self.context.get("candidate_multiplier", 3.0))
expand_links: bool = bool(self.context.get("expand_links", True))
max_links_per_direction: int = int(self.context.get("max_links_per_direction", 10))
assert query, "query cannot be empty"
assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be in [0, 1], got {vector_weight}"
assert limit > 0, f"limit must be positive, got {limit}"
candidates = min(_MAX_CANDIDATES, max(1, int(limit * candidate_multiplier)))
search_filter: dict = self.context.get("search_filter", {}) or {}
vector_results, keyword_results = await asyncio.gather(
self.file_store.vector_search(query, candidates, search_filter),
self.file_store.keyword_search(query, candidates, search_filter),
)
self.logger.info(
f"[{self.name}] query={query!r} candidates={candidates} "
f"vector_hits={len(vector_results)} keyword_hits={len(keyword_results)}",
)
hybrid = bool(vector_results) and bool(keyword_results)
if not vector_results and not keyword_results:
fused: list[FileChunk] = []
elif not keyword_results:
fused = vector_results
elif not vector_results:
fused = keyword_results
else:
fused = self._rrf_merge(vector_results, keyword_results, vector_weight)
if min_score > 0.0:
fused = [c for c in fused if c.score >= min_score]
fused = fused[:limit]
unique_paths = list(dict.fromkeys(c.path for c in fused))
link_expansion: dict[str, dict] = (
await self._expand_links(unique_paths, max_links_per_direction) if expand_links else {}
)
answer_lines: list[str] = []
for c in fused:
answer_lines.append(
f"========== {c.path}:{c.start_line}-{c.end_line} "
f"[{self._format_scores(c.scores, hybrid)}] ==========\n{c.text}",
)
answer_lines.extend(self._render_expansion_lines(link_expansion.get(c.path, {})))
self.context.response.answer = "\n".join(answer_lines)
self.context.response.metadata["results"] = [
c.model_dump(exclude_none=True, exclude={"embedding"}) for c in fused
]
self.context.response.metadata["link_expansion"] = link_expansion
self.context.response.metadata["counts"] = {
"vector": len(vector_results),
"keyword": len(keyword_results),
"returned": len(fused),
"hybrid": hybrid,
}
return self.context.response

View file

@ -0,0 +1,42 @@
"""Streaming demo steps: step1 prepares text, step2 streams it char-by-char."""
import asyncio
from ..base_step import BaseStep
from ...components import R
from ...enumeration import ChunkEnum
@R.register("stream_demo_step1")
class StreamDemoStep1(BaseStep):
"""Read query from context, repeat it 10x, write back for Step2 to stream."""
async def execute(self):
assert self.context is not None
query = self.context.get("query", "")
repeat = int(self.context.get("repeat", 10))
stream_text = (query * repeat) if query else ""
self.logger.info(f"[{self.name}] query={query!r}, repeat={repeat}, len={len(stream_text)}")
self.context["stream_text"] = stream_text
return self.context.response
@R.register("stream_demo_step2")
class StreamDemoStep2(BaseStep):
"""Stream stream_text char-by-char as CONTENT chunks with 0.1s pacing."""
async def execute(self):
assert self.context is not None
stream_text: str = self.context.get("stream_text", "")
interval = float(self.context.get("interval", 0.1))
self.logger.info(f"[{self.name}] streaming {len(stream_text)} chars, interval={interval}s")
for ch in stream_text:
await self.context.add_stream_string(ch, ChunkEnum.CONTENT)
await asyncio.sleep(interval)
return self.context.response

View file

@ -0,0 +1,19 @@
"""Return the package version."""
from ..base_step import BaseStep
from ...components import R
@R.register("version_step")
class VersionStep(BaseStep):
"""Emit reme4.__version__ as the response answer."""
async def execute(self):
assert self.context is not None
from ... import __version__
self.logger.info(f"[{self.name}] version={__version__}")
self.context.response.answer = __version__
self.context.response.metadata["version"] = __version__
return self.context.response

View file

@ -0,0 +1,7 @@
"""CRUD steps for markdown files under the working_dir."""
from .read import ReadStep
__all__ = [
"ReadStep",
]

View file

@ -0,0 +1,109 @@
"""Shared filesystem helpers for CRUD steps (path gating, safe read, truncation)."""
from pathlib import Path
import aiofiles
import aiofiles.os
from ...constants import DEFAULT_MAX_BYTES, MAX_FILE_READ_BYTES, TRUNCATION_NOTICE_MARKER
from ...utils import get_logger
logger = get_logger()
def resolve_path(working_path: Path, raw: str) -> tuple[Path | None, str | None]:
"""Resolve a relative `path=` argument under self.working_path.
Rules:
- the caller supplies the full relative path under ``self.working_path``;
absolute paths are rejected.
Returns ``(abs_path, None)`` on success, or ``(None, error_message)`` on failure.
Filetype-specific gating (e.g. markdown-only / suffix auto-append) is
layered on top by callers — see ``reme4/steps/crud/_file_io.py::gate_md``.
"""
if not raw or not str(raw).strip():
return None, "`path` is required"
s = str(raw).strip()
p = Path(s)
if p.is_absolute():
logger.info("absolute path detected, recommmending relative paths")
return p, None
return working_path / p, None
def gate_md(target: Path, raw: str) -> tuple[Path | None, str | None]:
"""Markdown-only gate: auto-append `.md` when no suffix; reject any non-`.md` suffix.
Layered on top of ``BaseStep.resolve_path`` to keep filetype-specific rules
out of the generic path resolver.
"""
if target.suffix == "":
return target.with_suffix(".md"), None
if target.suffix.lower() != ".md":
return None, (f"path {raw!r} is not a markdown file; this command only supports .md files")
return target, None
async def read_file_safe(file_path, max_bytes: int = MAX_FILE_READ_BYTES) -> str:
"""Read file with utf-8-sig (BOM-tolerant), fallback to errors='ignore'."""
stat = await aiofiles.os.stat(str(file_path))
read_size = min(stat.st_size, max_bytes)
try:
async with aiofiles.open(str(file_path), "r", encoding="utf-8-sig") as f:
return await f.read(read_size)
except UnicodeDecodeError:
async with aiofiles.open(
str(file_path),
"r",
encoding="utf-8-sig",
errors="ignore",
) as f:
return await f.read(read_size)
def truncate_text_output(
text: str,
*,
start_line: int = 1,
total_lines: int = 0,
max_bytes: int = DEFAULT_MAX_BYTES,
file_path: str | None = None,
encoding: str = "utf-8",
) -> str:
"""Truncate text by bytes preserving line integrity; append a continuation notice.
See qwenpaw `tools/utils.py` for the same semantics. Returns text unchanged when
it fits within max_bytes, when max_bytes <= 0, or when the last line itself
exceeds max_bytes (unhandled edge case).
"""
if not text or max_bytes <= 0:
return text
try:
text_bytes = text.encode(encoding)
if len(text_bytes) <= max_bytes:
return text
truncated = text_bytes[:max_bytes]
result = truncated.decode(encoding, errors="ignore")
newline_count = result.count("\n")
next_line = start_line + max(1, newline_count)
if next_line <= total_lines:
read_from = next_line
elif start_line < total_lines:
read_from = total_lines
else:
return result
notice = (
TRUNCATION_NOTICE_MARKER + f"\nThe output above was truncated."
f"\nThe full content is saved to the file and contains {total_lines} lines in total."
f"\nThis excerpt starts at line {start_line} and covers the next {max_bytes} bytes."
f"\nIf the current content is not enough, call `read` with file={file_path or ''} "
f"start_line={read_from} to read more."
)
return result + notice
except Exception:
logger.warning("truncate_text_output failed, returning original text", exc_info=True)
return text

91
reme4/steps/crud/read.py Normal file
View file

@ -0,0 +1,91 @@
"""Read a markdown file from the vault, with line-range slicing and byte-truncation."""
from ._file_io import (
gate_md,
resolve_path,
read_file_safe,
truncate_text_output,
)
from ..base_step import BaseStep
from ...components import R
@R.register("read_step")
class ReadStep(BaseStep):
"""Read a markdown file. Optional `start_line`/`end_line` for ranged reads."""
def _fail(self, message: str, **meta) -> None:
assert self.context is not None
self.context.response.success = False
self.context.response.answer = f"Error: {message}"
if meta:
self.context.response.metadata.update(meta)
async def execute(self): # pylint: disable=too-many-return-statements
assert self.context is not None
raw = str(self.context.get("path") or "")
start_line = self.context.get("start_line")
end_line = self.context.get("end_line")
target, err = resolve_path(self.working_path, raw)
if err:
self._fail(err)
return None
target, err = gate_md(target, raw)
if err:
self._fail(err)
return None
for label, value in (("start_line", start_line), ("end_line", end_line)):
if value is None:
continue
try:
int(value)
except (TypeError, ValueError):
self._fail(f"{label} must be an integer, got {value!r}")
return None
if not target.exists():
self._fail(f"file {target} does not exist", path=str(target))
return None
if not target.is_file():
self._fail(f"path {target} is not a file", path=str(target))
return None
try:
content = await read_file_safe(target)
except Exception as e:
self._fail(f"read failed: {e}", path=str(target))
return None
all_lines = content.split("\n")
total = len(all_lines)
s = max(1, int(start_line) if start_line is not None else 1)
e = min(total, int(end_line) if end_line is not None else total)
if s > total:
self._fail(
f"start_line {s} exceeds file length ({total} lines)",
path=str(target),
total_lines=total,
)
return None
if s > e:
self._fail(f"start_line ({s}) > end_line ({e})", path=str(target))
return None
selected = "\n".join(all_lines[s - 1 : e])
text = truncate_text_output(
selected,
start_line=s,
total_lines=total,
file_path=str(target),
)
self.context.response.success = True
self.context.response.answer = text
self.logger.info(
f"[{self.name}] read path={target} lines={s}-{e}/{total} bytes={len(text.encode('utf-8'))}",
)
return self.context.response

31
reme4/utils/__init__.py Normal file
View file

@ -0,0 +1,31 @@
"""Utility modules."""
from .common_utils import (
hash_text,
execute_stream_task,
mock_reme_server,
call_action,
call_and_check,
)
from .env_utils import load_env
from .logger_utils import get_logger
from .logo_utils import print_logo
from .service_utils import find_reme, locate_reme, precheck_start, cli_find_reme
from .similarity_utils import cosine_similarity, batch_cosine_similarity
__all__ = [
"hash_text",
"execute_stream_task",
"mock_reme_server",
"call_action",
"call_and_check",
"load_env",
"get_logger",
"print_logo",
"find_reme",
"locate_reme",
"precheck_start",
"cli_find_reme",
"cosine_similarity",
"batch_cosine_similarity",
]

249
reme4/utils/common_utils.py Normal file
View file

@ -0,0 +1,249 @@
"""Common utilities: hashing and async stream task execution."""
import asyncio
import hashlib
import json
import socket
import subprocess
import sys
import time
from collections.abc import AsyncGenerator, Callable
from contextlib import asynccontextmanager
from typing import Any, Literal
from .logger_utils import get_logger
from ..constants import REME_DEFAULT_HOST, REME_DEFAULT_PORT
from ..enumeration import ChunkEnum
from ..schema import StreamChunk
def hash_text(text: str, encoding: str = "utf-8") -> str:
"""Return SHA-256 hex digest of text."""
return hashlib.sha256(text.encode(encoding)).hexdigest()
def _format_chunk(
chunk: StreamChunk,
output_format: Literal["str", "bytes", "chunk"],
) -> str | bytes | StreamChunk:
"""Render a StreamChunk in the requested transport format."""
if output_format == "chunk":
return chunk
data = "data:[DONE]\n\n" if chunk.done else f"data:{chunk.model_dump_json()}\n\n"
return data.encode() if output_format == "bytes" else data
async def execute_stream_task(
stream_queue: asyncio.Queue[StreamChunk],
task: asyncio.Task[Any],
task_name: str | None = None,
output_format: Literal["str", "bytes", "chunk"] = "str",
) -> AsyncGenerator[str | bytes | StreamChunk, None]:
"""Yield chunks from stream_queue while monitoring task; cancels task on exit.
output_format: "str"/"bytes" emit SSE frames, "chunk" emits raw StreamChunk.
"""
logger = get_logger()
consumer: asyncio.Task[StreamChunk] | None = None
try:
while True:
consumer = get_chunk = asyncio.create_task(stream_queue.get())
done, _pending = await asyncio.wait({get_chunk, task}, return_when=asyncio.FIRST_COMPLETED)
# Producer still running — relay the next chunk and continue.
if task not in done:
chunk = get_chunk.result()
yield _format_chunk(chunk, output_format)
if chunk.done:
return
continue
# Producer finished. Capture any pending chunk, then stop the consumer wait
# so we can inspect task state safely.
pending_chunk: StreamChunk | None = None
if get_chunk in done:
pending_chunk = get_chunk.result()
else:
get_chunk.cancel()
try:
await get_chunk
except asyncio.CancelledError:
pass
# Surface task failure first — an exception trumps trailing data.
if task.cancelled():
msg = f"Task cancelled: {task_name}" if task_name else "Task cancelled"
raise asyncio.CancelledError(msg)
exc = task.exception()
if exc is not None:
log_msg = f"Task error in {task_name}: {exc}" if task_name else f"Task error: {exc}"
logger.error(log_msg, exc_info=exc)
raise exc
# Producer ended cleanly — flush pending + drain queue so no chunk is lost,
# then emit the terminal sentinel.
if pending_chunk is not None:
yield _format_chunk(pending_chunk, output_format)
if pending_chunk.done:
return
while not stream_queue.empty():
chunk = stream_queue.get_nowait()
yield _format_chunk(chunk, output_format)
if chunk.done:
return
yield _format_chunk(StreamChunk(chunk_type=ChunkEnum.DONE, chunk="", done=True), output_format)
return
finally:
# Cancel consumer wait if still pending (e.g. on consumer aclose).
if consumer is not None and not consumer.done():
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
# Cancel producer task if still running to avoid resource leaks.
if not task.done():
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
def _pick_free_port(host: str = REME_DEFAULT_HOST) -> int:
"""Bind to port 0 and return the OS-assigned free port."""
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind((host, 0))
return s.getsockname()[1]
async def _wait_reme_ready(host: str, port: int, timeout: float) -> None:
"""Poll find_reme until it reports 'reme' or timeout elapses."""
from .service_utils import find_reme
deadline = time.time() + timeout
while time.time() < deadline:
status = await find_reme(host, port)
if status == "reme":
return
await asyncio.sleep(0.2)
raise TimeoutError(f"ReMe service did not become ready at {host}:{port} within {timeout}s")
@asynccontextmanager
async def mock_reme_server(
host: str = REME_DEFAULT_HOST,
port: int | None = None,
config: str | None = None,
extra_args: list[str] | None = None,
startup_timeout: float = 30.0,
shutdown_timeout: float = 10.0,
log_to_file: bool = False,
enable_logo: bool = False,
):
"""Spawn `reme4 start` as a subprocess and yield (host, port) once ready.
Auto-picks a free port when port is None. Subprocess is terminated on exit.
"""
logger = get_logger()
if port is None:
port = _pick_free_port(host)
cmd: list[str] = [
sys.executable,
"-m",
"reme4.reme",
"start",
f"service.host={host}",
f"service.port={port}",
f"log_to_file={'true' if log_to_file else 'false'}",
f"enable_logo={'true' if enable_logo else 'false'}",
]
if config:
cmd.append(f"config={config}")
if extra_args:
cmd.extend(extra_args)
logger.info(f"Launching mock reme server: {' '.join(cmd)}")
proc = subprocess.Popen(
cmd,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
)
try:
await _wait_reme_ready(host, port, startup_timeout)
yield host, port
except Exception:
# Capture early-exit output for diagnostics.
if proc.poll() is not None and proc.stdout is not None:
tail = proc.stdout.read()
logger.error(f"reme server exited early. output:\n{tail}")
raise
finally:
if proc.poll() is None:
proc.terminate()
try:
proc.wait(timeout=shutdown_timeout)
except subprocess.TimeoutExpired:
logger.warning("reme server did not terminate gracefully, killing")
proc.kill()
proc.wait(timeout=shutdown_timeout)
if proc.stdout is not None:
try:
proc.stdout.close()
except Exception:
pass
async def call_action(
action: str,
host: str = REME_DEFAULT_HOST,
port: int = REME_DEFAULT_PORT,
timeout: float = 30.0,
**kwargs,
) -> dict | str:
"""POST to /{action}; return parsed JSON (dict) for JSON endpoints, raw text for SSE."""
from ..components.client.http_client import HttpClient
pieces: list[str] = []
async with HttpClient(action=action, host=host, port=port, timeout=timeout, **kwargs) as client:
async for chunk in client.stream_chunks():
payload = chunk.chunk
pieces.append(payload if isinstance(payload, str) else json.dumps(payload, ensure_ascii=False))
raw = "".join(pieces)
try:
return json.loads(raw)
except (ValueError, json.JSONDecodeError):
return raw
async def call_and_check(
action: str,
host: str = REME_DEFAULT_HOST,
port: int = REME_DEFAULT_PORT,
validator: Callable[[Any], bool] | None = None,
expected: Any = None,
timeout: float = 30.0,
**kwargs,
) -> Any:
"""Call action and verify response. Raises AssertionError on mismatch.
- validator(result) -> bool: custom predicate.
- expected: deep-equality target (compared to result, or to result[key] when expected is dict).
"""
result = await call_action(action, host=host, port=port, timeout=timeout, **kwargs)
if validator is not None and not validator(result):
raise AssertionError(f"validator rejected response for action={action!r}: {result!r}")
if expected is not None:
if isinstance(expected, dict) and isinstance(result, dict):
for k, v in expected.items():
if result.get(k) != v:
raise AssertionError(
f"action={action!r} expected {k}={v!r}, got {result.get(k)!r} (full: {result!r})",
)
elif result != expected:
raise AssertionError(f"action={action!r} expected {expected!r}, got {result!r}")
return result

36
reme4/utils/env_utils.py Normal file
View file

@ -0,0 +1,36 @@
"""Load .env files into os.environ (idempotent)."""
import os
from pathlib import Path
_LOADED = False
def _parse(path: Path) -> None:
for line in path.read_text(encoding="utf-8").splitlines():
line = line.strip()
if not line or line.startswith("#") or "=" not in line:
continue
key, value = line.split("=", 1)
os.environ[key.strip()] = value.strip().strip("'\"")
def load_env(path: str | Path | None = None) -> None:
"""Load .env from given path, or search cwd and up to 5 parents."""
global _LOADED
if _LOADED:
return
if path:
path = Path(path)
if path.exists():
_parse(path)
_LOADED = True
return
for directory in [Path.cwd(), *Path.cwd().parents[:5]]:
env_path = directory / ".env"
if env_path.exists():
_parse(env_path)
_LOADED = True
return

108
reme4/utils/logger_utils.py Normal file
View file

@ -0,0 +1,108 @@
"""Logger utilities supporting both loguru and standard logging backends."""
import logging
import os
import sys
from datetime import datetime
from logging.handlers import TimedRotatingFileHandler
_logger = None
_LOGURU_FORMAT = "{time:YYYY-MM-DD HH:mm:ss} | {level} | {file}:{line} | {function} | {message}"
_STDLIB_FORMAT = "%(asctime)s | %(levelname)s | %(filename)s:%(lineno)d | %(funcName)s | %(message)s"
_STDLIB_DATEFMT = "%Y-%m-%d %H:%M:%S"
def _enable_loguru() -> bool:
return os.getenv("REME_DISABLE_LOGURU", "").lower() != "true"
def _init_loguru(log_dir: str, level: str, log_to_console: bool, log_to_file: bool):
from loguru import logger
logger.remove()
if log_to_console:
logger.add(
sink=sys.stdout,
level=level,
format=_LOGURU_FORMAT,
colorize=True,
)
if log_to_file:
try:
os.makedirs(log_dir, exist_ok=True)
current_ts = datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
log_filepath = os.path.join(log_dir, f"{current_ts}.log")
logger.add(
log_filepath,
level=level,
rotation="00:00",
retention="7 days",
compression="zip",
encoding="utf-8",
format=_LOGURU_FORMAT,
)
except Exception as e:
logger.error(f"Error configuring file logging: {e}")
return logger
def _init_stdlib(log_dir: str, level: str, log_to_console: bool, log_to_file: bool):
logger = logging.getLogger("reme")
logger.setLevel(level)
logger.propagate = False
for handler in list(logger.handlers):
logger.removeHandler(handler)
formatter = logging.Formatter(_STDLIB_FORMAT, datefmt=_STDLIB_DATEFMT)
if log_to_console:
console_handler = logging.StreamHandler(sys.stdout)
console_handler.setLevel(level)
console_handler.setFormatter(formatter)
logger.addHandler(console_handler)
if log_to_file:
try:
os.makedirs(log_dir, exist_ok=True)
current_ts = datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
log_filepath = os.path.join(log_dir, f"{current_ts}.log")
file_handler = TimedRotatingFileHandler(
log_filepath,
when="midnight",
backupCount=7,
encoding="utf-8",
)
file_handler.setLevel(level)
file_handler.setFormatter(formatter)
logger.addHandler(file_handler)
except Exception as e:
logger.error(f"Error configuring file logging: {e}")
return logger
def get_logger(
log_dir: str = "logs",
level: str = "INFO",
log_to_console: bool = True,
log_to_file: bool = True,
force_init: bool = False,
):
"""Return the global logger, initializing sinks on first call (or when force_init)."""
global _logger
if _logger is not None and not force_init:
return _logger
if _enable_loguru():
_logger = _init_loguru(log_dir, level, log_to_console, log_to_file)
else:
_logger = _init_stdlib(log_dir, level, log_to_console, log_to_file)
return _logger

103
reme4/utils/logo_utils.py Normal file
View file

@ -0,0 +1,103 @@
"""Startup banner with ASCII logo and service metadata."""
import colorsys
import importlib.metadata
import random
from typing import TYPE_CHECKING
from rich.console import Console, Group
from rich.panel import Panel
from rich.table import Table
from rich.text import Text
if TYPE_CHECKING:
from ..schema import ApplicationConfig
def get_version(package_name: str) -> str:
"""Return installed package version, or empty string if not installed."""
try:
return importlib.metadata.version(package_name)
except importlib.metadata.PackageNotFoundError:
return ""
def _hsv_rgb(h: float, s: float = 0.85, v: float = 0.98) -> tuple[int, int, int]:
"""HSV → 0-255 RGB tuple. High saturation+value keeps colors vibrant."""
r, g, b = colorsys.hsv_to_rgb(h % 1.0, s, v)
return int(r * 255), int(g * 255), int(b * 255)
def print_logo(app_config: "ApplicationConfig"):
"""Print rainbow ASCII logo and runtime config (backend, URL, versions).
Color: each startup picks a random hue rotation; both horizontal
(across each line) and vertical (line-to-line) sweep ~half the
hue wheel, so the banner shows a fresh multi-color rainbow gradient
every run.
"""
ascii_art = [
r" ██████╗ ███████╗ ███╗ ███╗ ███████╗ ",
r" ██╔══██╗ ██╔════╝ ████╗ ████║ ██╔════╝ ",
r" ██████╔╝ █████╗ ██╔████╔██║ █████╗ ",
r" ██╔══██╗ ██╔══╝ ██║╚██╔╝██║ ██╔══╝ ",
r" ██║ ██║ ███████╗ ██║ ╚═╝ ██║ ███████╗ ",
r" ╚═╝ ╚═╝ ╚══════╝ ╚═╝ ╚═╝ ╚══════╝ ",
]
hue_base = random.random() # random starting hue per startup
horizontal_span = 0.5 # half the wheel left-to-right
vertical_shift = 0.08 # small per-line nudge for 2D rainbow
logo_text = Text()
for line_idx, line in enumerate(ascii_art):
line_len = max(1, len(line) - 1)
line_hue_start = hue_base + line_idx * vertical_shift
for i, char in enumerate(line):
ratio = i / line_len
r, g, b = _hsv_rgb(line_hue_start + horizontal_span * ratio)
logo_text.append(char, style=f"bold rgb({r},{g},{b})")
logo_text.append("\n")
info_table = Table.grid(padding=(0, 1))
info_table.add_column(style="bold", justify="center")
info_table.add_column(style="bold cyan", justify="left")
info_table.add_column(style="white", justify="left")
# service is a ComponentConfig with extra="allow"; backend-specific fields live in model_extra.
service = app_config.service
backend = service.backend
extra = service.model_extra or {}
info_table.add_row("📦", "Backend:", backend)
match backend:
case "http":
host = extra.get("host", "localhost")
port = extra.get("port", 8000)
info_table.add_row("🔗", "URL:", f"http://{host}:{port}")
info_table.add_row("📚", "FastAPI:", Text(get_version("fastapi"), style="dim"))
case "mcp":
transport = extra.get("transport", "stdio")
info_table.add_row("🚌", "Transport:", transport)
if transport != "stdio":
host = extra.get("host", "localhost")
port = extra.get("port", 8000)
url = f"http://{host}:{port}"
if transport == "sse":
url += "/sse"
info_table.add_row("🔗", "URL:", url)
info_table.add_row("📚", "FastMCP:", Text(get_version("fastmcp"), style="dim"))
info_table.add_row("🚀", "ReMe:", Text(get_version("reme-ai"), style="dim"))
panel = Panel(
Group(logo_text, info_table),
title=app_config.app_name,
title_align="left",
border_style="dim",
padding=(1, 4),
expand=False,
)
Console().print(Group("\n", panel, "\n"))

View file

@ -0,0 +1,96 @@
"""Service discovery utilities."""
import asyncio
import socket
import subprocess
import sys
from ..constants import REME_DEFAULT_HOST, REME_DEFAULT_PORT
async def find_reme(host: str, port: int) -> str:
"""Probe host:port. Returns 'reme', 'occupied', or 'free'."""
from ..components.client.http_client import HttpClient
try:
async with HttpClient(action="health_check", host=host, port=port, timeout=2.0) as client:
async for _ in client():
break
return "reme"
except Exception:
pass
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
try:
s.bind((host, port))
return "free"
except OSError:
return "occupied"
def _sh(cmd: list[str]) -> str:
"""Run cmd; return stdout, or '' on failure."""
try:
return subprocess.check_output(cmd, stderr=subprocess.DEVNULL, text=True)
except (subprocess.CalledProcessError, FileNotFoundError):
return ""
def _pid_on_port(port: int) -> int | None:
"""PID listening on TCP port, or None."""
out = _sh(["lsof", "-nP", f"-iTCP:{port}", "-sTCP:LISTEN", "-t"]).strip()
return int(out.splitlines()[0]) if out else None
def _scan_reme_procs() -> list[tuple[int, str, int]]:
"""List running 'reme ... start' processes as (pid, host, port)."""
procs: list[tuple[int, str, int]] = []
for line in _sh(["pgrep", "-af", "reme.* start"]).splitlines():
parts = line.split()
if not parts or not parts[0].isdigit():
continue
host, port = REME_DEFAULT_HOST, REME_DEFAULT_PORT
for t in parts[1:]:
if t.startswith("service.host="):
host = t.split("=", 1)[1]
elif t.startswith("service.port=") and t.split("=", 1)[1].isdigit():
port = int(t.split("=", 1)[1])
procs.append((int(parts[0]), host, port))
return procs
async def locate_reme() -> tuple[str, int, int | None] | None:
"""Find a running reme: try default port, then scanned processes."""
if await find_reme(REME_DEFAULT_HOST, REME_DEFAULT_PORT) == "reme":
return REME_DEFAULT_HOST, REME_DEFAULT_PORT, _pid_on_port(REME_DEFAULT_PORT)
for pid, host, port in _scan_reme_procs():
if await find_reme(host, port) == "reme":
return host, port, pid
return None
def precheck_start(svc_config: dict | None) -> bool:
"""Pre-flight check for `start`: False if reme is up, exits 1 on port conflict."""
host = (svc_config or {}).get("host") or REME_DEFAULT_HOST
port = (svc_config or {}).get("port") or REME_DEFAULT_PORT
status = asyncio.run(find_reme(host, port))
if status == "reme":
print(f"reme already running at {host}:{port}")
return False
if status == "occupied":
print(
f"port {port} occupied. Start on another port: reme4 start service.port=<other_port>",
file=sys.stderr,
)
sys.exit(1)
return True
def cli_find_reme() -> None:
"""Handle `reme find_reme`: print HOST/PORT/PID or a hint to start reme."""
found = asyncio.run(locate_reme())
if not found:
print("reme not started. Try: reme start", file=sys.stderr)
sys.exit(1)
host, port, pid = found
print(f"HOST={host} PORT={port} PID={pid or 'unknown'}")

View file

@ -0,0 +1,35 @@
"""Cosine similarity for single vectors and batched matrices."""
import numpy as np
def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
"""Cosine similarity of two equal-length vectors; returns 0.0 if either has zero norm."""
if len(vec1) != len(vec2):
raise ValueError(f"Vectors must have same length: {len(vec1)} != {len(vec2)}")
dot_product = sum(a * b for a, b in zip(vec1, vec2))
magnitude1 = sum(a * a for a in vec1) ** 0.5
magnitude2 = sum(b * b for b in vec2) ** 0.5
if magnitude1 == 0 or magnitude2 == 0:
return 0.0
return dot_product / (magnitude1 * magnitude2)
def batch_cosine_similarity(nd_array1: np.ndarray, nd_array2: np.ndarray) -> np.ndarray:
"""Pairwise cosine similarity matrix between two batches; output shape (N1, N2)."""
if nd_array1.shape[1] != nd_array2.shape[1]:
raise ValueError(
f"Embedding dimensions must match: {nd_array1.shape[1]} != {nd_array2.shape[1]}",
)
dot_products = np.dot(nd_array1, nd_array2.T)
norms1 = np.linalg.norm(nd_array1, axis=1)
norms2 = np.linalg.norm(nd_array2, axis=1)
norm_products = np.outer(norms1, norms2)
# Guard against zero-norm rows to keep division finite.
norm_products = np.where(norm_products == 0, 1e-10, norm_products)
return dot_products / norm_products

View file

@ -79,6 +79,15 @@ def temp_nested_dir(temp_dir: Path):
yield temp_dir
def make_mock_file_store():
"""Create an async-compatible mock file store."""
mock_file_store = MagicMock()
mock_file_store.clear_all = AsyncMock()
mock_file_store.list_files = AsyncMock(return_value=[])
mock_file_store.get_file_chunks = AsyncMock(return_value=[])
return mock_file_store
# ==================== Test Existing Paths ====================
@ -147,6 +156,21 @@ class TestExistingPaths:
await watcher.close()
@pytest.mark.asyncio
async def test_restart_resets_stop_event(self, temp_dir: Path):
"""Test restarting watcher resets the previous stop signal."""
watcher = BaseFileWatcher(watch_paths=str(temp_dir))
await watcher.start()
await watcher.close()
assert watcher._stop_event.is_set() is True
await watcher.start()
assert watcher._stop_event.is_set() is False
await watcher.close()
@pytest.mark.asyncio
async def test_multiple_start_calls(self, temp_dir: Path):
"""Test that multiple start calls don't create multiple tasks."""
@ -381,9 +405,7 @@ class TestRebuildIndexOnStart:
callback_called.append(changes)
# Create mock file_store
mock_file_store = MagicMock()
mock_file_store.list_files = AsyncMock(return_value=[])
mock_file_store.get_file_chunks = AsyncMock(return_value=[])
mock_file_store = make_mock_file_store()
watcher = BaseFileWatcher(
watch_paths=str(temp_dir),
@ -408,9 +430,7 @@ class TestRebuildIndexOnStart:
callback_called.append(changes)
# Create mock file_store
mock_file_store = MagicMock()
mock_file_store.list_files = AsyncMock(return_value=[])
mock_file_store.get_file_chunks = AsyncMock(return_value=[])
mock_file_store = make_mock_file_store()
watcher = BaseFileWatcher(
watch_paths=str(temp_dir),
@ -441,9 +461,7 @@ class TestRebuildIndexOnStart:
async def callback(changes):
callback_called.append(changes)
mock_file_store = MagicMock()
mock_file_store.list_files = AsyncMock(return_value=[])
mock_file_store.get_file_chunks = AsyncMock(return_value=[])
mock_file_store = make_mock_file_store()
watcher = BaseFileWatcher(
watch_paths=str(temp_dir),
@ -474,9 +492,7 @@ class TestRebuildIndexOnStart:
async def callback(changes):
callback_called.append(changes)
mock_file_store = MagicMock()
mock_file_store.list_files = AsyncMock(return_value=[])
mock_file_store.get_file_chunks = AsyncMock(return_value=[])
mock_file_store = make_mock_file_store()
watcher = BaseFileWatcher(
watch_paths=str(temp_nested_dir),
@ -510,9 +526,7 @@ class TestRebuildIndexOnStart:
async def callback(changes):
callback_called.append(changes)
mock_file_store = MagicMock()
mock_file_store.list_files = AsyncMock(return_value=[])
mock_file_store.get_file_chunks = AsyncMock(return_value=[])
mock_file_store = make_mock_file_store()
watcher = BaseFileWatcher(
watch_paths=str(temp_nested_dir),
@ -545,9 +559,7 @@ class TestRebuildIndexOnStart:
async def callback(changes):
callback_called.append(changes)
mock_file_store = MagicMock()
mock_file_store.list_files = AsyncMock(return_value=[])
mock_file_store.get_file_chunks = AsyncMock(return_value=[])
mock_file_store = make_mock_file_store()
watcher = BaseFileWatcher(
watch_paths="/nonexistent/path",
@ -698,9 +710,7 @@ class TestEdgeCases:
"""Test watching a single file instead of directory."""
file_path = temp_files["txt_0"]
mock_file_store = MagicMock()
mock_file_store.list_files = AsyncMock(return_value=[])
mock_file_store.get_file_chunks = AsyncMock(return_value=[])
mock_file_store = make_mock_file_store()
callback_called = []
@ -731,9 +741,7 @@ class TestEdgeCases:
empty_dir = temp_dir / "empty"
empty_dir.mkdir()
mock_file_store = MagicMock()
mock_file_store.list_files = AsyncMock(return_value=[])
mock_file_store.get_file_chunks = AsyncMock(return_value=[])
mock_file_store = make_mock_file_store()
callback_called = []
@ -774,7 +782,7 @@ class TestEdgeCases:
unicode_dir.mkdir()
file_path = unicode_dir / "文件.txt"
file_path.write_text("内容")
file_path.write_text("内容", encoding="utf-8")
watcher = BaseFileWatcher(
watch_paths=str(unicode_dir),

View file

@ -0,0 +1,67 @@
"""
Tests for the default watch path construction in ``ReMeLight``.
Verifies that the built-in watch list picks a single ``MEMORY.md`` /
``memory.md`` spelling so the file is not indexed twice on case-insensitive
filesystems (Windows NTFS, macOS APFS/HFS+). See agentscope-ai/ReMe#228.
"""
# pylint: disable=redefined-outer-name,protected-access,missing-function-docstring,missing-class-docstring
import tempfile
from pathlib import Path
from unittest.mock import patch
import pytest
from reme.reme_light import ReMeLight
@pytest.fixture
def temp_working_dir():
with tempfile.TemporaryDirectory() as tmp:
yield tmp
def _captured_watch_paths(working_dir: str, *, default_file_watcher_config=None):
"""Capture the ``watch_paths`` that ``ReMeLight`` would forward to its
parent ``Application.__init__``, without spinning up the full app stack."""
captured: dict = {}
def _capture(*_args, **kwargs):
captured.update(kwargs)
with patch("reme.reme_light.Application.__init__", _capture):
ReMeLight(
working_dir=working_dir,
default_file_watcher_config=default_file_watcher_config,
)
return list((captured.get("default_file_watcher_config") or {}).get("watch_paths", []))
class TestDefaultWatchPaths:
def test_defaults_to_uppercase_memory_md_when_neither_exists(self, temp_working_dir):
paths = _captured_watch_paths(temp_working_dir)
memory_dir = str(Path(temp_working_dir).absolute() / "memory")
assert paths == [str(Path(temp_working_dir).absolute() / "MEMORY.md"), memory_dir]
def test_picks_lowercase_memory_md_when_only_it_exists(self, temp_working_dir):
(Path(temp_working_dir) / "memory.md").write_text("")
if (Path(temp_working_dir) / "MEMORY.md").exists():
pytest.skip("case-insensitive filesystem treats both spellings as one file")
paths = _captured_watch_paths(temp_working_dir)
assert paths[0] == str(Path(temp_working_dir).absolute() / "memory.md")
def test_prefers_uppercase_memory_md_when_it_exists(self, temp_working_dir):
(Path(temp_working_dir) / "MEMORY.md").write_text("")
paths = _captured_watch_paths(temp_working_dir)
assert paths[0] == str(Path(temp_working_dir).absolute() / "MEMORY.md")
def test_user_provided_watch_paths_pass_through(self, temp_working_dir):
custom = [str(Path(temp_working_dir) / "notes.md")]
paths = _captured_watch_paths(
temp_working_dir,
default_file_watcher_config={"watch_paths": custom},
)
assert paths == custom

View file

@ -0,0 +1,337 @@
"""BM25Index performance tests for add_docs and retrieve."""
import asyncio
import os
import random
import tempfile
import time
from reme4.components.keyword_index import BM25Index
from reme4.components.tokenizer import RegexTokenizer
# A small vocab of realistic-looking words for generating random text
_VOCAB = [
"algorithm",
"data",
"machine",
"learning",
"model",
"network",
"neural",
"training",
"optimization",
"gradient",
"loss",
"function",
"parameter",
"weight",
"bias",
"layer",
"activation",
"relu",
"sigmoid",
"softmax",
"backpropagation",
"forward",
"pass",
"batch",
"epoch",
"iteration",
"convergence",
"divergence",
"regularization",
"dropout",
"attention",
"transformer",
"encoder",
"decoder",
"embedding",
"token",
"vector",
"matrix",
"tensor",
"computation",
"graph",
"node",
"edge",
"vertex",
"path",
"search",
"retrieval",
"index",
"query",
"document",
"corpus",
"term",
"frequency",
"inverse",
"score",
"rank",
"relevance",
"precision",
"recall",
"f1",
"metric",
"evaluation",
"benchmark",
"dataset",
"sample",
"feature",
"label",
"class",
"predict",
"classification",
"regression",
"clustering",
"dimension",
"reduction",
"pca",
"tsne",
"visualization",
"matplotlib",
"plot",
"chart",
"histogram",
"scatter",
"line",
"bar",
"database",
"sql",
"query",
"table",
"row",
"column",
"index",
"primary",
"foreign",
"key",
"constraint",
"schema",
"migration",
"version",
"control",
"git",
"commit",
"branch",
"merge",
"conflict",
"resolution",
"review",
"approve",
"reject",
"pull",
"request",
"issue",
"bug",
"fix",
"feature",
"enhancement",
"refactor",
"test",
"deploy",
"production",
"staging",
"development",
"environment",
"configuration",
"setting",
"variable",
"constant",
"global",
"local",
"scope",
"closure",
"callback",
"promise",
"async",
"await",
"synchronous",
"asynchronous",
"concurrent",
"parallel",
"thread",
"process",
"memory",
"cache",
"buffer",
"queue",
"stack",
"heap",
"pool",
]
class temp_chdir:
"""Context manager to temporarily chdir into a path and restore on exit."""
def __init__(self, path):
self.path = path
self.old = None
def __enter__(self):
self.old = os.getcwd()
os.chdir(self.path)
return self
def __exit__(self, *exc):
os.chdir(self.old)
def _gen_random_text(n_tokens: int) -> str:
"""Generate random text with approximately n_tokens words."""
words = random.choices(_VOCAB, k=n_tokens)
return " ".join(words)
def _gen_random_query(n_words: int) -> str:
"""Generate a random query with n_words words."""
words = random.choices(_VOCAB, k=n_words)
return " ".join(words)
async def _make_index() -> BM25Index:
"""Create and start a BM25Index using cwd as working dir, with non-filtering tokenizer."""
index = BM25Index()
tokenizer = RegexTokenizer(filter_stopwords=False)
index.tokenizer = tokenizer
index._owned.append(tokenizer) # pylint: disable=protected-access
await index.start()
return index
async def _setup_index_for_retrieve(n_docs: int = 100, doc_tokens: int = 1000) -> BM25Index:
"""Build an index with n_docs medium-sized docs in cwd."""
index = await _make_index()
docs = {f"doc_{i}": _gen_random_text(doc_tokens) for i in range(n_docs)}
await index.add_docs(docs)
return index
def test_add_docs_small():
"""Add 100 small docs (~100 tokens each)."""
async def run():
docs = {f"doc_{i}": _gen_random_text(100) for i in range(100)}
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
index = await _make_index()
t0 = time.perf_counter()
await index.add_docs(docs)
elapsed = time.perf_counter() - t0
print(f" add_docs (100 docs x ~100 tokens): {elapsed:.4f}s")
await index.close()
asyncio.run(run())
def test_add_docs_medium():
"""Add 100 medium docs (~1000 tokens each)."""
async def run():
docs = {f"doc_{i}": _gen_random_text(1000) for i in range(100)}
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
index = await _make_index()
t0 = time.perf_counter()
await index.add_docs(docs)
elapsed = time.perf_counter() - t0
print(f" add_docs (100 docs x ~1000 tokens): {elapsed:.4f}s")
await index.close()
asyncio.run(run())
def test_add_docs_large():
"""Add 100 large docs (~10000 tokens each)."""
async def run():
docs = {f"doc_{i}": _gen_random_text(10000) for i in range(100)}
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
index = await _make_index()
t0 = time.perf_counter()
await index.add_docs(docs)
elapsed = time.perf_counter() - t0
print(f" add_docs (100 docs x ~10000 tokens): {elapsed:.4f}s")
await index.close()
asyncio.run(run())
def test_retrieve_short_query():
"""Retrieve with 1-word query."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
index = await _setup_index_for_retrieve()
query = _gen_random_query(1)
t0 = time.perf_counter()
await index.retrieve(query, limit=10)
elapsed = time.perf_counter() - t0
print(f" retrieve (1-word query, 100 docs): {elapsed:.6f}s")
await index.close()
asyncio.run(run())
def test_retrieve_medium_query():
"""Retrieve with 5-word query."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
index = await _setup_index_for_retrieve()
query = _gen_random_query(5)
t0 = time.perf_counter()
await index.retrieve(query, limit=10)
elapsed = time.perf_counter() - t0
print(f" retrieve (5-word query, 100 docs): {elapsed:.6f}s")
await index.close()
asyncio.run(run())
def test_retrieve_long_query():
"""Retrieve with 20-word query."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
index = await _setup_index_for_retrieve()
query = _gen_random_query(20)
t0 = time.perf_counter()
await index.retrieve(query, limit=10)
elapsed = time.perf_counter() - t0
print(f" retrieve (20-word query, 100 docs): {elapsed:.6f}s")
await index.close()
asyncio.run(run())
def test_retrieve_very_long_query():
"""Retrieve with 100-word query."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
index = await _setup_index_for_retrieve()
query = _gen_random_query(100)
t0 = time.perf_counter()
await index.retrieve(query, limit=10)
elapsed = time.perf_counter() - t0
print(f" retrieve (100-word query, 100 docs): {elapsed:.6f}s")
await index.close()
asyncio.run(run())
if __name__ == "__main__":
random.seed(42)
print("=== BM25Index Performance Tests ===\n")
print("[add_docs]")
test_add_docs_small()
test_add_docs_medium()
test_add_docs_large()
print("\n[retrieve]")
test_retrieve_short_query()
test_retrieve_medium_query()
test_retrieve_long_query()
test_retrieve_very_long_query()
print("\nDone.")

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