mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-06 02:48:22 +00:00
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:
parent
83bfddb4a4
commit
a4efc0f776
102 changed files with 5560 additions and 4820 deletions
|
|
@ -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/
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"] = {}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
9
reme4/components/file_catalog/__init__.py
Normal file
9
reme4/components/file_catalog/__init__.py
Normal file
|
|
@ -0,0 +1,9 @@
|
|||
"""file catalog"""
|
||||
|
||||
from .base_file_catalog import BaseFileCatalog
|
||||
from .local_file_catalog import LocalFileCatalog
|
||||
|
||||
__all__ = [
|
||||
"BaseFileCatalog",
|
||||
"LocalFileCatalog",
|
||||
]
|
||||
39
reme4/components/file_catalog/base_file_catalog.py
Normal file
39
reme4/components/file_catalog/base_file_catalog.py
Normal 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."""
|
||||
62
reme4/components/file_catalog/local_file_catalog.py
Normal file
62
reme4/components/file_catalog/local_file_catalog.py
Normal 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)
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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*."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -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)})"
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
45
reme4/config/demo.yaml
Normal 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
474
reme4/config/qwenpaw.yaml
Normal 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
|
||||
|
|
@ -22,6 +22,8 @@ class ComponentEnum(str, Enum):
|
|||
|
||||
FILE_GRAPH = "file_graph"
|
||||
|
||||
FILE_CATALOG = "file_catalog"
|
||||
|
||||
KEYWORD_INDEX = "keyword_index"
|
||||
|
||||
SERVICE = "service"
|
||||
|
|
|
|||
|
|
@ -13,5 +13,7 @@ class LinkScopeEnum(str, Enum):
|
|||
"""
|
||||
|
||||
REAL = "real"
|
||||
|
||||
VIRTUAL = "virtual"
|
||||
|
||||
ALL = "all"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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)})
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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)})
|
||||
|
|
@ -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
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
@ -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,
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
0
reme4/steps/file_io/__init__.py
Normal file
0
reme4/steps/file_io/__init__.py
Normal file
492
reme4/steps/file_io/_file_io.py
Normal file
492
reme4/steps/file_io/_file_io.py
Normal 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,
|
||||
}
|
||||
94
reme4/steps/file_io/daily_create.py
Normal file
94
reme4/steps/file_io/daily_create.py
Normal 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
|
||||
|
|
@ -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})
|
||||
|
|
@ -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))
|
||||
|
|
@ -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:
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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
104
reme4/steps/file_io/list.py
Normal 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
|
||||
|
|
@ -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
115
reme4/steps/file_io/read.py
Normal 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
|
||||
|
|
@ -25,7 +25,6 @@ from pathlib import Path
|
|||
import frontmatter
|
||||
|
||||
from ..base_step import BaseStep
|
||||
|
||||
from ...components import R
|
||||
|
||||
|
||||
|
|
@ -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")
|
||||
|
|
@ -1,7 +0,0 @@
|
|||
"""Graph steps."""
|
||||
|
||||
from .traverse import GraphTraverseStep
|
||||
|
||||
__all__ = [
|
||||
"GraphTraverseStep",
|
||||
]
|
||||
|
|
@ -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)})
|
||||
0
reme4/steps/index/__init__.py
Normal file
0
reme4/steps/index/__init__.py
Normal file
30
reme4/steps/index/clear_and_scan.py
Normal file
30
reme4/steps/index/clear_and_scan.py
Normal 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
|
||||
|
|
@ -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")
|
||||
|
||||
|
|
@ -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"] = [
|
||||
111
reme4/steps/index/traverse.py
Normal file
111
reme4/steps/index/traverse.py
Normal 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
|
||||
94
reme4/steps/index/update_catalog.py
Normal file
94
reme4/steps/index/update_catalog.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
0
reme4/steps/transfer/__init__.py
Normal file
0
reme4/steps/transfer/__init__.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
129
reme4/utils/link_expansion.py
Normal file
129
reme4/utils/link_expansion.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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` /
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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所有测试通过!")
|
||||
|
|
@ -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
|
|
@ -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!")
|
||||
|
|
|
|||
161
tests4/unittest/test_file_catalog.py
Normal file
161
tests4/unittest/test_file_catalog.py
Normal 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所有测试通过!")
|
||||
|
|
@ -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")])
|
||||
|
||||
|
|
|
|||
957
tests4/unittest/test_keyword_index.py
Normal file
957
tests4/unittest/test_keyword_index.py
Normal 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)} 项测试通过!")
|
||||
254
tests4/unittest/test_link_expansion.py
Normal file
254
tests4/unittest/test_link_expansion.py
Normal 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
Loading…
Add table
Reference in a new issue