mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-28 01:31:46 +00:00
Merge remote-tracking branch 'origin/main'
This commit is contained in:
commit
d27ca21386
110 changed files with 12528 additions and 43 deletions
43
.github/workflows/unittest.yml
vendored
Normal file
43
.github/workflows/unittest.yml
vendored
Normal 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
2
.gitignore
vendored
|
|
@ -42,4 +42,4 @@ meta_memory/*
|
|||
**/data/*.json
|
||||
*.db
|
||||
memories/*
|
||||
.reme/*
|
||||
.reme/*
|
||||
|
|
|
|||
183
docs4/reme_design.md
Normal file
183
docs4/reme_design.md
Normal 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
7
docs4/todo.md
Normal 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
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
26
reme4/__init__.py
Normal 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
174
reme4/application.py
Normal 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)
|
||||
43
reme4/components/__init__.py
Normal file
43
reme4/components/__init__.py
Normal 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",
|
||||
]
|
||||
28
reme4/components/application_context.py
Normal file
28
reme4/components/application_context.py
Normal 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] = {}
|
||||
53
reme4/components/as_llm/__init__.py
Normal file
53
reme4/components/as_llm/__init__.py
Normal 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",
|
||||
]
|
||||
44
reme4/components/as_llm_formatter/__init__.py
Normal file
44
reme4/components/as_llm_formatter/__init__.py
Normal 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",
|
||||
]
|
||||
141
reme4/components/as_llm_formatter/reme_openai_chat_formatter.py
Normal file
141
reme4/components/as_llm_formatter/reme_openai_chat_formatter.py
Normal 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
|
||||
35
reme4/components/as_token_counter/__init__.py
Normal file
35
reme4/components/as_token_counter/__init__.py
Normal 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",
|
||||
]
|
||||
21
reme4/components/as_token_counter/estimate_token_counter.py
Normal file
21
reme4/components/as_token_counter/estimate_token_counter.py
Normal 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)
|
||||
185
reme4/components/base_component.py
Normal file
185
reme4/components/base_component.py
Normal 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()
|
||||
7
reme4/components/client/__init__.py
Normal file
7
reme4/components/client/__init__.py
Normal 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"]
|
||||
41
reme4/components/client/base_client.py
Normal file
41
reme4/components/client/base_client.py
Normal 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
|
||||
143
reme4/components/client/http_client.py
Normal file
143
reme4/components/client/http_client.py
Normal 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
|
||||
125
reme4/components/client/mcp_client.py
Normal file
125
reme4/components/client/mcp_client.py
Normal 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)
|
||||
77
reme4/components/component_registry.py
Normal file
77
reme4/components/component_registry.py
Normal 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()
|
||||
6
reme4/components/embedding/__init__.py
Normal file
6
reme4/components/embedding/__init__.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
"""Embedding model implementations."""
|
||||
|
||||
from .base_embedding_model import BaseEmbeddingModel
|
||||
from .openai_embedding_model import OpenAIEmbeddingModel
|
||||
|
||||
__all__ = ["BaseEmbeddingModel", "OpenAIEmbeddingModel"]
|
||||
214
reme4/components/embedding/base_embedding_model.py
Normal file
214
reme4/components/embedding/base_embedding_model.py
Normal 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")
|
||||
52
reme4/components/embedding/openai_embedding_model.py
Normal file
52
reme4/components/embedding/openai_embedding_model.py
Normal 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
|
||||
8
reme4/components/file_graph/__init__.py
Normal file
8
reme4/components/file_graph/__init__.py
Normal 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"]
|
||||
53
reme4/components/file_graph/base_file_graph.py
Normal file
53
reme4/components/file_graph/base_file_graph.py
Normal 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*."""
|
||||
138
reme4/components/file_graph/local_file_graph.py
Normal file
138
reme4/components/file_graph/local_file_graph.py
Normal 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
|
||||
]
|
||||
450
reme4/components/file_graph/neo4j_file_graph.py
Normal file
450
reme4/components/file_graph/neo4j_file_graph.py
Normal 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),
|
||||
)
|
||||
122
reme4/components/file_graph/nx_file_graph.py
Normal file
122
reme4/components/file_graph/nx_file_graph.py
Normal 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]
|
||||
8
reme4/components/file_parser/__init__.py
Normal file
8
reme4/components/file_parser/__init__.py
Normal 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"]
|
||||
22
reme4/components/file_parser/bare_file_parser.py
Normal file
22
reme4/components/file_parser/bare_file_parser.py
Normal 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=[]), []
|
||||
30
reme4/components/file_parser/base_file_parser.py
Normal file
30
reme4/components/file_parser/base_file_parser.py
Normal 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)."""
|
||||
123
reme4/components/file_parser/default_file_parser.py
Normal file
123
reme4/components/file_parser/default_file_parser.py
Normal 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
|
||||
729
reme4/components/file_parser/linked_file_parser.py
Normal file
729
reme4/components/file_parser/linked_file_parser.py
Normal 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()
|
||||
14
reme4/components/file_store/__init__.py
Normal file
14
reme4/components/file_store/__init__.py
Normal 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",
|
||||
]
|
||||
100
reme4/components/file_store/base_file_store.py
Normal file
100
reme4/components/file_store/base_file_store.py
Normal 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)
|
||||
177
reme4/components/file_store/local_file_store.py
Normal file
177
reme4/components/file_store/local_file_store.py
Normal 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
|
||||
9
reme4/components/file_watcher/__init__.py
Normal file
9
reme4/components/file_watcher/__init__.py
Normal 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",
|
||||
]
|
||||
118
reme4/components/file_watcher/base_file_watcher.py
Normal file
118
reme4/components/file_watcher/base_file_watcher.py
Normal 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)."""
|
||||
129
reme4/components/file_watcher/lite_file_watcher.py
Normal file
129
reme4/components/file_watcher/lite_file_watcher.py
Normal 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)
|
||||
6
reme4/components/job/__init__.py
Normal file
6
reme4/components/job/__init__.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
"""Job components for executing workflows."""
|
||||
|
||||
from .base_job import BaseJob
|
||||
from .stream_job import StreamJob
|
||||
|
||||
__all__ = ["BaseJob", "StreamJob"]
|
||||
54
reme4/components/job/base_job.py
Normal file
54
reme4/components/job/base_job.py
Normal 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
|
||||
21
reme4/components/job/stream_job.py
Normal file
21
reme4/components/job/stream_job.py
Normal 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()
|
||||
6
reme4/components/keyword_index/__init__.py
Normal file
6
reme4/components/keyword_index/__init__.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
"""Keyword index components."""
|
||||
|
||||
from .base_keyword_index import BaseKeywordIndex
|
||||
from .bm25_index import BM25Index
|
||||
|
||||
__all__ = ["BaseKeywordIndex", "BM25Index"]
|
||||
70
reme4/components/keyword_index/base_keyword_index.py
Normal file
70
reme4/components/keyword_index/base_keyword_index.py
Normal 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."""
|
||||
206
reme4/components/keyword_index/bm25_index.py
Normal file
206
reme4/components/keyword_index/bm25_index.py
Normal 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 = {}
|
||||
125
reme4/components/prompt_handler.py
Normal file
125
reme4/components/prompt_handler.py
Normal 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)})"
|
||||
87
reme4/components/runtime_context.py
Normal file
87
reme4/components/runtime_context.py
Normal 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
|
||||
11
reme4/components/service/__init__.py
Normal file
11
reme4/components/service/__init__.py
Normal 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",
|
||||
]
|
||||
48
reme4/components/service/base_service.py
Normal file
48
reme4/components/service/base_service.py
Normal 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)
|
||||
99
reme4/components/service/http_service.py
Normal file
99
reme4/components/service/http_service.py
Normal 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)
|
||||
71
reme4/components/service/mcp_service.py
Normal file
71
reme4/components/service/mcp_service.py
Normal 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)
|
||||
11
reme4/components/tokenizer/__init__.py
Normal file
11
reme4/components/tokenizer/__init__.py
Normal 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",
|
||||
]
|
||||
44
reme4/components/tokenizer/base_tokenizer.py
Normal file
44
reme4/components/tokenizer/base_tokenizer.py
Normal 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."""
|
||||
27
reme4/components/tokenizer/jieba_tokenizer.py
Normal file
27
reme4/components/tokenizer/jieba_tokenizer.py
Normal 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
|
||||
31
reme4/components/tokenizer/regex_tokenizer.py
Normal file
31
reme4/components/tokenizer/regex_tokenizer.py
Normal 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
|
||||
1395
reme4/components/tokenizer/stopwords
Normal file
1395
reme4/components/tokenizer/stopwords
Normal file
File diff suppressed because it is too large
Load diff
8
reme4/config/__init__.py
Normal file
8
reme4/config/__init__.py
Normal file
|
|
@ -0,0 +1,8 @@
|
|||
"""Config"""
|
||||
|
||||
from .config_parser import parse_args, resolve_app_config
|
||||
|
||||
__all__ = [
|
||||
"parse_args",
|
||||
"resolve_app_config",
|
||||
]
|
||||
219
reme4/config/config_parser.py
Normal file
219
reme4/config/config_parser.py
Normal 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
189
reme4/config/default.yaml
Normal 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
12
reme4/constants.py
Normal 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>>"
|
||||
9
reme4/enumeration/__init__.py
Normal file
9
reme4/enumeration/__init__.py
Normal file
|
|
@ -0,0 +1,9 @@
|
|||
"""Enumeration"""
|
||||
|
||||
from .chunk_enum import ChunkEnum
|
||||
from .component_enum import ComponentEnum
|
||||
|
||||
__all__ = [
|
||||
"ChunkEnum",
|
||||
"ComponentEnum",
|
||||
]
|
||||
21
reme4/enumeration/chunk_enum.py
Normal file
21
reme4/enumeration/chunk_enum.py
Normal 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"
|
||||
37
reme4/enumeration/component_enum.py
Normal file
37
reme4/enumeration/component_enum.py
Normal 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
43
reme4/reme.py
Normal 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
25
reme4/schema/__init__.py
Normal 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",
|
||||
]
|
||||
45
reme4/schema/application_config.py
Normal file
45
reme4/schema/application_config.py
Normal 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
34
reme4/schema/emb_node.py
Normal 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()
|
||||
26
reme4/schema/file_chunk.py
Normal file
26
reme4/schema/file_chunk.py
Normal 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
|
||||
24
reme4/schema/file_front_matter.py
Normal file
24
reme4/schema/file_front_matter.py
Normal 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
18
reme4/schema/file_link.py
Normal 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
16
reme4/schema/file_node.py
Normal 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
11
reme4/schema/request.py
Normal 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
15
reme4/schema/response.py
Normal 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")
|
||||
14
reme4/schema/stream_chunk.py
Normal file
14
reme4/schema/stream_chunk.py
Normal 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
11
reme4/steps/__init__.py
Normal 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
188
reme4/steps/base_step.py
Normal 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,
|
||||
)
|
||||
21
reme4/steps/common/__init__.py
Normal file
21
reme4/steps/common/__init__.py
Normal 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",
|
||||
]
|
||||
53
reme4/steps/common/demo.py
Normal file
53
reme4/steps/common/demo.py
Normal 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
|
||||
164
reme4/steps/common/health_check.py
Normal file
164
reme4/steps/common/health_check.py
Normal 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
|
||||
42
reme4/steps/common/help.py
Normal file
42
reme4/steps/common/help.py
Normal 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
|
||||
24
reme4/steps/common/reindex.py
Normal file
24
reme4/steps/common/reindex.py
Normal 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
|
||||
227
reme4/steps/common/search.py
Normal file
227
reme4/steps/common/search.py
Normal 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
|
||||
42
reme4/steps/common/stream_demo.py
Normal file
42
reme4/steps/common/stream_demo.py
Normal 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
|
||||
19
reme4/steps/common/version.py
Normal file
19
reme4/steps/common/version.py
Normal 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
|
||||
7
reme4/steps/crud/__init__.py
Normal file
7
reme4/steps/crud/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""CRUD steps for markdown files under the working_dir."""
|
||||
|
||||
from .read import ReadStep
|
||||
|
||||
__all__ = [
|
||||
"ReadStep",
|
||||
]
|
||||
109
reme4/steps/crud/_file_io.py
Normal file
109
reme4/steps/crud/_file_io.py
Normal 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
91
reme4/steps/crud/read.py
Normal 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
31
reme4/utils/__init__.py
Normal 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
249
reme4/utils/common_utils.py
Normal 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
36
reme4/utils/env_utils.py
Normal 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
108
reme4/utils/logger_utils.py
Normal 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
103
reme4/utils/logo_utils.py
Normal 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"))
|
||||
96
reme4/utils/service_utils.py
Normal file
96
reme4/utils/service_utils.py
Normal 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'}")
|
||||
35
reme4/utils/similarity_utils.py
Normal file
35
reme4/utils/similarity_utils.py
Normal 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
|
||||
|
|
@ -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),
|
||||
|
|
|
|||
67
tests/test_reme_light_watch_paths.py
Normal file
67
tests/test_reme_light_watch_paths.py
Normal 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
|
||||
337
tests4/unittest/test_bm25_index_perf.py
Normal file
337
tests4/unittest/test_bm25_index_perf.py
Normal 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
Loading…
Add table
Reference in a new issue