refactor(reme4): restructure steps packages (#258)

* fix(bm25_index): 修正BM25索引计算中的文档长度归一化问题

修复了在计算BM25相似度时对文档长度进行不正确归一化的bug,确保所有查询都能得到准确的相关性评分。

* up

* up

* up

* up

* up

* up

* up

* up

* up

* up

* up

* up

* up

* up

* up

* up

* up

* refactor(steps): Rename and adjust indexing step logic

- Rename `scan_changes.py` and `reindex.py` to `clear_and_scan.py`
- Update implementation details of `ScanChangesStep` and `ClearAndScanStep`
- Modify the scheduling mechanism in `WatchChangesStep`
- Adjust step registration and parameter configuration in config files
- Update related tests to align with the new interface changes

* up

* feat(daily): replace daily CRUD operations with slug provisioning approach

* refactor(tests): migrate CRUD step tests from HTTP server to direct LocalFileStore

* up

* up

* up

* up

---------

Co-authored-by: huangsen <huangsen.huang@alibaba-inc.com>
This commit is contained in:
jinliyl 2026-05-28 14:30:30 +08:00 • committed by GitHub
parent 83bfddb4a4
commit a4efc0f776
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
102 changed files with 5560 additions and 4820 deletions

View file

@ -108,6 +108,8 @@ ReMe 不把对话一股脑塞进数据库,而是按"原始 → 加工"分两
### 3.3 目录约定(用户可读、可备份、可迁移)
skills 建议使用index 看 whole picture~
```
~/reme_workspace/
├── resource/ # 原始素材(按日期归档)
@ -122,6 +124,8 @@ ReMe 不把对话一股脑塞进数据库,而是按"原始 → 加工"分两
│ └── reading-paper.md
└── digest/ # 固定四个子目录:personal / knowledge / procedural / proactive
├── personal/ # 个性化记忆:偏好、习惯、个人事件
.change_log.md difference
_moc.md
│ └── 用户偏好.md
├── knowledge/ # 知识类记忆:用户自定义二级目录(work / financial / ...)
│ ├── work/

View file

@ -44,6 +44,7 @@ dependencies = [
"fastmcp>=2.14.1",
"httpx>=0.28.1",
"jieba>=0.42.1",
"rjieba>=0.1.11",
"loguru>=0.7.3",
"mcp>=1.25.0",
"networkx>=3.4",

View file

@ -3,96 +3,131 @@
import asyncio
import heapq
from pathlib import Path
from typing import AsyncGenerator
from typing import AsyncGenerator, TypeVar
from .components import BaseComponent, ApplicationContext
from .components.job import BaseJob
from .components.service import BaseService
from .enumeration import ComponentEnum
from .schema import Response, StreamChunk
from .schema import ComponentConfig, Response, StreamChunk
from .utils import execute_stream_task, print_logo, get_logger
T = TypeVar("T", bound=BaseComponent)
_NodeKey = tuple[ComponentEnum, str]
class Application(BaseComponent):
"""Main application: initializes components, resolves dependencies, runs jobs."""
"""Wires components from config and runs jobs against them."""
def __init__(self, **kwargs) -> None:
self.context = ApplicationContext(**kwargs)
self._started_components: list[BaseComponent] = []
vault_path = Path(self.config.vault_dir).absolute()
vault_path.mkdir(parents=True, exist_ok=True)
(vault_path / self.config.metadata_dir).mkdir(parents=True, exist_ok=True)
(vault_path / self.config.daily_dir).mkdir(parents=True, exist_ok=True)
(vault_path / self.config.digest_dir).mkdir(parents=True, exist_ok=True)
(vault_path / self.config.resource_dir).mkdir(parents=True, exist_ok=True)
self._setup_vault_directories()
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)
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)
self._init_service()
self._init_components()
self._init_jobs()
@property
def config(self):
"""Application configuration."""
"""Typed view onto the application config held by the context."""
return self.context.app_config
# ----- Wiring (called once during __init__) --------------------------
def _setup_vault_directories(self) -> None:
"""Ensure the vault root and configured subdirectories exist on disk."""
cfg = self.config
vault_path = Path(cfg.vault_dir).absolute()
vault_path.mkdir(parents=True, exist_ok=True)
for subdir in [cfg.metadata_dir, cfg.daily_dir, cfg.digest_dir, cfg.resource_dir]:
if subdir:
(vault_path / subdir).mkdir(parents=True, exist_ok=True)
def _init_service(self) -> None:
"""Instantiate the single service backend declared in config.service."""
self.context.service = self._instantiate(
ComponentEnum.SERVICE,
self.config.service,
label="Service",
expected_type=BaseService,
)
def _init_components(self) -> None:
"""Instantiate every component declared under config.components."""
for ctype, group in self.config.components.items():
self.context.components[ctype] = {}
for name, cfg in group.items():
self.context.components[ctype][name] = self._instantiate(
ctype,
cfg,
label=f"Component '{name}'",
expected_type=BaseComponent,
name=name,
)
def _init_jobs(self) -> None:
"""Instantiate every job declared under config.jobs."""
for name, cfg in self.config.jobs.items():
self.context.jobs[name] = self._instantiate(
ComponentEnum.JOB,
cfg,
label=f"Job '{name}'",
expected_type=BaseJob,
name=name,
)
def _instantiate(
self,
ctype: ComponentEnum,
cfg: ComponentConfig,
*,
label: str,
expected_type: type[T],
name: str | None = None,
) -> T:
"""Resolve cfg.backend through the registry and construct the instance.
`label` is the human-readable identifier used only in error messages.
`expected_type` narrows the return type and guards against a backend
registered under the wrong ComponentEnum.
`name` is forwarded to the constructor for named components/jobs;
leave it None for the service, which is keyed solely by type.
"""
# Lazy import: the registry self-populates as component modules load.
from .components import R
if not cfg.backend:
raise ValueError(f"{label} is missing the required 'backend' field")
backend_cls = R.get(ctype, cfg.backend)
if backend_cls is None:
raise ValueError(f"Unregistered backend '{cfg.backend}' for {label}")
params = cfg.model_dump()
params["app_context"] = self.context
if name is not None:
params.setdefault("name", name)
instance = backend_cls(**params)
if not isinstance(instance, expected_type):
got, want = type(instance).__name__, expected_type.__name__
raise TypeError(f"{label} backend '{cfg.backend}' produced {got}, expected {want} subclass")
return instance
# ----- Dependency ordering ------------------------------------------
def _topological_order(self) -> list[BaseComponent]:
"""Kahn's algorithm. Raises on missing required dep or cycle."""
nodes: dict[tuple[ComponentEnum, str], BaseComponent] = {
"""Return components in dependency order via Kahn's algorithm; raise on missing dep or cycle."""
nodes: dict[_NodeKey, 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",
)
in_degree, dependents = self._build_dependency_graph(nodes)
ready = [k for k, d in in_degree.items() if d == 0]
heapq.heapify(ready)
@ -110,25 +145,49 @@ class Application(BaseComponent):
raise ValueError(f"Circular dependency detected among: {unresolved}")
return ordered
@staticmethod
def _build_dependency_graph(
nodes: dict[_NodeKey, BaseComponent],
) -> tuple[dict[_NodeKey, int], dict[_NodeKey, list[_NodeKey]]]:
"""Compute in-degree and adjacency lists; raise if a required dep is missing."""
in_degree: dict[_NodeKey, int] = dict.fromkeys(nodes, 0)
dependents: dict[_NodeKey, list[_NodeKey]] = {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 unregistered {dep.ctype.value}:{dep.name}",
)
return in_degree, dependents
# ----- Lifecycle -----------------------------------------------------
async def _start(self) -> None:
"""Start components, then regular jobs, then background jobs; record order for reverse close."""
"""Start components in dependency order, then jobs (background last)."""
components = self._topological_order()
jobs = list(self.context.jobs.values())
sequence = (
components + [j for j in jobs if j.backend != "background"] + [j for j in jobs if j.backend == "background"]
)
# Background jobs come last so they observe a fully wired system.
foreground = [j for j in jobs if j.backend != "background"]
background = [j for j in jobs if j.backend == "background"]
for c in components + foreground + background:
await self._start_one(c)
for c in sequence:
try:
if c.backend == "background":
self.logger.info(f"Starting background job: {c.name}")
await c.start()
self._started_components.append(c)
except Exception as e:
self.logger.exception(f"Failed to start {c.component_type.value}:{c.name}: {e}")
async def _start_one(self, c: BaseComponent) -> None:
"""Start one component and record it for ordered shutdown; log and swallow failures."""
try:
if c.backend == "background":
self.logger.info(f"Starting background job: {c.name}")
await c.start()
self._started_components.append(c)
except Exception as e:
self.logger.exception(f"Failed to start {c.component_type.value}:{c.name}: {e}")
async def _close(self) -> None:
"""Close in reverse order of successful start."""
"""Close in reverse start order so every peer outlives its dependents."""
for c in reversed(self._started_components):
try:
await c.close()
@ -136,19 +195,20 @@ class Application(BaseComponent):
self.logger.exception(f"Failed to close {c.component_type.value}:{c.name}: {e}")
self._started_components.clear()
# ----- Job execution -------------------------------------------------
async def run_job(self, name: str, /, **kwargs) -> Response:
"""Execute a registered job by name."""
"""Execute a registered job by name and return its final Response."""
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."""
"""Execute a streaming job, yielding chunks as they are produced."""
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))
stream_queue: asyncio.Queue = asyncio.Queue()
task = asyncio.create_task(self.context.jobs[name](stream_queue=stream_queue, **kwargs))
async for chunk in execute_stream_task(
stream_queue=stream_queue,
task=task,
@ -159,8 +219,6 @@ class Application(BaseComponent):
yield chunk
def run_app(self):
"""Start the service and serve the application."""
from .components.service import BaseService
"""Serve the application through the configured service backend."""
assert isinstance(self.context.service, BaseService)
self.context.service.run_app(app=self)

View file

@ -5,6 +5,7 @@ from . import as_llm_formatter
from . import as_token_counter
from . import client
from . import embedding
from . import file_catalog
from . import file_graph
from . import file_parser
from . import file_store
@ -31,6 +32,7 @@ __all__ = [
"as_token_counter",
"client",
"embedding",
"file_catalog",
"file_graph",
"file_parser",
"file_store",

View file

@ -1,28 +1,29 @@
"""Application context: shared state container for components, jobs, and service."""
from typing import TYPE_CHECKING
from ..enumeration import ComponentEnum
from ..schema import ApplicationConfig
if TYPE_CHECKING:
from .base_component import BaseComponent
from .job import BaseJob
from .service import BaseService
class ApplicationContext:
"""Holds the parsed config and instantiated components, jobs, and service.
"""Passive state container holding parsed config and wired components.
Acts as a passive state container. The actual wiring (resolving backends from
the registry and instantiating each component) is performed by Application.
The Application class performs the actual wiring (registry lookups and
component instantiation); this class only stores the results so that
components, jobs, and the service can find each other at runtime.
"""
def __init__(self, **kwargs):
# Parse and validate raw config kwargs into a typed ApplicationConfig.
# Parse raw kwargs into a typed, validated config object.
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] = {}
# Populated by Application during initialization.
self.service: "BaseService | None" = None
self.components: dict[ComponentEnum, dict[str, "BaseComponent"]] = {}
self.jobs: dict[str, "BaseJob"] = {}

View file

@ -1,4 +1,4 @@
"""Base class for components."""
"""Base class for components with async lifecycle and dependency injection."""
import asyncio
from abc import ABC
@ -15,7 +15,11 @@ T = TypeVar("T", bound="BaseComponent")
class Dependency:
"""Declared dependency: bind() return value, instance attribute placeholder, and topological-sort edge."""
"""Placeholder returned by ``BaseComponent.bind`` for an unresolved dependency.
Resolved into a real component (or None) when the owning component starts.
Accessing any attribute before resolution raises a clear error.
"""
__slots__ = ("ctype", "name", "default_factory", "optional")
@ -36,9 +40,9 @@ class Dependency:
return f"<unresolved {self.ctype.value}:{self.name}{suffix}>"
def __getattr__(self, item: str) -> Any:
# Guard against using the dependency before start() resolves it.
# Catches accidental use of the placeholder before start() resolves it.
raise RuntimeError(
f"Dependency {self.ctype.value}:{self.name} accessed before start() (attribute '{item}')",
f"Dependency {self.ctype.value}:{self.name} accessed before start() " f"(attribute '{item}')",
)
@ -58,13 +62,14 @@ class BaseComponent(ABC):
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)
logger = get_logger()
self.logger = logger.bind(component=self.name) if hasattr(logger, "bind") else logger
self._is_started: bool = False
self._lock: asyncio.Lock = asyncio.Lock()
# Components created from bind() default_factory in standalone mode (auto-managed lifecycle).
# Components created via bind() default_factory in standalone mode;
# their lifecycle is owned by this component.
self._owned: list["BaseComponent"] = []
@property
@ -82,82 +87,110 @@ class BaseComponent(ABC):
default_factory: Callable[[], T] | None = None,
optional: bool = True,
) -> T | None:
"""Declare a dependency on another component; resolved at start(). Empty name → None."""
"""Declare a dependency on another component.
Returns a ``Dependency`` placeholder resolved into the real component
(or None / a factory-produced instance) when ``start`` runs. An empty
`name` short-circuits to None so callers can skip optional wiring.
"""
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'")
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."""
"""All unresolved dependency placeholders 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.")
"""Replace every ``Dependency`` attribute with its resolved target."""
for attr, dep in list(self.__dict__.items()):
if isinstance(dep, Dependency):
self._resolve_one(attr, dep)
# ----- Lookup --------------------------------------------------------
def _resolve_one(self, attr: str, dep: Dependency) -> None:
"""Resolve a single dependency, dispatching by mode."""
if self.app_context is None:
self._resolve_standalone(attr, dep)
else:
self._resolve_from_context(attr, dep)
def _resolve_standalone(self, attr: str, dep: Dependency) -> None:
"""Standalone mode: use default_factory, or fall back to None when optional.
Required dependencies without a factory keep the placeholder so any
attribute access surfaces a clear error at the call site.
"""
if dep.default_factory is not None:
instance = dep.default_factory()
setattr(self, attr, instance)
if isinstance(instance, BaseComponent):
self._owned.append(instance)
elif dep.optional:
setattr(self, attr, None)
def _resolve_from_context(self, attr: str, dep: Dependency) -> None:
"""Context-bound mode: look up the component from ``app_context.components``."""
target = self.app_context.components.get(dep.ctype, {}).get(dep.name)
if target is not None:
setattr(self, attr, target)
elif dep.optional:
setattr(self, attr, None)
else:
raise ValueError(f"{dep.ctype.value} '{dep.name}' not found.")
# ----- Vault path helpers --------------------------------------------
@property
def vault_path(self) -> Path:
"""Resolved vault root path from app context or cwd."""
"""Absolute vault root directory (cwd when no app_context is attached)."""
if self.app_context is None:
return Path.cwd()
return Path(self.app_context.app_config.vault_dir).absolute()
@property
def vault_metadata_path(self) -> Path:
"""Resolved metadata directory: vault_path / metadata_dir, or absolute metadata_dir."""
"""Vault metadata directory: ``<vault>/<metadata_dir>``."""
if self.app_context is None:
return Path.cwd() / "metadata"
return self.vault_path / self.app_context.app_config.metadata_dir
@property
def component_metadata_path(self) -> Path:
"""Per-component metadata directory under the vault."""
return self.vault_metadata_path / self.component_type.value
def to_vault_relative(self, path: str | Path) -> str:
"""Return path relative to vault_path; absolute path string if outside."""
"""Convert `path` to a vault-relative string; return absolute path when outside."""
abs_path = Path(path).absolute()
try:
return str(abs_path.relative_to(self.vault_path))
except ValueError:
return str(abs_path)
# ----- Lifecycle -----------------------------------------------------
# ----- Lifecycle hooks (override in subclasses) ----------------------
async def _start(self) -> None:
"""Subclass hook: start logic."""
"""Subclass hook called once after dependencies are resolved."""
async def _close(self) -> None:
"""Subclass hook: close logic."""
"""Subclass hook called once during ``close``."""
async def dump(self) -> None:
"""Persist in-memory state to disk. Override in subclasses that need persistence."""
"""Persist in-memory state to disk. Override when persistence is needed."""
async def load(self) -> None:
"""Restore in-memory state from disk. Override in subclasses that need persistence."""
"""Restore in-memory state from disk. Override when persistence is needed."""
# ----- Lifecycle control --------------------------------------------
async def start(self) -> None:
"""Resolve bindings → start owned fallbacks → _start(). No-op if already started."""
"""Start the component once: resolve deps → start owned → run _start."""
async with self._lock:
if self._is_started:
return
@ -168,7 +201,7 @@ class BaseComponent(ABC):
self._is_started = True
async def close(self) -> None:
"""_close() → close owned fallbacks in reverse. No-op if not started."""
"""Close the component once: run _close → close owned in reverse order."""
async with self._lock:
if not self._is_started:
return
@ -178,7 +211,7 @@ class BaseComponent(ABC):
self._is_started = False
async def restart(self) -> None:
"""Close then start."""
"""Close then start the component."""
await self.close()
await self.start()

View file

@ -1,4 +1,4 @@
"""Global registry mapping (ComponentEnum, name) -> component class."""
"""Global registry mapping ``(ComponentEnum, name) -> component class``."""
from typing import Callable, TypeVar, cast
@ -10,7 +10,7 @@ T = TypeVar("T", bound=BaseComponent)
class ComponentRegistry:
"""Two-level registry: component_type -> name -> class.
"""Two-level registry: ``component_type -> name -> class``.
Supports both direct calls — ``R.register(MyClass, "name")`` — and
decorator usage — ``@R.register("name")``.
@ -21,16 +21,20 @@ class ComponentRegistry:
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."""
"""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")
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")
self.logger.warning(
f"Component '{name}' already registered for {component_type}, overwriting",
)
group[name] = cls
return cls
@ -40,16 +44,19 @@ class ComponentRegistry:
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.
# Direct call: register(MyClass) or register(MyClass, "alias").
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__)
cls = cast(type[T], cls_or_name)
return self._do_register(cls, name if name is not None else cls.__name__)
# Decorator mode: first arg is the registration name.
# Decorator call: @R.register("alias") — must receive a string name.
if not isinstance(cls_or_name, str):
raise TypeError(f"Expected a class or string, got {type(cls_or_name).__name__}")
registration_name = cls_or_name
def decorator(decorated_cls: type[T]) -> type[T]:
return self._do_register(decorated_cls, cls_or_name)
return self._do_register(decorated_cls, registration_name)
return decorator

View file

@ -13,9 +13,11 @@ from ..base_component import BaseComponent
from ...enumeration import ComponentEnum
from ...schema import EmbNode
Miss = tuple[int, str, str] # (result_index, text, cache_key)
class BaseEmbeddingModel(BaseComponent):
"""Embedding model with LRU cache and disk persistence."""
"""Embedding model with LRU cache, disk persistence, and concurrent batching."""
component_type = ComponentEnum.EMBEDDING_MODEL
@ -29,6 +31,7 @@ class BaseEmbeddingModel(BaseComponent):
max_batch_size: int = 10,
max_input_length: int = 8192,
max_cache_size: int = 10000,
max_concurrency: int = 2,
enable_cache: bool = True,
cache_version: str = "v1",
max_retries: int = 3,
@ -43,21 +46,25 @@ class BaseEmbeddingModel(BaseComponent):
self.max_batch_size = max_batch_size
self.max_input_length = max_input_length
self.max_cache_size = max_cache_size
self.max_concurrency = max_concurrency
self.enable_cache = enable_cache
self.cache_version = cache_version
self.max_retries = max_retries
self._embedding_cache: OrderedDict[str, np.ndarray] = OrderedDict()
self._cache: OrderedDict[str, np.ndarray] = OrderedDict()
self._key_suffix = f"|{model_name}|{dimensions}".encode()
self.is_healthy: bool = True
@property
def cache_path(self) -> Path:
"""Disk path for the embedding cache file."""
"""Path of the persisted embedding cache, namespaced by name and version."""
return self.vault_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 _close(self) -> None:
await self.dump()
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}"
@ -75,71 +82,23 @@ class BaseEmbeddingModel(BaseComponent):
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."""
"""Embed a single text; returns None if the provider yields nothing."""
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)
"""Get embeddings for texts. Cache hits return immediately; misses run concurrently."""
texts = [self._truncate(t) for t in input_text]
results, misses = self._partition_by_cache(texts)
if misses:
await self._fill_misses(misses, results, **kwargs)
return results
async def get_node_embeddings(self, nodes: list[EmbNode], **kwargs) -> list[EmbNode]:
"""Compute and assign embeddings for EmbNode objects."""
"""Embed each node's text in-place and return the same list."""
embeddings = await self.get_embeddings([n.text for n in nodes], **kwargs)
if len(embeddings) == len(nodes):
for node, vec in zip(nodes, embeddings):
@ -151,64 +110,134 @@ class BaseEmbeddingModel(BaseComponent):
async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float] | None]:
"""Get raw embeddings from the underlying provider."""
# -- Cache Operations --
# -- Batching --
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 _truncate(self, text: str) -> str:
return text if len(text) <= self.max_input_length else text[: self.max_input_length]
def _put_to_cache(self, text: str, embedding: np.ndarray) -> None:
"""Insert into LRU cache, evicting oldest if full."""
def _partition_by_cache(self, texts: list[str]) -> tuple[list[np.ndarray | None], list[Miss]]:
"""Split texts into pre-filled results (hits) and a miss list to compute."""
results: list[np.ndarray | None] = [None] * len(texts)
misses: list[Miss] = []
for idx, text in enumerate(texts):
key = self._cache_key(text)
hit = self._cache_get(key)
if hit is not None:
results[idx] = hit
else:
misses.append((idx, text, key))
return results, misses
async def _fill_misses(self, misses: list[Miss], results: list[np.ndarray | None], **kwargs) -> None:
"""Compute miss embeddings in concurrent batches and write into results + cache."""
size = self.max_batch_size
batches = [misses[i : i + size] for i in range(0, len(misses), size)]
sem = asyncio.Semaphore(self.max_concurrency)
async def run(batch: list[Miss]) -> list[tuple[int, str, np.ndarray]]:
async with sem:
return await self._compute_batch(batch, **kwargs)
for done in await asyncio.gather(*(run(b) for b in batches)):
for idx, key, emb in done:
results[idx] = emb
self._cache_put(key, emb)
async def _compute_batch(self, batch: list[Miss], **kwargs) -> list[tuple[int, str, np.ndarray]]:
"""Call provider for one batch with retry; returns [(idx, key, embedding)]."""
texts = [text for _, text, _ in batch]
embeddings = await self._call_with_retry(texts, **kwargs)
if not embeddings or len(embeddings) != len(texts):
return []
out: list[tuple[int, str, np.ndarray]] = []
for (idx, _text, key), raw in zip(batch, embeddings):
if raw is None:
continue
emb = self._normalize_dim(np.asarray(raw, dtype=np.float16))
out.append((idx, key, emb))
return out
async def _call_with_retry(self, texts: list[str], **kwargs) -> list[list[float] | None] | None:
"""Call provider with exponential backoff on transient errors."""
for attempt in range(self.max_retries):
try:
result = await self._get_embeddings(texts, **kwargs)
if result and len(result) == len(texts):
return result
except (TimeoutError, ConnectionError, OSError):
if attempt < self.max_retries - 1:
await asyncio.sleep(2**attempt)
except Exception:
self.logger.exception("Embedding request failed")
return None
return None
def _normalize_dim(self, emb: np.ndarray) -> np.ndarray:
if len(emb) == self.dimensions:
return emb
if len(emb) < self.dimensions:
return np.pad(emb, (0, self.dimensions - len(emb)))
return emb[: self.dimensions]
# -- Cache --
def _cache_key(self, text: str) -> str:
return hashlib.sha256(text.encode() + self._key_suffix).hexdigest()
def _cache_get(self, key: str) -> np.ndarray | None:
if not self.enable_cache or key not in self._cache:
return None
self._cache.move_to_end(key)
return self._cache[key]
def _cache_put(self, key: str, embedding: np.ndarray) -> None:
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)
cache = self._cache
if key in cache:
cache.move_to_end(key)
cache[key] = embedding
return
if len(cache) >= self.max_cache_size:
cache.popitem(last=False)
cache[key] = embedding
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 --
# -- Persistence --
async def load(self) -> None:
"""Load cached embeddings from disk (npz format); replaces in-memory cache."""
self._embedding_cache.clear()
"""Load cached embeddings from disk (npz); replaces in-memory cache."""
self._cache.clear()
if not self.enable_cache or not self.cache_path.exists():
return
await asyncio.to_thread(self._load_sync)
def _load_sync(self) -> None:
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:
if len(self._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}")
self._cache[str(key)] = emb.astype(np.float16)
self.logger.info(f"Loaded {len(self._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:
"""Persist in-memory cache to disk (npz)."""
if not self.enable_cache or not self._cache:
return
await asyncio.to_thread(self._dump_sync)
def _dump_sync(self) -> None:
self.cache_path.parent.mkdir(parents=True, exist_ok=True)
keys = list(self._embedding_cache.keys())
embeddings = np.stack(list(self._embedding_cache.values()))
keys = np.array(list(self._cache.keys()), dtype=str)
embeddings = np.stack(list(self._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}")
np.savez(self.cache_path, keys=keys, embeddings=embeddings)
self.logger.info(f"Saved {len(self._cache)} embeddings to {self.cache_path}")
except Exception:
self.logger.exception("Failed to save embedding cache")

View file

@ -0,0 +1,9 @@
"""file catalog"""
from .base_file_catalog import BaseFileCatalog
from .local_file_catalog import LocalFileCatalog
__all__ = [
"BaseFileCatalog",
"LocalFileCatalog",
]

View file

@ -0,0 +1,39 @@
"""Abstract base class for file catalog backends."""
from abc import abstractmethod
from ..base_component import BaseComponent
from ...enumeration import ComponentEnum
from ...schema import FileNode
class BaseFileCatalog(BaseComponent):
"""File-catalog backend recording FileNode entries keyed by path."""
component_type = ComponentEnum.FILE_CATALOG
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 persisted state. No-op without local files."""
async def dump(self) -> None:
"""Persist state. No-op without local files."""
@abstractmethod
async def upsert(self, nodes: list[FileNode]) -> None:
"""Insert or update nodes keyed by path."""
@abstractmethod
async def delete(self, path: str | list[str]) -> None:
"""Delete nodes by path; missing paths are skipped."""
@abstractmethod
async def get_nodes(self, paths: list[str] | None = None) -> list[FileNode]:
"""Return nodes by paths; None = all; missing paths are skipped."""

View file

@ -0,0 +1,62 @@
"""Local file catalog backend: in-memory dict persisted as JSONL."""
import aiofiles
from .base_file_catalog import BaseFileCatalog
from ..component_registry import R
from ...schema import FileNode
@R.register("local")
class LocalFileCatalog(BaseFileCatalog):
"""Dict-backed catalog persisted as JSONL."""
def __init__(self, encoding: str = "utf-8", **kwargs):
super().__init__(**kwargs)
self.encoding = encoding
self._nodes: dict[str, FileNode] = {}
self.component_metadata_path.mkdir(parents=True, exist_ok=True)
self._catalog_file = self.component_metadata_path / f"{self.name}.jsonl"
async def load(self) -> None:
if not self._catalog_file.exists():
return
try:
await self._read_jsonl()
self.logger.info(f"Loaded {len(self._nodes)} nodes from {self._catalog_file}")
except Exception as e:
self.logger.exception(f"Failed to load {self._catalog_file}: {e}")
async def dump(self) -> None:
try:
await self._write_jsonl()
self.logger.info(f"Saved {len(self._nodes)} nodes to {self._catalog_file}")
except Exception as e:
self.logger.exception(f"Failed to write {self._catalog_file}: {e}")
async def upsert(self, nodes: list[FileNode]) -> None:
for node in nodes:
self._nodes[node.path] = node
async def delete(self, path: str | list[str]) -> None:
paths = [path] if isinstance(path, str) else path
for p in paths:
self._nodes.pop(p, None)
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 _read_jsonl(self) -> None:
async with aiofiles.open(self._catalog_file, encoding=self.encoding) as f:
async for line in f:
if stripped := line.strip():
node = FileNode.model_validate_json(stripped)
self._nodes[node.path] = node
async def _write_jsonl(self) -> None:
tmp = self._catalog_file.with_suffix(".tmp")
async with aiofiles.open(tmp, "w", encoding=self.encoding) as f:
await f.write("\n".join(n.model_dump_json() for n in self._nodes.values()))
tmp.replace(self._catalog_file)

View file

@ -5,4 +5,9 @@ from .local_file_graph import LocalFileGraph
from .neo4j_file_graph import Neo4jFileGraph
from .nx_file_graph import NxFileGraph
__all__ = ["BaseFileGraph", "LocalFileGraph", "Neo4jFileGraph", "NxFileGraph"]
__all__ = [
"BaseFileGraph",
"LocalFileGraph",
"Neo4jFileGraph",
"NxFileGraph",
]

View file

@ -1,7 +1,6 @@
"""Abstract base for file-graph backends."""
from abc import abstractmethod
from pathlib import Path
from ..base_component import BaseComponent
from ...enumeration import ComponentEnum, LinkScopeEnum
@ -9,17 +8,17 @@ from ...schema import FileLink, FileNode
class BaseFileGraph(BaseComponent):
"""Abstract base for file-graph backends."""
"""Abstract base for file-graph backends.
Link scope (``get_outlinks`` / ``get_inlinks``):
REAL edges touching an indexed node
VIRTUAL edges touching a dangling placeholder
(referenced but never upserted, or already deleted)
ALL both
"""
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.vault_metadata_path / self.component_type.value
self.graph_path.mkdir(parents=True, exist_ok=True)
# -- Lifecycle ---------------------------------------------------------
async def _start(self) -> None:
@ -31,13 +30,7 @@ class BaseFileGraph(BaseComponent):
await super()._close()
async def load(self) -> None:
"""Load persisted state. No-op for backends without local files.
Called at the end of ``_start()`` after base resources are ready
but before subclass-specific resources are initialised. Backends
that need their own resources for loading should override
``_start()`` instead of this hook.
"""
"""Restore persisted state. No-op for backends without local files."""
async def dump(self) -> None:
"""Persist state. No-op for backends without local files."""
@ -46,15 +39,15 @@ class BaseFileGraph(BaseComponent):
@abstractmethod
async def upsert_nodes(self, nodes: list[FileNode]) -> None:
"""Insert or update nodes in the graph."""
"""Insert or update nodes."""
@abstractmethod
async def delete_nodes(self, paths: list[str]) -> None:
"""Delete nodes by path."""
"""Remove nodes by path."""
@abstractmethod
async def get_nodes(self, paths: list[str] | None = None) -> list[FileNode]:
"""Return nodes by paths; None = all real nodes; [] = []."""
"""Return nodes by paths; ``None`` = all real nodes."""
@abstractmethod
async def rebuild_links(self) -> None:
@ -67,30 +60,9 @@ class BaseFileGraph(BaseComponent):
# -- Link access -------------------------------------------------------
@abstractmethod
async def get_outlinks(
self,
path: str,
scope: LinkScopeEnum = LinkScopeEnum.REAL,
) -> list[FileLink]:
"""Return outgoing links for *path*.
``scope=REAL`` (default) → edges whose target is an indexed
(real) node. ``scope=VIRTUAL`` → only dangling edges (target
was referenced but never upserted, or was deleted). ``ALL``
→ both.
"""
async def get_outlinks(self, path: str, scope: LinkScopeEnum = LinkScopeEnum.REAL) -> list[FileLink]:
"""Outgoing links from *path*."""
@abstractmethod
async def get_inlinks(
self,
path: str,
scope: LinkScopeEnum = LinkScopeEnum.REAL,
) -> list[FileLink]:
"""Return incoming links for *path*.
``scope=REAL`` (default) → returns the inbound edges only when
*path* itself is a real node. ``scope=VIRTUAL`` → returns
inbound edges only when *path* is virtual (useful for retarget
/ lint that must see references to non-existent targets).
``ALL`` → both.
"""
async def get_inlinks(self, path: str, scope: LinkScopeEnum = LinkScopeEnum.REAL) -> list[FileLink]:
"""Inbound links to *path*."""

View file

@ -15,9 +15,10 @@ class LocalFileGraph(BaseFileGraph):
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"
self._inverse: dict[str, set[str]] = {} # real target → sources
self._pending: dict[str, set[str]] = {} # virtual target → sources
self.component_metadata_path.mkdir(parents=True, exist_ok=True)
self._graph_file: Path = self.component_metadata_path / f"{self.name}.jsonl"
# -- Lifecycle ---------------------------------------------------------
@ -26,20 +27,19 @@ class LocalFileGraph(BaseFileGraph):
await self.rebuild_links()
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)]
)
for line in f:
if line.strip():
node = FileNode.model_validate_json(line)
self._nodes[node.path] = node
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:
@ -49,22 +49,29 @@ class LocalFileGraph(BaseFileGraph):
except Exception as e:
self.logger.exception(f"Failed to write {self._graph_file}: {e}")
# -- Edge bookkeeping --------------------------------------------------
# -- Internals ---------------------------------------------------------
@staticmethod
def _targets(node: FileNode) -> list[str]:
return [lnk.target_path for lnk in node.links if lnk.target_path]
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]
if srcs and src in srcs:
srcs.discard(src)
if not srcs:
del bucket[target]
def _scope_match(self, target: str, scope: LinkScopeEnum) -> bool:
if scope is LinkScopeEnum.ALL:
return True
is_real = target in self._nodes
return is_real if scope is LinkScopeEnum.REAL else not is_real
# -- Node CRUD ---------------------------------------------------------
@ -73,14 +80,11 @@ class LocalFileGraph(BaseFileGraph):
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)
for target in self._targets(old):
self._remove_edge(path, target)
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.
for target in self._targets(node):
self._add_edge(path, target)
promoted = self._pending.pop(path, None)
if promoted:
self._inverse.setdefault(path, set()).update(promoted)
@ -90,10 +94,8 @@ class LocalFileGraph(BaseFileGraph):
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).
for target in self._targets(node):
self._remove_edge(path, target)
demoted = self._inverse.pop(path, None)
if demoted:
self._pending.setdefault(path, set()).update(demoted)
@ -104,13 +106,11 @@ class LocalFileGraph(BaseFileGraph):
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)
for target in self._targets(node):
self._add_edge(src, target)
async def clear(self):
self._nodes.clear()
@ -120,26 +120,13 @@ class LocalFileGraph(BaseFileGraph):
# -- Link access -------------------------------------------------------
async def get_outlinks(
self,
path: str,
scope: LinkScopeEnum = LinkScopeEnum.REAL,
) -> list[FileLink]:
# Source must be real (only real nodes carry a ``links`` payload).
# Targets may be real or virtual; ``scope`` selects which to surface.
async def get_outlinks(self, path: str, scope: LinkScopeEnum = LinkScopeEnum.REAL) -> list[FileLink]:
node = self._nodes.get(path)
if node is None:
return []
return [lnk for lnk in node.links if lnk.target_path and _match_target(lnk.target_path, self._nodes, scope)]
return [lnk for lnk in node.links if lnk.target_path and self._scope_match(lnk.target_path, scope)]
async def get_inlinks(
self,
path: str,
scope: LinkScopeEnum = LinkScopeEnum.REAL,
) -> list[FileLink]:
# ``_inverse`` keys real targets; ``_pending`` keys virtual ones.
# The queried ``path`` lives in at most one bucket, so ``scope``
# is satisfied by selecting which bucket to read.
async def get_inlinks(self, path: str, scope: LinkScopeEnum = LinkScopeEnum.REAL) -> list[FileLink]:
sources: set[str] = set()
if scope in (LinkScopeEnum.REAL, LinkScopeEnum.ALL):
sources |= self._inverse.get(path, set())
@ -148,11 +135,3 @@ class LocalFileGraph(BaseFileGraph):
return [
link for src in sources if src in self._nodes for link in self._nodes[src].links if link.target_path == path
]
def _match_target(target_path: str, nodes: dict, scope: LinkScopeEnum) -> bool:
"""Whether an edge into ``target_path`` should be surfaced under ``scope``."""
if scope is LinkScopeEnum.ALL:
return True
is_real = target_path in nodes
return is_real if scope is LinkScopeEnum.REAL else not is_real

View file

@ -103,7 +103,7 @@ class Neo4jFileGraph(BaseFileGraph):
)
real, virtual, edges = await self._counts(session)
self.logger.info(
f"Neo4jFileGraph '{self.graph_name}' connected at "
f"Neo4jFileGraph '{self.name}' connected at "
f"{self._uri}/{self._database}: "
f"{real} nodes, {edges} edges, {virtual} virtual",
)

View file

@ -16,121 +16,109 @@ from ...schema import FileLink, FileNode
@R.register("nx")
class NxFileGraph(BaseFileGraph):
"""Networkx-backed file graph; uses FileLink.target_path for adjacency."""
"""Networkx-backed file graph; uses FileLink.target_path for adjacency.
Real node carries ``node`` attr; virtual (dangling target) does not.
"""
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"
self.component_metadata_path.mkdir(parents=True, exist_ok=True)
self._graph_file: Path = self.component_metadata_path / f"{self.name}.pkl"
# -- Lifecycle ---------------------------------------------------------
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}")
self.logger.info(f"Loaded {self._real_count()} 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}")
self.logger.info(f"Saved {self._real_count()} nodes to {self._graph_file}")
except Exception as e:
self.logger.exception(f"Failed to write {self._graph_file}: {e}")
# -- Internals ---------------------------------------------------------
def _real_count(self) -> int:
return sum(1 for _, d in self._graph.nodes(data=True) if "node" in d)
def _is_real(self, key: str) -> bool:
return "node" in self._graph.nodes[key]
@staticmethod
def _edges_from(src: str, node: FileNode):
return ((src, lnk.target_path, {"link": lnk}) for lnk in node.links if lnk.target_path)
def _scope_match(self, key: str, scope: LinkScopeEnum) -> bool:
if scope is LinkScopeEnum.ALL:
return True
is_real = self._is_real(key)
return is_real if scope is LinkScopeEnum.REAL else not is_real
# -- 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)
self._graph.add_node(path, node=node) # promotes virtual placeholder
self._graph.add_edges_from(self._edges_from(path, node))
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)
self._graph.nodes[path].pop("node", None) # demote to virtual
if self._graph.in_degree(path) == 0:
self._graph.remove_node(path) # remove orphan virtual node
self._graph.remove_node(path)
async def get_nodes(self, paths: list[str] | None = None) -> list[FileNode]:
nodes_view = self._graph.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]]
return [d["node"] for _, d in view(data=True) if "node" in d]
return [view[p]["node"] for p in paths if p in view and "node" in view[p]]
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
)
for path, data in list(self._graph.nodes(data=True)):
self._graph.add_edges_from(self._edges_from(path, data["node"]))
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,
scope: LinkScopeEnum = LinkScopeEnum.REAL,
) -> list[FileLink]:
# Source must be real; targets may be virtual placeholders.
# ``scope`` picks real / virtual / both targets.
nodes_view = self._graph.nodes
if path not in nodes_view or "node" not in nodes_view[path]:
async def get_outlinks(self, path: str, scope: LinkScopeEnum = LinkScopeEnum.REAL) -> list[FileLink]:
view = self._graph.nodes
if path not in view or "node" not in view[path]:
return []
return [
d["link"]
for _, tgt, d in self._graph.out_edges(path, data=True)
if "link" in d and _match_node(nodes_view, tgt, scope)
if "link" in d and self._scope_match(tgt, scope)
]
async def get_inlinks(
self,
path: str,
scope: LinkScopeEnum = LinkScopeEnum.REAL,
) -> list[FileLink]:
# ``path`` is a single node — its realness selects which scope
# produces a non-empty result (REAL ↔ real node, VIRTUAL ↔
# virtual placeholder; ALL is always allowed).
nodes_view = self._graph.nodes
if path not in nodes_view or not _match_node(nodes_view, path, scope):
async def get_inlinks(self, path: str, scope: LinkScopeEnum = LinkScopeEnum.REAL) -> list[FileLink]:
view = self._graph.nodes
if path not in view or not self._scope_match(path, scope):
return []
return [d["link"] for _, _, d in self._graph.in_edges(path, data=True) if "link" in d]
def _match_node(nodes_view, key: str, scope: LinkScopeEnum) -> bool:
"""Whether ``key`` satisfies ``scope`` under the nx node-realness convention."""
if scope is LinkScopeEnum.ALL:
return True
is_real = "node" in nodes_view[key]
return is_real if scope is LinkScopeEnum.REAL else not is_real

View file

@ -9,7 +9,7 @@ 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]``. Wikilink extraction is
delegated to :class:`reme4.utils.wikilink_handler.WikilinkHandler` —
delegated to :class:`reme.utils.wikilink_handler.WikilinkHandler` —
the single source of truth for ``[[...]]`` syntax (including
Dataview-style typed predicates).
"""

View file

@ -10,34 +10,33 @@ from ...schema import FileChunk, FileLink, FileNode
class BaseFileStore(BaseComponent):
"""Abstract base for file store backends.
Defines the *semantic* contract a file store must offer: write (upsert / delete / clear),
retrieve (vector / keyword), and graph queries (nodes / links). Sub-component composition
(embedding model, keyword index, file graph) is each backend's implementation choice and
is not part of the base contract.
Defines the *semantic* contract a file store must offer: write (upsert / delete /
clear), retrieve (vector / keyword), and graph queries (nodes / links). How the
backend composes sub-components (embedding model, keyword index, file graph) is
an implementation detail outside this contract.
"""
component_type = ComponentEnum.FILE_STORE
def __init__(self, store_name: str, store_version: str = "v1", **kwargs):
super().__init__(**kwargs)
self.store_name = store_name or self.name
self.store_version = store_version
self.store_path = self.vault_metadata_path / self.component_type.value / self.store_name
self.store_path.mkdir(parents=True, exist_ok=True)
# -- CRUD ------------------------------------------------------------
# -- CRUD -----------------------------------------------------------------
@abstractmethod
async def upsert(self, files: list[tuple[FileNode, list[FileChunk]]]) -> None:
"""Upsert files and their chunks into the store."""
"""Upsert files and their chunks; existing chunks for the same path are replaced."""
@abstractmethod
async def delete(self, path: str | list[str]) -> None:
"""Delete files by path from the store."""
"""Delete the given path(s) and all their chunks; unknown paths are skipped."""
@abstractmethod
async def clear(self) -> None:
"""Drop every file and chunk in the store."""
# -- graph queries --------------------------------------------------------
@abstractmethod
async def get_nodes(self, paths: list[str] | None = None) -> list[FileNode]:
"""Return file nodes; None = all nodes; missing paths are skipped."""
"""Return file nodes; ``None`` = all; missing paths are skipped."""
@abstractmethod
async def get_outlinks(
@ -45,7 +44,7 @@ class BaseFileStore(BaseComponent):
path: str,
scope: LinkScopeEnum = LinkScopeEnum.REAL,
) -> list[FileLink]:
"""Return outgoing links for *path*. See ``BaseFileGraph.get_outlinks`` for scope semantics."""
"""Outgoing links for *path*; scope semantics match ``BaseFileGraph.get_outlinks``."""
@abstractmethod
async def get_inlinks(
@ -53,18 +52,14 @@ class BaseFileStore(BaseComponent):
path: str,
scope: LinkScopeEnum = LinkScopeEnum.REAL,
) -> list[FileLink]:
"""Return incoming links for *path*. See ``BaseFileGraph.get_inlinks`` for scope semantics."""
"""Incoming links for *path*; scope semantics match ``BaseFileGraph.get_inlinks``."""
@abstractmethod
async def clear(self) -> None:
"""Clear the store of all files and chunks."""
# -- Search -----------------------------------------------------------
# -- search ---------------------------------------------------------------
@abstractmethod
async def vector_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]:
"""Perform vector similarity search."""
"""Vector similarity search over chunk embeddings."""
@abstractmethod
async def keyword_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]:
"""Perform full-text keyword search."""
"""Full-text keyword search over chunk text."""

View file

@ -20,7 +20,7 @@ class FaissLocalFileStore(LocalFileStore):
remains the source of truth.
faiss is imported lazily inside ``__init__`` so that merely importing this
module (e.g. via ``reme4 version``) does not trigger the SWIG bindings and
module (e.g. via ``reme version``) does not trigger the SWIG bindings and
their associated DeprecationWarnings.
"""
@ -31,21 +31,25 @@ class FaissLocalFileStore(LocalFileStore):
**kwargs,
):
super().__init__(**kwargs)
self._faiss = self._import_faiss()
self.normalize = normalize
self.max_tombstones = max_tombstones
self.faiss_path = self.component_metadata_path / f"faiss_index_{self.name}_{self.store_version}.bin"
self.faiss_idmap_path = self.component_metadata_path / f"faiss_idmap_{self.name}_{self.store_version}.json"
self._faiss_index = None # faiss.Index | None
self._id_map: list[str] = [] # row -> chunk_id
self._id_to_row: dict[str, int] = {} # chunk_id -> row (live entries only)
self._tombstones: set[int] = set() # rows whose chunk_id was deleted
@staticmethod
def _import_faiss():
try:
import faiss
except ImportError as e:
raise ImportError(
"faiss is required for FaissLocalFileStore. Install with `pip install faiss-cpu`.",
) from e
self._faiss = faiss
self.normalize = normalize
self.max_tombstones = max_tombstones
self.faiss_path = self.store_path / f"faiss_index_{self.store_version}.bin"
self.faiss_idmap_path = self.store_path / f"faiss_idmap_{self.store_version}.json"
self._faiss_index = None # faiss.Index | None
self._id_map: list[str] = [] # row -> chunk_id
self._id_to_row: dict[str, int] = {} # chunk_id -> row
self._tombstones: set[int] = set() # rows whose chunk_id was deleted
return faiss
# -- helpers ----------------------------------------------------------
@ -105,57 +109,64 @@ class FaissLocalFileStore(LocalFileStore):
# -- persistence ------------------------------------------------------
async def load(self) -> None:
"""Load chunks (parent), then load FAISS sidecar; rebuild from chunks on miss/corruption."""
"""Load chunks via the parent, then attach FAISS state (sidecar or rebuild)."""
await super().load()
if self.embedding_model is None or self._dim == 0:
self._faiss_index = None
return
loaded = False
if self.faiss_path.exists() and self.faiss_idmap_path.exists():
try:
index = self._faiss.read_index(str(self.faiss_path))
if index.d != self._dim:
raise ValueError(f"FAISS dim {index.d} != embedding dim {self._dim}")
async with aiofiles.open(self.faiss_idmap_path, encoding=self.encoding) as f:
data = json.loads(await f.read())
id_map = list(data.get("id_map", []))
if len(id_map) != index.ntotal:
raise ValueError(f"id_map size {len(id_map)} != index ntotal {index.ntotal}")
self._faiss_index = index
self._id_map = id_map
self._tombstones = set(data.get("tombstones", []))
self._id_to_row = {cid: i for i, cid in enumerate(self._id_map) if i not in self._tombstones}
self.logger.info(f"Loaded FAISS index: {index.ntotal} vectors from {self.faiss_path}")
loaded = True
except Exception as e:
self.logger.exception(f"Failed to load FAISS index, will rebuild: {e}")
self.faiss_path.unlink(missing_ok=True)
self.faiss_idmap_path.unlink(missing_ok=True)
if not loaded:
if not await self._try_load_sidecar():
self._rebuild_index()
async def _try_load_sidecar(self) -> bool:
"""Read the binary index plus id-map sidecar. On any mismatch or read error,
wipe the partial files so the caller can rebuild from chunks cleanly.
"""
if not (self.faiss_path.exists() and self.faiss_idmap_path.exists()):
return False
try:
index = self._faiss.read_index(str(self.faiss_path))
if index.d != self._dim:
raise ValueError(f"FAISS dim {index.d} != embedding dim {self._dim}")
async with aiofiles.open(self.faiss_idmap_path, encoding=self.encoding) as f:
data = json.loads(await f.read())
id_map = list(data.get("id_map", []))
if len(id_map) != index.ntotal:
raise ValueError(f"id_map size {len(id_map)} != index ntotal {index.ntotal}")
self._faiss_index = index
self._id_map = id_map
self._tombstones = set(data.get("tombstones", []))
self._id_to_row = {cid: i for i, cid in enumerate(self._id_map) if i not in self._tombstones}
self.logger.info(f"Loaded FAISS index: {index.ntotal} vectors from {self.faiss_path}")
return True
except Exception as e:
self.logger.exception(f"Failed to load FAISS index, will rebuild: {e}")
self.faiss_path.unlink(missing_ok=True)
self.faiss_idmap_path.unlink(missing_ok=True)
return False
async def dump(self) -> None:
"""Persist chunks JSONL (parent) plus FAISS sidecar via atomic rename."""
"""Persist chunks JSONL via the parent, then write the FAISS sidecar atomically."""
await super().dump()
if self._faiss_index is None or self.embedding_model is None:
return
try:
self._compact_if_needed()
tmp_index = self.faiss_path.with_suffix(".tmp")
self._faiss.write_index(self._faiss_index, str(tmp_index))
tmp_index.replace(self.faiss_path)
tmp_idmap = self.faiss_idmap_path.with_suffix(".tmp")
payload = json.dumps({"id_map": self._id_map, "tombstones": sorted(self._tombstones)})
async with aiofiles.open(tmp_idmap, "w", encoding=self.encoding) as f:
await f.write(payload)
tmp_idmap.replace(self.faiss_idmap_path)
await self._write_sidecar()
self.logger.info(f"Saved FAISS index: {self._faiss_index.ntotal} vectors to {self.faiss_path}")
except Exception as e:
self.logger.exception(f"Failed to write FAISS index: {e}")
async def _write_sidecar(self) -> None:
tmp_index = self.faiss_path.with_suffix(".tmp")
self._faiss.write_index(self._faiss_index, str(tmp_index))
tmp_index.replace(self.faiss_path)
tmp_idmap = self.faiss_idmap_path.with_suffix(".tmp")
payload = json.dumps({"id_map": self._id_map, "tombstones": sorted(self._tombstones)})
async with aiofiles.open(tmp_idmap, "w", encoding=self.encoding) as f:
await f.write(payload)
tmp_idmap.replace(self.faiss_idmap_path)
# -- CRUD overrides ---------------------------------------------------
async def upsert(self, files: list[tuple[FileNode, list[FileChunk]]]) -> None:
@ -163,23 +174,27 @@ class FaissLocalFileStore(LocalFileStore):
return
assert self.file_graph is not None
# Snapshot the chunk_ids the file_graph currently holds for these paths,
# so we can compute add/delete deltas after super finishes.
# Snapshot pre-upsert chunk_ids so we can diff against the post-upsert state.
old_ids_by_path = {
n.path: set(n.chunk_ids) for n in await self.file_graph.get_nodes([node.path for node, _ in files])
}
await super().upsert(files)
if self._faiss_index is None or self.embedding_model is None:
return
self._sync_index_after_upsert(files, old_ids_by_path)
def _sync_index_after_upsert(
self,
files: list[tuple[FileNode, list[FileChunk]]],
old_ids_by_path: dict[str, set[str]],
) -> None:
"""Apply add/tombstone deltas to FAISS based on chunk_id set differences."""
existing = set(self._id_to_row)
to_add: list[FileChunk] = []
for node, _ in files:
new_ids = set(node.chunk_ids)
old_ids = old_ids_by_path.get(node.path, set())
for cid in old_ids - new_ids:
for cid in old_ids_by_path.get(node.path, set()) - new_ids:
self._tombstone(cid)
for cid in new_ids - existing:
chunk = self.file_chunks.get(cid)
@ -228,13 +243,16 @@ class FaissLocalFileStore(LocalFileStore):
if query_embedding is None:
return []
# Over-fetch by len(tombstones) so dropped rows can't starve the result set.
q = self._prepare(query_embedding)
# Over-fetch to compensate for tombstoned rows.
k = min(self._faiss_index.ntotal, limit + len(self._tombstones))
scores, rows = self._faiss_index.search(q, k)
return self._collect_hits(rows[0].tolist(), scores[0].tolist(), limit)
def _collect_hits(self, rows: list[int], scores: list[float], limit: int) -> list[FileChunk]:
"""Map raw FAISS rows back to chunks, skipping tombstones and stale ids."""
results: list[FileChunk] = []
for raw_row, score in zip(rows[0].tolist(), scores[0].tolist()):
for raw_row, score in zip(rows, scores):
row = int(raw_row)
if row < 0 or row in self._tombstones or row >= len(self._id_map):
continue

View file

@ -29,6 +29,7 @@ class LocalFileStore(BaseFileStore):
keyword_index: str = "default",
file_graph: str = "default",
encoding: str = "utf-8",
store_version: str = "v1",
**kwargs,
):
super().__init__(**kwargs)
@ -46,15 +47,17 @@ class LocalFileStore(BaseFileStore):
self.file_graph = self.bind(file_graph, BaseFileGraph, default_factory=LocalFileGraph)
self.encoding = encoding
self.store_version = store_version
self.component_metadata_path.mkdir(parents=True, exist_ok=True)
self.file_chunks: dict[str, FileChunk] = {}
self.chunks_path = self.store_path / f"file_chunks_{self.store_version}.jsonl"
self.chunks_path = self.component_metadata_path / f"file_chunks_{self.name}_{self.store_version}.jsonl"
# Lifecycle
# -- lifecycle ------------------------------------------------------------
async def _start(self) -> None:
await super()._start()
if self.embedding_model is not None and not await self.embedding_model.health_check():
self.logger.warning(f"{self.store_name}: embedding unhealthy, vector disabled")
self.logger.warning(f"{self.name}: embedding unhealthy, vector disabled")
self.embedding_model = None
await self.load()
@ -67,11 +70,13 @@ class LocalFileStore(BaseFileStore):
"""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.logger.error(f"{self.name}: embedding disabled, {reason}")
self.embedding_model = None
# -- persistence ----------------------------------------------------------
async def load(self) -> None:
"""Load chunks from JSONL file into memory."""
"""Load chunks from the JSONL file into memory; missing file is a no-op."""
if not self.chunks_path.exists():
return
try:
@ -86,7 +91,7 @@ class LocalFileStore(BaseFileStore):
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."""
"""Atomically rewrite the JSONL, then cascade dump into keyword_index and file_graph."""
assert self.file_graph is not None
try:
tmp = self.chunks_path.with_suffix(".tmp")
@ -100,7 +105,7 @@ class LocalFileStore(BaseFileStore):
await self.keyword_index.dump()
await self.file_graph.dump()
# CRUD
# -- CRUD -----------------------------------------------------------------
async def upsert(self, files: list[tuple[FileNode, list[FileChunk]]]) -> None:
if not files:
@ -108,40 +113,69 @@ class LocalFileStore(BaseFileStore):
assert self.file_graph is not None
old_map = {n.path: n for n in await self.file_graph.get_nodes([node.path for node, _ in files])}
new_nodes, needs_embed, keyword_docs = self._stage_upsert(files, old_map)
await self.file_graph.upsert_nodes(new_nodes)
await self._embed_pending(needs_embed)
if self.keyword_index and keyword_docs:
await self.keyword_index.add_docs(keyword_docs)
def _stage_upsert(
self,
files: list[tuple[FileNode, list[FileChunk]]],
old_map: dict[str, FileNode],
) -> tuple[list[FileNode], list[FileChunk], dict[str, str]]:
"""Mutate self.file_chunks for each file and collect the work the I/O step needs:
new graph nodes, chunks still needing an embedding, and keyword docs to index.
"""
new_nodes: list[FileNode] = []
needs_embed: list[FileChunk] = []
keyword_docs: dict[str, str] = {}
for node, chunks in files:
old_node: FileNode | None = old_map.get(node.path)
cached: dict = {}
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
cached = self._evict_prior_chunks(old_map.get(node.path))
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)
self._reuse_or_queue_embedding(c, cached, needs_embed)
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)
return new_nodes, needs_embed, keyword_docs
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)
def _evict_prior_chunks(self, old_node: FileNode | None) -> dict[str, np.ndarray]:
"""Drop chunks for the path being re-upserted; keep their embeddings around so
a new chunk reusing the same id avoids a redundant embedding call.
"""
cached: dict[str, np.ndarray] = {}
if not (old_node and self.embedding_model):
return cached
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
return cached
def _reuse_or_queue_embedding(
self,
chunk: FileChunk,
cached: dict[str, np.ndarray],
needs_embed: list[FileChunk],
) -> None:
if not self.embedding_model or chunk.embedding is not None:
return
if chunk.id in cached:
chunk.embedding = cached[chunk.id]
elif chunk.text:
needs_embed.append(chunk)
async def _embed_pending(self, chunks: list[FileChunk]) -> None:
if not (chunks and self.embedding_model):
return
try:
await self.embedding_model.get_node_embeddings(chunks)
except Exception as e:
self._disable_embedding(f"upsert: {type(e).__name__}: {e}")
async def delete(self, path: str | list[str]) -> None:
assert self.file_graph is not None
@ -184,7 +218,7 @@ class LocalFileStore(BaseFileStore):
await self.keyword_index.clear()
await self.file_graph.clear()
# Search
# -- search ---------------------------------------------------------------
async def vector_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]:
if self.embedding_model is None or not query:
@ -229,7 +263,7 @@ class LocalFileStore(BaseFileStore):
return results
# Extensions
# -- extensions -----------------------------------------------------------
async def rebuild_links(self) -> None:
"""Rebuild graph links via the underlying file graph."""

View file

@ -1,7 +1,9 @@
"""Long-running background job with optional supervisor."""
import asyncio
import contextlib
import random
import time
from .base_job import BaseJob
from ..component_registry import R
@ -18,6 +20,12 @@ class BackgroundJob(BaseJob):
exponential backoff (backoff_base * 2**attempt, capped at backoff_cap)
plus ±50% jitter. __call__ must NOT swallow exceptions, otherwise the
supervisor cannot trigger a restart.
On close, the stop_event is set and the task is given up to
``close_timeout`` seconds to exit gracefully; after that it is cancelled.
If a single run survives at least ``attempt_reset_after`` seconds before
crashing, the backoff attempt counter resets — so a long-stable job that
eventually crashes restarts quickly rather than at the capped delay.
"""
def __init__(
@ -25,13 +33,18 @@ class BackgroundJob(BaseJob):
supervisor: bool = True,
backoff_base: float = 1.0,
backoff_cap: float = 60.0,
close_timeout: float = 5.0,
attempt_reset_after: float = 60.0,
enable_serve: bool = False,
**kwargs,
):
super().__init__(**kwargs)
super().__init__(enable_serve=enable_serve, **kwargs)
self.supervisor: bool = supervisor
self.backoff_base: float = backoff_base
self.backoff_cap: float = backoff_cap
self._stop_event: asyncio.Event = asyncio.Event()
self.close_timeout: float = close_timeout
self.attempt_reset_after: float = attempt_reset_after
self._stop_event: asyncio.Event | None = None
self._task: asyncio.Task | None = None
async def _start(self) -> None:
@ -40,35 +53,63 @@ class BackgroundJob(BaseJob):
self._task = asyncio.create_task(self._run_with_supervisor())
async def _close(self) -> None:
self._stop_event.set()
if self._task is not None:
try:
await self._task
except Exception:
self.logger.exception(f"Background task '{self.name}' raised during close")
self._task = None
if self._stop_event is not None:
self._stop_event.set()
await self._shutdown_task()
await super()._close()
async def _shutdown_task(self) -> None:
"""Wait close_timeout for graceful exit, then force-cancel."""
if self._task is None:
return
try:
# shield prevents wait_for's cancellation from propagating to the task itself,
# so a timeout here truly times out instead of cancelling silently.
await asyncio.wait_for(asyncio.shield(self._task), timeout=self.close_timeout)
except asyncio.TimeoutError:
self._task.cancel()
with contextlib.suppress(BaseException):
await self._task
except Exception:
self.logger.exception(f"Background task '{self.name}' raised during close")
self._task = None
def _backoff_delay(self, attempt: int) -> float:
"""Exponential backoff with ±50% jitter, capped at backoff_cap."""
capped = min(self.backoff_base * (2**attempt), self.backoff_cap)
return min(capped * (0.5 + random.random()), self.backoff_cap)
async def _wait_or_stop(self, delay: float) -> None:
"""Sleep up to delay, returning immediately when stop_event is set."""
assert self._stop_event is not None
try:
await asyncio.wait_for(self._stop_event.wait(), timeout=delay)
except asyncio.TimeoutError:
pass
async def _run_with_supervisor(self) -> None:
assert self._stop_event is not None
attempt = 0
while not self._stop_event.is_set():
started_at = time.monotonic()
try:
await self()
return
except Exception as e:
if not self.supervisor:
raise
delay = min(self.backoff_base * (2**attempt), self.backoff_cap) * (0.5 + random.random())
# A long-stable run that just crashed restarts fresh rather than at the capped delay.
if time.monotonic() - started_at >= self.attempt_reset_after:
attempt = 0
delay = self._backoff_delay(attempt)
self.logger.exception(f"job body crashed, restart in {delay:.2f}s error={e}")
attempt += 1
try:
await asyncio.wait_for(self._stop_event.wait(), timeout=delay)
except asyncio.TimeoutError:
pass
await self._wait_or_stop(delay)
async def __call__(self, **kwargs) -> Response:
"""Default body: run step_components in order; errors propagate to supervisor."""
context = RuntimeContext(stop_event=self._stop_event, **self.kwargs)
for step in self.step_components:
"""Default body: run steps in order; errors propagate to supervisor."""
merged = {**self.kwargs, **kwargs}
context = RuntimeContext(stop_event=self._stop_event, **merged)
for step in self._build_steps():
await step(context)
return context.response

View file

@ -1,11 +1,16 @@
"""Base job component for sequential step execution."""
from typing import TYPE_CHECKING
from ..base_component import BaseComponent
from ..component_registry import R
from ..runtime_context import RuntimeContext
from ...enumeration import ComponentEnum
from ...schema import ComponentConfig, Response
if TYPE_CHECKING:
from ...steps import BaseStep
@R.register("base")
class BaseJob(BaseComponent):
@ -18,40 +23,47 @@ class BaseJob(BaseComponent):
description: str = "",
parameters: dict | None = None,
steps: list[ComponentConfig | dict] | None = None,
enable_serve: bool = True,
**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] = []
self.enable_serve = enable_serve
# Resolved at start: (cls, params) pairs. Steps are re-instantiated per call so they stay
# stateless across runs and concurrent invocations don't share mutable step state.
self.step_specs: list[tuple[type["BaseStep"], dict]] = []
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))
if self.app_context is None:
raise RuntimeError(f"app_context must be provided for job '{self.name}'")
self.step_specs = [self._resolve_step(raw) for raw in self.step_configs]
async def _close(self) -> None:
"""Release all step components."""
self.step_components.clear()
self.step_specs.clear()
def _resolve_step(self, raw: ComponentConfig | dict) -> tuple[type["BaseStep"], dict]:
"""Validate a step config and look up its class via the registry."""
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
return step_cls, params
def _build_steps(self) -> list["BaseStep"]:
# dict(params) copies kwargs so steps cannot mutate the shared spec.
return [step_cls(**dict(params)) for step_cls, params in self.step_specs]
async def __call__(self, **kwargs) -> Response:
"""Execute all steps in order and return the final response."""
"""Run all steps in order, capturing any failure into the response."""
context = RuntimeContext(**kwargs)
try:
for step in self.step_components:
for step in self._build_steps():
await step(context)
except Exception as e:
self.logger.exception(f"Failed to execute job: {e}")

View file

@ -11,11 +11,12 @@ 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."""
"""Run steps; emit failures as ERROR chunks, then a terminal DONE marker."""
context = RuntimeContext(**kwargs)
try:
for step in self.step_components:
for step in self._build_steps():
await step(context)
except Exception as e:
await context.add_stream_string(str(e), ChunkEnum.ERROR)
# Always emit DONE so consumers can detach even after an error.
await context.add_stream_done()

View file

@ -1,7 +1,6 @@
"""Abstract base class for keyword index implementations."""
"""Abstract base class for keyword indexes (BM25 and other lexical backends)."""
from abc import abstractmethod
from pathlib import Path
from ..base_component import BaseComponent
from ..tokenizer import BaseTokenizer
@ -9,62 +8,50 @@ from ...enumeration import ComponentEnum
class BaseKeywordIndex(BaseComponent):
"""Abstract base class for keyword index implementations."""
"""Common interface for keyword indexes (add / delete / retrieve / clear)."""
component_type = ComponentEnum.KEYWORD_INDEX
def __init__(self, tokenizer: str = "default", index_version: str = "v1", **kwargs):
def __init__(self, tokenizer: str = "default", **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.vault_metadata_path / self.component_type.value
self.index_path.mkdir(parents=True, exist_ok=True)
self.component_metadata_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."""
"""Tokenize a single text into a list of 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."""
"""Add or replace documents keyed by id."""
@abstractmethod
async def delete_docs(self, doc_ids: list[str]) -> None:
"""Remove documents by their IDs."""
"""Delete documents by id; missing ids are skipped."""
@abstractmethod
async def retrieve(self, query: str, limit: int = 3) -> dict[str, float]:
"""Search documents. Returns {doc_id: score} sorted descending."""
"""Return top-`limit` doc_id → score for the given query."""
@abstractmethod
async def clear(self) -> None:
"""Reset index to empty state."""
"""Wipe in-memory state and remove any persisted artifacts."""
async def reset_index(self, docs_dict: dict[str, str]) -> None:
"""Clear index, re-add all documents, and persist."""
"""Wipe the index, rebuild it from `docs_dict`, and persist the result."""
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."""
"""Compact or rebuild the index. No-op by default; override as needed."""

View file

@ -1,166 +1,326 @@
"""BM25 search engine with persistent index support.
"""BM25 inverted index with on-disk persistence.
Implements Okapi BM25 ranking with an inverted index for efficient
document lookup, incremental updates, and pickle-based persistence.
On-disk truth source (see `_snapshot` / `_restore`):
vocab : dict[token, token_id]
_doc_ids : list[doc_id], indexed by doc_idx
_doc_id_to_idx : dict[doc_id, doc_idx]
_doc_lens : np.int32[n], indexed by doc_idx
_deleted : np.bool[n], lazy-delete flag per doc_idx
_doc_token_ids : list[np.int32[]], unique token_ids per doc
_posting_doc_idxs : dict[token_id, np.int32[]], posting list (doc_idx)
_posting_tfs : dict[token_id, np.int32[]], aligned term frequencies
Deletion is lazy: setting `_deleted[idx] = True` retires the slot. The posting
lists keep the stale entries until `optimize_index` rewrites them. Updating an
existing doc_id retires the old slot first, then allocates a fresh idx.
"""
import math
import pickle
from collections import Counter
from typing import TypedDict
from pathlib import Path
import numpy as np
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.
"""BM25 inverted index with lazy deletion and on-disk 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):
def __init__(self, k1: float = 1.5, b: float = 0.75, index_version: str = "v1", **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.index_version = index_version
self.vocab: dict[str, int] = {}
self._doc_ids: list[str] = []
self._doc_id_to_idx: dict[str, int] = {}
self._doc_lens: np.ndarray = np.zeros(0, dtype=np.int32)
self._deleted: np.ndarray = np.zeros(0, dtype=bool)
self._doc_token_ids: list[np.ndarray] = []
self._posting_doc_idxs: dict[int, np.ndarray] = {}
self._posting_tfs: dict[int, np.ndarray] = {}
# IDF cache; invalidated whenever live-doc count or postings change.
self._idf_cache: dict[int, float] = {}
# -- Properties -----------------------------------------------------------
@property
def index_file(self) -> Path:
"""Path of the persisted index, namespaced by tokenizer and version."""
if self.tokenizer is None:
raise RuntimeError("Tokenizer not initialized. Call start() first.")
name = type(self.tokenizer).__name__.replace("Tokenizer", "").lower()
return self.component_metadata_path / f"bm25_{name}_{self.index_version}.pkl"
@property
def n_docs(self) -> int:
"""Number of indexed documents."""
return len(self.doc_meta)
"""Number of live (non-deleted) documents."""
return 0 if self._deleted.size == 0 else int((~self._deleted).sum())
@property
def total_len(self) -> int:
"""Sum of token counts across live documents."""
return 0 if self._deleted.size == 0 else int(self._doc_lens[~self._deleted].sum())
@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
"""Average length of live documents, used for BM25 length normalization."""
n = self.n_docs
return self.total_len / n if n > 0 else 0.0
@property
def doc_meta(self) -> dict[str, dict]:
"""Per-live-doc length and unique token_id set, keyed by doc_id."""
return {
self._doc_ids[idx]: {
"len": int(self._doc_lens[idx]),
"token_ids": {int(t) for t in self._doc_token_ids[idx]},
}
for idx in range(len(self._doc_ids))
if not self._deleted[idx]
}
@property
def inverted_index(self) -> dict[int, dict[str, int]]:
"""Readable view of postings: token_id -> {doc_id: tf}, deleted skipped."""
out: dict[int, dict[str, int]] = {}
for tid, doc_idxs in self._posting_doc_idxs.items():
tfs = self._posting_tfs[tid]
posting = {self._doc_ids[int(i)]: int(tf) for i, tf in zip(doc_idxs, tfs) if not self._deleted[int(i)]}
if posting:
out[tid] = posting
return out
# -- Internal helpers -----------------------------------------------------
def _tokens_to_ids(self, tokens: list[str]) -> list[int]:
"""Map tokens to integer IDs, assigning new IDs on first encounter."""
ids = []
"""Map tokens to ids, allocating a fresh id for any unseen token."""
vocab = self.vocab
ids: list[int] = []
for token in tokens:
token = token.strip()
if token:
ids.append(self.vocab.setdefault(token, len(self.vocab)))
if not token:
continue
tid = vocab.get(token)
if tid is None:
tid = len(vocab)
vocab[token] = tid
ids.append(tid)
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:
"""Lazy-delete a doc: flip `_deleted` and drop the id mapping."""
idx = self._doc_id_to_idx.get(doc_id)
if idx is None or self._deleted[idx]:
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]
self._deleted[idx] = True
self._doc_id_to_idx.pop(doc_id, None)
self._idf_cache = {}
def _get_idf(self, token_id: int) -> float:
"""Compute and cache IDF for a token ID."""
def _get_idf(self, token_id: int, n_docs: int | None = None) -> float:
"""Return the cached IDF for a token, computing it on miss."""
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]
doc_idxs = self._posting_doc_idxs.get(token_id)
if doc_idxs is None or doc_idxs.size == 0:
self._idf_cache[token_id] = 0.0
return 0.0
df = int((~self._deleted[doc_idxs]).sum())
if n_docs is None:
n_docs = self.n_docs
idf = math.log(1 + (n_docs - df + 0.5) / (df + 0.5)) if df else 0.0
self._idf_cache[token_id] = idf
return idf
def _prepare_doc(self, doc_id: str, content: str) -> tuple[np.ndarray, int, Counter] | None:
"""Tokenize and count terms; retire any prior version of `doc_id`."""
self._remove_doc(doc_id)
token_ids = self._tokens_to_ids(self._tokenize(content))
if not token_ids:
return None
counts = Counter(token_ids)
unique_tids = np.fromiter(counts.keys(), dtype=np.int32, count=len(counts))
return unique_tids, len(token_ids), counts
def _append_doc_arrays(
self,
new_doc_ids: list[str],
new_doc_lens: list[int],
new_doc_token_ids: list[np.ndarray],
) -> None:
"""Append metadata for a batch of new docs to the doc-level arrays."""
if not new_doc_ids:
return
self._doc_ids.extend(new_doc_ids)
self._doc_token_ids.extend(new_doc_token_ids)
self._doc_lens = np.concatenate([self._doc_lens, np.array(new_doc_lens, dtype=np.int32)])
self._deleted = np.concatenate([self._deleted, np.zeros(len(new_doc_ids), dtype=bool)])
def _extend_postings(self, pending: dict[int, list[tuple[int, int]]]) -> None:
"""Append pending (doc_idx, tf) pairs to each token's posting list."""
for tid, items in pending.items():
n = len(items)
new_idxs = np.fromiter((idx for idx, _ in items), dtype=np.int32, count=n)
new_tfs = np.fromiter((tf for _, tf in items), dtype=np.int32, count=n)
if tid in self._posting_doc_idxs:
self._posting_doc_idxs[tid] = np.concatenate([self._posting_doc_idxs[tid], new_idxs])
self._posting_tfs[tid] = np.concatenate([self._posting_tfs[tid], new_tfs])
else:
self._posting_doc_idxs[tid] = new_idxs
self._posting_tfs[tid] = new_tfs
def _encode_query(self, query: str) -> list[int]:
"""Tokenize query; drop OOV terms; deduplicate while preserving order."""
vocab = self.vocab
return list(dict.fromkeys(vocab[t] for t in self._tokenize(query) if t in vocab))
def _top_k(self, scores: np.ndarray, limit: int) -> np.ndarray:
"""Indices of the top `limit` strictly-positive scores, descending."""
if limit <= 0:
return np.empty(0, dtype=np.int64)
positive_count = int((scores > 0).sum())
if positive_count == 0:
return np.empty(0, dtype=np.int64)
k = min(limit, positive_count)
if k >= scores.size:
return np.argsort(-scores)[:k]
top = np.argpartition(-scores, k - 1)[:k]
return top[np.argsort(-scores[top])]
# -- Public API -----------------------------------------------------------
async def add_docs(self, docs_dict: dict[str, str]) -> None:
"""Index or update multiple documents. Mapping of doc_id to content."""
"""Add or replace documents in batch (existing doc_ids are overwritten)."""
if not docs_dict:
return
new_doc_ids: list[str] = []
new_doc_lens: list[int] = []
new_doc_token_ids: list[np.ndarray] = []
pending: dict[int, list[tuple[int, int]]] = {}
next_idx = len(self._doc_ids)
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:
prepared = self._prepare_doc(doc_id, content)
if prepared is None:
continue
token_ids = self._tokens_to_ids(tokens)
token_counts = Counter(token_ids)
unique_tids, n_tokens, token_counts = prepared
idx = next_idx
next_idx += 1
new_doc_ids.append(doc_id)
new_doc_lens.append(n_tokens)
new_doc_token_ids.append(unique_tids)
self._doc_id_to_idx[doc_id] = idx
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)
pending.setdefault(tid, []).append((idx, tf))
self._append_doc_arrays(new_doc_ids, new_doc_lens, new_doc_token_ids)
self._extend_postings(pending)
self._idf_cache = {}
async def delete_docs(self, doc_ids: list[str]) -> None:
"""Remove documents by their IDs."""
"""Lazy-delete a batch of doc_ids; physical reclaim happens in optimize_index."""
for doc_id in doc_ids:
self._remove_doc(doc_id)
self._idf_cache = {}
def _score_query(self, query_ids: list[int], n_docs: int) -> np.ndarray:
"""Compute BM25 scores across all docs; deleted docs zeroed out."""
avg_len = self.total_len / n_docs
k1, b = self.k1, self.b
denom_base = k1 * (1.0 - b)
denom_norm = k1 * b / avg_len if avg_len > 0 else 0.0
scores = np.zeros(self._doc_lens.size, dtype=np.float32)
for tid in query_ids:
doc_idxs = self._posting_doc_idxs.get(tid)
if doc_idxs is None or doc_idxs.size == 0:
continue
idf = self._get_idf(tid, n_docs)
if idf == 0.0:
continue
tfs = self._posting_tfs[tid].astype(np.float32)
d_lens = self._doc_lens[doc_idxs].astype(np.float32)
# Each doc_idx appears at most once per posting list (Counter dedups
# within a doc, and updates allocate fresh idxs), so fancy-index
# accumulation is safe here.
scores[doc_idxs] += idf * tfs * (k1 + 1.0) / (tfs + denom_base + denom_norm * d_lens)
if self._deleted.any():
scores[self._deleted] = 0.0
return scores
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:
"""BM25 retrieval; returns {doc_id: score} sorted by score descending."""
n_docs = self.n_docs
if n_docs == 0:
return {}
query_ids = self._encode_query(query)
if not query_ids:
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
scores = self._score_query(query_ids, n_docs)
top_idxs = self._top_k(scores, limit)
return {self._doc_ids[int(i)]: float(scores[int(i)]) for i in top_idxs}
return dict(sorted(scores.items(), key=lambda x: x[1], reverse=True)[:limit]) if scores else {}
# -- Persistence ----------------------------------------------------------
def _snapshot(self) -> dict:
"""Bundle every persistent field; mirrors `_restore`."""
return {
"vocab": self.vocab,
"doc_ids": self._doc_ids,
"doc_id_to_idx": self._doc_id_to_idx,
"doc_lens": self._doc_lens,
"deleted": self._deleted,
"doc_token_ids": self._doc_token_ids,
"posting_doc_idxs": self._posting_doc_idxs,
"posting_tfs": self._posting_tfs,
"k1": self.k1,
"b": self.b,
}
def _restore(self, data: dict) -> None:
"""Restore index state from a `_snapshot` dict."""
self.vocab = data["vocab"]
self._doc_ids = data["doc_ids"]
self._doc_id_to_idx = data["doc_id_to_idx"]
self._doc_lens = data["doc_lens"]
self._deleted = data["deleted"]
self._doc_token_ids = data["doc_token_ids"]
self._posting_doc_idxs = data["posting_doc_idxs"]
self._posting_tfs = data["posting_tfs"]
self.k1 = data.get("k1", 1.5)
self.b = data.get("b", 0.75)
self._idf_cache = {}
async def dump(self) -> None:
"""Persist index to disk via pickle (atomic rename)."""
"""Persist the index via temp file + atomic rename to avoid torn writes."""
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,
)
pickle.dump(self._snapshot(), 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."""
"""Load from disk; missing file is a no-op, corrupt file resets state."""
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._restore(data)
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}")
@ -168,39 +328,99 @@ class BM25Index(BaseKeywordIndex):
await self.clear()
async def clear(self) -> None:
"""Reset index to empty state and remove persisted file."""
"""Reset in-memory state and remove the persisted file."""
self.vocab = {}
self.inverted_index = {}
self.doc_meta = {}
self.total_len = 0
self._doc_ids = []
self._doc_id_to_idx = {}
self._doc_lens = np.zeros(0, dtype=np.int32)
self._deleted = np.zeros(0, dtype=bool)
self._doc_token_ids = []
self._posting_doc_idxs = {}
self._posting_tfs = {}
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
# -- Compaction -----------------------------------------------------------
# Build compact ID mapping
old_to_new: dict[int, int] = {}
def _build_idx_remap(self, active_mask: np.ndarray) -> tuple[np.ndarray, int]:
"""Build an old_idx -> new_idx array (-1 for retired slots)."""
active_old_idxs = np.where(active_mask)[0]
n_active = int(active_old_idxs.size)
remap = -np.ones(self._deleted.size, dtype=np.int32)
remap[active_old_idxs] = np.arange(n_active, dtype=np.int32)
return remap, n_active
def _compact_vocab(self, active_mask: np.ndarray) -> tuple[dict[str, int], dict[int, int]]:
"""Keep only tokens still referenced by a live doc; renumber contiguously."""
used_tids = {tid for tid, doc_idxs in self._posting_doc_idxs.items() if active_mask[doc_idxs].any()}
new_vocab: dict[str, int] = {}
old_to_new: dict[int, int] = {}
for token, old_tid in self.vocab.items():
if old_tid in used_token_ids:
if old_tid in used_tids:
new_tid = len(new_vocab)
new_vocab[token] = new_tid
old_to_new[old_tid] = new_tid
return new_vocab, old_to_new
# 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}
def _compact_postings(
self,
active_mask: np.ndarray,
old_to_new_idx: np.ndarray,
old_tid_to_new: dict[int, int],
) -> tuple[dict[int, np.ndarray], dict[int, np.ndarray]]:
"""Drop deleted entries and rewrite postings under new idx/tid numbering."""
new_idxs: dict[int, np.ndarray] = {}
new_tfs: dict[int, np.ndarray] = {}
for tid, doc_idxs in self._posting_doc_idxs.items():
if tid not in old_tid_to_new:
continue
mask = active_mask[doc_idxs]
new_tid = old_tid_to_new[tid]
new_idxs[new_tid] = old_to_new_idx[doc_idxs[mask]].astype(np.int32, copy=False)
new_tfs[new_tid] = self._posting_tfs[tid][mask].astype(np.int32, copy=False)
return new_idxs, new_tfs
def _compact_docs(
self,
active_mask: np.ndarray,
old_tid_to_new: dict[int, int],
) -> tuple[list[str], list[np.ndarray]]:
"""Rebuild doc_id list and unique-token arrays under the new vocab."""
active_old_idxs = np.where(active_mask)[0]
new_doc_ids = [self._doc_ids[int(i)] for i in active_old_idxs]
new_doc_token_ids = [
np.fromiter(
(old_tid_to_new[int(t)] for t in self._doc_token_ids[int(i)] if int(t) in old_tid_to_new),
dtype=np.int32,
)
for i in active_old_idxs
]
return new_doc_ids, new_doc_token_ids
async def optimize_index(self) -> None:
"""Physically reclaim deleted docs and unused vocab entries."""
if self._deleted.size == 0:
return
active_mask = ~self._deleted
if not active_mask.any():
await self.clear()
return
old_to_new_idx, n_active = self._build_idx_remap(active_mask)
new_vocab, old_tid_to_new = self._compact_vocab(active_mask)
new_posting_idxs, new_posting_tfs = self._compact_postings(
active_mask,
old_to_new_idx,
old_tid_to_new,
)
new_doc_ids, new_doc_token_ids = self._compact_docs(active_mask, old_tid_to_new)
self.vocab = new_vocab
self.inverted_index = new_inverted_index
self._doc_ids = new_doc_ids
self._doc_id_to_idx = {doc_id: i for i, doc_id in enumerate(new_doc_ids)}
self._doc_lens = self._doc_lens[active_mask].astype(np.int32, copy=True)
self._deleted = np.zeros(n_active, dtype=bool)
self._doc_token_ids = new_doc_token_ids
self._posting_doc_idxs = new_posting_idxs
self._posting_tfs = new_posting_tfs
self._idf_cache = {}

View file

@ -8,26 +8,27 @@ from string import Formatter
import yaml
# Matches a leading flag tag like "[verbose] some text".
# Matches a leading flag tag at line start: "[flag] rest of line".
_FLAG_PATTERN = re.compile(r"^\[(\w+)]")
class PromptHandler:
"""Loads prompts from YAML/JSON or class-adjacent files and formats them.
"""Loads prompts from YAML/JSON files and renders them with optional flags.
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.
Template keys may carry a language suffix (``key_en``, ``key_zh``); lookups
fall back to the bare key when no localized variant exists. Lines tagged
with ``[flag]`` are kept only when the matching boolean 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.
# Non-string kwargs are silently dropped — prompts must be strings.
self.data: dict[str, str] = {k: v for k, v in kwargs.items() if isinstance(v, str)}
self.language: str = language.strip()
# ----- Loading -------------------------------------------------------
def load_prompt_by_file(
self,
prompt_file_path: str | Path | None = None,
@ -41,13 +42,18 @@ class PromptHandler:
if not path.exists() or path.suffix.lower() not in self._SUPPORTED_EXTENSIONS:
return self
return self.load_prompt_dict(self._parse_prompt_file(path), overwrite)
@staticmethod
def _parse_prompt_file(path: Path) -> dict | None:
"""Parse a YAML or JSON prompt file; return None on any parse error."""
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)
if path.suffix.lower() in (".yaml", ".yml"):
return yaml.safe_load(f)
return json.load(f)
except (json.JSONDecodeError, yaml.YAMLError, OSError):
return self
return self.load_prompt_dict(prompt_dict, overwrite)
return None
def load_prompt_by_class(self, cls: type, overwrite: bool = True) -> "PromptHandler":
"""Load prompts from ``<class_module>.yaml`` (or ``.yml``) next to `cls`."""
@ -57,9 +63,9 @@ class PromptHandler:
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)
candidate = base_path.with_suffix(ext)
if candidate.exists():
return self.load_prompt_by_file(candidate, overwrite)
return self
def load_prompt_dict(self, prompt_dict: dict | None = None, overwrite: bool = True) -> "PromptHandler":
@ -70,21 +76,28 @@ class PromptHandler:
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
# ----- Lookup --------------------------------------------------------
def _candidate_keys(self, prompt_name: str) -> tuple[str, ...]:
"""Lookup order: localized key first when a language is set, then bare key."""
if self.language:
return (f"{prompt_name}_{self.language}", prompt_name)
return (prompt_name,)
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,):
for key in self._candidate_keys(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]}")
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)
return any(k in self.data for k in self._candidate_keys(prompt_name))
def list_prompts(self, language_filter: str | None = None) -> list[str]:
"""List all keys, optionally filtered to those ending with ``_<language>``."""
@ -93,33 +106,44 @@ class PromptHandler:
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.
# ----- Formatting ----------------------------------------------------
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``.
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 flag toggles; the rest are format variables.
With ``validate=True``, any missing ``{var}`` placeholder raises ``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)
prompt = self._apply_flag_filter(prompt, flags)
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)}")
self._check_required_vars(prompt, formats, prompt_name)
return prompt.format(**formats).strip() if formats else prompt
@staticmethod
def _apply_flag_filter(prompt: str, flags: dict[str, bool]) -> str:
"""Keep unflagged lines; keep flagged lines only when a matching flag is set."""
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)
return "\n".join(lines)
@staticmethod
def _check_required_vars(prompt: str, formats: dict, prompt_name: str) -> None:
"""Raise when any ``{var}`` placeholder lacks a corresponding kwarg."""
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)}",
)
def __repr__(self) -> str:
return f"PromptHandler(language='{self.language}', num_prompts={len(self.data)})"

View file

@ -9,8 +9,8 @@ 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.
Holds the response object, an optional stream queue, a stop event, and a
free-form data dict accessible via mapping-style operators (``ctx[key]``).
"""
def __init__(
@ -25,12 +25,14 @@ class RuntimeContext:
self.stop_event: asyncio.Event | None = stop_event
self.data: dict = kwargs
# ----- Data dict access ----------------------------------------------
def get(self, key: str, default=None):
"""Get a value from the data dict."""
"""Get a value from the data dict with an optional default."""
return self.data.get(key, default)
def update(self, data: dict) -> "RuntimeContext":
"""Merge data into the context."""
"""Merge `data` into the context and return self for chaining."""
self.data.update(data)
return self
@ -46,41 +48,40 @@ class RuntimeContext:
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.
"""Reuse `context` (merging kwargs into its data) or create a fresh one."""
if context is None:
return cls(**kwargs)
context.update(kwargs)
return context
return context.update(kwargs)
# ----- Streaming -----------------------------------------------------
@property
def stream(self) -> bool:
"""Whether a stream queue is attached (i.e., streaming is enabled)."""
return self.stream_queue is not None
async def _enqueue(self, chunk: StreamChunk) -> None:
"""Put a chunk on the stream queue."""
"""Put a chunk on the stream queue; raises when streaming is disabled."""
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
# ----- Misc ----------------------------------------------------------
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.
"""Copy ``data[source]`` to ``data[target]`` for each ``source: target`` pair."""
if not mapping:
return self
for source, target in mapping.items():

View file

@ -1,10 +1,14 @@
"""Base service class for exposing jobs via HTTP, MCP, etc."""
"""Base class for services that expose jobs over a network protocol."""
import json
import os
from abc import abstractmethod
from contextlib import asynccontextmanager
from typing import TYPE_CHECKING
from ..base_component import BaseComponent
from ..job.base_job import BaseJob
from ...constants import REME_SERVICE_INFO
from ...enumeration import ComponentEnum
if TYPE_CHECKING:
@ -12,30 +16,53 @@ if TYPE_CHECKING:
class BaseService(BaseComponent):
"""Base class for services that expose jobs via HTTP, MCP, etc."""
"""Skeleton for services (HTTP, MCP, ...) that turn jobs into endpoints or tools."""
component_type = ComponentEnum.SERVICE
def __init__(self, **kwargs):
super().__init__(**kwargs)
# Underlying framework instance (FastAPI, FastMCP, ...); populated by build_service().
self.service = None
# ----- Subclass contract ---------------------------------------------
@abstractmethod
def build_service(self, app: "Application") -> None:
"""Initialize the underlying service framework."""
"""Instantiate and configure the underlying server framework."""
@abstractmethod
def add_job(self, job: BaseJob) -> None:
"""Register a single job with the service."""
"""Register a single job as a callable endpoint or tool."""
@abstractmethod
def start_service(self, app: "Application") -> None:
"""Start serving requests."""
"""Block on serving requests until shutdown."""
# ----- Shared helpers ------------------------------------------------
def _lifespan(self, app: "Application", host: str, port: int):
"""Build an async-context lifespan that brackets the server with app start/close.
Publishes the bound address via the REME_SERVICE_INFO environment variable so
in-process clients can discover where this service is listening.
"""
@asynccontextmanager
async def lifespan(_):
await app.start()
service_info = json.dumps({"host": host, "port": port})
os.environ[REME_SERVICE_INFO] = service_info
self.logger.info(f"{self.name} started: {REME_SERVICE_INFO}={service_info}")
yield
await app.close()
return lifespan
def add_jobs(self, app: "Application") -> None:
"""Register all non-background jobs from the application context."""
"""Register every job whose enable_serve flag is True."""
for name, job in app.context.jobs.items():
if job.backend == "background":
if not job.enable_serve:
continue
try:
self.add_job(job)
@ -44,7 +71,7 @@ class BaseService(BaseComponent):
self.logger.error(f"Failed to add job {name}: {e}")
def run_app(self, app: "Application") -> None:
"""Build, populate, and start the service."""
"""Build the service, register jobs, then start serving (blocking)."""
self.build_service(app)
self.add_jobs(app)
self.start_service(app)

View file

@ -1,11 +1,8 @@
"""HTTP service implementation for ReMe."""
"""HTTP service: exposes jobs as FastAPI endpoints (JSON, or SSE for stream jobs)."""
import asyncio
import json
import os
import warnings
from collections.abc import AsyncGenerator
from contextlib import asynccontextmanager
from typing import TYPE_CHECKING
import uvicorn
@ -16,7 +13,7 @@ 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 ...constants import REME_DEFAULT_HOST, REME_DEFAULT_PORT
from ...schema import Request, Response
from ...utils import execute_stream_task
@ -24,27 +21,76 @@ if TYPE_CHECKING:
from ...application import Application
# uvicorn 0.41 still imports these deprecated websockets symbols on startup,
# even though we don't use WebSocket. Silence just those specific warnings.
_WEBSOCKET_DEPRECATION_PATTERNS = (
r".*websockets\.legacy is deprecated.*",
r".*WebSocketServerProtocol is deprecated.*",
)
@R.register("http")
class HttpService(BaseService):
"""HTTP service: normal jobs -> JSON endpoints, stream jobs -> SSE endpoints."""
"""Map non-stream jobs to JSON POST endpoints and StreamJobs to 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:
# ----- BaseService contract ------------------------------------------
def build_service(self, app: "Application") -> None:
"""Create the FastAPI app with permissive CORS and an app-managed lifespan."""
self.service = FastAPI(
title=app.config.app_name,
lifespan=self._lifespan(app, self.host, self.port),
)
self.service.add_middleware(
CORSMiddleware, # type: ignore[arg-type]
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
def add_job(self, job: BaseJob) -> None:
"""Dispatch to streaming or non-streaming registration based on job type."""
if isinstance(job, StreamJob):
self._add_stream_job(job)
else:
self._add_json_job(job)
def start_service(self, app: "Application") -> None:
"""Run uvicorn, suppressing unrelated websocket deprecation noise."""
for pattern in _WEBSOCKET_DEPRECATION_PATTERNS:
warnings.filterwarnings("ignore", category=DeprecationWarning, message=pattern)
uvicorn.run(self.service, host=self.host, port=self.port, **self.kwargs)
# ----- Endpoint factories --------------------------------------------
def _add_json_job(self, job: BaseJob) -> None:
"""Register a job as POST /{job.name} returning a JSON Response."""
async def 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)
self.service.post(
f"/{job.name}",
response_model=Response,
description=job.description,
)(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)))
"""Register a StreamJob as POST /{job.name} streaming chunks as text/event-stream."""
async def generate_stream() -> AsyncGenerator[bytes, None]:
async def endpoint(request: Request) -> StreamingResponse:
stream_queue: asyncio.Queue = asyncio.Queue()
task = asyncio.create_task(
job(stream_queue=stream_queue, **request.model_dump(exclude_none=True)),
)
async def body() -> AsyncGenerator[bytes, None]:
async for chunk in execute_stream_task(
stream_queue=stream_queue,
task=task,
@ -54,46 +100,6 @@ class HttpService(BaseService):
assert isinstance(chunk, bytes)
yield chunk
return StreamingResponse(generate_stream(), media_type="text/event-stream")
return StreamingResponse(body(), 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)
self.service.post(f"/{job.name}")(endpoint)

View file

@ -1,8 +1,5 @@
"""MCP (Model Context Protocol) service implementation."""
"""MCP (Model Context Protocol) service: exposes jobs as MCP tools."""
import json
import os
from contextlib import asynccontextmanager
from typing import TYPE_CHECKING
from fastmcp import FastMCP
@ -11,8 +8,8 @@ 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
from ..job import BaseJob, StreamJob
from ...constants import REME_DEFAULT_HOST, REME_DEFAULT_PORT
if TYPE_CHECKING:
from ...application import Application
@ -20,7 +17,7 @@ if TYPE_CHECKING:
@R.register("mcp")
class MCPService(BaseService):
"""Expose jobs as MCP (Model Context Protocol) tools."""
"""Expose non-stream jobs as MCP tools over stdio, SSE, or other supported transports."""
def __init__(
self,
@ -34,19 +31,17 @@ class MCPService(BaseService):
self.host: str = host
self.port: int = port
# ----- BaseService contract ------------------------------------------
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()
"""Create the FastMCP server with an app-managed lifespan."""
self.service = FastMCP(
name=app.config.app_name,
lifespan=self._lifespan(app, self.host, self.port),
)
self.service = FastMCP(name=app.config.app_name, lifespan=lifespan)
def add_job(self, job: "BaseJob") -> None:
def add_job(self, job: BaseJob) -> None:
"""Register a non-stream job as an MCP tool; StreamJobs are skipped (not supported)."""
if isinstance(job, StreamJob):
return
@ -64,7 +59,8 @@ class MCPService(BaseService):
)
def start_service(self, app: "Application") -> None:
transport_kwargs = {}
"""Run the MCP server; bind host/port only when the transport is network-based."""
transport_kwargs: dict = {}
if self.transport != "stdio":
transport_kwargs["host"] = self.host
transport_kwargs["port"] = self.port

View file

@ -10,18 +10,28 @@ from ...enumeration import ComponentEnum
class BaseTokenizer(BaseComponent):
"""Base tokenizer. Subclasses must implement `tokenize`. Loads stopwords on start."""
"""Tokenizer base class with shared stopword loading and post-processing.
Subclasses implement raw tokenization via `_tokenize_one`; lowercasing and
stopword filtering are handled here so every backend behaves consistently.
"""
component_type = ComponentEnum.TOKENIZER
DEFAULT_STOPWORDS_PATH = Path(__file__).parent / "stopwords"
def __init__(self, stopwords_path: str | Path | None = None, **kwargs):
def __init__(
self,
stopwords_path: str | Path | None = None,
filter_stopwords: bool = True,
**kwargs,
):
super().__init__(**kwargs)
self.stopwords_path = Path(stopwords_path) if stopwords_path else self.DEFAULT_STOPWORDS_PATH
self.filter_stopwords = filter_stopwords
self._stopwords: set[str] = set()
async def _start(self) -> None:
"""Load stopwords from file."""
# A missing file is non-fatal: tokenizers still work, just without filtering.
if not self.stopwords_path.exists():
self.logger.warning(f"Stopwords file not found: {self.stopwords_path}")
return
@ -31,14 +41,24 @@ class BaseTokenizer(BaseComponent):
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."""
"""Loaded stopwords (empty set if none were loaded)."""
return self._stopwords
def tokenize(self, texts: list[str], lower: bool = True, **kwargs) -> list[list[str]]:
"""Tokenize each text and apply shared post-processing."""
return [self._postprocess(self._tokenize_one(t, **kwargs), lower) for t in texts]
def _postprocess(self, tokens: list[str], lower: bool) -> list[str]:
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]
return tokens
@abstractmethod
def tokenize(self, texts: list[str], **kwargs) -> list[list[str]]:
"""Tokenize a list of texts."""
def _tokenize_one(self, text: str, **kwargs) -> list[str]:
"""Return raw tokens for one text; lowercasing/filtering happen upstream."""

View file

@ -1,27 +1,43 @@
"""Jieba tokenizer for Chinese text segmentation."""
from typing import Callable
from .base_tokenizer import BaseTokenizer
from ..component_registry import R
@R.register("jieba")
class JiebaTokenizer(BaseTokenizer):
"""Tokenizer using jieba for Chinese text segmentation."""
"""Tokenizer backed by jieba for Chinese word segmentation.
def __init__(self, filter_stopwords: bool = True, **kwargs):
`backend` selects the underlying implementation:
- "rjieba": Rust binding of jieba-rs, ~10-30x faster than pure Python (default).
- "jieba": Original pure-Python jieba, slowest but the reference.
"""
SUPPORTED_BACKENDS = ("rjieba", "jieba")
def __init__(self, backend: str = "rjieba", **kwargs):
super().__init__(**kwargs)
self.filter_stopwords = filter_stopwords
if backend not in self.SUPPORTED_BACKENDS:
raise ValueError(
f"Unknown jieba backend {backend!r}; expected one of {self.SUPPORTED_BACKENDS}",
)
self.backend = backend
self._cut: Callable[[str], list[str]] | None = None
def tokenize(self, texts: list[str], lower: bool = True, **kwargs) -> list[list[str]]:
"""Tokenize texts using jieba."""
import jieba
async def _start(self) -> None:
await super()._start()
# Resolve the backend once at startup so per-call overhead is just one attribute lookup.
if self.backend == "rjieba":
import rjieba
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
self._cut = rjieba.cut
else:
import jieba
self._cut = jieba.cut
self.logger.info(f"JiebaTokenizer using backend: {self.backend}")
def _tokenize_one(self, text: str, **kwargs) -> list[str]:
return list(self._cut(text))

View file

@ -1,31 +1,24 @@
"""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."""
"""Regex tokenizer: each CJK char is its own token, non-CJK uses word boundaries.
WORD_PATTERN = re.compile(r"(?u)\b\w\w+\b") # 2+ char words
CHINESE_PATTERN = re.compile(r"[一-鿿]") # single Chinese char
Treating CJK characters as individual tokens avoids needing a Chinese
segmenter while still giving BM25-style indexes useful unigrams.
"""
def __init__(self, filter_stopwords: bool = True, **kwargs):
super().__init__(**kwargs)
self.filter_stopwords = filter_stopwords
WORD_PATTERN = re.compile(r"(?u)\b\w\w+\b") # non-CJK words, 2+ chars
CHINESE_PATTERN = re.compile(r"[一-鿿]")
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
def _tokenize_one(self, text: str, **kwargs) -> list[str]:
# Pull CJK chars first, then strip them out so the word regex only sees the rest.
tokens = self.CHINESE_PATTERN.findall(text)
tokens.extend(self.WORD_PATTERN.findall(self.CHINESE_PATTERN.sub(" ", text)))
return tokens

View file

@ -1,89 +1,88 @@
service:
backend: http
# backend: mcp
# Default dev config points vault_dir at ./.reme so `python -m reme4 start`
# can be run from the repo root and exercise the full atomic-tool surface
# against the seeded test data. Override the `vault_dir=` CLI arg.
vault_dir: .reme
daily_dir: daily
digest_dir: digest
resource_dir: resource
resource_dir: ""
jobs:
# ════════════════════════════════════════════════════════════════════
# UTILITY — service introspection
# ════════════════════════════════════════════════════════════════════
- backend: base
name: version
description: "return reme4 package version"
update_store_index_loop:
backend: background
watch_paths: [ "daily", "digest" ]
suffix_filters: [ "md" ]
steps:
- backend: scan_changes_step
- backend: update_index_step
persist: true
- backend: watch_changes_step
dispatch_step: update_index_step
version:
backend: base
description: "return reme package version"
parameters:
type: object
properties: {}
properties: { }
steps:
- backend: version_step
- backend: base
name: health_check
description: "return a concise health-check snapshot of reme4 components"
health_check:
backend: base
description: "return a concise health-check snapshot of reme components"
parameters:
type: object
properties: {}
properties: { }
steps:
- backend: health_check_step
- backend: base
name: help
help:
backend: base
description: "list all registered jobs with their metadata"
parameters:
type: object
properties: {}
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: index_changes
description: "apply a batch of file changes (added/modified/deleted) into file_store"
traverse:
backend: base
description: "Walk the wikilink graph from a path."
parameters:
type: object
properties:
changes:
type: array
description: "list of change items"
items:
type: object
properties:
change:
type: string
enum: [added, modified, deleted]
description: "type of file change"
path:
type: string
description: "absolute file path"
required:
- change
- path
path:
type: string
description: "path"
depth:
type: integer
description: "hop limit"
default: 1
direction:
type: string
enum:
- forward
- backward
- both
default: both
required:
- changes
- path
steps:
- backend: index_changes_step
- backend: traverse_step
# ════════════════════════════════════════════════════════════════════
# ATOMIC TOOLS — same surface plugins/reme-{service,expert} expose
# ════════════════════════════════════════════════════════════════════
reindex:
backend: base
description: "wipe the file store and rebuild it from the existing files"
parameters:
type: object
properties: { }
steps:
- backend: clear_and_scan_step
- backend: update_index_step
persist: true
# ── Retrieve ───────────────────────────────────────────────────────
- backend: base
name: search
search:
backend: base
description: "Hybrid vault search (vector + BM25, RRF-fused)."
parameters:
type: object
@ -108,31 +107,118 @@ jobs:
expand_links: true
max_links_per_direction: 10
- backend: base
name: traverse
description: "Walk the wikilink graph from a seed path."
daily:create:
backend: base
description: "Provision a note slug under a daily folder: daily/<date>/<slug>.md"
parameters:
type: object
properties:
slug:
type: string
description: "the file stem"
date:
type: string
description: "YYYY-MM-DD; empty = today"
default: ""
required:
- slug
steps:
- backend: daily_create_step
daily:list:
backend: base
description: "List notes under a single day."
parameters:
type: object
properties:
date:
type: string
description: "YYYY-MM-DD; empty = today"
default: ""
steps:
- backend: daily_list_step
daily:reindex:
backend: base
description: "Rebuild the day-index page daily/<date>.md."
parameters:
type: object
properties:
date:
type: string
description: "YYYY-MM-DD; empty = today"
default: ""
steps:
- backend: daily_reindex_step
frontmatter:delete:
backend: base
description: "Drop keys from a file's frontmatter."
parameters:
type: object
properties:
path:
type: string
description: "seed path (vault-relative)"
depth:
type: integer
description: "hop limit"
default: 1
direction:
description: "vault-relative path"
keys:
type: array
description: "keys to remove"
items:
type: string
required:
- path
- keys
steps:
- backend: frontmatter_delete_step
frontmatter:read:
backend: base
description: "Read a file's frontmatter as a dict."
parameters:
type: object
properties:
path:
type: string
description: "forward / backward / both"
default: both
description: "vault-relative path"
required:
- path
steps:
- backend: traverse_step
- backend: frontmatter_read_step
# ── Read Operations ───────────────────────────────────────────────────────────
- backend: base
name: list
frontmatter:update:
backend: base
description: "Merge key-values into a file's frontmatter."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
metadata:
type: object
description: "key-values to merge"
required:
- path
- metadata
steps:
- backend: frontmatter_update_step
stat:
backend: base
description: "Stat path (size, mtime, exists, is_dir, is_file)."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
required:
- path
steps:
- backend: stat_step
list:
backend: base
description: "List files under a vault path."
parameters:
type: object
@ -152,164 +238,8 @@ jobs:
steps:
- backend: list_step
- backend: base
name: read
description: "Read a markdown file under the vault."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path; markdown only"
start_line:
type: integer
description: "first line (1-based, inclusive)"
end_line:
type: integer
description: "last line (1-based, inclusive)"
required:
- path
steps:
- backend: read_step
- backend: base
name: stat
description: "Stat a vault file (size, mtime, exists, is_dir, is_file)."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
required:
- path
steps:
- backend: stat_step
- backend: base
name: frontmatter:read
description: "Read a file's YAML frontmatter as a dict."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
required:
- path
steps:
- backend: frontmatter:read_step
# ── Write Operations──────────────────────────────────────────────────────────
- backend: base
name: write
description: "Write a markdown file (create or overwrite) with name/description frontmatter."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path; markdown only"
name:
type: string
description: "frontmatter name"
description:
type: string
description: "frontmatter description"
content:
type: string
description: "body"
required:
- path
- name
- description
- content
steps:
- backend: write_step
- backend: base
name: edit
description: "Find-and-replace in a markdown file (all occurrences)."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
old:
type: string
description: "text to find"
new:
type: string
description: "replacement"
default: ""
required:
- path
- old
- new
steps:
- backend: edit_step
- backend: base
name: append
description: "Append content to a markdown file."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
content:
type: string
description: "content to append"
required:
- path
- content
steps:
- backend: append_step
- backend: base
name: frontmatter:update
description: "Merge keys into a file's YAML frontmatter."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
metadata:
type: object
description: "keys to merge"
additionalProperties: true
required:
- path
- metadata
steps:
- backend: frontmatter_update_step
- backend: base
name: frontmatter:delete
description: "Drop keys from a file's YAML frontmatter."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
keys:
type: array
description: "keys to remove"
items:
type: string
required:
- path
- keys
steps:
- backend: frontmatter_delete_step
# ── File Operations (relocate / cross-realm) ──────────────────────────────
- backend: base
name: move
move:
backend: base
description: "Move / rename a vault file; rewrites inbound wikilinks by default."
parameters:
type: object
@ -334,8 +264,8 @@ jobs:
steps:
- backend: move_step
- backend: base
name: delete
delete:
backend: base
description: "Delete a vault file or folder; returns surviving inbound wikilinks."
parameters:
type: object
@ -348,165 +278,74 @@ jobs:
steps:
- backend: delete_step
- backend: base
name: upload
description: "Copy a host file into the vault at an explicit destination."
parameters:
type: object
properties:
src_path:
type: string
description: "host absolute path"
dst_path:
type: string
description: "vault-relative destination"
overwrite:
type: boolean
description: "overwrite if dst exists"
default: false
required:
- src_path
- dst_path
steps:
- backend: upload_step
- backend: base
name: upload_resource
description: "Ingest an external-channel asset into resource/<today>/ with provenance."
read:
backend: base
description: "Read a markdown file under the vault."
parameters:
type: object
properties:
path:
type: string
description: "host source path"
channel:
type: string
description: "channel id (wechat / email / browser / api / ...)"
description:
type: string
description: "what the asset is and how to interpret it"
metadata:
type: object
description: "extra provenance keys (e.g. source)"
default: {}
description: "vault-relative path; markdown only"
start_line:
type: integer
description: "first line (1-based, inclusive)"
end_line:
type: integer
description: "last line (1-based, inclusive)"
required:
- path
- channel
steps:
- backend: read_step
write:
backend: base
description: "Write a markdown file (create or overwrite) with name/description frontmatter."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path; markdown only"
name:
type: string
description: "frontmatter name"
description:
type: string
description: "frontmatter description"
content:
type: string
description: "body"
required:
- path
- name
- description
- content
steps:
- backend: upload_resource_step
- backend: write_step
- backend: base
name: download
description: "Copy a vault file out to the host filesystem."
edit:
backend: base
description: "Find-and-replace in a markdown file (all occurrences)."
parameters:
type: object
properties:
src_path:
path:
type: string
description: "vault-relative source"
dst_path:
description: "vault-relative path"
old:
type: string
description: "host absolute dest; empty = temp file"
default: ""
overwrite:
type: boolean
description: "overwrite if dst exists"
default: false
required:
- src_path
steps:
- backend: download_step
# ── Daily Operations (note CRUD + day-index rollup) ───────────────────
- backend: base
name: daily:read
description: "Read daily/<date>/<slug>.md (body + frontmatter)."
parameters:
type: object
properties:
slug:
description: "text to find"
new:
type: string
description: "note slug"
date:
type: string
description: "ISO date; empty = today"
description: "replacement"
default: ""
required:
- slug
- path
- old
- new
steps:
- backend: daily_read_step
- backend: base
name: daily:write
description: "Write daily/<date>/<slug>.md (body + frontmatter); refreshes the day index."
parameters:
type: object
properties:
slug:
type: string
description: "note slug"
body:
type: string
description: "note body"
default: ""
frontmatter:
type: object
description: "frontmatter dict; defaults to {name: <slug>}"
default: {}
date:
type: string
description: "ISO date; empty = today"
default: ""
overwrite:
type: boolean
description: "false = skip if exists; true = replace"
default: false
refresh_index:
type: boolean
description: "refresh daily/<date>.md after write"
default: true
required:
- slug
steps:
- backend: daily_write_step
- backend: base
name: daily:list
description: "List notes under a single day."
parameters:
type: object
properties:
date:
type: string
description: "ISO date; empty = today"
default: ""
steps:
- backend: daily_list_step
- backend: base
name: daily:reindex
description: "Rebuild the day-index page daily/<date>.md."
parameters:
type: object
properties:
date:
type: string
description: "ISO date; empty = today"
default: ""
steps:
- backend: daily_reindex_step
- backend: background
name: watch_file
watch_paths:
- MEMORY.md
- memory
suffix_filters:
- md
steps:
- backend: update_store_step
- backend: watch_changes_step
- backend: edit_step
components:
tokenizer:
@ -517,7 +356,7 @@ components:
default:
backend: ${EMBEDDING_BACKEND:-openai}
api_key: ${EMBEDDING_API_KEY:-}
base_url: ${EMBEDDING_BASE_URL:-https://api.openai.com/v1}
base_url: ${EMBEDDING_BASE_URL:-https://dashscope.aliyuncs.com/compatible-mode/v1}
model_name: ${EMBEDDING_MODEL_NAME:-text-embedding-v4}
dimensions: 1024
@ -528,16 +367,10 @@ components:
file_parser:
linked:
backend: linked
supported_extensions:
- md
supported_extensions: [ "md" ]
chunked:
backend: chunked
supported_extensions:
- txt
- html
- json
- yaml
- py
supported_extensions: [ "txt", "html", "json", "yaml", "py" ]
default:
backend: default
@ -550,21 +383,7 @@ components:
default:
backend: local
store_name: local
embedding_model: default
# embedding_model: default
embedding_model: ""
keyword_index: default
file_graph: default
# as_llm / formatter aren't required for atomic primitives; configure
# only if you'll invoke digester/synchronizer or other LLM-driven
# paths from the dev server.
# as_llm:
# default:
# backend: ${LLM_BACKEND:-openai}
# api_key: ${LLM_API_KEY:-}
# model_name: ${LLM_MODEL_NAME:-gpt-4o-mini}
# client_kwargs:
# base_url: ${LLM_BASE_URL:-https://api.openai.com/v1}
#
# as_llm_formatter:
# default:
# backend: ${LLM_BACKEND:-openai}

45
reme4/config/demo.yaml Normal file
View file

@ -0,0 +1,45 @@
service:
backend: http
jobs:
demo:
backend: base
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
stream_demo:
backend: stream
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

474
reme4/config/qwenpaw.yaml Normal file
View file

@ -0,0 +1,474 @@
service:
backend: http
daily_dir: daily
digest_dir: digest
resource_dir: ""
jobs:
version:
backend: base
description: "return reme package version"
parameters:
type: object
properties: {}
steps:
- backend: version_step
health_check:
backend: base
description: "return a concise health-check snapshot of reme components"
parameters:
type: object
properties: {}
steps:
- backend: health_check_step
help:
backend: base
description: "list all registered jobs with their metadata"
parameters:
type: object
properties: {}
steps:
- backend: help_step
traverse:
backend: base
description: "Walk the wikilink graph from a path."
parameters:
type: object
properties:
path:
type: string
description: "path"
depth:
type: integer
description: "hop limit"
default: 1
direction:
type: string
enum:
- forward
- backward
- both
default: both
required:
- path
steps:
- backend: traverse_step
reindex:
backend: base
description: "wipe the file store and rebuild it from the watcher's tracked files"
parameters:
type: object
properties: {}
steps:
- backend: clear_and_scan_step
- backend: update_index_step
persist: true
index_changes:
backend: base
description: "apply a batch of file changes (added/modified/deleted) into file_store"
parameters:
type: object
properties:
changes:
type: array
description: "list of change items"
items:
type: object
properties:
change:
type: string
enum: [added, modified, deleted]
description: "type of file change"
path:
type: string
description: "absolute file path"
required:
- change
- path
required:
- changes
steps:
- backend: index_changes_step
# ════════════════════════════════════════════════════════════════════
# ATOMIC TOOLS — same surface plugins/reme-{service,expert} expose
# ════════════════════════════════════════════════════════════════════
# ── Retrieve ───────────────────────────────────────────────────────
search:
backend: base
description: "Hybrid vault search (vector + BM25, RRF-fused)."
parameters:
type: object
properties:
query:
type: string
description: "search query"
limit:
type: integer
description: "max results"
default: 5
min_score:
type: number
description: "min fused score"
default: 0.0
required:
- query
steps:
- backend: search_step
vector_weight: 0.7
candidate_multiplier: 3.0
expand_links: true
max_links_per_direction: 10
# ── Read Operations ───────────────────────────────────────────────────────────
list:
backend: base
description: "List files under a vault path."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative dir; empty = root"
default: ""
recursive:
type: boolean
description: "recurse"
default: false
limit:
type: integer
description: "max results"
default: 100
steps:
- backend: list_step
read:
backend: base
description: "Read a markdown file under the vault."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path; markdown only"
start_line:
type: integer
description: "first line (1-based, inclusive)"
end_line:
type: integer
description: "last line (1-based, inclusive)"
required:
- path
steps:
- backend: read_step
stat:
backend: base
description: "Stat a vault file (size, mtime, exists, is_dir, is_file)."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
required:
- path
steps:
- backend: stat_step
frontmatter:read:
backend: base
description: "Read a file's YAML frontmatter as a dict."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
required:
- path
steps:
- backend: frontmatter:read_step
# ── Write Operations──────────────────────────────────────────────────────────
write:
backend: base
description: "Write a markdown file (create or overwrite) with name/description frontmatter."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path; markdown only"
name:
type: string
description: "frontmatter name"
description:
type: string
description: "frontmatter description"
content:
type: string
description: "body"
required:
- path
- name
- description
- content
steps:
- backend: write_step
edit:
backend: base
description: "Find-and-replace in a markdown file (all occurrences)."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
old:
type: string
description: "text to find"
new:
type: string
description: "replacement"
default: ""
required:
- path
- old
- new
steps:
- backend: edit_step
append:
backend: base
description: "Append content to a markdown file."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
content:
type: string
description: "content to append"
required:
- path
- content
steps:
- backend: append_step
frontmatter:update:
backend: base
description: "Merge keys into a file's YAML frontmatter."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
metadata:
type: object
description: "keys to merge"
additionalProperties: true
required:
- path
- metadata
steps:
- backend: frontmatter_update_step
frontmatter:delete:
backend: base
description: "Drop keys from a file's YAML frontmatter."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
keys:
type: array
description: "keys to remove"
items:
type: string
required:
- path
- keys
steps:
- backend: frontmatter_delete_step
delete:
backend: base
description: "Delete a vault file or folder; returns surviving inbound wikilinks."
parameters:
type: object
properties:
path:
type: string
description: "vault-relative path"
required:
- path
steps:
- backend: delete_step
upload_resource:
backend: base
description: "Ingest an external-channel asset into resource/<today>/ with provenance."
parameters:
type: object
properties:
path:
type: string
description: "host source path"
channel:
type: string
description: "channel id (wechat / email / browser / api / ...)"
description:
type: string
description: "what the asset is and how to interpret it"
metadata:
type: object
description: "extra provenance keys (e.g. source)"
default: {}
required:
- path
- channel
- description
steps:
- backend: upload_resource_step
download:
backend: base
description: "Copy a vault file out to the host filesystem."
parameters:
type: object
properties:
src_path:
type: string
description: "vault-relative source"
dst_path:
type: string
description: "host absolute dest; empty = temp file"
default: ""
overwrite:
type: boolean
description: "overwrite if dst exists"
default: false
required:
- src_path
steps:
- backend: download_step
# ── Daily Operations (slug provisioning + day-index rollup) ──────────
daily:create:
backend: base
description: "Provision daily/<date>/<slug>.md (empty body, frontmatter {name: slug}); idempotent; refreshes the day index."
parameters:
type: object
properties:
slug:
type: string
description: "note slug"
date:
type: string
description: "ISO date; empty = today"
default: ""
required:
- slug
steps:
- backend: daily_create_step
daily:list:
backend: base
description: "List notes under a single day."
parameters:
type: object
properties:
date:
type: string
description: "ISO date; empty = today"
default: ""
steps:
- backend: daily_list_step
daily:reindex:
backend: base
description: "Rebuild the day-index page daily/<date>.md."
parameters:
type: object
properties:
date:
type: string
description: "ISO date; empty = today"
default: ""
steps:
- backend: daily_reindex_step
watch_file:
backend: background
watch_paths:
- MEMORY.md
- memory
suffix_filters:
- md
steps:
- backend: scan_changes_step
- backend: update_index_step
persist: true
- backend: watch_changes_step
dispatch_step: update_index_step
components:
tokenizer:
default:
backend: regex
embedding_model:
default:
backend: ${EMBEDDING_BACKEND:-openai}
api_key: ${EMBEDDING_API_KEY}
base_url: ${EMBEDDING_BASE_URL:-https://api.openai.com/v1}
model_name: ${EMBEDDING_MODEL_NAME:-text-embedding-v4}
dimensions: 1024
file_graph:
default:
backend: local
file_parser:
linked:
backend: linked
supported_extensions:
- md
chunked:
backend: chunked
supported_extensions:
- txt
- html
- json
- yaml
- py
default:
backend: default
keyword_index:
default:
backend: bm25
tokenizer: default
file_store:
default:
backend: local
store_name: local
embedding_model: default
keyword_index: default
file_graph: default

View file

@ -22,6 +22,8 @@ class ComponentEnum(str, Enum):
FILE_GRAPH = "file_graph"
FILE_CATALOG = "file_catalog"
KEYWORD_INDEX = "keyword_index"
SERVICE = "service"

View file

@ -13,5 +13,7 @@ class LinkScopeEnum(str, Enum):
"""
REAL = "real"
VIRTUAL = "virtual"
ALL = "all"

View file

@ -16,12 +16,12 @@ class ComponentConfig(BaseModel):
class JobConfig(ComponentConfig):
"""Config for a job — an ordered sequence of step components."""
"""Config for a job — an ordered sequence of step components. Keyed by name in ApplicationConfig.jobs."""
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")
enable_serve: bool = Field(default=True, description="Whether to expose this job through the service layer")
class ApplicationConfig(BaseModel):
@ -42,7 +42,10 @@ class ApplicationConfig(BaseModel):
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")
jobs: dict[str, JobConfig] = Field(
default_factory=dict,
description="Job definitions keyed by job name",
)
components: dict[ComponentEnum, dict[str, ComponentConfig]] = Field(
default_factory=dict,
description="Component registry keyed by type then name",

View file

@ -1,46 +1,71 @@
"""steps — registers every BaseStep subclass at import time.
"""steps"""
Each submodule's ``@R.register`` decorators only fire when the module
is imported. Auto-importing them here means any config that names a
step backend (e.g. ``graph_traverse_step``, ``write``, ``digester``)
will find it in the registry without the caller having to remember
which submodule it lives in.
File-I/O is split by blast radius. The ``crud`` package covers both
opaque-byte ops (list / stat / move / delete / upload / download) and
whole-file text ops (read / write / append / edit) — they share the
same path-resolution helpers, so they live in one package.
``frontmatter`` is the one sliced surface that earns its own RUD
package (YAML is structured data — surgical key edits cannot be safely
emulated with string-substitution on the body). For mid-file body
edits, use ``edit`` (exact string replacement) or do a read + write
round-trip.
* ``common`` — search / health_check / help / reindex / version / graph_traverse
* ``crud`` — list / stat / move / delete / upload / download / read / write / append / edit
* ``frontmatter`` — markdown frontmatter slice RUD (frontmatter_read_step / update / delete)
* ``daily`` — note genesis / list / day-index reindex
* ``jobs`` — synchronizer / digester (LLM-driven orchestrators)
"""
from . import common # noqa: F401 -- registers common steps (search, version, graph_traverse, ...)
from . import crud # noqa: F401 -- registers list/stat/upload/download/move/delete/read/write/append/edit
from . import frontmatter # noqa: F401 -- registers frontmatter_read_step/update/delete
from . import (
daily,
) # noqa: F401 -- registers daily_read_step / daily_write_step / daily_list_step / daily_reindex_step
from . import background # noqa: F401
# from . import jobs # noqa: F401 -- registers synchronizer / digester
from .base_step import BaseStep
from . import graph # noqa: F401
from .common.demo import DemoEchoStep1, DemoEchoStep2
from .common.health_check import HealthCheckStep
from .common.help import HelpStep
from .common.stream_demo import StreamDemoStep1, StreamDemoStep2
from .common.version import VersionStep
from .file_io.daily_create import DailyCreateStep
from .file_io.daily_list import DailyListStep
from .file_io.daily_reindex import DailyReindexStep
from .file_io.delete import DeleteStep
from .file_io.edit import EditStep
from .file_io.frontmatter_delete import FrontmatterDeleteStep
from .file_io.frontmatter_read import FrontmatterReadStep
from .file_io.frontmatter_update import FrontmatterUpdateStep
from .file_io.list import ListStep
from .file_io.move import MoveStep
from .file_io.read import ReadStep
from .file_io.stat import StatStep
from .file_io.write import WriteStep
from .index.clear_and_scan import ClearAndScanStep
from .index.scan_changes import ScanChangesStep
from .index.search import SearchStep
from .index.traverse import TraverseStep
from .index.update_catalog import UpdateCatalogStep
from .index.update_index import UpdateIndexStep
from .index.watch_changes import WatchChangesStep
from .transfer.download import DownloadStep
from .transfer.ingest import IngestStep
from .transfer.upload import UploadStep
__all__ = [
"background",
"common",
"crud",
"graph",
"frontmatter",
"daily",
"BaseStep",
# common
"DemoEchoStep1",
"DemoEchoStep2",
"HealthCheckStep",
"HelpStep",
"StreamDemoStep1",
"StreamDemoStep2",
"VersionStep",
# file_io
"DeleteStep",
"EditStep",
"ListStep",
"MoveStep",
"ReadStep",
"StatStep",
"WriteStep",
# file_io (daily)
"DailyCreateStep",
"DailyListStep",
"DailyReindexStep",
# file_io.frontmatter
"FrontmatterDeleteStep",
"FrontmatterReadStep",
"FrontmatterUpdateStep",
# index
"ClearAndScanStep",
"ScanChangesStep",
"SearchStep",
"TraverseStep",
"UpdateCatalogStep",
"UpdateIndexStep",
"WatchChangesStep",
# transfer
"DownloadStep",
"IngestStep",
"UploadStep",
]

View file

@ -1,11 +0,0 @@
"""Background steps."""
from .index_changes import IndexChangesStep
from .update_store import UpdateStoreStep
from .watch_changes import WatchChangesStep
__all__ = [
"IndexChangesStep",
"UpdateStoreStep",
"WatchChangesStep",
]

View file

@ -1,17 +0,0 @@
"""Common steps."""
from .health_check import HealthCheckStep
from .help import HelpStep
from .reindex import ReindexStep
from .search import SearchStep
from .traverse import TraverseStep
from .version import VersionStep
__all__ = [
"HealthCheckStep",
"HelpStep",
"ReindexStep",
"SearchStep",
"TraverseStep",
"VersionStep",
]

View file

@ -1,76 +1,86 @@
"""Return a concise health check snapshot of ReMe runtime components."""
"""Concise health snapshot of ReMe runtime components."""
import sys
from collections.abc import Mapping
import numpy as np
from ...enumeration.component_enum import ComponentEnum
from ..base_step import BaseStep
from ... import __version__
from ...components import R
from ...enumeration import ComponentEnum
# ---------------------------------------------------------------------------
# Memory accounting
# ---------------------------------------------------------------------------
def _deep_size(obj, _seen: set | None = None) -> int:
"""Recursive sizeof; uses ndarray.nbytes for numpy and walks containers / __dict__."""
"""Recursive sizeof. Uses ndarray.nbytes; walks Mappings, sequences, __dict__."""
if _seen is None:
_seen = set()
obj_id = id(obj)
if obj_id in _seen:
if id(obj) in _seen:
return 0
_seen.add(obj_id)
_seen.add(id(obj))
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())
extra = 0
elif isinstance(obj, Mapping):
extra = 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)
extra = sum(_deep_size(item, _seen) for item in obj)
elif hasattr(obj, "__dict__"):
size += _deep_size(vars(obj), _seen)
extra = _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
extra = sum(_deep_size(getattr(obj, s), _seen) for s in obj.__slots__ if hasattr(obj, s))
else:
extra = 0
return size + extra
def _mb_str(*objs) -> str:
"""Return summed deep size of objs formatted as 'X.XX MB'."""
"""Sum deep size of objs and format as 'X.XX MB'."""
seen: set = set()
total = sum(_deep_size(o, seen) for o in objs)
return f"{total / (1024 * 1024):.2f} MB"
# ---------------------------------------------------------------------------
# Per-component status collectors
# ---------------------------------------------------------------------------
def _embedding_status(comp) -> dict:
cache = getattr(comp, "_embedding_cache", {}) or {}
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 {}),
"cache_size": len(cache),
"memory": _mb_str(cache),
}
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.
def _file_graph_nx_status(comp, graph) -> dict:
"""Networkx backend: virtuals are nodes without a 'node' payload."""
n_real = sum(1 for _, d in graph.nodes(data=True) if "node" in d)
return {
"is_started": comp.is_started,
"n_nodes": n_real,
"n_edges": graph.number_of_edges(),
"n_virtual": graph.number_of_nodes() - n_real,
"memory": _mb_str(graph),
}
def _file_graph_local_status(comp) -> dict:
"""Local backend: nodes / inverse edges / pending edges held as separate dicts."""
nodes = getattr(comp, "_nodes", {}) or {}
inverse = getattr(comp, "_inverse", {}) or {}
pending = getattr(comp, "_pending", {}) or {}
@ -83,6 +93,11 @@ def _file_graph_status(comp) -> dict:
}
def _file_graph_status(comp) -> dict:
graph = getattr(comp, "_graph", None)
return _file_graph_nx_status(comp, graph) if graph is not None else _file_graph_local_status(comp)
def _file_store_status(comp) -> dict:
chunks = getattr(comp, "file_chunks", {}) or {}
return {
@ -94,12 +109,13 @@ def _file_store_status(comp) -> dict:
def _keyword_index_status(comp) -> dict:
vocab = getattr(comp, "vocab", {}) or {}
return {
"is_started": comp.is_started,
"n_docs": getattr(comp, "n_docs", None),
"vocab_size": len(getattr(comp, "vocab", {}) or {}),
"vocab_size": len(vocab),
"memory": _mb_str(
getattr(comp, "vocab", {}) or {},
vocab,
getattr(comp, "inverted_index", {}) or {},
getattr(comp, "doc_meta", {}) or {},
getattr(comp, "_idf_cache", {}) or {},
@ -115,8 +131,13 @@ _HANDLERS = {
}
def _is_status_healthy(ctype: ComponentEnum, status: dict) -> bool:
"""Per-component health rule. Unstarted = unhealthy; type-specific extras checked."""
# ---------------------------------------------------------------------------
# Health rules and step entry point
# ---------------------------------------------------------------------------
def _is_healthy(ctype: ComponentEnum, status: dict) -> bool:
"""Unstarted = unhealthy; embedding model also requires is_healthy != False."""
if not status.get("is_started"):
return False
if ctype is ComponentEnum.EMBEDDING_MODEL and status.get("is_healthy") is False:
@ -124,30 +145,38 @@ def _is_status_healthy(ctype: ComponentEnum, status: dict) -> bool:
return True
def _collect_components(app_context) -> tuple[dict, bool]:
"""Walk every registered component type and produce {type: {name: status}}, plus overall flag."""
components: dict = {}
healthy = True
for ctype, handler in _HANDLERS.items():
bucket = {}
for name, comp in app_context.components.get(ctype, {}).items():
status = handler(comp)
bucket[name] = status
if not _is_healthy(ctype, status):
healthy = False
components[ctype.value] = bucket
return components, healthy
@R.register("health_check_step")
class HealthCheckStep(BaseStep):
"""Collect a concise health check snapshot of the relevant components."""
"""Collect a concise health 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
components, healthy = _collect_components(self.app_context)
else:
components, healthy = {}, True
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'}"
emoji = "✅" if healthy else "❌"
label = "healthy" if healthy else "unhealthy"
self.context.response.answer = f"{emoji} ReMe v{__version__} - {label}"
self.context.response.metadata["health"] = health
return self.context.response

View file

@ -4,26 +4,26 @@ from ..base_step import BaseStep
from ...components import R
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 name, schema in props.items():
ptype = schema.get("type", "any")
if name in required:
parts.append(f"{name}:{ptype}*")
elif "default" in schema:
parts.append(f"{name}:{ptype}={schema['default']}")
else:
parts.append(f"{name}:{ptype}")
return ", ".join(parts)
@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)
"""List all registered jobs (excluding self and non-servable) as compact one-liners for an LLM."""
async def execute(self):
assert self.context is not None
@ -31,12 +31,11 @@ class HelpStep(BaseStep):
lines = []
if self.app_context is not None:
for name, job in self.app_context.jobs.items():
if name == "help":
if name == "help" or not getattr(job, "enable_serve", True):
continue
lines.append(f"🛠️ `{name}` — {job.description} 📥 {self._format_params(job.parameters)}")
lines.append(f"🛠️ `{name}` — {job.description} 📥 {_format_params(job.parameters)}")
self.logger.info(f"[{self.name}] returning {len(lines)} jobs")
self.context.response.answer = "\n".join(lines)
self.context.response.metadata["job_count"] = len(lines)
return self.context.response

View file

@ -1,35 +0,0 @@
"""Wipe the file store and rebuild it by scanning the vault from disk."""
from ..base_step import BaseStep
from ...components import R
@R.register("reindex_step")
class ReindexStep(BaseStep):
"""Full re-index: clear store, walk vault, hand the file list to index_changes."""
async def execute(self):
assert self.context is not None
suffix_filters: list[str] = self.context.get("suffix_filters", ["md"])
suffixes = tuple("." + s.strip(".") for s in suffix_filters) if suffix_filters else None
await self.file_store.clear()
paths: list[str] = []
for p in self.vault_path.rglob("*"):
if not p.is_file():
continue
if suffixes and not str(p).endswith(suffixes):
continue
paths.append(str(p.absolute()))
if paths:
await self.run_job("index_changes", changes=[{"change": "added", "path": p} for p in paths])
await self.file_store.dump()
counts = {"added": len(paths), "modified": 0, "deleted": 0}
self.logger.info(f"[{self.name}] reindexed {counts}")
self.context.response.answer = f"🔄 Reindexed {counts['added']} file(s)"
self.context.response.metadata["counts"] = counts
return self.context.response

View file

@ -1,154 +0,0 @@
"""``traverse_step`` — BFS over wikilink edges from a seed file.
Single tool for relationship browsing. ``depth=1`` covers the trivial
"what does this link to / what links here" lookups (set ``direction``
accordingly); higher depth opens up multi-hop exploration.
Output is one record per edge traversed (not per node), so the same
target can appear multiple times if reached via different predicates
or paths — agents dedupe at the call site if they want a flat node
set. Each record carries ``via`` (the predecessor) and the link's
``predicate`` / ``anchor`` so the agent can reconstruct the path.
Adjacency is loaded once via ``file_graph.get_nodes(None)`` — every
real node arrives with its full ``links`` payload, and we build both
the outbound and the inbound index in a single pass. The BFS then
runs purely in memory: no per-frontier-node graph round-trips, no
filesystem walk. The ``get_inlinks`` / ``get_outlinks`` contract
methods stay unused here because they'd add network round-trips for
data we already have.
Direction vocabulary accepts both the standard convention
(``forward`` / ``backward`` / ``both``) and the engine convention
(``out`` / ``in`` / ``both``).
Seeds are paths relative to the vault used as-is — short-form resolution
is no longer attempted. Seeds that don't match any graph node yield
empty BFS results (no error).
"""
from collections import deque
from pathlib import Path
from ..base_step import BaseStep
from ...components import R
from ...schema import FileLink
_FORWARD = {"out", "forward"}
_BACKWARD = {"in", "backward"}
_BOTH = {"both"}
_VALID_DIRECTIONS = _FORWARD | _BACKWARD | _BOTH
async def _build_indexes(
file_store,
) -> tuple[
dict[str, list[tuple[str, FileLink]]],
dict[str, list[tuple[str, FileLink]]],
]:
"""One ``get_nodes(None)`` call → (outbound, inbound) adjacency dicts.
Each dict is keyed by node path; values are ``(neighbor_path, link)``
tuples. Source paths land in the inbound index alongside the link
object — solving the contract gap where ``get_inlinks`` returns
target-shaped FileLinks without source attribution.
"""
outbound: dict[str, list[tuple[str, FileLink]]] = {}
inbound: dict[str, list[tuple[str, FileLink]]] = {}
if not file_store.file_graph:
return outbound, inbound
for node in await file_store.file_graph.get_nodes():
for link in node.links:
if not link.target_path:
continue
outbound.setdefault(node.path, []).append((link.target_path, link))
inbound.setdefault(link.target_path, []).append((node.path, link))
return outbound, inbound
def _bfs(
seeds: list[str],
max_depth: int,
direction: str,
outbound: dict[str, list[tuple[str, FileLink]]],
inbound: dict[str, list[tuple[str, FileLink]]],
) -> list[dict]:
"""In-memory BFS. One record per edge traversed."""
walk_out = direction in _FORWARD or direction in _BOTH
walk_in = direction in _BACKWARD or direction in _BOTH
visited_edges: set[tuple[str, str, str | None]] = set()
results: list[dict] = []
queue: deque[tuple[str, int]] = deque((s, 0) for s in seeds)
while queue:
current, depth = queue.popleft()
if depth >= max_depth:
continue
edges: list[tuple[str, str | None, str | None]] = []
if walk_out:
for tgt, link in outbound.get(current, ()):
edges.append((tgt, link.predicate, link.target_anchor))
if walk_in:
for src, link in inbound.get(current, ()):
edges.append((src, link.predicate, link.target_anchor))
for next_path, pred, anchor in edges:
edge_key = (current, next_path, pred)
if edge_key in visited_edges:
continue
visited_edges.add(edge_key)
results.append(
{
"path": next_path,
"depth": depth + 1,
"via": current,
"predicate": pred,
"anchor": anchor,
},
)
if depth + 1 < max_depth:
queue.append((next_path, depth + 1))
return results
def _normalize_seeds(raw) -> list[str]:
"""Coerce raw seed input to a non-empty list of strings relative to the vault."""
if isinstance(raw, (str, Path)):
items = [raw]
else:
items = list(raw or [])
return [str(p) for p in items if p]
@R.register("traverse_step")
class TraverseStep(BaseStep):
"""BFS from a seed file to explore wikilink relationships.
Parameters:
path — single seed (str) or a list of seeds.
direction — ``forward`` / ``backward`` / ``both`` (or ``out`` / ``in`` / ``both``).
depth — hop limit (default 1 = immediate neighbors).
"""
async def execute(self):
assert self.context is not None
seeds_raw = self.context.get("path")
depth = int(self.context.get("depth") or 1)
direction = (self.context.get("direction") or "both").lower()
assert (
direction in _VALID_DIRECTIONS
), f"direction must be one of {sorted(_VALID_DIRECTIONS)}, got {direction!r}"
seeds = _normalize_seeds(seeds_raw)
assert seeds, "path is required"
outbound, inbound = await _build_indexes(self.file_store)
results = _bfs(seeds, depth, direction, outbound, inbound)
self.context.response.success = True
seed_label = seeds[0] if len(seeds) == 1 else f"{len(seeds)} seeds"
self.context.response.answer = f"Traversed {len(results)} edge(s) from {seed_label}"
self.context.response.metadata.update({"edges": results, "count": len(results)})

View file

@ -7,7 +7,7 @@ from ...components import R
@R.register("version_step")
class VersionStep(BaseStep):
"""Emit reme4.__version__ as the response answer."""
"""Emit reme.__version__ as the response answer."""
async def execute(self):
assert self.context is not None

View file

@ -1,40 +0,0 @@
"""File-level ops on vault_dir — both opaque-byte and text-content surfaces.
The package covers two related surfaces:
* **Opaque-byte ops** (don't care about file type): ``delete``,
``download``, ``list``, ``move``, ``stat``, ``upload``,
``upload_resource``.
* **Text-content ops** (markdown-aware; layered on the same path-
resolution helpers in ``_file_io.py``): ``read``, ``write``,
``append``, ``edit``.
For frontmatter slice RUD (YAML structured-data semantics) see
``reme4.steps.frontmatter``.
"""
from .read import ReadStep
from .edit import EditStep
from .delete import DeleteStep
from .write import WriteStep
from .append import AppendStep
from .move import MoveStep
from .stat import StatStep
from .download import DownloadStep
from .list import ListStep
from .upload import UploadStep
from .upload_resource import UploadResourceStep
__all__ = [
"DeleteStep",
"WriteStep",
"AppendStep",
"MoveStep",
"StatStep",
"DownloadStep",
"ListStep",
"UploadStep",
"UploadResourceStep",
"ReadStep",
"EditStep",
]

View file

@ -1,221 +0,0 @@
"""Shared filesystem helpers for CRUD steps (path gating, safe read, truncation)."""
from pathlib import Path
from typing import Iterable
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()
NON_MD_WARNING = (
"non-markdown file detected; CRUD operations are recommended on markdown files. "
"Operating in compatibility mode may carry risks of errors."
)
# Modern text formats — assume UTF-8 by convention.
_STANDARD_TEXT_EXTS = {
".md",
".py",
".js",
".ts",
".json",
".yaml",
".yml",
".html",
".css",
".xml",
".log",
".conf",
".ini",
".txt",
".sh",
}
# Legacy formats that may use ANSI/GBK on Chinese Windows systems.
_NON_STANDARD_EXTS = {".csv", ".bat", ".cmd", ".reg"}
def resolve_path(vault_path: Path, raw: str) -> tuple[Path | None, str | None]:
"""Resolve a `path=` argument against ``vault_path``.
Rules:
- Relative paths are joined under ``vault_path``.
- Absolute paths are accepted and returned as-is; a warning is logged
recommending relative paths, but the read still proceeds.
Returns ``(abs_path, None)`` on success, or ``(None, error_message)`` on failure
(currently only when ``raw`` is empty/blank).
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, recommending relative paths")
return p, None
return vault_path / p, None
def gate_md(target: Path) -> tuple[Path, bool]:
"""Markdown gate with compatibility fallback.
Returns ``(path, is_md)``:
- No suffix → auto-append `.md`, ``is_md=True``.
- `.md` suffix → ``is_md=True``.
- Any other suffix → ``is_md=False`` (caller handles degraded mode).
"""
if target.suffix == "":
return target.with_suffix(".md"), True
if target.suffix.lower() != ".md":
return target, False
return target, True
def _try_decode(data: bytes, encodings: Iterable[str]) -> tuple[str, str] | None:
"""Return ``(text, encoding)`` for the first encoding that decodes ``data`` cleanly."""
for enc in encodings:
try:
return data.decode(enc), enc
except (UnicodeDecodeError, LookupError):
continue
return None
def decode_known_file(data: bytes, file_extension: str) -> tuple[str, str]:
"""Decode file bytes using the extension as a hint. Returns ``(text, encoding)``.
Strategy:
1. BOM-based detection.
2. Extension-driven defaults:
3. Last resort → UTF-8 with ``errors='replace'`` so the function never raises.
"""
if data.startswith(b"\xef\xbb\xbf"):
return data.decode("utf-8-sig"), "utf-8-sig"
if data.startswith((b"\xff\xfe", b"\xfe\xff")):
try:
return data.decode("utf-16"), "utf-16"
except UnicodeDecodeError:
pass
ext = (file_extension or "").lower()
if ext in _STANDARD_TEXT_EXTS:
try:
return data.decode("utf-8-sig"), "utf-8"
except UnicodeDecodeError:
pass # fall through
if ext in _NON_STANDARD_EXTS:
result = _try_decode(data, ("utf-8-sig", "gbk"))
if result is not None:
text, enc = result
return text, "utf-8" if enc == "utf-8-sig" else enc
# Unknown extension or earlier strategies failed.
return data.decode("utf-8", errors="replace"), "utf-8"
async def read_file_safe(file_path, max_bytes: int = MAX_FILE_READ_BYTES) -> str:
"""Read file in byte mode and decode to string using extension-aware strategy."""
stat = await aiofiles.os.stat(str(file_path))
read_size = min(stat.st_size, max_bytes)
async with aiofiles.open(str(file_path), "rb") as f:
data = await f.read(read_size)
text, _ = decode_known_file(data, Path(file_path).suffix)
return text
async def detect_file_encoding(file_path, sniff_bytes: int = 8192) -> str:
"""Detect the encoding of an existing file so writes can preserve it.
Reads up to ``sniff_bytes`` from the head of the file (enough for BOM
detection and statistical analysis). Falls back to ``utf-8`` if the file
is unreadable.
"""
try:
async with aiofiles.open(str(file_path), "rb") as f:
data = await f.read(sniff_bytes)
except Exception: # pylint: disable=broad-except
return "utf-8"
_, enc = decode_known_file(data, Path(file_path).suffix)
return enc
async def write_file_safe(file_path: Path, content: str | bytes, encoding: str = "utf-8") -> None:
"""Write ``content`` to ``file_path`` in binary mode; creates parent dirs.
``str`` input is encoded with ``encoding`` (default UTF-8); callers wanting
to preserve a file's original encoding should pass the result of
:func:`detect_file_encoding`. If the requested ``encoding`` can't represent
some characters, falls back to UTF-8 to avoid data loss.
``bytes`` input is written verbatim — callers managing their own encoding
can pass raw bytes directly.
"""
file_path.parent.mkdir(parents=True, exist_ok=True)
if isinstance(content, str):
try:
payload = content.encode(encoding)
except (UnicodeEncodeError, LookupError):
logger.warning(
"write_file_safe: %r cannot encode all chars, falling back to utf-8",
encoding,
)
payload = content.encode("utf-8")
else:
payload = content
async with aiofiles.open(str(file_path), "wb") as f:
await f.write(payload)
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

View file

@ -1,75 +0,0 @@
"""Append content to the end of a file (auto-creates if missing)."""
import aiofiles
from ._file_io import NON_MD_WARNING, detect_file_encoding, gate_md, resolve_path
from ..base_step import BaseStep
from ...components import R
@R.register("append_step")
class AppendStep(BaseStep):
"""Append `content` to the target file. If the file does not exist it is
created (a system notice is appended to the answer in that case).
Content is appended verbatim — callers control whether a separator newline
is included in the input."""
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 "")
content = self.context.get("content")
content_str = "" if content is None else str(content)
target, err = resolve_path(self.vault_path, raw)
if err:
self._fail(err)
return None
target, is_md = gate_md(target)
if target.exists() and not target.is_file():
self._fail(f"path {target} is not a file", path=str(target))
return None
created = not target.exists()
# Preserve the existing file's encoding so appended bytes don't corrupt
# a non-UTF-8 file (e.g. GBK CSV). New files default to UTF-8.
encoding = "utf-8" if created else await detect_file_encoding(target)
try:
if created:
target.parent.mkdir(parents=True, exist_ok=True)
try:
payload = content_str.encode(encoding)
except (UnicodeEncodeError, LookupError):
self.logger.warning(
f"[{self.name}] cannot encode appended content as {encoding!r}, falling back to utf-8",
)
payload = content_str.encode("utf-8")
async with aiofiles.open(str(target), "ab") as f:
await f.write(payload)
except Exception as e: # pylint: disable=broad-except
self._fail(f"write failed: {e}", path=str(target))
return None
nbytes = len(payload)
self.context.response.success = True
if created:
answer = f"Appended {nbytes} bytes to {target} [system notice: file did not exist and was auto-created]"
else:
answer = f"Appended {nbytes} bytes to {target}"
if not is_md:
answer = f"{answer} [system notice: {NON_MD_WARNING}]"
self.context.response.answer = answer
self.logger.info(
f"[{self.name}] appended path={target} bytes={nbytes} encoding={encoding} "
f"created={created} is_md={is_md}",
)
return self.context.response

View file

@ -1,55 +0,0 @@
"""``file_list`` — enumerate files under a directory in the vault.
Reads directly from the filesystem (``Path.iterdir`` /
``Path.rglob``), **not** the file_store index. The store may lag
behind disk during indexing or after rapid mutations; for the most
current view, the on-disk walk is the source of truth.
Parameters:
path — directory to list under (relative to the vault or absolute).
Empty = vault root.
limit — cap the number of returned items.
recursive — descend into subdirectories. Default False = direct
children only.
No frontmatter is read — this is a plain directory walker. Callers
that need frontmatter-based filtering should iterate the result and
call ``frontmatter_read`` per candidate.
"""
from pathlib import Path
from ..base_step import BaseStep
from ...components import R
@R.register("list_step")
class ListStep(BaseStep):
"""Enumerate files under a directory in the vault."""
async def execute(self):
assert self.context is not None
path: str = self.context.get("path") or ""
recursive: bool = bool(self.context.get("recursive", False))
limit: int = int(self.context.get("limit") or 100)
vault_dir = Path(self.file_store.vault_path or ".").resolve()
target_dir = (vault_dir / (path or ".")).resolve()
items: list[str] = []
if target_dir.is_dir():
entries = target_dir.rglob("*") if recursive else target_dir.iterdir()
for entry in entries:
if not entry.is_file():
continue
try:
rel = str(entry.relative_to(vault_dir))
except ValueError:
rel = str(entry)
items.append(rel)
if len(items) >= limit:
break
self.context.response.success = True
self.context.response.answer = f"Listed {len(items)} file(s) under {path or '.'}"
self.context.response.metadata.update({"items": items, "count": len(items)})

View file

@ -1,91 +0,0 @@
"""Read a markdown file from vault_dir, with line-range slicing and byte-truncation."""
from ._file_io import (
NON_MD_WARNING,
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.vault_path, raw)
if err:
self._fail(err)
return None
target, is_md = gate_md(target)
if not is_md:
self.logger.info(f"[{self.name}] {NON_MD_WARNING} path={target}")
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

View file

@ -1,44 +0,0 @@
"""Daily-aware steps — CRUD on note md + day-level index.
A daily note is the single file ``daily/<YYYY-MM-DD>/<slug>.md``.
The day-level index ``daily/<YYYY-MM-DD>.md`` aggregates that day's
notes into a richer overview page (note list with name / description).
The index is a derived artifact — its source of truth lives in each
note's frontmatter; refreshes are idempotent and preserve manual
annotations in marker-delimited sections.
Tool boundary. The daily module exposes only the operations whose
shape is note- or day-specific:
* ``daily_read_step`` — read a note by ``slug + date``; returns the
body in ``answer`` and the parsed frontmatter as a dict in metadata,
so callers skip a separate ``frontmatter_read`` round-trip.
* ``daily_write_step`` — write the full body + frontmatter for
``daily/<date>/<slug>.md`` in one shot. Validates the slug, mkdirs
the day folder, refreshes the day index. ``mode="create"`` (default)
is idempotent skip-if-exists; ``mode="overwrite"`` is unconditional.
* ``daily_list_step`` — pure read of the notes under a single day
(defaults to today); returns ``{date, notes: [{path, slug, name,
description}, ...]}``. Does **not** touch the day index — call
``daily_reindex`` explicitly when the rollup page needs rebuilding.
* ``daily_reindex_step`` — explicit idempotent rebuild of a day's
index (historical backfill, drift recovery, batch-write reindex).
Body mid-edits / appends / arbitrary-path reads go through the generic
``read`` / ``write`` / ``append`` / ``edit`` steps. Frontmatter slice
mutations go through ``frontmatter_update`` / ``frontmatter_delete``.
The day-index is rebuilt explicitly via ``daily_reindex`` after a
batch of mutations.
"""
from .read import DailyReadStep
from .write import DailyWriteStep
from .list import DailyListStep
from .reindex import DailyReindexStep
__all__ = [
"DailyReadStep",
"DailyWriteStep",
"DailyListStep",
"DailyReindexStep",
]

View file

@ -1,306 +0,0 @@
"""Internal helpers for daily-aware steps — slug validation + day-index rebuild.
Two related concerns, both private to the ``daily`` package:
1. **Slug naming** — Windows-safe filename validation for the slug
that becomes the stem of ``daily/<YYYY-MM-DD>/<slug>.md``.
2. **Day-index** — the derived rollup page ``daily/<YYYY-MM-DD>.md``
listing every note under that date with name + description. The
index is auto-managed in marker-delimited sections; user-edited
manual sections are preserved verbatim across refreshes.
Public entry points:
* :func:`validate_slug` — return an error string, or ``None`` when the
slug is safe to use as a filename.
* :func:`scan_notes` — walk ``<daily_dir>/<date>/*.md`` and pull each
note's reserved frontmatter (``name`` / ``description``).
* :func:`refresh_day_index` — rebuild ``<daily_dir>/<date>.md`` from
the current state of its notes. Idempotent, safe to re-run.
"""
import re
from pathlib import Path
import frontmatter
# ---------------------------------------------------------------------------
# Slug validation
# ---------------------------------------------------------------------------
_INVALID_CHARS = re.compile(r'[<>:"/\\|?*\x00-\x1f]')
_RESERVED_NAMES = {
"CON",
"PRN",
"AUX",
"NUL",
*(f"COM{i}" for i in range(1, 10)),
*(f"LPT{i}" for i in range(1, 10)),
}
def validate_slug(slug: str) -> str | None:
"""Return an error message, or ``None`` when ``slug`` is a safe filename.
Rules (Windows is the strictest filesystem, so we validate to its bar):
- non-empty, no leading / trailing whitespace
- no reserved characters: ``< > : " / \\ | ? *`` or control chars (``\\x00-\\x1f``)
- no reserved device names: ``CON`` / ``PRN`` / ``AUX`` / ``NUL`` /
``COM1-9`` / ``LPT1-9`` (Windows reserves these with or without an
extension — ``CON.txt`` is also forbidden)
- no trailing ``.``
"""
if not slug:
return "slug is required"
if slug != slug.strip():
return f"slug cannot have leading or trailing whitespace: {slug!r}"
if _INVALID_CHARS.search(slug):
return f'slug contains invalid characters (one of < > : " / \\ | ? * ' f"or a control char): {slug!r}"
if slug.endswith("."):
return f"slug cannot end with '.': {slug!r}"
if slug.split(".", 1)[0].upper() in _RESERVED_NAMES:
return f"slug is a Windows-reserved device name: {slug!r}"
return None
# ---------------------------------------------------------------------------
# Day-index rebuild
# ---------------------------------------------------------------------------
#
# The day index is a derived artifact whose single job is **daily-note
# consolidation** — its source of truth lives in each note's
# frontmatter. The rebuild refreshes auto-managed sections while
# preserving any manual content the user has added between markers.
#
# Frontmatter shape — only the two reserved fields::
#
# name: <date>
# description: <one-line note-count digest>
#
# The note inventory lives in the body's ``<!-- notes:auto -->``
# wikilinks (graph edges feed off them). No bespoke status / lifecycle
# / scope / role / source / created axes — those are user-defined and
# intentionally absent from the auto-managed payload.
#
# Body auto sections (rebuilt on every refresh, marker-delimited):
#
# * ``notes`` — bulleted list of ``[[link]]\n name — description`` rows
#
# Manual sections live outside the auto markers and are preserved
# verbatim across refreshes. A fresh day file gets a ``## 备忘``
# section seeded as the manual scratch area.
# Marker syntax: HTML comments so they're invisible in rendered markdown
# but trivially detectable in source. Each block has a paired open/close.
_BLOCK_NAMES = ("notes",)
_BLOCK_OPEN = "<!-- {name}:auto -->"
_BLOCK_CLOSE = "<!-- /{name}:auto -->"
_HEADINGS = {
"notes": "## 今日笔记",
}
_MANUAL_HEADING = "## 备忘"
_MANUAL_STUB = "(人工记录区,刷新索引时不会动)"
def _block_re(name: str) -> re.Pattern:
"""Capturing regex for an auto block: heading + open marker + inner + close."""
return re.compile(
rf"(?P<heading>^{re.escape(_HEADINGS[name])}\s*\n)?"
rf"{re.escape(_BLOCK_OPEN.format(name=name))}"
r"(?P<inner>.*?)"
rf"{re.escape(_BLOCK_CLOSE.format(name=name))}",
re.DOTALL | re.MULTILINE,
)
def _count_digest(n: int) -> str:
"""One-line note count, used as the index ``description``."""
if n == 0:
return "本日暂无笔记。"
return f"今日 {n} 篇笔记。"
def scan_notes(vault_dir: Path, date: str, daily_dir: str) -> list[dict]:
"""Walk ``<daily_dir>/<date>/*.md`` and pull each note's frontmatter.
Returns one dict per note::
{"slug": str, "path": str, "name": str, "description": str}
Each ``.md`` directly under the day folder is a note; the file's
stem is the slug. Only reserved fields (name / description) are
read — user-defined frontmatter keys are ignored by the index.
"""
date_dir = vault_dir / daily_dir / date
if not date_dir.is_dir():
return []
out: list[dict] = []
for md_path in sorted(p for p in date_dir.iterdir() if p.is_file() and p.suffix == ".md"):
slug = md_path.stem
try:
post = frontmatter.loads(md_path.read_text(encoding="utf-8"))
except Exception: # pylint: disable=broad-except
continue
meta = post.metadata or {}
out.append(
{
"slug": slug,
"path": f"{daily_dir}/{date}/{slug}.md",
"name": str(meta.get("name") or slug),
"description": str(meta.get("description") or "").strip(),
},
)
return out
def _render_notes_block(notes: list[dict]) -> str:
"""Bulleted note digest: link on the bullet line, then an indented
``name — description`` summary so an agent can scan "what's
happening today" without opening each note.
The indented summary is omitted entirely when both name and
description add no information beyond the slug already shown in
the link.
"""
if not notes:
return "(无)"
lines: list[str] = []
for note in notes:
lines.append(f"- [[{note['path']}]]")
name = note["name"] if note["name"] and note["name"] != note["slug"] else ""
description = note["description"]
if name and description:
lines.append(f" {name} — {description}")
elif name:
lines.append(f" {name}")
elif description:
lines.append(f" {description}")
return "\n".join(lines)
def _wrap_block(name: str, inner: str) -> str:
"""Wrap rendered inner content with heading + auto markers."""
return f"{_HEADINGS[name]}\n" f"{_BLOCK_OPEN.format(name=name)}\n" f"{inner}\n" f"{_BLOCK_CLOSE.format(name=name)}"
def _replace_or_append(body: str, name: str, fresh_block: str) -> str:
"""Replace an existing auto block in-place; append at end if absent.
The replacement keeps the user's heading line if they renamed the
auto-heading (we only own the marker-wrapped inner). Appending uses
our canonical heading + markers so future refreshes find them.
"""
pattern = _block_re(name)
if pattern.search(body):
replacement = f"{_BLOCK_OPEN.format(name=name)}\n" f"{fresh_block}\n" f"{_BLOCK_CLOSE.format(name=name)}"
return pattern.sub(
lambda m: (m.group("heading") or "") + replacement,
body,
count=1,
)
suffix = _wrap_block(name, fresh_block)
return f"{body.rstrip()}\n\n{suffix}\n" if body.strip() else f"{suffix}\n"
def _seed_body(blocks: dict[str, str]) -> str:
"""Fresh-file body: all auto blocks in canonical order + manual stub."""
parts = [_wrap_block(name, blocks[name]) for name in _BLOCK_NAMES]
parts.append(f"{_MANUAL_HEADING}\n{_MANUAL_STUB}")
return "\n\n".join(parts) + "\n"
def _merge_blocks(body: str, blocks: dict[str, str]) -> str:
"""Refresh every auto block in-place; never touch manual content."""
for name in _BLOCK_NAMES:
body = _replace_or_append(body, name, blocks[name])
return body
def _frontmatter_payload(date: str, notes: list[dict]) -> dict:
"""Reserved-field-only frontmatter for the index page.
Emits ``name`` / ``description`` and nothing else — other axes
(status / lifecycle / scope / role / source / created) are
user-defined and belong in note bodies, not in this derived
aggregate.
"""
return {
"name": date,
"description": _count_digest(len(notes)),
}
async def refresh_day_index(file_store, date: str, daily_dir: str = "daily") -> dict:
"""Rebuild ``<daily_dir>/<date>.md`` from the current state of its notes.
Behaviour:
* No ``<daily_dir>/<date>/`` at all and no existing index file → no-op.
* Notes present → write the index file (create if missing,
otherwise merge auto blocks into the existing body, preserve
manual segments, refresh frontmatter).
* Notes directory empty but index file exists → rebuild with
empty auto blocks (keeps the file in sync with reality).
``daily_dir`` defaults to ``"daily"`` for tests / pure-helper
consumers; the registered steps pass the configured
``application_config.daily_dir`` so the on-disk layout always
matches what the index file claims.
Returns::
{
"date": str,
"path": "<daily_dir>/<date>.md",
"notes": [
{"path": "<daily_dir>/<date>/<slug>.md",
"name": str,
"description": str},
...
],
"created": bool, # True if index file was just written for the first time
}
The ``notes`` list mirrors the order rendered in the index body
(sorted by slug). The ``created`` field reflects index-page creation,
not note creation, so callers can log "index emerged" events
distinctly.
"""
vault_dir = Path(file_store.vault_path or ".").resolve()
index_rel = f"{daily_dir}/{date}.md"
index_abs = vault_dir / index_rel
notes = scan_notes(vault_dir, date, daily_dir)
notes_payload = [{"path": n["path"], "name": n["name"], "description": n["description"]} for n in notes]
if not notes and not index_abs.is_file():
return {
"date": date,
"path": index_rel,
"notes": notes_payload,
"created": False,
}
blocks = {"notes": _render_notes_block(notes)}
if index_abs.is_file():
post = frontmatter.loads(index_abs.read_text(encoding="utf-8"))
new_body = _merge_blocks(post.content, blocks)
was_created = False
else:
index_abs.parent.mkdir(parents=True, exist_ok=True)
new_body = _seed_body(blocks)
was_created = True
fm = _frontmatter_payload(date, notes)
out = frontmatter.Post(new_body, **fm)
index_abs.write_text(frontmatter.dumps(out), encoding="utf-8")
return {
"date": date,
"path": index_rel,
"notes": notes_payload,
"created": was_created,
}

View file

@ -1,91 +0,0 @@
"""``daily_read`` — read a daily note by slug + date; return body + parsed frontmatter.
Convenience wrapper around the generic ``read`` step for the
``daily/<YYYY-MM-DD>/<slug>.md`` path shape. The value over a raw
``file_read`` is two-fold:
* slug validation (Windows-safe filename rules) up front
* frontmatter parsed into a dict in metadata, so callers don't need a
separate ``frontmatter_read`` round-trip
Inputs:
slug (required, validated)
date (default today, ISO ``YYYY-MM-DD``)
Outputs:
answer = note body (frontmatter stripped)
metadata = {date, slug, path, exists, frontmatter: dict}
For arbitrary-path reads or ranged reads use the generic ``read`` step.
"""
from datetime import date as _date
import frontmatter
from ._daily_io import validate_slug
from ..crud._file_io import read_file_safe
from ..base_step import BaseStep
from ...components import R
@R.register("daily_read_step")
class DailyReadStep(BaseStep):
"""Read ``daily/<date>/<slug>.md`` → body + parsed frontmatter."""
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):
assert self.context is not None
slug: str = self.context.get("slug", "") or ""
day: str = (self.context.get("date") or "").strip() or _date.today().isoformat()
err = validate_slug(slug)
if err:
self._fail(err)
return None
daily_dir = self.app_context.app_config.daily_dir if self.app_context is not None else "daily"
path_rel = f"{daily_dir}/{day}/{slug}.md"
path_abs = (self.vault_path / path_rel).resolve()
if not path_abs.is_file():
self._fail(
f"note {path_rel} does not exist",
date=day,
slug=slug,
path=path_rel,
exists=False,
)
return None
try:
text = await read_file_safe(path_abs)
except Exception as e: # pylint: disable=broad-except
self._fail(f"read failed: {e}", date=day, slug=slug, path=path_rel)
return None
post = frontmatter.loads(text)
body = post.content
meta = dict(post.metadata or {})
self.context.response.success = True
self.context.response.answer = body
self.context.response.metadata.update(
{
"date": day,
"slug": slug,
"path": path_rel,
"exists": True,
"frontmatter": meta,
},
)
self.logger.info(
f"[{self.name}] read path={path_rel} " f"bytes={len(body.encode('utf-8'))} fm_keys={list(meta)}",
)
return self.context.response

View file

@ -1,131 +0,0 @@
"""``daily_write`` — write a daily note's full body + frontmatter; refresh day index.
Collapses the old ``daily_resolve`` → ``daily_create`` → ``file_write``
chain into a single call. Validates the slug, mkdirs the day folder,
writes body + frontmatter in one shot, refreshes ``daily/<date>.md``
index.
The ``overwrite`` flag picks between two behaviours:
* ``overwrite=False`` (default) — idempotent create: when the note
already exists returns ``{created: False, overwritten: False}``
without touching it (mirrors the old ``daily_resolve`` semantics).
Index still refreshes (siblings may have changed; cheap self-healing).
* ``overwrite=True`` — unconditional write. Preserves existing-file
encoding via ``detect_file_encoding`` (mirrors ``write_step``).
Frontmatter input is a dict; defaults to ``{name: slug}``. Empty /
None values are dropped (mirrors ``write_step``'s lenient frontmatter
handling). For partial frontmatter mutations on an existing note use
``frontmatter_update`` instead — ``daily_write`` is a full-file C/U.
"""
from datetime import date as _date
import frontmatter
from ._daily_io import refresh_day_index, validate_slug
from ..crud._file_io import detect_file_encoding, write_file_safe
from ..base_step import BaseStep
from ...components import R
@R.register("daily_write_step")
class DailyWriteStep(BaseStep):
"""Write ``daily/<date>/<slug>.md`` (create/overwrite); refresh day index."""
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
slug: str = self.context.get("slug", "") or ""
body: str = self.context.get("body", "") or ""
day: str = (self.context.get("date") or "").strip() or _date.today().isoformat()
overwrite: bool = bool(self.context.get("overwrite", False))
refresh_index: bool = bool(self.context.get("refresh_index", True))
fm_input = self.context.get("frontmatter")
if fm_input is None:
meta_in: dict = {"name": slug}
elif isinstance(fm_input, dict):
meta_in = dict(fm_input)
meta_in.setdefault("name", slug)
else:
self._fail(f"frontmatter must be a dict, got {type(fm_input).__name__}")
return None
err = validate_slug(slug)
if err:
self._fail(err)
return None
daily_dir = self.app_context.app_config.daily_dir if self.app_context is not None else "daily"
path_rel = f"{daily_dir}/{day}/{slug}.md"
path_abs = (self.vault_path / path_rel).resolve()
existed = path_abs.is_file()
# overwrite=False + file exists → idempotent skip (matches old daily_resolve semantics).
if existed and not overwrite:
payload: dict = {
"date": day,
"slug": slug,
"path": path_rel,
"created": False,
"overwritten": False,
}
if refresh_index:
payload["index"] = await refresh_day_index(self.file_store, day, daily_dir)
self.context.response.success = True
self.context.response.answer = f"Reused existing daily note {path_rel}"
self.context.response.metadata.update(payload)
return self.context.response
# Build the post. Drop empty / None values (write_step idiom).
clean_meta: dict = {}
for k, v in meta_in.items():
if v is None:
continue
if isinstance(v, str) and not v.strip():
continue
clean_meta[k] = v
post = frontmatter.Post(body, **clean_meta)
text = frontmatter.dumps(post)
if not text.endswith("\n"):
text += "\n"
# Preserve existing file encoding on overwrite; new files = UTF-8.
encoding = await detect_file_encoding(path_abs) if existed else "utf-8"
try:
await write_file_safe(path_abs, text, encoding=encoding)
except Exception as e: # pylint: disable=broad-except
self._fail(f"write failed: {e}", date=day, slug=slug, path=path_rel)
return None
payload = {
"date": day,
"slug": slug,
"path": path_rel,
"created": not existed,
"overwritten": existed,
}
if refresh_index:
payload["index"] = await refresh_day_index(self.file_store, day, daily_dir)
self.context.response.success = True
verb = "Wrote" if not existed else "Overwrote"
self.context.response.answer = f"{verb} daily note {path_rel}"
self.context.response.metadata.update(payload)
try:
nbytes = len(text.encode(encoding))
except (UnicodeEncodeError, LookupError):
nbytes = len(text.encode("utf-8"))
self.logger.info(
f"[{self.name}] wrote path={path_rel} bytes={nbytes} "
f"overwrite={overwrite} existed={existed} refresh_index={refresh_index}",
)
return self.context.response

View file

View file

@ -0,0 +1,492 @@
"""Shared filesystem helpers for CRUD steps.
Two related concerns, both private to the ``crud`` package:
1. **Generic file IO** — path gating, encoding-aware read/write, output
truncation (used by every CRUD step that touches the filesystem).
2. **Daily-note helpers** — slug validation + ``daily/<date>.md`` index
rebuild (used by the ``daily_*`` steps). The day index is a derived
rollup page auto-managed in marker-delimited sections; user-edited
manual sections are preserved verbatim across refreshes.
"""
import re
from pathlib import Path
from typing import Iterable
import aiofiles
import aiofiles.os
import frontmatter
from ...constants import DEFAULT_MAX_BYTES, MAX_FILE_READ_BYTES, TRUNCATION_NOTICE_MARKER
from ...utils import get_logger
logger = get_logger()
# ---------------------------------------------------------------------------
# Generic file IO
# ---------------------------------------------------------------------------
NON_MD_WARNING = (
"non-markdown file detected; CRUD operations are recommended on markdown files. "
"Operating in compatibility mode may carry risks of errors."
)
# Modern text formats — assume UTF-8 by convention.
_STANDARD_TEXT_EXTS = {
".md",
".py",
".js",
".ts",
".json",
".yaml",
".yml",
".html",
".css",
".xml",
".log",
".conf",
".ini",
".txt",
".sh",
}
# Legacy formats that may use ANSI/GBK on Chinese Windows systems.
_NON_STANDARD_EXTS = {".csv", ".bat", ".cmd", ".reg"}
# Path helpers
# ------------
# Filename validation. Windows is the strictest mainstream filesystem, so we
# validate to its bar — paths that pass here also work on macOS and Linux,
# and survive sync to a Windows machine or zip-and-share workflows.
_INVALID_CHARS = re.compile(r'[<>:"/\\|?*\x00-\x1f]')
_RESERVED_NAMES = {
"CON",
"PRN",
"AUX",
"NUL",
*(f"COM{i}" for i in range(1, 10)),
*(f"LPT{i}" for i in range(1, 10)),
}
def validate_filename_component(name: str, *, kind: str = "filename") -> str | None:
"""Return an error message, or ``None`` when ``name`` is a safe filename component.
A *component* is a single path segment — no ``/`` or ``\\`` allowed inside.
Used for both daily-note slugs and ``resolve_path`` per-component checks.
Rules:
- non-empty, no leading / trailing whitespace
- no reserved characters: ``< > : " / \\ | ? *`` or control chars (``\\x00-\\x1f``)
- no reserved device names: ``CON`` / ``PRN`` / ``AUX`` / ``NUL`` /
``COM1-9`` / ``LPT1-9`` (Windows reserves these with or without an
extension — ``CON.txt`` is also forbidden)
- no trailing ``.`` (also rejects ``..``, which doubles as path-traversal protection
for callers that validate per component)
``kind`` is the human-readable label inserted into error messages
(e.g. ``"slug"``, ``"path component"``).
"""
if not name:
return f"{kind} is required"
if name != name.strip():
return f"{kind} cannot have leading or trailing whitespace: {name!r}"
if _INVALID_CHARS.search(name):
return f'{kind} contains invalid characters (one of < > : " / \\ | ? * or a control char): {name!r}'
if name.endswith("."):
return f"{kind} cannot end with '.': {name!r}"
if name.split(".", 1)[0].upper() in _RESERVED_NAMES:
return f"{kind} is a Windows-reserved device name: {name!r}"
return None
def resolve_path(vault_path: Path, raw: str) -> tuple[Path | None, str | None]:
"""Resolve a `path=` argument against ``vault_path``.
Rules:
- Relative paths are joined under ``vault_path``.
- Absolute paths are accepted and returned as-is; a warning is logged
recommending relative paths, but the read still proceeds.
- Each path component is validated against the same Windows-strict
filename rules used for daily-note slugs — see
:func:`validate_filename_component`. ``..`` is rejected by the
trailing-``.`` rule, which doubles as path-traversal protection.
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 ``reme/steps/file_io/_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)
for part in p.parts:
if part == p.anchor:
continue
err = validate_filename_component(part, kind="path component")
if err:
return None, err
if p.is_absolute():
logger.info("absolute path detected, recommending relative paths")
return p, None
return vault_path / p, None
def gate_md(target: Path) -> tuple[Path, bool]:
"""Markdown gate with compatibility fallback.
Returns ``(path, is_md)``:
- No suffix → auto-append `.md`, ``is_md=True``.
- `.md` suffix → ``is_md=True``.
- Any other suffix → ``is_md=False`` (caller handles degraded mode).
"""
if target.suffix == "":
return target.with_suffix(".md"), True
if target.suffix.lower() != ".md":
return target, False
return target, True
# Encoding detection (private)
# ----------------------------
def _try_decode(data: bytes, encodings: Iterable[str]) -> tuple[str, str] | None:
"""Return ``(text, encoding)`` for the first encoding that decodes ``data`` cleanly."""
for enc in encodings:
try:
return data.decode(enc), enc
except (UnicodeDecodeError, LookupError):
continue
return None
def _decode_known_file(data: bytes, file_extension: str) -> tuple[str, str]:
"""Decode file bytes using the extension as a hint. Returns ``(text, encoding)``.
Strategy:
1. BOM-based detection.
2. Extension-driven defaults:
3. Last resort → UTF-8 with ``errors='replace'`` so the function never raises.
"""
if data.startswith(b"\xef\xbb\xbf"):
return data.decode("utf-8-sig"), "utf-8-sig"
if data.startswith((b"\xff\xfe", b"\xfe\xff")):
try:
return data.decode("utf-16"), "utf-16"
except UnicodeDecodeError:
pass
ext = (file_extension or "").lower()
if ext in _STANDARD_TEXT_EXTS:
try:
return data.decode("utf-8-sig"), "utf-8"
except UnicodeDecodeError:
pass # fall through
if ext in _NON_STANDARD_EXTS:
result = _try_decode(data, ("utf-8-sig", "gbk"))
if result is not None:
text, enc = result
return text, "utf-8" if enc == "utf-8-sig" else enc
# Unknown extension or earlier strategies failed.
return data.decode("utf-8", errors="replace"), "utf-8"
# File read / write
# -----------------
async def read_file_safe(file_path, max_bytes: int = MAX_FILE_READ_BYTES) -> tuple[str, str]:
"""Read file in byte mode and decode using extension-aware strategy.
Returns ``(text, encoding)``. Callers that need to write the file back
in its original encoding can pass ``encoding`` straight to
:func:`write_file_safe`, avoiding a second read via
:func:`detect_file_encoding`.
"""
stat = await aiofiles.os.stat(str(file_path))
read_size = min(stat.st_size, max_bytes)
async with aiofiles.open(str(file_path), "rb") as f:
data = await f.read(read_size)
return _decode_known_file(data, Path(file_path).suffix)
async def detect_file_encoding(file_path, sniff_bytes: int = 8192) -> str:
"""Detect the encoding of an existing file so writes can preserve it.
Reads up to ``sniff_bytes`` from the head of the file (enough for BOM
detection and statistical analysis). Falls back to ``utf-8`` if the file
is unreadable.
"""
try:
async with aiofiles.open(str(file_path), "rb") as f:
data = await f.read(sniff_bytes)
except Exception: # pylint: disable=broad-except
return "utf-8"
_, enc = _decode_known_file(data, Path(file_path).suffix)
return enc
async def write_file_safe(file_path: Path, content: str | bytes, encoding: str = "utf-8") -> None:
"""Write ``content`` to ``file_path`` in binary mode; creates parent dirs.
``str`` input is encoded with ``encoding`` (default UTF-8); callers wanting
to preserve a file's original encoding should pass the result of
:func:`detect_file_encoding`. If the requested ``encoding`` can't represent
some characters, falls back to UTF-8 to avoid data loss.
``bytes`` input is written verbatim — callers managing their own encoding
can pass raw bytes directly.
"""
file_path.parent.mkdir(parents=True, exist_ok=True)
if isinstance(content, str):
try:
payload = content.encode(encoding)
except (UnicodeEncodeError, LookupError):
logger.warning(
"write_file_safe: %r cannot encode all chars, falling back to utf-8",
encoding,
)
payload = content.encode("utf-8")
else:
payload = content
async with aiofiles.open(str(file_path), "wb") as f:
await f.write(payload)
# Output formatting
# -----------------
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
# ---------------------------------------------------------------------------
# Daily-note helpers: slug validation + day-index rebuild
# ---------------------------------------------------------------------------
# Slug validation
# ---------------
def validate_slug(slug: str) -> str | None:
"""Validate a daily-note slug. Thin wrapper over :func:`validate_filename_component`."""
return validate_filename_component(slug, kind="slug")
# Day-index rebuild
# -----------------
# The day index is a derived artifact whose single job is daily-note
# consolidation — its source of truth lives in each note's
# frontmatter. The rebuild refreshes the auto-managed notes block
# while preserving any user content sitting outside the markers.
#
# Frontmatter shape — only the two reserved fields:
# name: <date>
# description: <one-line note-count digest>
#
# The note inventory lives in the body's ``<!-- notes:auto -->``
# block: each note becomes a single line with its full frontmatter
# inlined (``- [[path]] name: ... description: ... <other keys>``),
# letting an agent scan the day at a glance. Content outside the
# auto markers is preserved verbatim across refreshes.
_NOTES_OPEN = "<!-- notes:auto -->"
_NOTES_CLOSE = "<!-- /notes:auto -->"
_NOTES_BLOCK_RE = re.compile(
rf"{re.escape(_NOTES_OPEN)}(?P<inner>.*?){re.escape(_NOTES_CLOSE)}",
re.DOTALL,
)
def _wrap_notes_block(inner: str) -> str:
return f"{_NOTES_OPEN}\n{inner}\n{_NOTES_CLOSE}"
# Content rendering
# -----------------
def _render_notes_block(notes: list[dict]) -> str:
"""Render each note as a single line with its full frontmatter inlined.
Format: ``- [[path]] key1: value1 key2: value2 ...``. ``name`` and
``description`` lead (when present) so columns line up across notes;
remaining keys follow in frontmatter insertion order. Empty / None
values are skipped; newlines in values collapse to spaces so the
single-line invariant holds.
"""
if not notes:
return "(none)"
lines: list[str] = []
for note in notes:
meta: dict = note["metadata"]
ordered_keys = [k for k in ("name", "description") if k in meta]
ordered_keys += [k for k in meta if k not in ("name", "description")]
parts = [f"- [[{note['path']}]]"]
for key in ordered_keys:
value = meta[key]
if value is None or value == "":
continue
value_str = str(value).replace("\r\n", " ").replace("\r", " ").replace("\n", " ")
parts.append(f"{key}: {value_str}")
lines.append(" ".join(parts))
return "\n".join(lines)
# Body manipulation
# -----------------
def _replace_or_append_notes(body: str, fresh_block: str) -> str:
"""Replace an existing notes auto block in-place; append at end if absent."""
if _NOTES_BLOCK_RE.search(body):
replacement = _wrap_notes_block(fresh_block)
return _NOTES_BLOCK_RE.sub(lambda m: replacement, body, count=1)
suffix = _wrap_notes_block(fresh_block)
return f"{body.rstrip()}\n\n{suffix}\n" if body.strip() else f"{suffix}\n"
# Public scan + rebuild
# ---------------------
def scan_notes(vault_dir: Path, date: str, daily_dir: str) -> list[dict]:
"""Walk ``<daily_dir>/<date>/*.md`` and pull each note's frontmatter.
Returns one dict per note::
{"slug": str, "path": str, "metadata": dict}
``metadata`` is the raw frontmatter dict (insertion-ordered);
consumers decide which keys to surface. Each ``.md`` directly
under the day folder is a note; the file's stem is the slug.
"""
date_dir = vault_dir / daily_dir / date
if not date_dir.is_dir():
return []
out: list[dict] = []
for md_path in sorted(p for p in date_dir.iterdir() if p.is_file() and p.suffix == ".md"):
slug = md_path.stem
try:
post = frontmatter.loads(md_path.read_text(encoding="utf-8"))
except Exception: # pylint: disable=broad-except
continue
out.append(
{
"slug": slug,
"path": f"{daily_dir}/{date}/{slug}.md",
"metadata": dict(post.metadata or {}),
},
)
return out
async def refresh_day_index(file_store, date: str, daily_dir: str) -> dict:
"""Rebuild ``<daily_dir>/<date>.md`` from the current state of its notes.
Behaviour:
* No ``<daily_dir>/<date>/`` at all and no existing index file → no-op.
* Notes present → write the index file (create if missing,
otherwise refresh the auto block in place, preserving content
outside the markers, refresh frontmatter).
* Notes directory empty but index file exists → rebuild with
empty auto block (keeps the file in sync with reality).
Returns ``{date, path, notes, created}``. Each row in ``notes`` is
``{path, slug, metadata}`` with the raw frontmatter dict.
"""
vault_dir = Path(file_store.vault_path or ".").resolve()
index_rel = f"{daily_dir}/{date}.md"
index_abs = vault_dir / index_rel
notes = scan_notes(vault_dir, date, daily_dir)
notes_payload = [{"path": n["path"], "slug": n["slug"], "metadata": n["metadata"]} for n in notes]
if not notes and not index_abs.is_file():
return {
"date": date,
"path": index_rel,
"notes": notes_payload,
"created": False,
}
notes_block = _render_notes_block(notes)
n = len(notes)
fm = {"name": date, "description": "No notes today." if n == 0 else f"{n} note(s) today."}
if index_abs.is_file():
post = frontmatter.loads(index_abs.read_text(encoding="utf-8"))
new_body = _replace_or_append_notes(post.content, notes_block)
merged = dict(post.metadata or {})
for key, value in fm.items():
if not merged.get(key):
merged[key] = value
fm = merged
was_created = False
else:
index_abs.parent.mkdir(parents=True, exist_ok=True)
new_body = _wrap_notes_block(notes_block) + "\n"
was_created = True
out = frontmatter.Post(new_body, **fm)
index_abs.write_text(frontmatter.dumps(out), encoding="utf-8")
return {
"date": date,
"path": index_rel,
"notes": notes_payload,
"created": was_created,
}

View file

@ -0,0 +1,94 @@
"""``daily_create`` — provision a note slug under a daily folder: ``daily/<date>/<slug>.md``.
Minimal slug provisioner. Validates the slug, mkdirs the day folder,
writes an empty-body note with frontmatter ``{name: slug}`` if (and
only if) the file does not already exist, refreshes the day index,
and returns the vault-relative path.
Idempotent: when the note already exists this is a no-op write (the
day index still refreshes — siblings may have changed; cheap
self-healing). The caller fills the body via ``file_write`` /
``file_edit`` / ``file_append`` (or a native editor); ``daily_create``
deliberately does not accept a body.
Inputs:
slug (required, validated) — the note's name (also the file stem)
date (optional, ``YYYY-MM-DD``; empty = today)
Outputs:
answer = one-line human-readable status
metadata = {date, slug, path, created, index?}
"""
from datetime import date as _date
from pathlib import Path
import frontmatter
from ._file_io import refresh_day_index, validate_slug, write_file_safe
from ..base_step import BaseStep
from ...components import R
@R.register("daily_create_step")
class DailyCreateStep(BaseStep):
"""Provision ``daily/<date>/<slug>.md`` (idempotent); refresh day index."""
def _fail(self, message: str, **meta) -> None:
"""Mark response failed; copy ``meta`` into ``response.metadata``."""
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)
def _collect_params(self) -> tuple[str, str, str]:
"""Read ``slug`` + ``date`` from context; default ``date`` today, ``daily_dir`` from app config."""
assert self.context is not None
slug = self.context.get("slug", "")
day = self.context.get("date", "") or _date.today().strftime("%Y-%m-%d")
daily_dir = self.app_context.app_config.daily_dir if self.app_context is not None else "daily"
return slug, day, daily_dir
@staticmethod
def _empty_note_text(slug: str) -> str:
"""Serialize an empty-body markdown note with frontmatter ``{name: slug}``; trailing newline guaranteed."""
text = frontmatter.dumps(frontmatter.Post("", name=slug))
return text if text.endswith("\n") else text + "\n"
async def _create_if_missing(self, path_abs: Path, slug: str) -> bool:
"""Write the empty note only when the file is absent. Returns ``True`` iff a new file was created."""
if path_abs.is_file():
return False
await write_file_safe(path_abs, self._empty_note_text(slug), encoding="utf-8")
return True
def _set_success(self, payload: dict, created: bool) -> None:
"""Stamp the response with success + human-readable answer + metadata payload."""
assert self.context is not None
self.context.response.success = True
self.context.response.answer = f"{'Created' if created else 'Reused existing'} daily note {payload['path']}"
self.context.response.metadata.update(payload)
async def execute(self):
"""Validate the slug, provision the note file, refresh the day index, stamp the response."""
assert self.context is not None
slug, day, daily_dir = self._collect_params()
err = validate_slug(slug)
if err:
self._fail(err)
return None
path_rel = f"{daily_dir}/{day}/{slug}.md"
path_abs = (self.vault_path / path_rel).resolve()
try:
created = await self._create_if_missing(path_abs, slug)
except Exception as e: # pylint: disable=broad-except
self._fail(f"create failed: {e}", date=day, slug=slug, path=path_rel)
return None
index = await refresh_day_index(self.file_store, day, daily_dir)
self._set_success({"date": day, "slug": slug, "path": path_rel, "created": created, "index": index}, created)
self.logger.info(f"[{self.name}] {'created' if created else 'reused'} path={path_rel}")
return self.context.response

View file

@ -1,23 +1,22 @@
"""``daily_list`` — list the notes under a single day (pure read, no side effects).
Returns one row per ``daily/<date>/<slug>.md`` note file with its
vault-relative ``path`` plus ``slug`` / ``name`` / ``description``
from frontmatter. Sorted by slug for stable output.
vault-relative ``path``, ``slug``, and the raw ``metadata`` dict
(full frontmatter). Sorted by slug for stable output.
**Does NOT refresh** ``daily/<date>.md`` — call ``daily_reindex``
explicitly when the index page needs to be rebuilt. Decoupling
read from write keeps each step's effect predictable.
Input is a single optional ``date`` (ISO ``YYYY-MM-DD``); falls back
Input is a single optional ``date`` (``YYYY-MM-DD``); falls back
to today.
"""
from datetime import date as _date
from pathlib import Path
from ._daily_io import scan_notes
from ._file_io import scan_notes
from ..base_step import BaseStep
from ...components import R
@ -25,23 +24,24 @@ from ...components import R
class DailyListStep(BaseStep):
"""List the notes under a single day. Pure read — no index refresh."""
async def execute(self):
def _collect_params(self) -> tuple[str, str, Path]:
"""Read ``date`` (default today, ``YYYY-MM-DD``), resolve ``daily_dir``, locate the vault root on disk."""
assert self.context is not None
day: str = (self.context.get("date") or "").strip() or _date.today().isoformat()
day = self.context.get("date", "") or _date.today().strftime("%Y-%m-%d")
daily_dir = self.app_context.app_config.daily_dir if self.app_context is not None else "daily"
vault_dir = Path(self.file_store.vault_path or ".").resolve()
return day, daily_dir, vault_dir
scanned = scan_notes(vault_dir, day, daily_dir)
notes = [
{
"path": n["path"],
"slug": n["slug"],
"name": n["name"],
"description": n["description"],
}
for n in scanned
]
@staticmethod
def _project(note: dict) -> dict:
"""Keep only the user-facing keys (drop internal scan_notes fields, if any)."""
return {"path": note["path"], "slug": note["slug"], "metadata": note["metadata"]}
async def execute(self):
"""Scan ``<daily_dir>/<date>/`` and emit one projected record per note."""
assert self.context is not None
day, daily_dir, vault_dir = self._collect_params()
notes = [self._project(n) for n in scan_notes(vault_dir, day, daily_dir)]
self.context.response.success = True
self.context.response.answer = f"Listed {len(notes)} note(s) for {day}"
self.context.response.metadata.update({"date": day, "notes": notes})

View file

@ -2,18 +2,18 @@
The day index ``daily/<date>.md`` is a derived artifact whose job is to
list and describe every note file under ``daily/<date>/``. It is auto-
refreshed by ``daily_write`` after every body write. Generic ops like
``file_write`` / ``file_append`` / ``frontmatter_update`` leave it
stale — this step is the standalone writer to call after batch flows
(historical backfill, drift recovery, end-of-batch consolidation, or
a ``frontmatter_update`` that touched ``name`` / ``description``).
refreshed by ``daily_create``. Generic ops like ``file_write`` /
``file_append`` / ``frontmatter_update`` leave it stale — this step
is the standalone writer to call after batch flows (historical
backfill, drift recovery, end-of-batch consolidation, or a
``frontmatter_update`` that touched ``name`` / ``description``).
This is the **write view**: it reports the index-page path and a
``created`` flag (true when the file was just emitted for the first
time), which is what a caller running a rebuild wants to confirm. For
the per-note inventory use ``daily_list``.
Input is a single optional ``date`` (ISO ``YYYY-MM-DD``); falls back to
Input is a single optional ``date`` (``YYYY-MM-DD``); falls back to
today.
Always idempotent and safe to re-run.
@ -21,9 +21,8 @@ Always idempotent and safe to re-run.
from datetime import date as _date
from ._daily_io import refresh_day_index
from ._file_io import refresh_day_index
from ..base_step import BaseStep
from ...components import R
@ -31,11 +30,16 @@ from ...components import R
class DailyReindexStep(BaseStep):
"""Rebuild ``daily/<date>.md`` from the current state of its notes."""
async def execute(self):
def _collect_params(self) -> tuple[str, str]:
"""Read ``date`` (default today) and ``daily_dir`` (default ``daily``) from context/app config."""
assert self.context is not None
day: str = (self.context.get("date") or "").strip() or _date.today().isoformat()
day = self.context.get("date", "") or _date.today().strftime("%Y-%m-%d")
daily_dir = self.app_context.app_config.daily_dir if self.app_context is not None else "daily"
refreshed = await refresh_day_index(self.file_store, day, daily_dir)
return day, daily_dir
def _apply_result(self, refreshed: dict) -> None:
"""Mirror the rebuild outcome (surfaced error or success payload) onto the response object."""
assert self.context is not None
if "error" in refreshed:
self.context.response.success = False
self.context.response.answer = f"Error: {refreshed['error']}"
@ -52,3 +56,8 @@ class DailyReindexStep(BaseStep):
"notes_count": notes_count,
},
)
async def execute(self):
"""Trigger the index rebuild and stamp the response."""
day, daily_dir = self._collect_params()
self._apply_result(await refresh_day_index(self.file_store, day, daily_dir))

View file

@ -25,9 +25,8 @@ import shutil
from pathlib import Path
from ..base_step import BaseStep
from ...utils.wikilink_handler import WikilinkHandler
from ...components import R
from ...utils.wikilink_handler import WikilinkHandler
def _is_inside(rel: str, folder_rel: str) -> bool:

View file

@ -2,7 +2,7 @@
import frontmatter
from ._file_io import NON_MD_WARNING, detect_file_encoding, gate_md, read_file_safe, resolve_path, write_file_safe
from ._file_io import NON_MD_WARNING, gate_md, read_file_safe, resolve_path, write_file_safe
from ..base_step import BaseStep
from ...components import R
@ -52,7 +52,7 @@ class EditStep(BaseStep):
return None
try:
raw_text = await read_file_safe(target)
raw_text, encoding = await read_file_safe(target)
except Exception as e: # pylint: disable=broad-except
self._fail(f"read failed: {e}", path=str(target))
return None
@ -87,9 +87,8 @@ class EditStep(BaseStep):
else:
new_text = new_body
# Preserve the file's original encoding so edits don't silently re-encode
# non-UTF-8 files (e.g. GBK CSV) to UTF-8.
encoding = await detect_file_encoding(target)
# Preserve the file's original encoding (returned by read_file_safe above)
# so edits don't silently re-encode non-UTF-8 files (e.g. GBK CSV) to UTF-8.
try:
await write_file_safe(target, new_text, encoding=encoding)
except Exception as e: # pylint: disable=broad-except

View file

@ -14,13 +14,12 @@ from pathlib import Path
import frontmatter
from ..base_step import BaseStep
from ...components import R
@R.register("frontmatter_delete_step")
class FrontmatterDeleteStep(BaseStep):
"""Remove keys from a markdown file's frontmatter."""
"""Remove keys from a Markdown file's frontmatter."""
async def execute(self):
assert self.context is not None

View file

@ -13,13 +13,12 @@ from pathlib import Path
import frontmatter
from ..base_step import BaseStep
from ...components import R
@R.register("frontmatter_read_step")
class FrontmatterReadStep(BaseStep):
"""Read a markdown file's frontmatter (YAML metadata only)."""
"""Read a Markdown file's frontmatter (YAML metadata only)."""
async def execute(self):
assert self.context is not None

View file

@ -17,7 +17,6 @@ from pathlib import Path
import frontmatter
from ..base_step import BaseStep
from ...components import R

104
reme4/steps/file_io/list.py Normal file
View file

@ -0,0 +1,104 @@
"""``file_list`` — enumerate files under a directory in the vault.
Reads directly from the filesystem (``Path.iterdir`` / ``Path.rglob``),
**not** the file_store index. The store may lag behind disk during
indexing or after rapid mutations; the on-disk walk is the source of truth.
Parameters:
path — dir to list under (relative to the vault or absolute). Empty = vault root.
limit — cap on the number of returned items (default 100, must be > 0).
recursive — descend into subdirectories. Default False = direct children only.
No frontmatter is read. Callers needing frontmatter-based filtering
should iterate the result and call ``frontmatter_read`` per candidate.
"""
from pathlib import Path
from typing import Iterable
from ..base_step import BaseStep
from ...components import R
# Default cap on returned items so huge vaults don't blow up the response.
DEFAULT_LIMIT = 100
@R.register("list_step")
class ListStep(BaseStep):
"""Enumerate files under a directory in the vault."""
def _fail(self, message: str, **meta) -> None:
"""Set a failed response (matches the read/edit/... fail envelope)."""
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)
def _collect_params(self) -> tuple[str, bool, int]:
"""Read ``path`` / ``recursive`` / ``limit`` from context; coerce permissively."""
assert self.context is not None
path = str(self.context.get("path") or "")
recursive = bool(self.context.get("recursive", False))
raw_limit = self.context.get("limit")
# Strings like "50" are accepted; bad/non-positive values fall back to default.
try:
limit = int(raw_limit) if raw_limit is not None else DEFAULT_LIMIT
except (TypeError, ValueError):
limit = DEFAULT_LIMIT
return path, recursive, limit if limit > 0 else DEFAULT_LIMIT
@staticmethod
def _resolve_target_dir(vault_dir: Path, path: str) -> Path:
"""Empty → vault root; absolute → as-is; relative → joined under vault_dir."""
if not path:
return vault_dir
candidate = Path(path)
return candidate.resolve() if candidate.is_absolute() else (vault_dir / candidate).resolve()
@staticmethod
def _walk_files(target_dir: Path, recursive: bool, limit: int) -> list[Path]:
"""Return up to ``limit`` regular files under ``target_dir``; short-circuits at the cap."""
entries: Iterable[Path] = target_dir.rglob("*") if recursive else target_dir.iterdir()
files: list[Path] = []
for entry in entries:
if not entry.is_file(): # skip dirs, sockets, broken links, etc.
continue
files.append(entry)
if len(files) >= limit:
break
return files
@staticmethod
def _format_relative(files: list[Path], vault_dir: Path) -> list[str]:
"""Render as vault-relative paths; fall back to absolute when outside the vault."""
out: list[str] = []
for entry in files:
try:
out.append(str(entry.relative_to(vault_dir)))
except ValueError:
out.append(str(entry))
return out
async def execute(self):
assert self.context is not None
path, recursive, limit = self._collect_params()
vault_dir = Path(self.file_store.vault_path or ".").resolve()
target_dir = self._resolve_target_dir(vault_dir, path)
if not target_dir.exists():
self._fail(f"directory {target_dir} does not exist", path=str(target_dir))
return None
if not target_dir.is_dir():
self._fail(f"path {target_dir} is not a directory", path=str(target_dir))
return None
items = self._format_relative(self._walk_files(target_dir, recursive, limit), vault_dir)
self.context.response.success = True
self.context.response.answer = f"Listed {len(items)} file(s) under {path or '.'}"
self.context.response.metadata.update({"items": items, "count": len(items)})
self.logger.info(
f"[{self.name}] listed dir={target_dir} recursive={recursive} count={len(items)} limit={limit}",
)
return self.context.response

View file

@ -33,9 +33,8 @@ import shutil
from pathlib import Path
from ..base_step import BaseStep
from ...utils.wikilink_handler import WikilinkHandler
from ...components import R
from ...utils.wikilink_handler import WikilinkHandler
@R.register("move_step")
@ -44,8 +43,8 @@ class MoveStep(BaseStep):
async def execute(self):
assert self.context is not None
src_path: str = self.context.get("src_path", "") or ""
dst_path: str = self.context.get("dst_path", "") or ""
src_path: str = self.context.get("src_path", "")
dst_path: str = self.context.get("dst_path", "")
overwrite: bool = bool(self.context.get("overwrite", False))
retarget: bool = bool(self.context.get("retarget", True))
assert src_path and dst_path, "src_path and dst_path are required"

115
reme4/steps/file_io/read.py Normal file
View file

@ -0,0 +1,115 @@
"""Read a markdown file from vault_dir, with line-range slicing and byte-truncation."""
from pathlib import Path
from ._file_io import NON_MD_WARNING, gate_md, read_file_safe, resolve_path, 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:
"""Mark the response failed and stash a human-readable error."""
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)
def _resolve_target(self, raw: str) -> Path | None:
"""Resolve ``raw`` under vault and gate the markdown suffix.
Non-md suffixes only warn (compatibility mode), not fail. Returns
the absolute path, or ``None`` when ``raw`` is empty/invalid.
"""
target, err = resolve_path(self.vault_path, raw)
if err:
self._fail(err)
return None
target, is_md = gate_md(target)
if not is_md:
self.logger.info(f"[{self.name}] {NON_MD_WARNING} path={target}")
return target
def _validate_line_args(self, start_line, end_line) -> bool:
"""Accept ``None`` or any value that parses via ``int()`` (JSON/CLI often stringify)."""
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 False
return True
def _check_file(self, target: Path) -> bool:
"""Confirm ``target`` exists and is a regular file."""
if not target.exists():
self._fail(f"file {target} does not exist", path=str(target))
return False
if not target.is_file():
self._fail(f"path {target} is not a file", path=str(target))
return False
return True
def _resolve_range(self, total: int, start_line, end_line, target: Path) -> tuple[int, int] | None:
"""Normalize 1-based inclusive ``[s, e]``; reject past-EOF or inverted ranges."""
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
return s, e
async def _load_content(self, target: Path) -> str | None:
"""Read via the encoding-aware helper; convert exceptions to ``_fail``."""
try:
content, _ = await read_file_safe(target)
return content
except Exception as e: # pylint: disable=broad-except
self._fail(f"read failed: {e}", path=str(target))
return None
async def execute(self):
assert self.context is not None
raw = str(self.context.get("path") or "")
start_line, end_line = self.context.get("start_line"), self.context.get("end_line")
# Validate inputs and target before touching the filesystem twice.
target = self._resolve_target(raw)
if target is None:
return None
if not self._validate_line_args(start_line, end_line):
return None
if not self._check_file(target):
return None
content = await self._load_content(target)
if content is None:
return None
all_lines = content.split("\n")
total = len(all_lines)
bounds = self._resolve_range(total, start_line, end_line, target)
if bounds is None:
return None
s, e = bounds
text = truncate_text_output(
"\n".join(all_lines[s - 1 : e]),
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

View file

@ -25,7 +25,6 @@ from pathlib import Path
import frontmatter
from ..base_step import BaseStep
from ...components import R

View file

@ -1,24 +0,0 @@
"""Frontmatter steps — RUD on the YAML frontmatter slice of a markdown file.
Three Steps:
frontmatter_read_step — return the frontmatter dict
frontmatter_update_step — merge a patch into the frontmatter
frontmatter_delete_step — drop the listed keys
Operates only on ``.md`` files; non-markdown targets get
``error="not markdown"`` and the call is **not** executed.
Body content stays untouched; the sibling ``file`` package owns every
other file-level surface — opaque-byte ops (list / stat / move /
delete / upload / download) and whole-file text ops (read / write /
append / edit). For mid-body edits, use ``file.edit`` (exact string
replacement) or do a read + write round-trip.
Each Step here is a pure disk read-modify-write — the watcher / parser
notices the change and refreshes the projections asynchronously.
"""
from . import read # noqa: F401 -- @R.register("frontmatter_read_step")
from . import update # noqa: F401 -- @R.register("frontmatter_update_step")
from . import delete # noqa: F401 -- @R.register("frontmatter_delete_step")

View file

@ -1,7 +0,0 @@
"""Graph steps."""
from .traverse import GraphTraverseStep
__all__ = [
"GraphTraverseStep",
]

View file

@ -1,119 +0,0 @@
"""``graph_traverse_step`` — BFS over wikilink edges from a seed file.
Single tool for relationship browsing. ``depth=1`` covers the trivial
"what does this link to / what links here" lookups (set ``direction``
accordingly); higher depth opens up multi-hop exploration.
Output is one record per edge traversed (not per node), so the same
target can appear multiple times if reached via different predicates
or paths — agents dedupe at the call site if they want a flat node
set. Each record carries ``via`` (the predecessor) and the link's
``predicate`` / ``anchor`` so the agent can reconstruct the path.
Adjacency is loaded once via ``file_graph.get_nodes(None)`` — every
real node arrives with its full ``links`` payload, and we build both
the outbound and the inbound index in a single pass. The BFS then
runs purely in memory: no per-frontier-node graph round-trips, no
filesystem walk. The ``get_inlinks`` / ``get_outlinks`` contract
methods stay unused here because they'd add network round-trips for
data we already have.
Direction vocabulary accepts both the standard convention
(``forward`` / ``backward`` / ``both``) and the engine convention
(``out`` / ``in`` / ``both``).
The seed ``path`` is taken as-is (vault-relative). A seed that doesn't
match any graph node yields an empty result (no error).
"""
from collections import deque
from ..base_step import BaseStep
from ...components import R
from ...schema import FileLink
_FORWARD = {"out", "forward"}
_BACKWARD = {"in", "backward"}
_BOTH = {"both"}
_VALID_DIRECTIONS = _FORWARD | _BACKWARD | _BOTH
@R.register("graph_traverse_step")
class GraphTraverseStep(BaseStep):
"""BFS from a seed file to explore wikilink relationships.
Parameters:
path — seed path (vault-relative).
direction — ``forward`` / ``backward`` / ``both`` (or ``out`` / ``in`` / ``both``).
depth — hop limit (default 1 = immediate neighbors).
predicate — optional edge-type filter; ``None`` = no filter.
"""
async def execute(self):
"""BFS from ``path`` and emit one record per traversed edge."""
assert self.context is not None
seed = str(self.context.get("path") or "").strip()
assert seed, "path is required"
max_depth = int(self.context.get("depth") or 1)
direction = (self.context.get("direction") or "both").lower()
predicate = self.context.get("predicate")
assert (
direction in _VALID_DIRECTIONS
), f"direction must be one of {sorted(_VALID_DIRECTIONS)}, got {direction!r}"
# Build outbound / inbound adjacency in one pass over all nodes.
outbound: dict[str, list[tuple[str, FileLink]]] = {}
inbound: dict[str, list[tuple[str, FileLink]]] = {}
if self.file_store.file_graph:
for node in await self.file_store.file_graph.get_nodes():
for link in node.links:
if not link.target_path:
continue
outbound.setdefault(node.path, []).append((link.target_path, link))
inbound.setdefault(link.target_path, []).append((node.path, link))
walk_out = direction in _FORWARD or direction in _BOTH
walk_in = direction in _BACKWARD or direction in _BOTH
visited_edges: set[tuple[str, str, str | None]] = set()
results: list[dict] = []
queue: deque[tuple[str, int]] = deque([(seed, 0)])
while queue:
current, depth = queue.popleft()
if depth >= max_depth:
continue
edges: list[tuple[str, str | None, str | None]] = []
if walk_out:
for tgt, link in outbound.get(current, ()):
if predicate is not None and link.predicate != predicate:
continue
edges.append((tgt, link.predicate, link.target_anchor))
if walk_in:
for src, link in inbound.get(current, ()):
if predicate is not None and link.predicate != predicate:
continue
edges.append((src, link.predicate, link.target_anchor))
for next_path, pred, anchor in edges:
edge_key = (current, next_path, pred)
if edge_key in visited_edges:
continue
visited_edges.add(edge_key)
results.append(
{
"path": next_path,
"depth": depth + 1,
"via": current,
"predicate": pred,
"anchor": anchor,
},
)
if depth + 1 < max_depth:
queue.append((next_path, depth + 1))
self.context.response.success = True
self.context.response.answer = f"Traversed {len(results)} edge(s) from {seed}"
self.context.response.metadata.update({"edges": results, "count": len(results)})

View file

View file

@ -0,0 +1,30 @@
"""Wipe the file store and emit every vault file as an ``added`` change.
Designed to be chained before ``update_index_step`` so that the second step
performs the actual re-indexing and persistence.
"""
from ..base_step import BaseStep
from ...components import R
@R.register("clear_and_scan_step")
class ClearAndScanStep(BaseStep):
"""Clear the file store, walk the vault, and write changes into the context."""
async def execute(self):
assert self.context is not None
suffixes = tuple("." + s.strip(".") for s in self.context.get("suffix_filters", ["md"]))
await self.file_store.clear()
paths = [
str(p.absolute())
for p in self.vault_path.rglob("*")
if p.is_file() and (not suffixes or str(p).endswith(suffixes))
]
self.context["changes"] = [{"change": "added", "path": p} for p in paths]
counts = {"added": len(paths), "modified": 0, "deleted": 0}
self.context.response.metadata["counts"] = counts
self.logger.info(f"[{self.name}] cleared store and scanned {len(paths)} file(s)")
return self.context.response

View file

@ -1,4 +1,8 @@
"""Initial sync: diff watch_paths vs file_store, then index the diff."""
"""One-shot scan: diff watch_paths vs file_store and write changes into context.
Designed to be chained before ``update_index_step`` so that the second step
performs the actual writes and persistence.
"""
from pathlib import Path
@ -6,14 +10,13 @@ from ..base_step import BaseStep
from ...components import R
@R.register("update_store_step")
class UpdateStoreStep(BaseStep):
"""One-shot sync: compute added/modified/deleted vs file_store and index."""
@R.register("scan_changes_step")
class ScanChangesStep(BaseStep):
"""One-shot scan: compute added/modified/deleted vs file_store and write to context."""
def __init__(self, recursive: bool = True, dump: bool = True, **kwargs):
def __init__(self, recursive: bool = True, **kwargs):
super().__init__(**kwargs)
self.recursive: bool = recursive
self.dump: bool = dump
async def execute(self):
assert self.context is not None
@ -54,11 +57,9 @@ class UpdateStoreStep(BaseStep):
)
counts = {"added": len(to_add), "modified": len(to_modify), "deleted": len(to_delete)}
self.context["changes"] = changes
if changes:
self.logger.info(f"[{self.name}] initial sync: {counts}")
await self.run_job("index_changes", changes=changes)
if self.dump:
await self.file_store.dump()
self.logger.info(f"[{self.name}] scan: {counts}")
else:
self.logger.info(f"[{self.name}] store is up to date")

View file

@ -4,7 +4,8 @@ import asyncio
from ..base_step import BaseStep
from ...components import R
from ...schema import FileChunk, FileLink, FileNode
from ...schema import FileChunk
from ...utils import expand_links, render_expansion_lines
_RRF_K = 60
_MAX_CANDIDATES = 200
@ -58,105 +59,6 @@ class SearchStep(BaseStep):
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 (name/description) from a FileNode."""
if node is None:
return {}
fm = node.front_matter
meta: dict = {}
if fm.name:
meta["name"] = fm.name
if fm.description:
meta["description"] = fm.description
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 "name" in meta:
parts.append(f'name="{meta["name"]}"')
if "description" in meta:
parts.append(f'description="{meta["description"]}"')
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()
@ -164,7 +66,7 @@ class SearchStep(BaseStep):
min_score: float = float(self.context.get("min_score", 0.0))
vector_weight: float = float(self.kwargs.get("vector_weight", 0.7))
candidate_multiplier: float = float(self.kwargs.get("candidate_multiplier", 3.0))
expand_links: bool = bool(self.kwargs.get("expand_links", True))
expand_links_enabled: bool = bool(self.kwargs.get("expand_links", True))
max_links_per_direction: int = int(self.kwargs.get("max_links_per_direction", 10))
if not query:
@ -203,7 +105,7 @@ class SearchStep(BaseStep):
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 {}
await expand_links(self.file_store, unique_paths, max_links_per_direction) if expand_links_enabled else {}
)
answer_lines: list[str] = []
@ -212,7 +114,7 @@ class SearchStep(BaseStep):
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, {})))
answer_lines.extend(render_expansion_lines(link_expansion.get(c.path, {})))
self.context.response.answer = "\n".join(answer_lines)
self.context.response.metadata["results"] = [

View file

@ -0,0 +1,111 @@
"""BFS over wikilink edges from one or more seed files.
One record per traversed *edge* (not per node): the same target can repeat
if reached via different predicates or paths. Each record carries the
predecessor plus the link's predicate/anchor so callers can reconstruct
the path. Adjacency is built once via a single ``file_store.get_nodes()``
call — BFS then runs purely in memory with no per-frontier round-trips.
"""
from collections import deque
from pathlib import Path
from ..base_step import BaseStep
from ...components import R
from ...schema import FileLink
_OUT = {"out", "forward", "both"}
_IN = {"in", "backward", "both"}
_VALID = _OUT | _IN
# source path -> list of (neighbor path, link)
Adjacency = dict[str, list[tuple[str, FileLink]]]
async def _build_adjacency(file_store) -> tuple[Adjacency, Adjacency]:
"""Single ``get_nodes()`` pass → (outbound, inbound) adjacency maps.
Inbound stores the source path next to each link so BFS can attribute
inbound edges back to their origin — ``get_inlinks`` alone returns
target-shaped FileLinks without source attribution.
"""
outbound: Adjacency = {}
inbound: Adjacency = {}
for node in await file_store.get_nodes():
for link in node.links:
if link.target_path:
outbound.setdefault(node.path, []).append((link.target_path, link))
inbound.setdefault(link.target_path, []).append((node.path, link))
return outbound, inbound
def _bfs(
seeds: list[str],
max_depth: int,
direction: str,
outbound: Adjacency,
inbound: Adjacency,
) -> list[dict]:
"""In-memory BFS; emits one record per unique (src, dst, predicate) edge."""
sources: list[Adjacency] = []
if direction in _OUT:
sources.append(outbound)
if direction in _IN:
sources.append(inbound)
visited: set[tuple[str, str, str | None]] = set()
results: list[dict] = []
queue: deque[tuple[str, int]] = deque((s, 0) for s in seeds)
while queue:
current, depth = queue.popleft()
if depth >= max_depth:
continue
for src in sources:
for next_path, link in src.get(current, ()):
key = (current, next_path, link.predicate)
if key in visited:
continue
visited.add(key)
results.append(
{
"path": next_path,
"depth": depth + 1,
"via": current,
"predicate": link.predicate,
"anchor": link.target_anchor,
},
)
if depth + 1 < max_depth:
queue.append((next_path, depth + 1))
return results
@R.register("traverse_step")
class TraverseStep(BaseStep):
"""BFS from one or more seed files to explore wikilink relationships.
Parameters:
path — single seed (str) or list of seeds (vault-relative).
direction — ``forward`` / ``backward`` / ``both`` (or ``out`` / ``in`` / ``both``).
depth — hop limit (default 1 = immediate neighbors).
"""
async def execute(self):
assert self.context is not None
raw = self.context.get("path")
items = [raw] if isinstance(raw, (str, Path)) else list(raw or [])
seeds = [str(p) for p in items if p]
assert seeds, "path is required"
depth = int(self.context.get("depth") or 1)
direction = (self.context.get("direction") or "both").lower()
assert direction in _VALID, f"direction must be one of {sorted(_VALID)}, got {direction!r}"
outbound, inbound = await _build_adjacency(self.file_store)
results = _bfs(seeds, depth, direction, outbound, inbound)
label = seeds[0] if len(seeds) == 1 else f"{len(seeds)} seeds"
self.context.response.success = True
self.context.response.answer = f"Traversed {len(results)} edge(s) from {label}"
self.context.response.metadata.update({"edges": results, "count": len(results)})
return self.context.response

View file

@ -0,0 +1,94 @@
"""Update file catalog with a batch of file changes."""
from pathlib import Path
from watchfiles import Change
from ..base_step import BaseStep
from ...components import R
from ...components.file_catalog import BaseFileCatalog
from ...enumeration import ComponentEnum
from ...schema import FileNode
@R.register("update_catalog_step")
class UpdateCatalogStep(BaseStep):
"""Classify raw watcher changes and update the file_catalog."""
@property
def file_catalog(self) -> BaseFileCatalog:
"""Return the file catalog component."""
return self._resolve("file_catalog", BaseFileCatalog, ComponentEnum.FILE_CATALOG)
def _to_vault_relative(self, path: str | Path) -> str:
abs_path = Path(path).absolute()
try:
return str(abs_path.relative_to(self.vault_path))
except ValueError:
return str(abs_path)
async def execute(self):
assert self.context is not None
# Each item: {"change": Change | "added"|"modified"|"deleted", "path": absolute path}
changes: list[dict] = self.context.get("changes") or []
persist: bool = bool(self.context.get("persist", False))
buckets: dict[Change, list[str]] = {Change.added: [], Change.modified: [], Change.deleted: []}
for item in changes:
c = item["change"]
if isinstance(c, str):
c = Change.__members__.get(c)
if isinstance(c, Change) and c in buckets:
buckets[c].append(item["path"])
results: list[dict] = []
for change, action in ((Change.added, "Adding"), (Change.modified, "Updating")):
paths = buckets[change]
if not paths:
continue
self.logger.info(f"Detected {len(paths)} {change.name} file(s)")
nodes: list[FileNode] = []
ok_paths: list[str] = []
for path in paths:
abs_path = Path(path)
if not abs_path.is_file():
results.append({"change": change.name, "path": path, "success": False, "error": "not a file"})
continue
self.logger.info(f"{action} file: {path}")
try:
stat = abs_path.stat()
nodes.append(FileNode(path=self._to_vault_relative(abs_path), st_mtime=stat.st_mtime))
ok_paths.append(path)
except Exception as e:
self.logger.exception(f"Failed to stat {path}")
results.append({"change": change.name, "path": path, "success": False, "error": str(e)})
if nodes:
try:
await self.file_catalog.delete([n.path for n in nodes])
await self.file_catalog.upsert(nodes)
results.extend({"change": change.name, "path": p, "success": True} for p in ok_paths)
except Exception as e:
self.logger.exception(f"Failed to upsert {len(nodes)} {change.name} file(s)")
results.extend(
{"change": change.name, "path": p, "success": False, "error": str(e)} for p in ok_paths
)
if deleted := buckets[Change.deleted]:
if self.file_catalog is None:
raise RuntimeError("file_catalog is not initialized!")
self.logger.info(f"Detected {len(deleted)} deleted file(s)")
rel_deleted = [self._to_vault_relative(p) for p in deleted]
try:
await self.file_catalog.delete(rel_deleted)
results.extend({"change": "deleted", "path": p, "success": True} for p in deleted)
except Exception as e:
self.logger.exception(f"Failed to delete {len(deleted)} file(s)")
results.extend({"change": "deleted", "path": p, "success": False, "error": str(e)} for p in deleted)
if persist and results:
await self.file_catalog.dump()
self.context.response.answer = results
self.context.response.success = all(r["success"] for r in results) if results else True
return self.context.response

View file

@ -1,4 +1,4 @@
"""Index a batch of file changes into file_store."""
"""Update index with a batch of file changes."""
from pathlib import Path
@ -9,9 +9,13 @@ from ...components import R
from ...schema import FileChunk, FileNode
@R.register("index_changes_step")
class IndexChangesStep(BaseStep):
"""Classify raw watcher changes and index them into file_store."""
@R.register("update_index_step")
class UpdateIndexStep(BaseStep):
"""Classify raw watcher changes and update the file_store index."""
def __init__(self, persist: bool = False, **kwargs):
super().__init__(**kwargs)
self.persist: bool = persist
async def execute(self):
assert self.context is not None
@ -76,6 +80,9 @@ class IndexChangesStep(BaseStep):
self.logger.exception(f"Failed to delete {len(deleted)} file(s)")
results.extend({"change": "deleted", "path": p, "success": False, "error": str(e)} for p in deleted)
if self.persist and results:
await self.file_store.dump()
self.context.response.answer = results
self.context.response.success = all(r["success"] for r in results) if results else True
return self.context.response

View file

@ -1,16 +1,17 @@
"""Long-running awatch loop: convert raw changes into index_changes calls."""
"""Long-running awatch loop: convert raw changes into update_index calls."""
import asyncio
from watchfiles import Change, awatch
from ..base_step import BaseStep
from ...components import R
from ...components import R, BaseComponent
from ...enumeration import ComponentEnum
@R.register("watch_changes_step")
class WatchChangesStep(BaseStep):
"""Watch files and forward each batch of raw changes to the index_changes job."""
"""Watch files and forward each batch of raw changes to a downstream step."""
def __init__(
self,
@ -18,6 +19,7 @@ class WatchChangesStep(BaseStep):
force_polling: bool = True,
debounce: int = 2000,
poll_delay_ms: int = 2000,
dispatch_step: str = "",
**kwargs,
):
super().__init__(**kwargs)
@ -25,6 +27,7 @@ class WatchChangesStep(BaseStep):
self.force_polling: bool = force_polling
self.debounce: int = debounce
self.poll_delay_ms: int = poll_delay_ms
self.dispatch_step: str = dispatch_step
def _filter(self, _change: Change, path: str) -> bool:
suffixes = (self.context.get("suffix_filters") if self.context else None) or ["md"]
@ -43,6 +46,12 @@ class WatchChangesStep(BaseStep):
if not valid_paths:
raise RuntimeError(f"No valid watch paths under {self.vault_path}: {paths}")
dispatch_step_cls: type[BaseComponent] | None = None
if self.dispatch_step:
dispatch_step_cls = R.get(ComponentEnum.STEP, self.dispatch_step)
if dispatch_step_cls is None:
raise RuntimeError(f"Unregistered step '{self.dispatch_step}'")
self.logger.info(f"Watching: {[str(p) for p in valid_paths]}")
async for raw_changes in awatch(
*valid_paths,
@ -62,6 +71,8 @@ class WatchChangesStep(BaseStep):
]
if changes:
self.logger.info(f"Detected {len(changes)} change(s)")
await self.run_job("index_changes", changes=changes)
if dispatch_step_cls is not None:
step = dispatch_step_cls(app_context=self.app_context)
await step(changes=changes)
return self.context.response

View file

View file

@ -1,8 +1,8 @@
"""``upload_resource`` — copy an externally-received asset into ``resource/<date>/``.
"""``ingest`` — capture an externally-received asset into ``resource/<date>/``.
This step is the **passive** ingest entry point: when an external
channel (wechat group, email, browser save, API push, ...) hands
the agent a file, ``upload_resource`` lands it in ``resource/<YYYY-MM-DD>/``
This step is the dedicated information-capture interface: when an
external channel (wechat group, email, browser save, API push, ...)
hands the agent a file, ``ingest`` lands it in ``resource/<YYYY-MM-DD>/``
keyed by the day it was received (always today, local time),
alongside a ``meta.json`` row recording provenance.
@ -40,7 +40,7 @@ are reported as an error — the step never silently dedupes, so
callers see the conflict and can decide whether to retry, rename
upstream, or skip.
Each ``upload_resource`` call:
Each ``ingest`` call:
1. Resolves the bucket date as today (local time).
2. Validates the inputs: ``path`` exists, ``channel`` matches the
@ -101,9 +101,9 @@ _CHANNEL_RE = re.compile(r"^[a-z0-9][a-z0-9-]*$")
_RESERVED_METADATA_KEYS = frozenset({"name", "channel", "received_at", "description"})
@R.register("upload_resource_step")
class UploadResourceStep(BaseStep):
"""Land an external asset in ``resource/<date>/`` and update the day's meta + index."""
@R.register("ingest_step")
class IngestStep(BaseStep):
"""Capture an external asset into ``resource/<date>/`` and update the day's meta + index."""
async def execute(self):
assert self.context is not None
@ -129,7 +129,7 @@ class UploadResourceStep(BaseStep):
"description": description,
},
)
except _DuplicateUpload as e:
except _DuplicateIngest as e:
self._fail({"error": str(e)})
return
except Exception as e:
@ -137,7 +137,7 @@ class UploadResourceStep(BaseStep):
return
self.context.response.success = True
self.context.response.answer = f"Uploaded {outcome['name']} to {outcome['path']}"
self.context.response.answer = f"Ingested {outcome['name']} to {outcome['path']}"
self.context.response.metadata.update(outcome)
# ------------------------------------------------------------------
@ -145,7 +145,7 @@ class UploadResourceStep(BaseStep):
def _fail(self, payload: dict) -> None:
assert self.context is not None
self.context.response.success = False
self.context.response.answer = f"Error: {payload.get('error', 'upload failed')}"
self.context.response.answer = f"Error: {payload.get('error', 'ingest failed')}"
self.context.response.metadata.update(payload)
def _resource_dir_name(self) -> str:
@ -176,7 +176,7 @@ class UploadResourceStep(BaseStep):
on_disk = {p.name for p in bucket.iterdir() if p.is_file()}
existing_names = {Path(e.path).name for e in existing_entries} | on_disk
if final_name in existing_names:
raise _DuplicateUpload(
raise _DuplicateIngest(
f"duplicate: {final_name!r} already exists in {resource_dir}/{date}/",
)
@ -206,7 +206,7 @@ class UploadResourceStep(BaseStep):
}
class _DuplicateUpload(Exception):
class _DuplicateIngest(Exception):
"""Raised when the derived name already exists in the bucket."""
@ -362,7 +362,7 @@ def _read_meta(meta_path: Path) -> list[FileNode]:
def _atomic_write_text(target: Path, text: str) -> None:
"""Atomic text write via tempfile + os.replace in the same directory."""
target.parent.mkdir(parents=True, exist_ok=True)
tmp_fd, tmp_path = tempfile.mkstemp(prefix=".upload-", dir=target.parent)
tmp_fd, tmp_path = tempfile.mkstemp(prefix=".ingest-", dir=target.parent)
try:
with os.fdopen(tmp_fd, "w", encoding="utf-8") as f:
f.write(text)

View file

@ -14,7 +14,7 @@ callers must opt in to clobber an existing destination.
For the resource-bucket ingest path (channel-tagged, dated under
``resource/<YYYY-MM-DD>/`` with provenance metadata) use
``upload_resource`` instead.
``ingest`` instead.
"""
import mimetypes

View file

@ -8,6 +8,7 @@ from .common_utils import (
call_and_check,
)
from .env_utils import load_env
from .link_expansion import expand_links, render_expansion_lines
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
@ -20,6 +21,8 @@ __all__ = [
"call_action",
"call_and_check",
"load_env",
"expand_links",
"render_expansion_lines",
"get_logger",
"print_logo",
"find_reme",

View file

@ -138,12 +138,12 @@ async def mock_reme_server(
port: int | None = None,
config: str | None = None,
extra_args: list[str] | None = None,
startup_timeout: float = 30.0,
startup_timeout: float = 120.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.
"""Spawn `reme start` as a subprocess and yield (host, port) once ready.
Auto-picks a free port when port is None. Subprocess is terminated on exit.
"""
@ -154,7 +154,7 @@ async def mock_reme_server(
cmd: list[str] = [
sys.executable,
"-m",
"reme4.reme",
"reme.reme",
"start",
f"service.host={host}",
f"service.port={port}",

View file

@ -0,0 +1,129 @@
"""Expand a file's wikilink neighbors and render them as indented text.
Used by :class:`~reme.steps.index.search.SearchStep` for per-hit
context expansion. Pure helper — no step state, only ``file_store`` is
required.
Two-layer split so callers can pick what they need:
* :func:`expand_links` — data layer. Returns a structured dict keyed
by source path, each value carrying its outlinks / inlinks with
neighbor meta and per-edge predicate/anchor.
* :func:`render_expansion_lines` — view layer. Turns one path's
expansion sub-dict into the same `` → path name=… description=…``
block ``SearchStep`` has historically printed.
"""
import asyncio
from ..schema import FileLink, FileNode
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
def _node_meta(node: FileNode | None) -> dict:
"""Extract a compact meta dict (name/description) from a FileNode."""
if node is None:
return {}
fm = node.front_matter
meta: dict = {}
if fm.name:
meta["name"] = fm.name
if fm.description:
meta["description"] = fm.description
return meta
def _format_meta_inline(meta: dict) -> str:
"""One-line render of node meta for the answer; '(no meta)' when empty."""
parts = []
if "name" in meta:
parts.append(f'name="{meta["name"]}"')
if "description" in meta:
parts.append(f'description="{meta["description"]}"')
return " ".join(parts) if parts else "(no meta)"
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(
file_store,
paths: list[str],
max_per_direction: int = 10,
) -> dict[str, dict]:
"""Fetch out/in links for each path and attach neighbor meta.
Returns ``{path: {"outlinks": [...], "inlinks": [...]}, ...}`` where
each list item is ``{"path": str, "meta": {...}, "edges": [{"predicate", "anchor"}, ...]}``.
Empty input returns ``{}``. ``max_per_direction`` caps the neighbor
list per direction *before* meta lookup so we don't fetch nodes
that won't be displayed.
"""
if not paths:
return {}
out_lists, in_lists = await asyncio.gather(
asyncio.gather(*(file_store.get_outlinks(p) for p in paths)),
asyncio.gather(*(file_store.get_inlinks(p) for p in paths)),
)
out_grouped = [
dict(list(_group_by_neighbor(outs, "target_path").items())[:max_per_direction]) for outs in out_lists
]
in_grouped = [dict(list(_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 file_store.get_nodes(neighbor_paths) if neighbor_paths else []
meta_by_path = {n.path: _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 {p: {"outlinks": _attach(og), "inlinks": _attach(ig)} for p, og, ig in zip(paths, out_grouped, in_grouped)}
def render_expansion_lines(expansion: dict, indent: str = " ") -> list[str]:
"""Render one path's expansion sub-dict as indented lines.
``expansion`` is one value from :func:`expand_links` — i.e.
``{"outlinks": [...], "inlinks": [...]}``. Returns ``[]`` when both
directions are empty (caller decides whether to append a blank
line). ``indent`` controls the leading indent of the direction
header; neighbor lines and per-edge ``via`` lines nest further.
"""
lines: list[str] = []
inner = indent + " "
edge_indent = indent + " "
for direction, arrow, items in (
("outlinks", "→", expansion.get("outlinks") or []),
("inlinks", "←", expansion.get("inlinks") or []),
):
if not items:
continue
lines.append(f"{indent}{direction} ({len(items)}):")
for item in items:
lines.append(f"{inner}{arrow} {item['path']} {_format_meta_inline(item['meta'])}")
for edge in item["edges"]:
lines.append(f"{edge_indent}via {_format_via(edge)}")
return lines

View file

@ -79,7 +79,7 @@ def precheck_start(svc_config: dict | None) -> bool:
return False
if status == "occupied":
print(
f"port {port} occupied. Start on another port: reme4 start service.port=<other_port>",
f"port {port} occupied. Start on another port: reme start service.port=<other_port>",
file=sys.stderr,
)
sys.exit(1)

View file

@ -4,7 +4,7 @@ One class, :class:`WikilinkHandler`, owning every wikilink concern:
* **Pure text** — regex, Dataview predicate inference, validation:
:meth:`~WikilinkHandler.extract_links` (used by
:mod:`reme4.components.file_parser.linked_file_parser`),
:mod:`reme.components.file_parser.linked_file_parser`),
:meth:`~WikilinkHandler.scan_and_rewrite`,
:meth:`~WikilinkHandler.validate_src_dst` /
:meth:`~WikilinkHandler.validate_scope` /

View file

@ -1,10 +1,11 @@
"""Tests for background steps: UpdateStoreStep + WatchChangesStep.
"""Tests for background steps: ScanChangesStep + WatchChangesStep.
Both steps are subclasses of BaseStep. To exercise them without spinning up the
full ApplicationContext / index_changes job, we:
* pass real (started) file_store/file_parser via the step's kwargs (so the
BaseStep _resolve() machinery returns them);
* stub run_job() with a small recorder that captures the changes payload.
full ApplicationContext, we pass real (started) file_store/file_parser via the
step's kwargs (so the BaseStep _resolve() machinery returns them).
ScanChangesStep writes its result into ``context["changes"]`` for a downstream
``update_index_step`` to consume; tests assert against that key directly.
"""
# pylint: disable=protected-access
@ -14,15 +15,13 @@ import os
import tempfile
import warnings
from pathlib import Path
from typing import Any
from watchfiles import Change
from reme4.components.file_parser import ChunkedFileParser
from reme4.components.file_store import LocalFileStore
from reme4.components.runtime_context import RuntimeContext
from reme4.schema import Response
from reme4.steps.background import UpdateStoreStep, WatchChangesStep
from reme4.steps import ScanChangesStep, WatchChangesStep
warnings.filterwarnings("ignore", category=DeprecationWarning, module="jieba")
warnings.filterwarnings("ignore", category=DeprecationWarning, module="pkg_resources")
@ -52,52 +51,24 @@ def write_file(path: Path, content: str = "x") -> Path:
# ---------------------------------------------------------------------------
# UpdateStoreStep
# ScanChangesStep
# ---------------------------------------------------------------------------
class _RecorderStep:
"""Mixin: replaces run_job with a recorder that captures the changes payload."""
recorded: list[dict]
dispatched: int
def install_recorder(self):
"""Install a fake run_job that records dispatched 'index_changes' payloads."""
self.recorded = []
self.dispatched = 0
async def fake_run_job(name: str, **kwargs: Any):
assert name == "index_changes"
self.recorded = kwargs.get("changes") or []
self.dispatched += 1
return Response()
# pylint: disable-next=attribute-defined-outside-init
self.run_job = fake_run_job # type: ignore[assignment]
class _RecordingUpdateStoreStep(UpdateStoreStep, _RecorderStep):
pass
async def _make_update_step(
async def _make_scan_step(
watch_paths: list[str] | str = "vault",
suffix_filters: list[str] | None = None,
recursive: bool = True,
dump: bool = True,
) -> tuple[_RecordingUpdateStoreStep, RuntimeContext, LocalFileStore, ChunkedFileParser]:
fs = LocalFileStore(store_name="test_store", embedding_model="")
) -> tuple[ScanChangesStep, RuntimeContext, LocalFileStore, ChunkedFileParser]:
fs = LocalFileStore(name="test_store", embedding_model="")
parser = ChunkedFileParser()
await fs.start()
await parser.start()
step = _RecordingUpdateStoreStep(
step = ScanChangesStep(
recursive=recursive,
dump=dump,
file_store=fs,
file_parser=parser,
)
step.install_recorder()
context = RuntimeContext(
watch_paths=watch_paths,
suffix_filters=suffix_filters or ["md"],
@ -110,7 +81,7 @@ async def _teardown(fs: LocalFileStore, parser: ChunkedFileParser) -> None:
await fs.close()
def test_update_store_initial_all_added():
def test_scan_changes_initial_all_added():
"""First run on a fresh store emits 'added' for every existing file (abs paths)."""
async def run():
@ -121,49 +92,49 @@ def test_update_store_initial_all_added():
vault = cwd / "vault"
write_file(vault / "a.md", "alpha")
write_file(vault / "b.md", "beta")
step, ctx, fs, parser = await _make_update_step()
step, ctx, fs, parser = await _make_scan_step()
try:
resp = await step(ctx)
counts = resp.metadata["counts"]
assert counts == {"added": 2, "modified": 0, "deleted": 0}
assert step.dispatched == 1
kinds = sorted(item["change"] for item in step.recorded)
paths = sorted(item["path"] for item in step.recorded)
changes = ctx["changes"]
kinds = sorted(item["change"] for item in changes)
paths = sorted(item["path"] for item in changes)
assert kinds == ["added", "added"]
expected = sorted([str(cwd / "vault/a.md"), str(cwd / "vault/b.md")])
assert paths == expected
finally:
await _teardown(fs, parser)
print("✓ test_update_store_initial_all_added passed")
print("✓ test_scan_changes_initial_all_added passed")
asyncio.run(run())
def test_update_store_no_changes_skips_dispatch():
"""A second run over an unchanged store reports zero counts and does not dispatch."""
def test_scan_changes_no_changes_emits_empty_list():
"""A second run over an unchanged store reports zero counts and empty changes."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
cwd = Path.cwd()
vault = cwd / "vault"
a = write_file(vault / "a.md", "alpha")
seed_step, ctx, fs, parser = await _make_update_step()
step, ctx, fs, parser = await _make_scan_step()
try:
node, chunks = await parser.parse(a)
await fs.upsert([(node, chunks)])
resp = await seed_step(ctx)
resp = await step(ctx)
counts = resp.metadata["counts"]
assert counts == {"added": 0, "modified": 0, "deleted": 0}
assert seed_step.dispatched == 0
assert ctx["changes"] == []
finally:
await _teardown(fs, parser)
print("✓ test_update_store_no_changes_skips_dispatch passed")
print("✓ test_scan_changes_no_changes_emits_empty_list passed")
asyncio.run(run())
def test_update_store_detects_modify_and_delete():
def test_scan_changes_detects_modify_and_delete():
"""Second pass distinguishes added/modified/deleted; paths are absolute."""
async def run():
@ -172,7 +143,7 @@ def test_update_store_detects_modify_and_delete():
vault = cwd / "vault"
a = write_file(vault / "a.md", "alpha")
b = write_file(vault / "b.md", "beta")
step, ctx, fs, parser = await _make_update_step()
step, ctx, fs, parser = await _make_scan_step()
try:
# Seed via direct parse/upsert.
for p in (a, b):
@ -188,31 +159,31 @@ def test_update_store_detects_modify_and_delete():
resp = await step(ctx)
counts = resp.metadata["counts"]
assert counts == {"added": 1, "modified": 1, "deleted": 1}
by_kind = {item["change"]: item["path"] for item in step.recorded}
by_kind = {item["change"]: item["path"] for item in ctx["changes"]}
assert by_kind["added"] == str(c)
assert by_kind["modified"] == str(a)
assert by_kind["deleted"] == str(b)
finally:
await _teardown(fs, parser)
print("✓ test_update_store_detects_modify_and_delete passed")
print("✓ test_scan_changes_detects_modify_and_delete passed")
asyncio.run(run())
def test_update_store_missing_watch_path_silently_skipped():
def test_scan_changes_missing_watch_path_silently_skipped():
"""Non-existent watch_paths entries are dropped silently."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
(Path(tmpdir) / "vault").mkdir()
step, ctx, fs, parser = await _make_update_step(watch_paths=["vault", "ghost"])
step, ctx, fs, parser = await _make_scan_step(watch_paths=["vault", "ghost"])
try:
resp = await step(ctx)
assert resp.metadata["counts"] == {"added": 0, "modified": 0, "deleted": 0}
assert step.dispatched == 0
assert ctx["changes"] == []
finally:
await _teardown(fs, parser)
print("✓ test_update_store_missing_watch_path_silently_skipped passed")
print("✓ test_scan_changes_missing_watch_path_silently_skipped passed")
asyncio.run(run())
@ -222,16 +193,11 @@ def test_update_store_missing_watch_path_silently_skipped():
# ---------------------------------------------------------------------------
class _RecordingWatchChangesStep(WatchChangesStep, _RecorderStep):
pass
def test_watch_changes_requires_stop_event():
"""Missing stop_event in context raises a clear error."""
async def run():
step = _RecordingWatchChangesStep()
step.install_recorder()
step = WatchChangesStep()
step.context = RuntimeContext(watch_paths=["vault"], suffix_filters=["md"])
try:
await step.execute()
@ -249,8 +215,7 @@ def test_watch_changes_raises_when_no_valid_paths():
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
step = _RecordingWatchChangesStep()
step.install_recorder()
step = WatchChangesStep()
stop = asyncio.Event()
step.context = RuntimeContext(
stop_event=stop,
@ -264,7 +229,6 @@ def test_watch_changes_raises_when_no_valid_paths():
assert "No valid watch paths" in str(e)
else:
raise AssertionError("expected RuntimeError")
assert step.dispatched == 0
print("✓ test_watch_changes_raises_when_no_valid_paths passed")
asyncio.run(run())
@ -273,8 +237,7 @@ def test_watch_changes_raises_when_no_valid_paths():
def test_watch_changes_filter_only_passes_md():
"""The internal filter pulls suffix_filters from runtime context."""
step = _RecordingWatchChangesStep()
step.install_recorder()
step = WatchChangesStep()
step.context = RuntimeContext(suffix_filters=["md"])
assert step._filter(Change.added, "/x/foo.md")
assert not step._filter(Change.added, "/x/foo.txt")
@ -283,11 +246,11 @@ def test_watch_changes_filter_only_passes_md():
if __name__ == "__main__":
print("\n=== Background Steps Tests ===")
# UpdateStoreStep
test_update_store_initial_all_added()
test_update_store_no_changes_skips_dispatch()
test_update_store_detects_modify_and_delete()
test_update_store_missing_watch_path_silently_skipped()
# ScanChangesStep
test_scan_changes_initial_all_added()
test_scan_changes_no_changes_emits_empty_list()
test_scan_changes_detects_modify_and_delete()
test_scan_changes_missing_watch_path_silently_skipped()
# WatchChangesStep
test_watch_changes_requires_stop_event()
test_watch_changes_raises_when_no_valid_paths()

View file

@ -1,586 +0,0 @@
"""Tests for BM25Index search engine."""
# pylint: disable=protected-access
import asyncio
import os
import tempfile
import warnings
from reme4.components.keyword_index import BM25Index
from reme4.components.tokenizer import RegexTokenizer
# Filter jieba/pkg_resources deprecation warnings
warnings.filterwarnings("ignore", category=DeprecationWarning, module="jieba")
warnings.filterwarnings("ignore", category=DeprecationWarning, module="pkg_resources")
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)
async def create_bm25(k1: float = 1.5, b: float = 0.75) -> BM25Index:
"""Create and start a BM25Index in cwd with a non-filtering tokenizer.
The non-filtering tokenizer keeps short test texts (e.g. "hello world") visible,
since several common test words ("hello", "我", "的") are in the default stopwords.
"""
bm25 = BM25Index(k1=k1, b=b)
# Replace the unresolved Dependency placeholder with a real tokenizer instance.
tokenizer = RegexTokenizer(filter_stopwords=False)
bm25.tokenizer = tokenizer
bm25._owned.append(tokenizer)
await bm25.start()
return bm25
def test_basic_init():
"""Test BM25Index initialization."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = BM25Index()
assert bm25.k1 == 1.5
assert bm25.b == 0.75
assert bm25.vocab == {}
assert bm25.inverted_index == {}
assert bm25.doc_meta == {}
assert bm25.n_docs == 0
assert bm25.avg_len == 0.0
print("✓ test_basic_init passed")
asyncio.run(run())
def test_start_with_tokenizer():
"""Test BM25Index starts and initializes tokenizer."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
assert bm25.tokenizer is not None
assert bm25.is_started
await bm25.close()
assert not bm25.is_started
print("✓ test_start_with_tokenizer passed")
asyncio.run(run())
def test_add_single_doc():
"""Test adding a single document."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
await bm25.add_docs({"doc1": "hello world"})
assert bm25.n_docs == 1
assert bm25.total_len > 0
assert "doc1" in bm25.doc_meta
await bm25.close()
print("✓ test_add_single_doc passed")
asyncio.run(run())
def test_add_multiple_docs():
"""Test adding multiple documents."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
docs = {
"doc1": "hello world",
"doc2": "hello python",
"doc3": "world python",
}
await bm25.add_docs(docs)
assert bm25.n_docs == 3
assert len(bm25.vocab) > 0
assert len(bm25.inverted_index) > 0
await bm25.close()
print("✓ test_add_multiple_docs passed")
asyncio.run(run())
def test_retrieve_basic():
"""Test basic retrieval functionality."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
docs = {
"doc1": "python programming language",
"doc2": "java programming language",
"doc3": "python data analysis",
}
await bm25.add_docs(docs)
results = await bm25.retrieve("python", limit=3)
assert len(results) <= 3
assert "doc1" in results or "doc3" in results
await bm25.close()
print("✓ test_retrieve_basic passed")
asyncio.run(run())
def test_retrieve_with_limit():
"""Test retrieval with result limit."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
docs = {f"doc{i}": f"python programming {i}" for i in range(10)}
await bm25.add_docs(docs)
results = await bm25.retrieve("python", limit=3)
assert len(results) == 3
results = await bm25.retrieve("python", limit=5)
assert len(results) == 5
await bm25.close()
print("✓ test_retrieve_with_limit passed")
asyncio.run(run())
def test_retrieve_empty_query():
"""Test retrieval with empty or unknown query."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
docs = {"doc1": "hello world"}
await bm25.add_docs(docs)
results = await bm25.retrieve("", limit=3)
assert results == {}
results = await bm25.retrieve("unknownxyz", limit=3)
assert results == {}
await bm25.close()
print("✓ test_retrieve_empty_query passed")
asyncio.run(run())
def test_retrieve_empty_index():
"""Test retrieval from empty index."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
results = await bm25.retrieve("python", limit=3)
assert results == {}
await bm25.close()
print("✓ test_retrieve_empty_index passed")
asyncio.run(run())
def test_update_doc():
"""Test updating an existing document."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
await bm25.add_docs({"doc1": "hello world python"})
old_len = bm25.total_len
await bm25.add_docs({"doc1": "java"})
assert bm25.n_docs == 1
assert bm25.total_len != old_len
results = await bm25.retrieve("java", limit=1)
assert "doc1" in results
await bm25.close()
print("✓ test_update_doc passed")
asyncio.run(run())
def test_remove_doc():
"""Test removing a document."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
docs = {
"doc1": "hello world",
"doc2": "hello python",
}
await bm25.add_docs(docs)
assert bm25.n_docs == 2
bm25._remove_doc("doc1")
assert bm25.n_docs == 1
assert "doc1" not in bm25.doc_meta
results = await bm25.retrieve("hello", limit=2)
assert "doc1" not in results
assert "doc2" in results
await bm25.close()
print("✓ test_remove_doc passed")
asyncio.run(run())
def test_remove_nonexistent_doc():
"""Test removing a nonexistent document (should be no-op)."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
await bm25.add_docs({"doc1": "hello world"})
bm25._remove_doc("nonexistent")
assert bm25.n_docs == 1
await bm25.close()
print("✓ test_remove_nonexistent_doc passed")
asyncio.run(run())
def test_clear():
"""Test clearing the index."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
await bm25.add_docs(
{
"doc1": "hello world",
"doc2": "hello python",
},
)
assert bm25.n_docs == 2
await bm25.clear()
assert bm25.n_docs == 0
assert bm25.vocab == {}
assert bm25.inverted_index == {}
assert bm25.doc_meta == {}
assert bm25.total_len == 0
assert bm25._idf_cache == {}
await bm25.close()
print("✓ test_clear passed")
asyncio.run(run())
def test_optimize_index():
"""Test optimize_index functionality to compact vocab."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
await bm25.add_docs({"doc1": "hello world"})
bm25._remove_doc("doc1")
assert bm25.n_docs == 0
assert len(bm25.vocab) > 0
await bm25.optimize_index()
assert bm25.vocab == {}
assert bm25.inverted_index == {}
await bm25.close()
print("✓ test_optimize_index passed")
asyncio.run(run())
def test_optimize_index_with_docs():
"""Test optimize_index with remaining documents."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
await bm25.add_docs(
{
"doc1": "hello world",
"doc2": "hello python",
},
)
old_vocab = bm25.vocab.copy()
bm25._remove_doc("doc1")
await bm25.optimize_index()
assert bm25.n_docs == 1
assert "doc2" in bm25.doc_meta
assert len(bm25.vocab) < len(old_vocab)
results = await bm25.retrieve("hello", limit=1)
assert "doc2" in results
await bm25.close()
print("✓ test_optimize_index_with_docs passed")
asyncio.run(run())
def test_persistence():
"""Test dump and load persistence."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
docs = {
"doc1": "hello world",
"doc2": "hello python",
"doc3": "programming language",
}
await bm25.add_docs(docs)
old_vocab = bm25.vocab.copy()
old_doc_meta = {k: dict(v) for k, v in bm25.doc_meta.items()}
await bm25.dump()
await bm25.close()
bm25_new = await create_bm25()
assert bm25_new.vocab == old_vocab
assert bm25_new.n_docs == 3
for doc_id in old_doc_meta:
assert doc_id in bm25_new.doc_meta
results = await bm25_new.retrieve("hello", limit=2)
assert "doc1" in results or "doc2" in results
await bm25_new.close()
print("✓ test_persistence passed")
asyncio.run(run())
def test_custom_params():
"""Test custom k1 and b parameters."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25(k1=2.0, b=0.5)
assert bm25.k1 == 2.0
assert bm25.b == 0.5
await bm25.add_docs({"doc1": "test document"})
results = await bm25.retrieve("test", limit=1)
assert "doc1" in results
await bm25.close()
print("✓ test_custom_params passed")
asyncio.run(run())
def test_chinese_text():
"""Test with Chinese text."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
docs = {
"doc1": "我爱北京天安门",
"doc2": "北京是中国的首都",
"doc3": "上海的天气很好",
}
await bm25.add_docs(docs)
results = await bm25.retrieve("北", limit=2)
assert len(results) <= 2
assert "doc1" in results or "doc2" in results
await bm25.close()
print("✓ test_chinese_text passed")
asyncio.run(run())
def test_mixed_chinese_english():
"""Test with mixed Chinese and English text."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
docs = {
"doc1": "Python 是一种编程语言",
"doc2": "Java 编程语言",
"doc3": "Python 数据分析",
}
await bm25.add_docs(docs)
results = await bm25.retrieve("Python", limit=3)
assert len(results) > 0
results = await bm25.retrieve("编", limit=2)
assert len(results) > 0
await bm25.close()
print("✓ test_mixed_chinese_english passed")
asyncio.run(run())
def test_idf_cache():
"""Test IDF cache functionality."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
await bm25.add_docs(
{
"doc1": "hello world",
"doc2": "hello python",
},
)
token = "hello"
if token in bm25.vocab:
tid = bm25.vocab[token]
idf1 = bm25._get_idf(tid)
assert tid in bm25._idf_cache
idf2 = bm25._get_idf(tid)
assert idf1 == idf2
await bm25.close()
print("✓ test_idf_cache passed")
asyncio.run(run())
def test_avg_len():
"""Test average document length calculation."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
assert bm25.avg_len == 0.0
await bm25.add_docs({"doc1": "hello world python"})
assert bm25.avg_len > 0
await bm25.add_docs({"doc2": "test"})
new_avg = bm25.avg_len
assert new_avg > 0
await bm25.close()
print("✓ test_avg_len passed")
asyncio.run(run())
def test_score_ordering():
"""Test that results are ordered by score descending."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
docs = {
"doc1": "python python python",
"doc2": "python python",
"doc3": "python",
}
await bm25.add_docs(docs)
results = await bm25.retrieve("python", limit=3)
scores = list(results.values())
for i in range(len(scores) - 1):
assert scores[i] >= scores[i + 1]
await bm25.close()
print("✓ test_score_ordering passed")
asyncio.run(run())
def test_empty_doc():
"""Test adding empty document."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
bm25 = await create_bm25()
await bm25.add_docs({"doc1": ""})
assert bm25.n_docs == 0
await bm25.add_docs({"doc2": " "})
assert bm25.n_docs == 0
await bm25.close()
print("✓ test_empty_doc passed")
asyncio.run(run())
if __name__ == "__main__":
print("\n=== BM25Index Tests ===")
test_basic_init()
test_start_with_tokenizer()
test_add_single_doc()
test_add_multiple_docs()
test_retrieve_basic()
test_retrieve_with_limit()
test_retrieve_empty_query()
test_retrieve_empty_index()
test_update_doc()
test_remove_doc()
test_remove_nonexistent_doc()
test_clear()
test_optimize_index()
test_optimize_index_with_docs()
test_persistence()
test_custom_params()
test_chinese_text()
test_mixed_chinese_english()
test_idf_cache()
test_avg_len()
test_score_ordering()
test_empty_doc()
print("\n所有测试通过!")

View file

@ -2,10 +2,10 @@
Two surfaces share this file:
* **HTTP / MCP E2E tests** (top half) spawn ``reme4 start`` via
``mock_reme_server`` and drive ``version`` / ``help`` / ``search`` /
``init`` / ``demo`` over the wire. Each test uses an isolated cwd so
the vault (``.reme`` by default) does not collide.
* **In-process job tests** (top half) build an ``Application`` from the
default config and call ``run_job`` directly — no subprocess, no HTTP.
Each test uses an isolated cwd so the vault (``.reme`` by default)
does not collide.
* **Direct unit tests** (bottom half) exercise ``TraverseStep``
(registered as ``traverse_step``) — BFS over wikilink edges from a
seed file, forward / backward / both — against a freshly built
@ -19,11 +19,12 @@ import os
import tempfile
import warnings
from reme4 import __version__ as REME_VERSION
from reme4 import Application, __version__ as REME_VERSION
from reme4.components.file_store import LocalFileStore
from reme4.config import resolve_app_config
from reme4.schema import FileLink, FileNode
from reme4.steps.common import traverse as traverse_mod
from reme4.utils import call_action, call_and_check, mock_reme_server
from reme4.steps.index import traverse as traverse_mod
from reme4.utils import load_env
warnings.filterwarnings("ignore", category=DeprecationWarning, module="jieba")
warnings.filterwarnings("ignore", category=DeprecationWarning, module="pkg_resources")
@ -50,6 +51,15 @@ def _run(coro):
asyncio.run(coro)
async def _make_app() -> Application:
"""Build and start an Application with the default config, logging silenced."""
load_env()
cfg = resolve_app_config(log_to_console=False, log_to_file=False, enable_logo=False)
app = Application(**cfg)
await app.start()
return app
def _node(path: str, links: list[tuple[str, str | None, str | None]] | None = None) -> FileNode:
"""Build a FileNode with (target_path, target_anchor, predicate) outgoing edges."""
return FileNode(
@ -61,7 +71,7 @@ def _node(path: str, links: list[tuple[str, str | None, str | None]] | None = No
async def _make_store(nodes: list[FileNode]) -> LocalFileStore:
"""LocalFileStore seeded with the given graph nodes (no files on disk)."""
store = LocalFileStore(store_name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_model="")
await store.start()
if nodes:
await store.file_graph.upsert_nodes(nodes)
@ -73,7 +83,7 @@ def _edges(step) -> list[dict]:
# ===========================================================================
# HTTP / MCP E2E tests: version / help / search / init / demo
# In-process job tests: version / help / health_check / search / reindex
# ===========================================================================
@ -82,18 +92,14 @@ def test_version_job():
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
async with mock_reme_server() as (host, port):
await call_and_check(
"version",
host=host,
port=port,
validator=lambda r: (
isinstance(r, dict)
and r.get("success") is True
and r.get("answer") == REME_VERSION
and r.get("metadata", {}).get("version") == REME_VERSION
),
)
app = await _make_app()
try:
resp = await app.run_job("version")
assert resp.success is True
assert resp.answer == REME_VERSION
assert resp.metadata.get("version") == REME_VERSION
finally:
await app.close()
print("✓ test_version_job passed")
_run(run())
@ -104,24 +110,17 @@ def test_help_job():
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
async with mock_reme_server() as (host, port):
result = await call_and_check(
"help",
host=host,
port=port,
validator=lambda r: (
isinstance(r, dict)
and r.get("success") is True
and isinstance(r.get("answer"), str)
and r.get("metadata", {}).get("job_count", 0) > 0
and "`help`" not in r["answer"]
),
)
# Spot-check that a couple of known jobs appear in the listing.
answer = result["answer"]
app = await _make_app()
try:
resp = await app.run_job("help")
assert resp.success is True
assert isinstance(resp.answer, str)
assert resp.metadata.get("job_count", 0) > 0
assert "`help`" not in resp.answer
for expected_job in ("version", "health_check", "search"):
if expected_job not in answer:
raise AssertionError(f"help output missing job {expected_job!r}: {answer!r}")
assert expected_job in resp.answer, f"help missing {expected_job!r}: {resp.answer!r}"
finally:
await app.close()
print("✓ test_help_job passed")
_run(run())
@ -132,92 +131,61 @@ def test_search_job_empty_store():
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
async with mock_reme_server() as (host, port):
await call_and_check(
"search",
host=host,
port=port,
query="hello world",
limit=5,
validator=lambda r: (
isinstance(r, dict)
and r.get("success") is True
and isinstance(r.get("metadata"), dict)
and isinstance(r["metadata"].get("counts"), dict)
and r["metadata"]["counts"].get("returned", -1) == 0
),
)
app = await _make_app()
try:
resp = await app.run_job("search", query="hello world", limit=5)
assert resp.success is True
counts = resp.metadata.get("counts", {})
assert isinstance(counts, dict)
assert counts.get("returned", -1) == 0
finally:
await app.close()
print("✓ test_search_job_empty_store passed")
_run(run())
def test_search_job_missing_query():
"""search without a query should surface the assertion error in `answer`."""
"""search with empty query returns success=False and a query-related error in answer."""
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
async with mock_reme_server() as (host, port):
result = await call_action("search", host=host, port=port, query="")
if not isinstance(result, dict):
raise AssertionError(f"expected dict response, got {result!r}")
if "query" not in str(result.get("answer", "")).lower():
raise AssertionError(f"expected query-related error in answer, got {result!r}")
app = await _make_app()
try:
resp = await app.run_job("search", query="")
assert resp.success is False
assert "query" in str(resp.answer).lower()
finally:
await app.close()
print("✓ test_search_job_missing_query passed")
_run(run())
# -- aggregate: reuse one server instance for all jobs -------------------
def test_all_jobs_one_server():
"""Run every common job against a single shared server for efficiency."""
def test_all_jobs_single_app():
"""Run every common job against one shared in-process Application for efficiency."""
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
async with mock_reme_server() as (host, port):
# version
await call_and_check(
"version",
host=host,
port=port,
validator=lambda r: isinstance(r, dict) and r.get("answer") == REME_VERSION,
)
# help
await call_and_check(
"help",
host=host,
port=port,
validator=lambda r: isinstance(r, dict) and r.get("metadata", {}).get("job_count", 0) > 0,
)
# health_check
await call_and_check(
"health_check",
host=host,
port=port,
validator=lambda r: isinstance(r, dict)
and isinstance(
r.get("metadata", {}).get("health"),
dict,
),
)
# search (empty store)
await call_and_check(
"search",
host=host,
port=port,
query="anything",
validator=lambda r: isinstance(r, dict) and r.get("success") is True,
)
# reindex
await call_and_check(
"reindex",
host=host,
port=port,
validator=lambda r: isinstance(r, dict) and isinstance(r.get("metadata", {}).get("counts"), dict),
)
print("✓ test_all_jobs_one_server passed")
app = await _make_app()
try:
resp = await app.run_job("version")
assert resp.answer == REME_VERSION
resp = await app.run_job("help")
assert resp.metadata.get("job_count", 0) > 0
resp = await app.run_job("health_check")
assert isinstance(resp.metadata.get("health"), dict)
resp = await app.run_job("search", query="anything")
assert resp.success is True
resp = await app.run_job("reindex")
assert isinstance(resp.metadata.get("counts"), dict)
finally:
await app.close()
print("✓ test_all_jobs_single_app passed")
_run(run())
@ -365,12 +333,12 @@ def test_traverse_both_directions():
if __name__ == "__main__":
print("\n=== reme4 common steps E2E tests ===")
print("\n=== reme4 common steps in-process tests ===")
test_version_job()
test_help_job()
test_search_job_empty_store()
test_search_job_missing_query()
test_all_jobs_one_server()
test_all_jobs_single_app()
print("\n=== traverse step tests ===")
test_traverse_forward_depth_1()
test_traverse_backward_returns_inlinks()

File diff suppressed because it is too large Load diff

View file

@ -1,24 +1,23 @@
"""Tests for daily-aware steps: daily_read / daily_write / daily_list / daily_reindex.
"""Tests for daily-aware steps: daily_create / daily_list / daily_reindex.
Sets up a small ``daily/`` tree with mixed dates and exercises note
read / write / list / index-rebuild operations. Arbitrary body
mid-edits, plain appends, and frontmatter mutations are generic CRUD
(covered in test_crud_steps and test_property_steps).
provision / listing / index-rebuild operations. Body authoring and
frontmatter mutations are generic CRUD (covered in test_crud_steps
and test_property_steps).
A daily note is the single file ``daily/<YYYY-MM-DD>/<slug>.md`` (no
folder, no sibling materials). ``daily_read`` returns body in answer
and parsed frontmatter as a dict in metadata; ``daily_write`` writes
body + frontmatter in one shot — ``overwrite=False`` (default) is an
idempotent skip-if-exists, ``overwrite=True`` is unconditional. Both
validate the slug up-front (Windows-safe filename rules) so the
path-shape contract is enforced at the daily boundary, not inside
generic CRUD.
folder, no sibling materials). ``daily_create`` is a minimal slug
provisioner: it validates the slug, writes an empty-body note with
default ``{name: slug}`` frontmatter when the file is absent, and
refreshes the day index. When the file already exists it is a no-op
write (``created=False``) — the body is filled in afterwards via
``file_write`` / ``file_edit`` / ``frontmatter_update`` or a native
editor.
``daily_list`` is now a **pure read** — it no longer triggers index
refresh. Use ``daily_reindex`` explicitly when the index page needs
to be rebuilt. ``daily_write`` auto-refreshes the index by default;
``frontmatter_update`` / ``file_append`` flows leave it stale and
require an explicit ``daily_reindex``.
``daily_list`` is a **pure read** — it never refreshes the index.
Use ``daily_reindex`` explicitly when the index page needs to be
rebuilt (e.g. after batch flows or a ``frontmatter_update`` that
touched ``name`` / ``description``).
Note: status / lifecycle / scope / role / source are no longer
core-reserved fields — the reme schema reserves only name /
@ -37,11 +36,10 @@ from pathlib import Path
import warnings
from reme4.components.file_store import LocalFileStore
from reme4.steps.daily import (
read as daily_read_step,
write as daily_write_step,
list as daily_list_step,
reindex as daily_reindex_step,
from reme4.steps.file_io import (
daily_create as daily_create_step,
daily_list as daily_list_step,
daily_reindex as daily_reindex_step,
)
warnings.filterwarnings("ignore", category=DeprecationWarning, module="jieba")
@ -75,7 +73,7 @@ async def _make_store_with_dailies(entries: list[tuple[str, str, str]]) -> Local
``daily/<date>/<slug>.md`` with a minimal ``name``-only
frontmatter — no opinionated status / lifecycle axes.
"""
store = LocalFileStore(store_name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_model="")
await store.start()
for day, slug, body in entries:
day_dir = Path.cwd() / "daily" / day
@ -153,12 +151,12 @@ def test_daily_list_filters_by_date():
asyncio.run(run())
def test_daily_list_returns_path_slug_name_description():
"""Each note row exposes path / slug / name / description (and nothing else)."""
def test_daily_list_returns_path_slug_metadata():
"""Each note row exposes path / slug / metadata (full frontmatter dict)."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = LocalFileStore(store_name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_model="")
await store.start()
await _seed_note(
"2026-05-18",
@ -173,12 +171,11 @@ def test_daily_list_returns_path_slug_name_description():
{
"path": "daily/2026-05-18/alpha.md",
"slug": "alpha",
"name": "Alpha Project",
"description": "JWT auth migration",
"metadata": {"name": "Alpha Project", "description": "JWT auth migration"},
},
]
await store.close()
print("✓ test_daily_list_returns_path_slug_name_description passed")
print("✓ test_daily_list_returns_path_slug_metadata passed")
asyncio.run(run())
@ -216,7 +213,7 @@ def test_daily_list_empty_when_no_daily_dir():
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = LocalFileStore(store_name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_model="")
await store.start()
step = daily_list_step.DailyListStep(file_store=store)
await step(date="2026-05-18")
@ -272,165 +269,42 @@ def test_daily_list_response_shape():
asyncio.run(run())
# -- daily_read_step ----------------------------------------------------------
# -- daily_create_step --------------------------------------------------------
def test_daily_read_returns_body_and_frontmatter_dict():
"""daily_read on an existing note returns body in answer and parsed frontmatter dict."""
def test_daily_create_provisions_note_and_refreshes_index():
"""Fresh slug ⇒ empty-body note with ``{name: slug}`` + day index refreshed."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = LocalFileStore(store_name="t", embedding_model="")
await store.start()
day_dir = Path(tmp) / "daily" / "2026-05-18"
day_dir.mkdir(parents=True, exist_ok=True)
(day_dir / "alpha.md").write_text(
"---\nname: Alpha Project\ndescription: JWT migration\n---\n## Objective\nfoo\n",
encoding="utf-8",
)
step = daily_read_step.DailyReadStep(file_store=store)
await step(slug="alpha", date="2026-05-18")
store = await _make_store_with_dailies([])
step = daily_create_step.DailyCreateStep(file_store=store)
await step(slug="kickoff", date="2026-05-18")
payload = _metadata(step)
assert step.context.response.success is True
assert "## Objective\nfoo" in step.context.response.answer
assert "---" not in step.context.response.answer # frontmatter stripped
assert payload["date"] == "2026-05-18"
assert payload["slug"] == "alpha"
assert payload["path"] == "daily/2026-05-18/alpha.md"
assert payload["exists"] is True
assert payload["frontmatter"] == {
"name": "Alpha Project",
"description": "JWT migration",
}
await store.close()
print("✓ test_daily_read_returns_body_and_frontmatter_dict passed")
asyncio.run(run())
def test_daily_read_default_date_is_today():
"""Omitted ``date`` ⇒ today's folder."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies(
[(_today(), "live", "current body")],
)
step = daily_read_step.DailyReadStep(file_store=store)
await step(slug="live")
payload = _metadata(step)
assert payload["date"] == _today()
assert payload["path"] == f"daily/{_today()}/live.md"
assert payload["exists"] is True
assert "current body" in step.context.response.answer
await store.close()
print("✓ test_daily_read_default_date_is_today passed")
asyncio.run(run())
def test_daily_read_missing_file_reports_exists_false():
"""Note absent ⇒ success=False, payload carries exists=False + path."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies([])
step = daily_read_step.DailyReadStep(file_store=store)
await step(slug="nothing-here", date="2026-05-18")
payload = _metadata(step)
assert step.context.response.success is False
assert payload["exists"] is False
assert payload["date"] == "2026-05-18"
assert payload["slug"] == "nothing-here"
assert payload["path"] == "daily/2026-05-18/nothing-here.md"
await store.close()
print("✓ test_daily_read_missing_file_reports_exists_false passed")
asyncio.run(run())
def test_daily_read_rejects_invalid_slug():
"""Slug validation (Windows-safe filename rules) runs up-front."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies([])
step = daily_read_step.DailyReadStep(file_store=store)
for bad in ("foo/bar", "foo:bar", "foo*bar", "CON", "lpt9", "foo.", " bar"):
await step(slug=bad)
assert step.context.response.success is False, f"expected reject for {bad!r}"
await store.close()
print("✓ test_daily_read_rejects_invalid_slug passed")
asyncio.run(run())
def test_daily_read_empty_frontmatter_dict():
"""No frontmatter ⇒ ``frontmatter`` key is the empty dict."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = LocalFileStore(store_name="t", embedding_model="")
await store.start()
day_dir = Path(tmp) / "daily" / "2026-05-18"
day_dir.mkdir(parents=True, exist_ok=True)
(day_dir / "plain.md").write_text("just body\n", encoding="utf-8")
step = daily_read_step.DailyReadStep(file_store=store)
await step(slug="plain", date="2026-05-18")
payload = _metadata(step)
assert payload["frontmatter"] == {}
assert step.context.response.answer.strip() == "just body"
await store.close()
print("✓ test_daily_read_empty_frontmatter_dict passed")
asyncio.run(run())
# -- daily_write_step ---------------------------------------------------------
def test_daily_write_creates_note_and_refreshes_index():
"""Fresh slug ⇒ note file written + day index refreshed by default."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies([])
step = daily_write_step.DailyWriteStep(file_store=store)
await step(
slug="kickoff",
date="2026-05-18",
body="## Plan\nfirst pass\n",
frontmatter={"name": "Kickoff Task", "description": "Day-one plan"},
)
payload = _metadata(step)
assert payload["created"] is True
assert payload["overwritten"] is False
assert payload["date"] == "2026-05-18"
assert payload["slug"] == "kickoff"
assert payload["path"] == "daily/2026-05-18/kickoff.md"
note = Path(tmp) / "daily" / "2026-05-18" / "kickoff.md"
text = note.read_text(encoding="utf-8")
assert "name: Kickoff Task" in text
assert "description: Day-one plan" in text
assert "## Plan\nfirst pass" in text
assert "name: kickoff" in text
# Body is empty — file is frontmatter + trailing newline.
assert text.rstrip().endswith("---")
index = Path(tmp) / "daily" / "2026-05-18.md"
assert index.is_file()
assert "[[daily/2026-05-18/kickoff.md]]" in index.read_text(encoding="utf-8")
await store.close()
print("✓ test_daily_write_creates_note_and_refreshes_index passed")
print("✓ test_daily_create_provisions_note_and_refreshes_index passed")
asyncio.run(run())
def test_daily_write_create_mode_is_idempotent():
"""``overwrite=False`` + file exists ⇒ created=False, file untouched, index still refreshed."""
def test_daily_create_is_idempotent_on_existing():
"""File exists ⇒ ``created=False``; the file body is NOT touched; index still refreshes."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
@ -440,228 +314,113 @@ def test_daily_write_create_mode_is_idempotent():
file_path = Path(tmp) / "daily" / "2026-05-18" / "ongoing.md"
before = file_path.read_text(encoding="utf-8")
step = daily_write_step.DailyWriteStep(file_store=store)
await step(slug="ongoing", date="2026-05-18", body="ignored new body")
step = daily_create_step.DailyCreateStep(file_store=store)
await step(slug="ongoing", date="2026-05-18")
payload = _metadata(step)
assert step.context.response.success is True
assert payload["created"] is False
assert payload["overwritten"] is False
assert payload["path"] == "daily/2026-05-18/ongoing.md"
assert file_path.read_text(encoding="utf-8") == before
assert payload["index"]["path"] == "daily/2026-05-18.md"
await store.close()
print("✓ test_daily_write_create_mode_is_idempotent passed")
print("✓ test_daily_create_is_idempotent_on_existing passed")
asyncio.run(run())
def test_daily_write_overwrite_mode_replaces_existing():
"""``overwrite=True`` ⇒ unconditional rewrite; overwritten=True."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies(
[("2026-05-18", "live", "stale text")],
)
step = daily_write_step.DailyWriteStep(file_store=store)
await step(
slug="live",
date="2026-05-18",
body="## Updated\nfresh body\n",
frontmatter={"name": "Live", "description": "post-merge"},
overwrite=True,
)
payload = _metadata(step)
assert payload["created"] is False
assert payload["overwritten"] is True
note = Path(tmp) / "daily" / "2026-05-18" / "live.md"
text = note.read_text(encoding="utf-8")
assert "stale text" not in text
assert "## Updated\nfresh body" in text
assert "description: post-merge" in text
await store.close()
print("✓ test_daily_write_overwrite_mode_replaces_existing passed")
asyncio.run(run())
def test_daily_write_default_frontmatter_uses_slug_as_name():
"""Omitted ``frontmatter`` ⇒ defaults to ``{name: slug}``."""
def test_daily_create_default_date_is_today():
"""Omitted ``date`` ⇒ today's folder."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies([])
step = daily_write_step.DailyWriteStep(file_store=store)
step = daily_create_step.DailyCreateStep(file_store=store)
await step(slug="today-task")
payload = _metadata(step)
assert payload["date"] == _today()
assert payload["path"] == f"daily/{_today()}/today-task.md"
assert payload["created"] is True
await store.close()
print("✓ test_daily_create_default_date_is_today passed")
asyncio.run(run())
def test_daily_create_default_frontmatter_uses_slug_as_name():
"""The provisioned note's frontmatter is ``{name: slug}`` (no body)."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies([])
step = daily_create_step.DailyCreateStep(file_store=store)
await step(slug="auth-refactor", date="2026-05-18")
note = Path(tmp) / "daily" / "2026-05-18" / "auth-refactor.md"
assert "name: auth-refactor" in note.read_text(encoding="utf-8")
await store.close()
print("✓ test_daily_write_default_frontmatter_uses_slug_as_name passed")
asyncio.run(run())
def test_daily_write_default_body_is_empty():
"""Omitted ``body`` ⇒ empty body, frontmatter-only note."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies([])
step = daily_write_step.DailyWriteStep(file_store=store)
await step(slug="stub", date="2026-05-18")
note = Path(tmp) / "daily" / "2026-05-18" / "stub.md"
text = note.read_text(encoding="utf-8")
# Just frontmatter + a single newline after the closing ---.
assert text.startswith("---\n")
assert "name: stub" in text
# Body section is empty: the post body resolves to "".
assert text.rstrip().endswith("---")
await store.close()
print("✓ test_daily_write_default_body_is_empty passed")
asyncio.run(run())
def test_daily_write_drops_empty_frontmatter_values():
"""Empty / None frontmatter values are dropped (write_step idiom)."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies([])
step = daily_write_step.DailyWriteStep(file_store=store)
await step(
slug="trim",
date="2026-05-18",
frontmatter={
"name": "Trim",
"description": " ", # whitespace-only → drop
"extra": None, # None → drop
"kept": "value",
},
)
note = Path(tmp) / "daily" / "2026-05-18" / "trim.md"
text = note.read_text(encoding="utf-8")
assert "name: Trim" in text
assert "kept: value" in text
assert "name: auth-refactor" in text
# No description in default frontmatter.
assert "description:" not in text
assert "extra:" not in text
await store.close()
print("✓ test_daily_write_drops_empty_frontmatter_values passed")
print("✓ test_daily_create_default_frontmatter_uses_slug_as_name passed")
asyncio.run(run())
def test_daily_write_rejects_invalid_slug():
"""Slug validation runs before any IO."""
def test_daily_create_rejects_invalid_slug():
"""Slug validation runs before any IO; no day folder is created on reject."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies([])
step = daily_write_step.DailyWriteStep(file_store=store)
step = daily_create_step.DailyCreateStep(file_store=store)
for bad in ("foo/bar", "foo:bar", "CON", "lpt9", "foo.", " bar"):
await step(slug=bad, date="2026-05-18", body="x")
await step(slug=bad, date="2026-05-18")
assert step.context.response.success is False, f"expected reject for {bad!r}"
# No day folder should be created on rejection.
assert not (Path(tmp) / "daily" / "2026-05-18").exists()
await store.close()
print("✓ test_daily_write_rejects_invalid_slug passed")
print("✓ test_daily_create_rejects_invalid_slug passed")
asyncio.run(run())
def test_daily_write_overwrite_default_is_false_create_then_skip():
"""First call (file absent) creates; second call without overwrite skips — proves the
overwrite=False default mirrors the old daily_resolve idempotent probe."""
def test_daily_create_rejects_empty_slug():
"""Empty / missing slug is rejected with a clear message."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies([])
step = daily_write_step.DailyWriteStep(file_store=store)
step = daily_create_step.DailyCreateStep(file_store=store)
await step(slug="", date="2026-05-18")
assert step.context.response.success is False
assert "slug" in (step.context.response.answer or "").lower()
await store.close()
print("✓ test_daily_create_rejects_empty_slug passed")
await step(slug="probe", date="2026-05-18", body="original")
asyncio.run(run())
def test_daily_create_then_skip_round_trip():
"""First call provisions, second call is an idempotent no-op write."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies([])
step = daily_create_step.DailyCreateStep(file_store=store)
await step(slug="probe", date="2026-05-18")
first = _metadata(step)
assert first["created"] is True
assert first["overwritten"] is False
await step(slug="probe", date="2026-05-18", body="ignored")
second = _metadata(step)
assert second["created"] is False
assert second["overwritten"] is False
note = Path(tmp) / "daily" / "2026-05-18" / "probe.md"
assert "original" in note.read_text(encoding="utf-8")
assert "ignored" not in note.read_text(encoding="utf-8")
before = note.read_text(encoding="utf-8")
await step(slug="probe", date="2026-05-18")
second = _metadata(step)
assert second["created"] is False
assert note.read_text(encoding="utf-8") == before
await store.close()
print("✓ test_daily_write_overwrite_default_is_false_create_then_skip passed")
asyncio.run(run())
def test_daily_write_rejects_non_dict_frontmatter():
"""``frontmatter`` must be a dict when supplied."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies([])
step = daily_write_step.DailyWriteStep(file_store=store)
await step(slug="ok", date="2026-05-18", frontmatter="not a dict")
assert step.context.response.success is False
assert "dict" in (step.context.response.answer or "")
await store.close()
print("✓ test_daily_write_rejects_non_dict_frontmatter passed")
asyncio.run(run())
def test_daily_write_refresh_index_can_be_disabled():
"""``refresh_index=False`` ⇒ note written but day index untouched."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies([])
step = daily_write_step.DailyWriteStep(file_store=store)
await step(
slug="solo",
date="2026-05-18",
body="x",
refresh_index=False,
)
note = Path(tmp) / "daily" / "2026-05-18" / "solo.md"
assert note.is_file()
assert "index" not in _metadata(step)
assert not (Path(tmp) / "daily" / "2026-05-18.md").exists()
await store.close()
print("✓ test_daily_write_refresh_index_can_be_disabled passed")
asyncio.run(run())
def test_daily_write_round_trips_with_daily_read():
"""A note written via daily_write must be retrievable via daily_read with the same data."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _make_store_with_dailies([])
body = "## Plan\nstep one\nstep two\n"
fm = {"name": "Round-trip", "description": "CRUD smoke"}
await daily_write_step.DailyWriteStep(file_store=store)(
slug="round-trip",
date="2026-05-18",
body=body,
frontmatter=fm,
)
read_step = daily_read_step.DailyReadStep(file_store=store)
await read_step(slug="round-trip", date="2026-05-18")
assert read_step.context.response.answer.strip() == body.strip()
assert _metadata(read_step)["frontmatter"] == fm
await store.close()
print("✓ test_daily_write_round_trips_with_daily_read passed")
print("✓ test_daily_create_then_skip_round_trip passed")
asyncio.run(run())
@ -678,7 +437,7 @@ def test_day_index_lists_each_note():
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = LocalFileStore(store_name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_model="")
await store.start()
await _seed_note("2026-05-18", "alpha", name="Alpha Project")
await _seed_note("2026-05-18", "beta", name="Beta Project")
@ -696,13 +455,11 @@ def test_day_index_lists_each_note():
def test_day_index_includes_note_descriptions():
"""Note ``description`` fields land in the rendered block so the
index reads as a one-glance "what's happening today" summary.
"""
"""Each note line inlines the full frontmatter (single-line, key: value pairs)."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = LocalFileStore(store_name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_model="")
await store.start()
cases = [
("alpha", "Alpha Project", "实现 JWT auth 中间件,迁移 session middleware"),
@ -714,9 +471,14 @@ def test_day_index_includes_note_descriptions():
await daily_reindex_step.DailyReindexStep(file_store=store)(date="2026-05-18")
text = _day_index_text(tmp, "2026-05-18")
assert "Alpha Project — 实现 JWT auth 中间件" in text
assert "调研增值税新政对 SaaS 的影响" in text
assert " Gamma\n" in text or text.rstrip().endswith("Gamma")
# name + description inline on the same line as the wikilink
assert "[[daily/2026-05-18/alpha.md]] name: Alpha Project description: 实现 JWT auth 中间件" in text
assert "[[daily/2026-05-18/beta.md]] name: beta description: 调研增值税新政对 SaaS 的影响" in text
# gamma has no description → only name is emitted, no trailing `description:` cruft
assert "[[daily/2026-05-18/gamma.md]] name: Gamma\n" in text or text.rstrip().endswith(
"[[daily/2026-05-18/gamma.md]] name: Gamma",
)
assert "description:" not in text.split("[[daily/2026-05-18/gamma.md]]")[1].split("\n")[0]
await store.close()
print("✓ test_day_index_includes_note_descriptions passed")
@ -739,19 +501,20 @@ def test_day_index_description_is_note_count():
text = _day_index_text(tmp, "2026-05-18")
assert "description:" in text
assert "2 篇笔记" in text
assert "2 note(s) today." in text
await store.close()
print("✓ test_day_index_description_is_note_count passed")
asyncio.run(run())
def test_day_index_preserves_manual_segment():
"""The ``## 备忘`` (manual) segment is preserved across refreshes."""
def test_day_index_preserves_user_content_outside_marker():
"""Any user-authored content sitting outside the auto markers is
preserved verbatim across refreshes."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = LocalFileStore(store_name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_model="")
await store.start()
await _seed_note("2026-05-18", "alpha")
reindex = daily_reindex_step.DailyReindexStep(file_store=store)
@ -759,20 +522,19 @@ def test_day_index_preserves_manual_segment():
index_path = Path(tmp) / "daily" / "2026-05-18.md"
text = index_path.read_text(encoding="utf-8")
patched = text.replace(
"(人工记录区,刷新索引时不会动)",
"MY HAND-WRITTEN NOTE\n这是我手写的备忘,不该被覆盖",
)
index_path.write_text(patched, encoding="utf-8")
# Append user content AFTER the auto block; it should survive refresh.
user_block = "\n\n## 我的笔记\nMY HAND-WRITTEN NOTE\n这是我手写的备忘,不该被覆盖\n"
index_path.write_text(text.rstrip() + user_block, encoding="utf-8")
await _seed_note("2026-05-18", "beta")
await reindex(date="2026-05-18")
after = index_path.read_text(encoding="utf-8")
assert "MY HAND-WRITTEN NOTE" in after
assert "这是我手写的备忘" in after
assert "## 我的笔记" in after
assert "[[daily/2026-05-18/beta.md]]" in after
await store.close()
print("✓ test_day_index_preserves_manual_segment passed")
print("✓ test_day_index_preserves_user_content_outside_marker passed")
asyncio.run(run())
@ -840,31 +602,22 @@ if __name__ == "__main__":
print("\n=== Daily step tests ===")
test_daily_list_default_date_is_today()
test_daily_list_filters_by_date()
test_daily_list_returns_path_slug_name_description()
test_daily_list_returns_path_slug_metadata()
test_daily_list_ignores_subdirectories()
test_daily_list_empty_when_no_daily_dir()
test_daily_list_does_not_refresh_index()
test_daily_list_response_shape()
test_daily_read_returns_body_and_frontmatter_dict()
test_daily_read_default_date_is_today()
test_daily_read_missing_file_reports_exists_false()
test_daily_read_rejects_invalid_slug()
test_daily_read_empty_frontmatter_dict()
test_daily_write_creates_note_and_refreshes_index()
test_daily_write_create_mode_is_idempotent()
test_daily_write_overwrite_mode_replaces_existing()
test_daily_write_default_frontmatter_uses_slug_as_name()
test_daily_write_default_body_is_empty()
test_daily_write_drops_empty_frontmatter_values()
test_daily_write_rejects_invalid_slug()
test_daily_write_overwrite_default_is_false_create_then_skip()
test_daily_write_rejects_non_dict_frontmatter()
test_daily_write_refresh_index_can_be_disabled()
test_daily_write_round_trips_with_daily_read()
test_daily_create_provisions_note_and_refreshes_index()
test_daily_create_is_idempotent_on_existing()
test_daily_create_default_date_is_today()
test_daily_create_default_frontmatter_uses_slug_as_name()
test_daily_create_rejects_invalid_slug()
test_daily_create_rejects_empty_slug()
test_daily_create_then_skip_round_trip()
test_day_index_lists_each_note()
test_day_index_includes_note_descriptions()
test_day_index_description_is_note_count()
test_day_index_preserves_manual_segment()
test_day_index_preserves_user_content_outside_marker()
test_daily_reindex_returns_write_view()
test_daily_reindex_created_flag_flips_on_rerun()
print("\nAll tests passed!")

View file

@ -0,0 +1,161 @@
"""Tests for FileCatalog backends."""
# pylint: disable=protected-access
import asyncio
import os
import tempfile
import pytest
from reme4.components.file_catalog import LocalFileCatalog
from reme4.schema import FileNode
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 make_node(path: str, mtime: float = 1.0) -> FileNode:
"""Build a FileNode with the given path and mtime for fixture use."""
return FileNode(path=path, st_mtime=mtime)
# All backends should satisfy the same BaseFileCatalog contract.
BACKENDS = [LocalFileCatalog]
@pytest.mark.parametrize("backend_cls", BACKENDS)
def test_upsert_and_get_nodes(backend_cls):
"""upsert stores nodes; get_nodes returns by path or all."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
catalog = backend_cls()
await catalog.start()
await catalog.upsert([make_node("a.md"), make_node("b.md")])
got_all = await catalog.get_nodes()
assert {n.path for n in got_all} == {"a.md", "b.md"}
got_one = await catalog.get_nodes(["a.md"])
assert len(got_one) == 1
assert got_one[0].path == "a.md"
got_missing = await catalog.get_nodes(["nope.md"])
assert got_missing == []
await catalog.close()
print(f"✓ test_upsert_and_get_nodes[{backend_cls.__name__}] passed")
asyncio.run(run())
@pytest.mark.parametrize("backend_cls", BACKENDS)
def test_upsert_replaces_existing(backend_cls):
"""Re-upserting a node with the same path overwrites the prior entry."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
catalog = backend_cls()
await catalog.start()
await catalog.upsert([make_node("a.md", mtime=1.0)])
await catalog.upsert([make_node("a.md", mtime=2.0)])
nodes = await catalog.get_nodes(["a.md"])
assert len(nodes) == 1
assert nodes[0].st_mtime == 2.0
await catalog.close()
print(f"✓ test_upsert_replaces_existing[{backend_cls.__name__}] passed")
asyncio.run(run())
@pytest.mark.parametrize("backend_cls", BACKENDS)
def test_delete_single_and_list(backend_cls):
"""delete accepts both a single path and a list; missing paths are no-ops."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
catalog = backend_cls()
await catalog.start()
await catalog.upsert([make_node("a.md"), make_node("b.md"), make_node("c.md")])
await catalog.delete("a.md")
assert {n.path for n in await catalog.get_nodes()} == {"b.md", "c.md"}
await catalog.delete(["b.md", "ghost.md"])
assert {n.path for n in await catalog.get_nodes()} == {"c.md"}
await catalog.close()
print(f"✓ test_delete_single_and_list[{backend_cls.__name__}] passed")
asyncio.run(run())
@pytest.mark.parametrize("backend_cls", BACKENDS)
def test_get_nodes_empty_inputs(backend_cls):
"""get_nodes([]) returns []; get_nodes(None) returns all."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
catalog = backend_cls()
await catalog.start()
await catalog.upsert([make_node("a.md")])
assert await catalog.get_nodes([]) == []
assert len(await catalog.get_nodes(None)) == 1
await catalog.close()
print(f"✓ test_get_nodes_empty_inputs[{backend_cls.__name__}] passed")
asyncio.run(run())
@pytest.mark.parametrize("backend_cls", BACKENDS)
def test_persistence_roundtrip(backend_cls):
"""close() dumps; a fresh instance loads the same nodes from disk."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
c1 = backend_cls()
await c1.start()
await c1.upsert([make_node("a.md", mtime=10.0), make_node("b.md", mtime=20.0)])
await c1.close()
c2 = backend_cls()
await c2.start()
nodes = sorted(await c2.get_nodes(), key=lambda n: n.path)
assert [n.path for n in nodes] == ["a.md", "b.md"]
assert [n.st_mtime for n in nodes] == [10.0, 20.0]
await c2.close()
print(f"✓ test_persistence_roundtrip[{backend_cls.__name__}] passed")
asyncio.run(run())
if __name__ == "__main__":
print("\n=== FileCatalog Tests ===")
for backend in BACKENDS:
test_upsert_and_get_nodes(backend)
test_upsert_replaces_existing(backend)
test_delete_single_and_list(backend)
test_get_nodes_empty_inputs(backend)
test_persistence_roundtrip(backend)
print("\n所有测试通过!")

View file

@ -60,7 +60,7 @@ class temp_chdir:
async def make_store(store_name: str = "test_store", **kwargs) -> LocalFileStore:
"""Build a started LocalFileStore with embedding disabled (no OpenAI dep)."""
store = LocalFileStore(store_name=store_name, embedding_model="", **kwargs)
store = LocalFileStore(name=store_name, embedding_model="", **kwargs)
await store.start()
return store
@ -352,7 +352,7 @@ def _skip_if_no_faiss(name: str) -> bool:
async def make_faiss_store(store_name: str = "test_faiss", **kwargs) -> FaissLocalFileStore:
"""Build a started FaissLocalFileStore wired to FakeEmbeddingModel (no API calls)."""
store = FaissLocalFileStore(store_name=store_name, embedding_model="fake", **kwargs)
store = FaissLocalFileStore(name=store_name, embedding_model="fake", **kwargs)
fake = FakeEmbeddingModel()
# Replace the unresolved Dependency placeholder with a concrete instance and
# let start() cascade lifecycle to it via _owned.
@ -531,7 +531,7 @@ def test_faiss_disabled_without_embedding():
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = FaissLocalFileStore(store_name="disabled", embedding_model="")
store = FaissLocalFileStore(name="disabled", embedding_model="")
await store.start()
await store.upsert([make_file("a.md", "alpha")])

View file

@ -0,0 +1,957 @@
"""Tests for BaseKeywordIndex implementations (currently: BM25Index).
Covers full lifecycle, CRUD, retrieval, persistence, optimize and — as the focus
of this file — Chinese / English / mixed-language behaviour driven by the
default RegexTokenizer (Chinese split per char, English words lowercased,
single-char ASCII words dropped).
"""
# pylint: disable=protected-access
import asyncio
import os
import tempfile
import warnings
from reme4.components.keyword_index import BM25Index
from reme4.components.tokenizer import RegexTokenizer
warnings.filterwarnings("ignore", category=DeprecationWarning, module="jieba")
warnings.filterwarnings("ignore", category=DeprecationWarning, module="pkg_resources")
# --------------------------------------------------------------------------- #
# Helpers #
# --------------------------------------------------------------------------- #
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)
async def create_bm25(
k1: float = 1.5,
b: float = 0.75,
filter_stopwords: bool = False,
) -> BM25Index:
"""Create and start a BM25Index in cwd with a non-filtering RegexTokenizer.
Stopword filtering is off so short test words ("hello", "我", "的") survive.
"""
bm25 = BM25Index(k1=k1, b=b)
tokenizer = RegexTokenizer(filter_stopwords=filter_stopwords)
bm25.tokenizer = tokenizer
bm25._owned.append(tokenizer)
await bm25.start()
return bm25
def run(coro):
"""Tiny shorthand to avoid repeating asyncio.run wrappers."""
return asyncio.run(coro)
# --------------------------------------------------------------------------- #
# Initialisation & lifecycle #
# --------------------------------------------------------------------------- #
def test_basic_init():
"""Default constructor produces empty BM25 state."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = BM25Index()
assert bm25.k1 == 1.5
assert bm25.b == 0.75
assert bm25.index_version == "v1"
assert bm25.vocab == {}
assert not bm25.inverted_index
assert bm25.doc_meta == {}
assert bm25.n_docs == 0
assert bm25.total_len == 0
assert bm25.avg_len == 0.0
run(go())
def test_custom_params():
"""k1, b and index_version are honoured."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = BM25Index(k1=2.0, b=0.5, index_version="v2")
assert bm25.k1 == 2.0
assert bm25.b == 0.5
assert bm25.index_version == "v2"
run(go())
def test_index_file_raises_when_tokenizer_is_none():
"""index_file must raise when tokenizer is explicitly None."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = BM25Index()
bm25.tokenizer = None
try:
_ = bm25.index_file
except RuntimeError:
return
raise AssertionError("expected RuntimeError when tokenizer is None")
run(go())
def test_start_close_lifecycle():
"""start/close toggles is_started and runs underlying tokenizer."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
assert bm25.is_started
assert bm25.tokenizer is not None
await bm25.close()
assert not bm25.is_started
run(go())
def test_index_file_path_includes_tokenizer_and_version():
"""index_file path embeds tokenizer name + index_version."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
path = str(bm25.index_file)
assert "bm25_regex_v1.pkl" in path
await bm25.close()
run(go())
# --------------------------------------------------------------------------- #
# add_docs / delete_docs / update #
# --------------------------------------------------------------------------- #
def test_add_empty_dict_noop():
"""Adding an empty dict must not touch state."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({})
assert bm25.n_docs == 0
assert bm25.vocab == {}
await bm25.close()
run(go())
def test_add_single_doc():
"""A single doc populates length, vocab and metadata."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "hello world"})
assert bm25.n_docs == 1
assert bm25.total_len == 2 # 'hello', 'world'
assert bm25.avg_len == 2.0
assert set(bm25.vocab) == {"hello", "world"}
assert "d1" in bm25.doc_meta
assert bm25.doc_meta["d1"]["len"] == 2
await bm25.close()
run(go())
def test_add_multiple_docs_and_inverted_index():
"""Inverted index lists postings for every term across multiple docs."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs(
{
"d1": "hello world",
"d2": "hello python",
"d3": "world python",
},
)
assert bm25.n_docs == 3
inv = bm25.inverted_index
tid_hello = bm25.vocab["hello"]
tid_world = bm25.vocab["world"]
tid_python = bm25.vocab["python"]
assert set(inv[tid_hello]) == {"d1", "d2"}
assert set(inv[tid_world]) == {"d1", "d3"}
assert set(inv[tid_python]) == {"d2", "d3"}
await bm25.close()
run(go())
def test_add_doc_empty_or_whitespace_is_skipped():
"""Empty / whitespace-only content yields no tokens and is silently dropped."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "", "d2": " ", "d3": "\n\t"})
assert bm25.n_docs == 0
await bm25.close()
run(go())
def test_update_existing_doc_swaps_content():
"""Re-adding same doc_id replaces tokens; old terms no longer match it."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "hello world python"})
old_len = bm25.total_len
await bm25.add_docs({"d1": "java"})
assert bm25.n_docs == 1
assert bm25.total_len != old_len
assert bm25.doc_meta["d1"]["len"] == 1
# Old term must no longer return d1.
assert "d1" not in await bm25.retrieve("hello", limit=5)
# New term does.
assert "d1" in await bm25.retrieve("java", limit=5)
await bm25.close()
run(go())
def test_delete_single_doc():
"""delete_docs removes a single doc from retrieval and meta."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "hello world", "d2": "hello python"})
assert bm25.n_docs == 2
await bm25.delete_docs(["d1"])
assert bm25.n_docs == 1
assert "d1" not in bm25.doc_meta
assert "d2" in bm25.doc_meta
results = await bm25.retrieve("hello", limit=2)
assert "d1" not in results
assert "d2" in results
await bm25.close()
run(go())
def test_delete_multiple_docs():
"""delete_docs handles a batch list."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({f"d{i}": "hello world" for i in range(5)})
assert bm25.n_docs == 5
await bm25.delete_docs(["d0", "d2", "d4"])
assert bm25.n_docs == 2
assert set(bm25.doc_meta) == {"d1", "d3"}
await bm25.close()
run(go())
def test_delete_nonexistent_is_noop():
"""Deleting unknown doc_ids must not raise."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "hello"})
await bm25.delete_docs(["nope", "still_nope"])
assert bm25.n_docs == 1
await bm25.close()
run(go())
def test_re_add_after_delete():
"""Adding a doc_id back after deletion yields a fresh idx and is retrievable."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "hello"})
await bm25.delete_docs(["d1"])
assert bm25.n_docs == 0
await bm25.add_docs({"d1": "world"})
assert bm25.n_docs == 1
assert "d1" in await bm25.retrieve("world", limit=1)
await bm25.close()
run(go())
# --------------------------------------------------------------------------- #
# Retrieval #
# --------------------------------------------------------------------------- #
def test_retrieve_empty_index():
"""Retrieving from an empty index returns {}."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
assert await bm25.retrieve("python", limit=3) == {}
await bm25.close()
run(go())
def test_retrieve_empty_or_unknown_query():
"""Empty queries and out-of-vocab queries both return {}."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "hello world"})
assert await bm25.retrieve("", limit=3) == {}
assert await bm25.retrieve("zzzunknownxyz", limit=3) == {}
await bm25.close()
run(go())
def test_retrieve_limit_caps_results():
"""retrieve honours `limit`."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({f"d{i}": f"python lang {i}" for i in range(10)})
assert len(await bm25.retrieve("python", limit=3)) == 3
assert len(await bm25.retrieve("python", limit=5)) == 5
# limit greater than matches: bounded by positive matches.
assert len(await bm25.retrieve("python", limit=99)) == 10
await bm25.close()
run(go())
def test_retrieve_score_ordering_by_tf():
"""A doc with higher term frequency for the query token outranks others."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs(
{
"high": "python python python",
"mid": "python python other",
"low": "python alpha beta",
},
)
results = await bm25.retrieve("python", limit=3)
assert list(results.keys()) == ["high", "mid", "low"]
scores = list(results.values())
assert scores[0] >= scores[1] >= scores[2]
await bm25.close()
run(go())
def test_retrieve_idf_favours_rare_terms():
"""In a query of {common, rare}, the doc containing the rare term wins."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
# 'common' appears everywhere → low IDF.
# 'rare' appears in only one doc → high IDF.
docs = {f"d{i}": "common filler text" for i in range(10)}
docs["target"] = "common rare term"
await bm25.add_docs(docs)
results = await bm25.retrieve("common rare", limit=3)
assert next(iter(results)) == "target"
await bm25.close()
run(go())
def test_retrieve_length_normalization():
"""With b=0.75 (default), a much longer doc with same tf scores lower."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs(
{
"short": "python",
"long": "python " + " ".join(f"w{i}" for i in range(50)),
},
)
results = await bm25.retrieve("python", limit=2)
assert results["short"] > results["long"]
await bm25.close()
run(go())
def test_retrieve_duplicate_query_tokens_dont_double_count():
"""Repeating the same query token should not boost its contribution."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "python rocks"})
once = await bm25.retrieve("python", limit=1)
many = await bm25.retrieve("python python python", limit=1)
assert once["d1"] == many["d1"]
await bm25.close()
run(go())
# --------------------------------------------------------------------------- #
# Chinese / English / mixed-language behaviour (focus) #
# --------------------------------------------------------------------------- #
def test_chinese_only_corpus():
"""Pure Chinese corpus indexes per-character and retrieves correctly."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs(
{
"d1": "我爱北京天安门",
"d2": "北京是中国的首都",
"d3": "上海的天气很好",
},
)
# Regex tokenizer splits Chinese per character.
assert "北" in bm25.vocab
assert "京" in bm25.vocab
# Query "北京" → two tokens, both d1 and d2 match; d3 does not.
results = await bm25.retrieve("北京", limit=3)
assert set(results) == {"d1", "d2"}
await bm25.close()
run(go())
def test_english_only_corpus_is_lowercased():
"""English tokens are lowercased so case-insensitive retrieval works."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs(
{
"d1": "Python Programming Language",
"d2": "Java Programming Language",
},
)
assert "python" in bm25.vocab
assert "Python" not in bm25.vocab
r_upper = await bm25.retrieve("PYTHON", limit=2)
r_lower = await bm25.retrieve("python", limit=2)
assert r_upper == r_lower
assert "d1" in r_upper
await bm25.close()
run(go())
def test_single_char_english_dropped():
"""RegexTokenizer's \\w\\w+ pattern drops single-letter ASCII tokens like 'I'."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "I love Beijing"})
# 'I' must not appear; 'love' and 'beijing' must.
assert "i" not in bm25.vocab
assert "love" in bm25.vocab
assert "beijing" in bm25.vocab
# Querying with just "I" returns nothing.
assert await bm25.retrieve("I", limit=1) == {}
await bm25.close()
run(go())
def test_mixed_doc_chinese_query_matches():
"""A Chinese query hits docs containing those Chinese chars even when mixed."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs(
{
"d1": "Python 是一种编程语言",
"d2": "Java 编程语言",
"d3": "Python 数据分析",
},
)
results = await bm25.retrieve("编程", limit=3)
assert set(results) >= {"d1", "d2"}
assert "d3" not in results
await bm25.close()
run(go())
def test_mixed_doc_english_query_matches():
"""An English query hits the right mixed-language docs."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs(
{
"d1": "Python 是一种编程语言",
"d2": "Java 编程语言",
"d3": "Python 数据分析",
},
)
results = await bm25.retrieve("python", limit=3)
assert set(results) == {"d1", "d3"}
assert "d2" not in results
await bm25.close()
run(go())
def test_mixed_query_combines_chinese_and_english_signal():
"""A query mixing English and Chinese aggregates IDF×tf contributions."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs(
{
"py_cn": "Python 编程", # matches both 'python' and '编','程'
"py_only": "Python tutorial", # matches only 'python'
"cn_only": "编程入门", # matches only '编','程'
# Avoid Chinese chars that the query splits into ('编','程') — '教程' would
# leak '程' into 'other' and pollute IDF, so use unrelated chars only.
"other": "Java 教学",
},
)
results = await bm25.retrieve("Python 编程", limit=4)
# py_cn should rank highest because it matches both branches.
assert next(iter(results)) == "py_cn"
# 'other' should not appear.
assert "other" not in results
# Both unimodal matches should still appear.
assert "py_only" in results and "cn_only" in results
await bm25.close()
run(go())
def test_mixed_doc_more_matches_outrank_fewer():
"""Doc covering more query tokens (Chinese+English) outranks partial matches."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs(
{
"full": "machine learning 机器 学习",
"en_only": "machine learning algorithm",
"cn_only": "机器 学习 算法",
},
)
results = await bm25.retrieve("machine 机器", limit=3)
# full has both English and Chinese hits → highest score.
assert next(iter(results)) == "full"
await bm25.close()
run(go())
def test_unicode_word_with_digits_preserved():
"""Alphanumeric tokens like 'iphone15' stay whole; trailing Chinese still split."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs(
{
"d1": "iPhone15 Pro 售价 9999 元",
"d2": "Android 旗舰 999 元",
},
)
assert "iphone15" in bm25.vocab
assert "9999" in bm25.vocab
assert "元" in bm25.vocab
r1 = await bm25.retrieve("iphone15", limit=2)
assert list(r1) == ["d1"]
r2 = await bm25.retrieve("元", limit=2)
assert set(r2) == {"d1", "d2"}
await bm25.close()
run(go())
def test_chinese_punctuation_ignored():
"""CJK punctuation should not produce tokens."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "你好,世界!这是 Python。"})
for sym in [",", "!", "。"]:
assert sym not in bm25.vocab
assert "你" in bm25.vocab
assert "python" in bm25.vocab
await bm25.close()
run(go())
def test_mixed_persistence_roundtrip():
"""A mixed-language index round-trips through dump/load with identical retrieval."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs(
{
"d1": "Python 编程语言",
"d2": "Java 编程",
"d3": "数据分析 with Python",
},
)
before = await bm25.retrieve("Python 编程", limit=3)
await bm25.close() # close triggers dump
bm25_2 = await create_bm25() # start triggers load
assert bm25_2.n_docs == 3
after = await bm25_2.retrieve("Python 编程", limit=3)
assert before == after
await bm25_2.close()
run(go())
# --------------------------------------------------------------------------- #
# Persistence #
# --------------------------------------------------------------------------- #
def test_dump_load_roundtrip_preserves_state():
"""dump → fresh instance → load reconstructs vocab, postings and params."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25(k1=2.0, b=0.4)
await bm25.add_docs(
{
"d1": "hello world",
"d2": "hello python",
"d3": "programming language",
},
)
old_vocab = dict(bm25.vocab)
old_meta = {k: dict(v) for k, v in bm25.doc_meta.items()}
await bm25.dump()
await bm25.close()
bm25_2 = await create_bm25() # default k1/b — load must overwrite
assert bm25_2.vocab == old_vocab
assert bm25_2.n_docs == 3
assert set(bm25_2.doc_meta) == set(old_meta)
assert bm25_2.k1 == 2.0
assert bm25_2.b == 0.4
await bm25_2.close()
run(go())
def test_load_missing_file_keeps_empty_state():
"""Calling load() with no file on disk is a no-op."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
# No add_docs, nothing persisted.
assert not bm25.index_file.exists()
await bm25.load()
assert bm25.n_docs == 0
await bm25.close()
run(go())
def test_load_corrupt_file_resets_index():
"""A corrupt pickle on disk is reported, deleted, and the index is cleared."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "hello"})
await bm25.dump()
# Corrupt the file.
bm25.index_file.write_bytes(b"not a pickle")
await bm25.load()
assert bm25.n_docs == 0
assert bm25.vocab == {}
assert not bm25.index_file.exists()
await bm25.close()
run(go())
# --------------------------------------------------------------------------- #
# clear / optimize / reset_index #
# --------------------------------------------------------------------------- #
def test_clear_wipes_everything():
"""clear() empties state and removes the index file."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "hello", "d2": "world"})
await bm25.dump()
assert bm25.index_file.exists()
await bm25.clear()
assert bm25.n_docs == 0
assert bm25.vocab == {}
assert bm25.inverted_index == {}
assert bm25.doc_meta == {}
assert bm25.total_len == 0
assert bm25._idf_cache == {}
assert not bm25.index_file.exists()
await bm25.close()
run(go())
def test_optimize_drops_deleted_only_terms():
"""After deleting, optimize_index drops vocab entries that no live doc uses."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs(
{
"d1": "alpha unique_to_d1",
"d2": "alpha beta",
},
)
assert "unique_to_d1" in bm25.vocab
await bm25.delete_docs(["d1"])
await bm25.optimize_index()
assert bm25.n_docs == 1
assert "d2" in bm25.doc_meta
# Term that only existed in d1 is gone.
assert "unique_to_d1" not in bm25.vocab
# Shared/own terms of d2 survive.
assert "alpha" in bm25.vocab and "beta" in bm25.vocab
# Retrieval still works correctly.
assert "d2" in await bm25.retrieve("alpha", limit=1)
await bm25.close()
run(go())
def test_optimize_when_all_deleted_clears_index():
"""optimize_index on a fully-deleted state collapses to empty index."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "hello world"})
await bm25.delete_docs(["d1"])
await bm25.optimize_index()
assert bm25.n_docs == 0
assert bm25.vocab == {}
assert bm25.inverted_index == {}
await bm25.close()
run(go())
def test_optimize_noop_when_no_deletions():
"""With nothing deleted, optimize_index leaves state intact."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "hello world"})
vocab_before = dict(bm25.vocab)
await bm25.optimize_index()
assert bm25.vocab == vocab_before
assert bm25.n_docs == 1
await bm25.close()
run(go())
def test_reset_index_replaces_all_docs():
"""reset_index (inherited from base) wipes and re-adds in one call."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "old content"})
await bm25.reset_index({"d2": "new content"})
assert bm25.n_docs == 1
assert "d2" in bm25.doc_meta
assert "d1" not in bm25.doc_meta
await bm25.close()
run(go())
# --------------------------------------------------------------------------- #
# Internal invariants #
# --------------------------------------------------------------------------- #
def test_idf_cache_populates_and_invalidates():
"""_get_idf caches results; add/delete clear the cache."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "hello world", "d2": "hello python"})
tid_hello = bm25.vocab["hello"]
idf1 = bm25._get_idf(tid_hello)
assert tid_hello in bm25._idf_cache
assert bm25._get_idf(tid_hello) == idf1
# Mutating the index must invalidate the cache.
await bm25.add_docs({"d3": "hello there"})
assert bm25._idf_cache == {}
await bm25.delete_docs(["d1"])
# delete_docs also clears cache; populate again then trigger via add.
_ = bm25._get_idf(bm25.vocab["hello"])
assert bm25._idf_cache # non-empty now
await bm25.add_docs({"d4": "x y z"})
assert bm25._idf_cache == {}
await bm25.close()
run(go())
def test_avg_len_tracks_live_docs_only():
"""avg_len excludes deleted docs.
Note: RegexTokenizer's `\\w\\w+` pattern drops 1-letter words, so we
use multi-letter tokens to keep length math predictable.
"""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
assert bm25.avg_len == 0.0
await bm25.add_docs({"d1": "alpha beta gamma delta"}) # 4 tokens
await bm25.add_docs({"d2": "alpha beta"}) # 2 tokens
assert bm25.avg_len == 3.0
await bm25.delete_docs(["d1"])
assert bm25.avg_len == 2.0
await bm25.close()
run(go())
def test_deleted_docs_excluded_from_scoring():
"""A deleted doc must score 0 and never appear in retrieve()."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "python", "d2": "python", "d3": "python"})
await bm25.delete_docs(["d2"])
results = await bm25.retrieve("python", limit=10)
assert set(results) == {"d1", "d3"}
assert all(s > 0 for s in results.values())
await bm25.close()
run(go())
def test_inverted_index_hides_deleted_postings():
"""inverted_index view skips postings whose doc is deleted."""
async def go():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
bm25 = await create_bm25()
await bm25.add_docs({"d1": "alpha", "d2": "alpha beta"})
tid_alpha = bm25.vocab["alpha"]
await bm25.delete_docs(["d1"])
inv = bm25.inverted_index
# 'alpha' posting now contains only the live doc.
assert tid_alpha in inv
assert set(inv[tid_alpha]) == {"d2"}
await bm25.close()
run(go())
# --------------------------------------------------------------------------- #
# Manual runner #
# --------------------------------------------------------------------------- #
if __name__ == "__main__":
import inspect
import sys
mod = sys.modules[__name__]
tests = [(name, obj) for name, obj in inspect.getmembers(mod, inspect.isfunction) if name.startswith("test_")]
print(f"\n=== BaseKeywordIndex / BM25Index Tests ({len(tests)}) ===\n")
failed = []
for name, fn in tests:
try:
fn()
print(f"✓ {name}")
except Exception as exc: # noqa: BLE001
print(f"✗ {name}: {exc!r}")
failed.append(name)
print()
if failed:
print(f"FAILED: {len(failed)} / {len(tests)}")
for n in failed:
print(f" - {n}")
sys.exit(1)
print(f"所有 {len(tests)} 项测试通过!")

View file

@ -0,0 +1,254 @@
"""Tests for ``reme4.utils.link_expansion``.
Two pure helpers:
* ``expand_links(file_store, paths, max_per_direction)`` — fetch
outlinks / inlinks for each path with neighbor meta attached.
* ``render_expansion_lines(expansion)`` — turn one path's expansion
dict into the indented ``→`` / ``←`` block used by SearchStep
answers.
"""
# pylint: disable=protected-access
import asyncio
import os
import tempfile
import warnings
from pathlib import Path
from reme4.components.file_store import LocalFileStore
from reme4.schema import FileFrontMatter, FileNode
from reme4.utils.link_expansion import expand_links, render_expansion_lines
from reme4.utils.wikilink_handler import WikilinkHandler
warnings.filterwarnings("ignore", category=DeprecationWarning, module="jieba")
warnings.filterwarnings("ignore", category=DeprecationWarning, module="pkg_resources")
class temp_chdir:
"""Test helper: chdir to ``path`` on enter, restore previous cwd 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)
async def _store_with(files: dict[str, dict]) -> LocalFileStore:
"""LocalFileStore seeded with files + parsed wikilinks + optional frontmatter meta.
Each value: ``{"body": str, "name": str?, "description": str?}``. The
body is written to disk and wikilinks are extracted; ``name`` /
``description`` populate FileFrontMatter so neighbor meta lookups
have something to surface.
"""
store = LocalFileStore(name="t", embedding_model="")
await store.start()
nodes: list[FileNode] = []
root = Path.cwd()
for rel, spec in files.items():
body = spec["body"]
abs_path = root / rel
abs_path.parent.mkdir(parents=True, exist_ok=True)
abs_path.write_text(body, encoding="utf-8")
fm = FileFrontMatter(
name=spec.get("name", ""),
description=spec.get("description", ""),
)
nodes.append(
FileNode(
path=rel,
st_mtime=abs_path.stat().st_mtime,
links=WikilinkHandler.extract_links(body, rel),
front_matter=fm,
),
)
if nodes:
await store.file_graph.upsert_nodes(nodes)
return store
# -- expand_links -------------------------------------------------------------
def test_expand_links_empty_paths_short_circuits():
"""No paths ⇒ empty dict, no file_store calls needed."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = LocalFileStore(name="t", embedding_model="")
await store.start()
result = await expand_links(store, [])
assert result == {}
await store.close()
print("✓ test_expand_links_empty_paths_short_circuits passed")
asyncio.run(run())
def test_expand_links_returns_outlinks_and_inlinks_with_meta():
"""A.md links to B.md ⇒ A has B as outlink, B has A as inlink, meta surfaced."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _store_with(
{
"A.md": {"body": "See [[B.md]] for details.", "name": "A Doc", "description": "alpha"},
"B.md": {"body": "End node.", "name": "B Doc", "description": "beta"},
},
)
result = await expand_links(store, ["A.md", "B.md"])
assert set(result.keys()) == {"A.md", "B.md"}
a_out = result["A.md"]["outlinks"]
assert len(a_out) == 1
assert a_out[0]["path"] == "B.md"
assert a_out[0]["meta"] == {"name": "B Doc", "description": "beta"}
assert a_out[0]["edges"] == [{"predicate": None, "anchor": None}]
assert result["A.md"]["inlinks"] == []
b_in = result["B.md"]["inlinks"]
assert len(b_in) == 1
assert b_in[0]["path"] == "A.md"
assert b_in[0]["meta"] == {"name": "A Doc", "description": "alpha"}
assert result["B.md"]["outlinks"] == []
await store.close()
print("✓ test_expand_links_returns_outlinks_and_inlinks_with_meta passed")
asyncio.run(run())
def test_expand_links_max_per_direction_caps_neighbors():
"""max_per_direction=2 ⇒ only first two distinct neighbors per direction kept."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _store_with(
{
"hub.md": {
"body": "[[a.md]] [[b.md]] [[c.md]] [[d.md]]",
},
"a.md": {"body": "a"},
"b.md": {"body": "b"},
"c.md": {"body": "c"},
"d.md": {"body": "d"},
},
)
result = await expand_links(store, ["hub.md"], max_per_direction=2)
out = result["hub.md"]["outlinks"]
assert len(out) == 2
assert [n["path"] for n in out] == ["a.md", "b.md"]
await store.close()
print("✓ test_expand_links_max_per_direction_caps_neighbors passed")
asyncio.run(run())
def test_expand_links_node_without_meta_returns_empty_meta_dict():
"""Neighbor with no frontmatter name/description ⇒ meta = {}."""
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = await _store_with(
{
"src.md": {"body": "[[dst.md]]"},
"dst.md": {"body": "no meta"},
},
)
result = await expand_links(store, ["src.md"])
assert result["src.md"]["outlinks"][0]["meta"] == {}
await store.close()
print("✓ test_expand_links_node_without_meta_returns_empty_meta_dict passed")
asyncio.run(run())
# -- render_expansion_lines ---------------------------------------------------
def test_render_expansion_lines_empty_input_yields_empty_list():
"""Both directions empty ⇒ no lines."""
assert not render_expansion_lines({})
assert not render_expansion_lines({"outlinks": [], "inlinks": []})
print("✓ test_render_expansion_lines_empty_input_yields_empty_list passed")
def test_render_expansion_lines_outlinks_only():
"""Single outlink with meta + plain edge renders as 3 lines."""
expansion = {
"outlinks": [
{
"path": "B.md",
"meta": {"name": "B", "description": "beta"},
"edges": [{"predicate": None, "anchor": None}],
},
],
"inlinks": [],
}
lines = render_expansion_lines(expansion)
assert lines == [
" outlinks (1):",
' → B.md name="B" description="beta"',
" via plain",
]
print("✓ test_render_expansion_lines_outlinks_only passed")
def test_render_expansion_lines_inlinks_only_with_predicate_and_anchor():
"""Inlink edge with predicate + anchor renders via descriptor."""
expansion = {
"outlinks": [],
"inlinks": [
{
"path": "src.md",
"meta": {},
"edges": [{"predicate": "references", "anchor": "intro"}],
},
],
}
lines = render_expansion_lines(expansion)
assert lines == [
" inlinks (1):",
" ← src.md (no meta)",
" via predicate=references, anchor=#intro",
]
print("✓ test_render_expansion_lines_inlinks_only_with_predicate_and_anchor passed")
def test_render_expansion_lines_both_directions_in_order():
"""outlinks block precedes inlinks block."""
expansion = {
"outlinks": [
{"path": "out.md", "meta": {"name": "Out"}, "edges": [{"predicate": None, "anchor": None}]},
],
"inlinks": [
{"path": "in.md", "meta": {"description": "incoming"}, "edges": [{"predicate": None, "anchor": None}]},
],
}
lines = render_expansion_lines(expansion)
assert lines[0] == " outlinks (1):"
assert lines[3] == " inlinks (1):"
assert lines[1].lstrip().startswith("→")
assert lines[4].lstrip().startswith("←")
print("✓ test_render_expansion_lines_both_directions_in_order passed")
if __name__ == "__main__":
test_expand_links_empty_paths_short_circuits()
test_expand_links_returns_outlinks_and_inlinks_with_meta()
test_expand_links_max_per_direction_caps_neighbors()
test_expand_links_node_without_meta_returns_empty_meta_dict()
test_render_expansion_lines_empty_input_yields_empty_list()
test_render_expansion_lines_outlinks_only()
test_render_expansion_lines_inlinks_only_with_predicate_and_anchor()
test_render_expansion_lines_both_directions_in_order()

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