From a4efc0f7761d1e66095b71634cdc42abc5eed991 Mon Sep 17 00:00:00 2001 From: jinliyl <6469360+jinliyl@users.noreply.github.com> Date: Thu, 28 May 2026 14:30:30 +0800 Subject: [PATCH] refactor(reme4): restructure steps packages (#258) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * 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 --- docs4/reme4_report.md | 4 + pyproject.toml | 1 + reme4/application.py | 226 +++-- reme4/components/__init__.py | 2 + reme4/components/application_context.py | 31 +- reme4/components/base_component.py | 121 ++- reme4/components/component_registry.py | 25 +- .../embedding/base_embedding_model.py | 215 ++-- reme4/components/file_catalog/__init__.py | 9 + .../file_catalog/base_file_catalog.py | 39 + .../file_catalog/local_file_catalog.py | 62 ++ reme4/components/file_graph/__init__.py | 7 +- .../components/file_graph/base_file_graph.py | 60 +- .../components/file_graph/local_file_graph.py | 89 +- .../components/file_graph/neo4j_file_graph.py | 2 +- reme4/components/file_graph/nx_file_graph.py | 96 +- .../file_parser/linked_file_parser.py | 2 +- .../components/file_store/base_file_store.py | 43 +- .../file_store/faiss_local_file_store.py | 122 ++- .../components/file_store/local_file_store.py | 94 +- reme4/components/job/background_job.py | 75 +- reme4/components/job/base_job.py | 52 +- reme4/components/job/stream_job.py | 5 +- .../keyword_index/base_keyword_index.py | 35 +- reme4/components/keyword_index/bm25_index.py | 444 ++++++-- reme4/components/prompt_handler.py | 98 +- reme4/components/runtime_context.py | 37 +- reme4/components/service/base_service.py | 43 +- reme4/components/service/http_service.py | 116 ++- reme4/components/service/mcp_service.py | 34 +- reme4/components/tokenizer/base_tokenizer.py | 34 +- reme4/components/tokenizer/jieba_tokenizer.py | 46 +- reme4/components/tokenizer/regex_tokenizer.py | 31 +- reme4/config/default.yaml | 609 ++++------- reme4/config/demo.yaml | 45 + reme4/config/qwenpaw.yaml | 474 +++++++++ reme4/enumeration/component_enum.py | 2 + reme4/enumeration/link_scope_enum.py | 2 + reme4/schema/application_config.py | 9 +- reme4/steps/__init__.py | 105 +- reme4/steps/background/__init__.py | 11 - reme4/steps/common/__init__.py | 17 - reme4/steps/common/health_check.py | 129 ++- reme4/steps/common/help.py | 41 +- reme4/steps/common/reindex.py | 35 - reme4/steps/common/traverse.py | 154 --- reme4/steps/common/version.py | 2 +- reme4/steps/crud/__init__.py | 40 - reme4/steps/crud/_file_io.py | 221 ---- reme4/steps/crud/append.py | 75 -- reme4/steps/crud/list.py | 55 - reme4/steps/crud/read.py | 91 -- reme4/steps/daily/__init__.py | 44 - reme4/steps/daily/_daily_io.py | 306 ------ reme4/steps/daily/read.py | 91 -- reme4/steps/daily/write.py | 131 --- reme4/steps/file_io/__init__.py | 0 reme4/steps/file_io/_file_io.py | 492 +++++++++ reme4/steps/file_io/daily_create.py | 94 ++ .../{daily/list.py => file_io/daily_list.py} | 34 +- .../reindex.py => file_io/daily_reindex.py} | 31 +- reme4/steps/{crud => file_io}/delete.py | 3 +- reme4/steps/{crud => file_io}/edit.py | 9 +- .../frontmatter_delete.py} | 3 +- .../read.py => file_io/frontmatter_read.py} | 3 +- .../frontmatter_update.py} | 1 - reme4/steps/file_io/list.py | 104 ++ reme4/steps/{crud => file_io}/move.py | 7 +- reme4/steps/file_io/read.py | 115 +++ reme4/steps/{crud => file_io}/stat.py | 1 - reme4/steps/{crud => file_io}/write.py | 0 reme4/steps/frontmatter/__init__.py | 24 - reme4/steps/graph/__init__.py | 7 - reme4/steps/graph/traverse.py | 119 --- reme4/steps/index/__init__.py | 0 reme4/steps/index/clear_and_scan.py | 30 + .../update_store.py => index/scan_changes.py} | 21 +- reme4/steps/{common => index}/search.py | 108 +- reme4/steps/index/traverse.py | 111 ++ reme4/steps/index/update_catalog.py | 94 ++ .../update_index.py} | 15 +- .../{background => index}/watch_changes.py | 19 +- reme4/steps/transfer/__init__.py | 0 reme4/steps/{crud => transfer}/download.py | 0 .../upload_resource.py => transfer/ingest.py} | 28 +- reme4/steps/{crud => transfer}/upload.py | 2 +- reme4/utils/__init__.py | 3 + reme4/utils/common_utils.py | 6 +- reme4/utils/link_expansion.py | 129 +++ reme4/utils/service_utils.py | 2 +- reme4/utils/wikilink_handler.py | 2 +- tests4/unittest/test_background_steps.py | 117 +-- tests4/unittest/test_bm25_lite.py | 586 ----------- tests4/unittest/test_common_steps.py | 186 ++-- tests4/unittest/test_crud_steps.py | 948 ++++++----------- tests4/unittest/test_daily_steps.py | 503 +++------ tests4/unittest/test_file_catalog.py | 161 +++ tests4/unittest/test_file_store.py | 6 +- tests4/unittest/test_keyword_index.py | 957 ++++++++++++++++++ tests4/unittest/test_link_expansion.py | 254 +++++ tests4/unittest/test_resource_steps.py | 52 +- tests4/unittest/test_wikilink_utils.py | 4 +- 102 files changed, 5560 insertions(+), 4820 deletions(-) create mode 100644 reme4/components/file_catalog/__init__.py create mode 100644 reme4/components/file_catalog/base_file_catalog.py create mode 100644 reme4/components/file_catalog/local_file_catalog.py create mode 100644 reme4/config/demo.yaml create mode 100644 reme4/config/qwenpaw.yaml delete mode 100644 reme4/steps/background/__init__.py delete mode 100644 reme4/steps/common/reindex.py delete mode 100644 reme4/steps/common/traverse.py delete mode 100644 reme4/steps/crud/__init__.py delete mode 100644 reme4/steps/crud/_file_io.py delete mode 100644 reme4/steps/crud/append.py delete mode 100644 reme4/steps/crud/list.py delete mode 100644 reme4/steps/crud/read.py delete mode 100644 reme4/steps/daily/__init__.py delete mode 100644 reme4/steps/daily/_daily_io.py delete mode 100644 reme4/steps/daily/read.py delete mode 100644 reme4/steps/daily/write.py create mode 100644 reme4/steps/file_io/__init__.py create mode 100644 reme4/steps/file_io/_file_io.py create mode 100644 reme4/steps/file_io/daily_create.py rename reme4/steps/{daily/list.py => file_io/daily_list.py} (51%) rename reme4/steps/{daily/reindex.py => file_io/daily_reindex.py} (59%) rename reme4/steps/{crud => file_io}/delete.py (99%) rename reme4/steps/{crud => file_io}/edit.py (91%) rename reme4/steps/{frontmatter/delete.py => file_io/frontmatter_delete.py} (97%) rename reme4/steps/{frontmatter/read.py => file_io/frontmatter_read.py} (96%) rename reme4/steps/{frontmatter/update.py => file_io/frontmatter_update.py} (99%) create mode 100644 reme4/steps/file_io/list.py rename reme4/steps/{crud => file_io}/move.py (97%) create mode 100644 reme4/steps/file_io/read.py rename reme4/steps/{crud => file_io}/stat.py (99%) rename reme4/steps/{crud => file_io}/write.py (100%) delete mode 100644 reme4/steps/frontmatter/__init__.py delete mode 100644 reme4/steps/graph/__init__.py delete mode 100644 reme4/steps/graph/traverse.py create mode 100644 reme4/steps/index/__init__.py create mode 100644 reme4/steps/index/clear_and_scan.py rename reme4/steps/{background/update_store.py => index/scan_changes.py} (79%) rename reme4/steps/{common => index}/search.py (53%) create mode 100644 reme4/steps/index/traverse.py create mode 100644 reme4/steps/index/update_catalog.py rename reme4/steps/{background/index_changes.py => index/update_index.py} (89%) rename reme4/steps/{background => index}/watch_changes.py (73%) create mode 100644 reme4/steps/transfer/__init__.py rename reme4/steps/{crud => transfer}/download.py (100%) rename reme4/steps/{crud/upload_resource.py => transfer/ingest.py} (95%) rename reme4/steps/{crud => transfer}/upload.py (99%) create mode 100644 reme4/utils/link_expansion.py delete mode 100644 tests4/unittest/test_bm25_lite.py create mode 100644 tests4/unittest/test_file_catalog.py create mode 100644 tests4/unittest/test_keyword_index.py create mode 100644 tests4/unittest/test_link_expansion.py diff --git a/docs4/reme4_report.md b/docs4/reme4_report.md index b16ca5a9..9b66aa28 100644 --- a/docs4/reme4_report.md +++ b/docs4/reme4_report.md @@ -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/ diff --git a/pyproject.toml b/pyproject.toml index 884b452c..19d853f5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/reme4/application.py b/reme4/application.py index dc2a3fd4..fb610868 100644 --- a/reme4/application.py +++ b/reme4/application.py @@ -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) diff --git a/reme4/components/__init__.py b/reme4/components/__init__.py index 336da644..7239c7d6 100644 --- a/reme4/components/__init__.py +++ b/reme4/components/__init__.py @@ -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", diff --git a/reme4/components/application_context.py b/reme4/components/application_context.py index 90a7f194..0feac294 100644 --- a/reme4/components/application_context.py +++ b/reme4/components/application_context.py @@ -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"] = {} diff --git a/reme4/components/base_component.py b/reme4/components/base_component.py index 9b59dcfe..bc96c64b 100644 --- a/reme4/components/base_component.py +++ b/reme4/components/base_component.py @@ -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"" 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: ``/``.""" 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() diff --git a/reme4/components/component_registry.py b/reme4/components/component_registry.py index c78ceb9c..8f141d43 100644 --- a/reme4/components/component_registry.py +++ b/reme4/components/component_registry.py @@ -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 diff --git a/reme4/components/embedding/base_embedding_model.py b/reme4/components/embedding/base_embedding_model.py index 64add5b5..a149cfcf 100644 --- a/reme4/components/embedding/base_embedding_model.py +++ b/reme4/components/embedding/base_embedding_model.py @@ -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") diff --git a/reme4/components/file_catalog/__init__.py b/reme4/components/file_catalog/__init__.py new file mode 100644 index 00000000..1ea9197f --- /dev/null +++ b/reme4/components/file_catalog/__init__.py @@ -0,0 +1,9 @@ +"""file catalog""" + +from .base_file_catalog import BaseFileCatalog +from .local_file_catalog import LocalFileCatalog + +__all__ = [ + "BaseFileCatalog", + "LocalFileCatalog", +] diff --git a/reme4/components/file_catalog/base_file_catalog.py b/reme4/components/file_catalog/base_file_catalog.py new file mode 100644 index 00000000..4b6612f5 --- /dev/null +++ b/reme4/components/file_catalog/base_file_catalog.py @@ -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.""" diff --git a/reme4/components/file_catalog/local_file_catalog.py b/reme4/components/file_catalog/local_file_catalog.py new file mode 100644 index 00000000..150be1bb --- /dev/null +++ b/reme4/components/file_catalog/local_file_catalog.py @@ -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) diff --git a/reme4/components/file_graph/__init__.py b/reme4/components/file_graph/__init__.py index 257bc7dc..3aebbb8b 100644 --- a/reme4/components/file_graph/__init__.py +++ b/reme4/components/file_graph/__init__.py @@ -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", +] diff --git a/reme4/components/file_graph/base_file_graph.py b/reme4/components/file_graph/base_file_graph.py index f68488fd..c34d53e0 100644 --- a/reme4/components/file_graph/base_file_graph.py +++ b/reme4/components/file_graph/base_file_graph.py @@ -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*.""" diff --git a/reme4/components/file_graph/local_file_graph.py b/reme4/components/file_graph/local_file_graph.py index f34293de..1a67c8a6 100644 --- a/reme4/components/file_graph/local_file_graph.py +++ b/reme4/components/file_graph/local_file_graph.py @@ -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 diff --git a/reme4/components/file_graph/neo4j_file_graph.py b/reme4/components/file_graph/neo4j_file_graph.py index 828d2776..ec22a7ef 100644 --- a/reme4/components/file_graph/neo4j_file_graph.py +++ b/reme4/components/file_graph/neo4j_file_graph.py @@ -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", ) diff --git a/reme4/components/file_graph/nx_file_graph.py b/reme4/components/file_graph/nx_file_graph.py index e41cafe7..5157d5d0 100644 --- a/reme4/components/file_graph/nx_file_graph.py +++ b/reme4/components/file_graph/nx_file_graph.py @@ -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 diff --git a/reme4/components/file_parser/linked_file_parser.py b/reme4/components/file_parser/linked_file_parser.py index 6fa46b3d..229b9d14 100644 --- a/reme4/components/file_parser/linked_file_parser.py +++ b/reme4/components/file_parser/linked_file_parser.py @@ -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). """ diff --git a/reme4/components/file_store/base_file_store.py b/reme4/components/file_store/base_file_store.py index 7788487e..afba93a3 100644 --- a/reme4/components/file_store/base_file_store.py +++ b/reme4/components/file_store/base_file_store.py @@ -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.""" diff --git a/reme4/components/file_store/faiss_local_file_store.py b/reme4/components/file_store/faiss_local_file_store.py index d98ef718..a0dd3c89 100644 --- a/reme4/components/file_store/faiss_local_file_store.py +++ b/reme4/components/file_store/faiss_local_file_store.py @@ -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 diff --git a/reme4/components/file_store/local_file_store.py b/reme4/components/file_store/local_file_store.py index 92c92e02..173ca97c 100644 --- a/reme4/components/file_store/local_file_store.py +++ b/reme4/components/file_store/local_file_store.py @@ -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.""" diff --git a/reme4/components/job/background_job.py b/reme4/components/job/background_job.py index a0a4efa3..c72de54f 100644 --- a/reme4/components/job/background_job.py +++ b/reme4/components/job/background_job.py @@ -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 diff --git a/reme4/components/job/base_job.py b/reme4/components/job/base_job.py index 4928f2d0..3a414bee 100644 --- a/reme4/components/job/base_job.py +++ b/reme4/components/job/base_job.py @@ -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}") diff --git a/reme4/components/job/stream_job.py b/reme4/components/job/stream_job.py index 89f1fab1..651d5488 100644 --- a/reme4/components/job/stream_job.py +++ b/reme4/components/job/stream_job.py @@ -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() diff --git a/reme4/components/keyword_index/base_keyword_index.py b/reme4/components/keyword_index/base_keyword_index.py index 86433e97..466f4c66 100644 --- a/reme4/components/keyword_index/base_keyword_index.py +++ b/reme4/components/keyword_index/base_keyword_index.py @@ -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.""" diff --git a/reme4/components/keyword_index/bm25_index.py b/reme4/components/keyword_index/bm25_index.py index 81899828..81c5de7a 100644 --- a/reme4/components/keyword_index/bm25_index.py +++ b/reme4/components/keyword_index/bm25_index.py @@ -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 = {} diff --git a/reme4/components/prompt_handler.py b/reme4/components/prompt_handler.py index a9ae203a..bbc4a046 100644 --- a/reme4/components/prompt_handler.py +++ b/reme4/components/prompt_handler.py @@ -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 ``.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 ``_``.""" @@ -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)})" diff --git a/reme4/components/runtime_context.py b/reme4/components/runtime_context.py index 8d8f4409..d9454900 100644 --- a/reme4/components/runtime_context.py +++ b/reme4/components/runtime_context.py @@ -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(): diff --git a/reme4/components/service/base_service.py b/reme4/components/service/base_service.py index 1c50abd3..2384fe71 100644 --- a/reme4/components/service/base_service.py +++ b/reme4/components/service/base_service.py @@ -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) diff --git a/reme4/components/service/http_service.py b/reme4/components/service/http_service.py index a1092947..95ed5e72 100644 --- a/reme4/components/service/http_service.py +++ b/reme4/components/service/http_service.py @@ -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) diff --git a/reme4/components/service/mcp_service.py b/reme4/components/service/mcp_service.py index 8f22450d..9203aff9 100644 --- a/reme4/components/service/mcp_service.py +++ b/reme4/components/service/mcp_service.py @@ -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 diff --git a/reme4/components/tokenizer/base_tokenizer.py b/reme4/components/tokenizer/base_tokenizer.py index 713c2ecb..8795f4a4 100644 --- a/reme4/components/tokenizer/base_tokenizer.py +++ b/reme4/components/tokenizer/base_tokenizer.py @@ -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.""" diff --git a/reme4/components/tokenizer/jieba_tokenizer.py b/reme4/components/tokenizer/jieba_tokenizer.py index 4391c89c..5e35fd23 100644 --- a/reme4/components/tokenizer/jieba_tokenizer.py +++ b/reme4/components/tokenizer/jieba_tokenizer.py @@ -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)) diff --git a/reme4/components/tokenizer/regex_tokenizer.py b/reme4/components/tokenizer/regex_tokenizer.py index 4379a60e..a5b8be11 100644 --- a/reme4/components/tokenizer/regex_tokenizer.py +++ b/reme4/components/tokenizer/regex_tokenizer.py @@ -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 diff --git a/reme4/config/default.yaml b/reme4/config/default.yaml index 73afdbe8..5d69b2db 100644 --- a/reme4/config/default.yaml +++ b/reme4/config/default.yaml @@ -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//.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/.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// 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//.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//.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: }" - 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/.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/.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} diff --git a/reme4/config/demo.yaml b/reme4/config/demo.yaml new file mode 100644 index 00000000..f765a194 --- /dev/null +++ b/reme4/config/demo.yaml @@ -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 diff --git a/reme4/config/qwenpaw.yaml b/reme4/config/qwenpaw.yaml new file mode 100644 index 00000000..96d30ad9 --- /dev/null +++ b/reme4/config/qwenpaw.yaml @@ -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// 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//.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/.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 diff --git a/reme4/enumeration/component_enum.py b/reme4/enumeration/component_enum.py index d0ae0ef1..315c760c 100644 --- a/reme4/enumeration/component_enum.py +++ b/reme4/enumeration/component_enum.py @@ -22,6 +22,8 @@ class ComponentEnum(str, Enum): FILE_GRAPH = "file_graph" + FILE_CATALOG = "file_catalog" + KEYWORD_INDEX = "keyword_index" SERVICE = "service" diff --git a/reme4/enumeration/link_scope_enum.py b/reme4/enumeration/link_scope_enum.py index c00163e6..86bdef63 100644 --- a/reme4/enumeration/link_scope_enum.py +++ b/reme4/enumeration/link_scope_enum.py @@ -13,5 +13,7 @@ class LinkScopeEnum(str, Enum): """ REAL = "real" + VIRTUAL = "virtual" + ALL = "all" diff --git a/reme4/schema/application_config.py b/reme4/schema/application_config.py index 6e82e527..3df6ff1d 100644 --- a/reme4/schema/application_config.py +++ b/reme4/schema/application_config.py @@ -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", diff --git a/reme4/steps/__init__.py b/reme4/steps/__init__.py index b039ba08..358ecd38 100644 --- a/reme4/steps/__init__.py +++ b/reme4/steps/__init__.py @@ -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", ] diff --git a/reme4/steps/background/__init__.py b/reme4/steps/background/__init__.py deleted file mode 100644 index c55db436..00000000 --- a/reme4/steps/background/__init__.py +++ /dev/null @@ -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", -] diff --git a/reme4/steps/common/__init__.py b/reme4/steps/common/__init__.py index 73248fd5..e69de29b 100644 --- a/reme4/steps/common/__init__.py +++ b/reme4/steps/common/__init__.py @@ -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", -] diff --git a/reme4/steps/common/health_check.py b/reme4/steps/common/health_check.py index 2fefd826..7bdf182d 100644 --- a/reme4/steps/common/health_check.py +++ b/reme4/steps/common/health_check.py @@ -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 diff --git a/reme4/steps/common/help.py b/reme4/steps/common/help.py index 6c850772..3ccf94d1 100644 --- a/reme4/steps/common/help.py +++ b/reme4/steps/common/help.py @@ -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 diff --git a/reme4/steps/common/reindex.py b/reme4/steps/common/reindex.py deleted file mode 100644 index 84f86939..00000000 --- a/reme4/steps/common/reindex.py +++ /dev/null @@ -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 diff --git a/reme4/steps/common/traverse.py b/reme4/steps/common/traverse.py deleted file mode 100644 index ae955acf..00000000 --- a/reme4/steps/common/traverse.py +++ /dev/null @@ -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)}) diff --git a/reme4/steps/common/version.py b/reme4/steps/common/version.py index e43aa698..b5ed0b29 100644 --- a/reme4/steps/common/version.py +++ b/reme4/steps/common/version.py @@ -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 diff --git a/reme4/steps/crud/__init__.py b/reme4/steps/crud/__init__.py deleted file mode 100644 index 78356a47..00000000 --- a/reme4/steps/crud/__init__.py +++ /dev/null @@ -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", -] diff --git a/reme4/steps/crud/_file_io.py b/reme4/steps/crud/_file_io.py deleted file mode 100644 index 8b02c4b4..00000000 --- a/reme4/steps/crud/_file_io.py +++ /dev/null @@ -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 diff --git a/reme4/steps/crud/append.py b/reme4/steps/crud/append.py deleted file mode 100644 index ed4c0913..00000000 --- a/reme4/steps/crud/append.py +++ /dev/null @@ -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 diff --git a/reme4/steps/crud/list.py b/reme4/steps/crud/list.py deleted file mode 100644 index 9b5d1b0c..00000000 --- a/reme4/steps/crud/list.py +++ /dev/null @@ -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)}) diff --git a/reme4/steps/crud/read.py b/reme4/steps/crud/read.py deleted file mode 100644 index 52cf8c98..00000000 --- a/reme4/steps/crud/read.py +++ /dev/null @@ -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 diff --git a/reme4/steps/daily/__init__.py b/reme4/steps/daily/__init__.py deleted file mode 100644 index 3f0da979..00000000 --- a/reme4/steps/daily/__init__.py +++ /dev/null @@ -1,44 +0,0 @@ -"""Daily-aware steps — CRUD on note md + day-level index. - -A daily note is the single file ``daily//.md``. -The day-level index ``daily/.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//.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", -] diff --git a/reme4/steps/daily/_daily_io.py b/reme4/steps/daily/_daily_io.py deleted file mode 100644 index 0751a3fe..00000000 --- a/reme4/steps/daily/_daily_io.py +++ /dev/null @@ -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//.md``. -2. **Day-index** — the derived rollup page ``daily/.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 ``//*.md`` and pull each - note's reserved frontmatter (``name`` / ``description``). -* :func:`refresh_day_index` — rebuild ``/.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: -# description: -# -# The note inventory lives in the body's ```` -# 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 = "" -_BLOCK_CLOSE = "" - -_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^{re.escape(_HEADINGS[name])}\s*\n)?" - rf"{re.escape(_BLOCK_OPEN.format(name=name))}" - r"(?P.*?)" - 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 ``//*.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 ``/.md`` from the current state of its notes. - - Behaviour: - * No ``//`` 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": "/.md", - "notes": [ - {"path": "//.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, - } diff --git a/reme4/steps/daily/read.py b/reme4/steps/daily/read.py deleted file mode 100644 index 8569f229..00000000 --- a/reme4/steps/daily/read.py +++ /dev/null @@ -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//.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//.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 diff --git a/reme4/steps/daily/write.py b/reme4/steps/daily/write.py deleted file mode 100644 index 4ccb6734..00000000 --- a/reme4/steps/daily/write.py +++ /dev/null @@ -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/.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//.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 diff --git a/reme4/steps/file_io/__init__.py b/reme4/steps/file_io/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/reme4/steps/file_io/_file_io.py b/reme4/steps/file_io/_file_io.py new file mode 100644 index 00000000..7a90acbd --- /dev/null +++ b/reme4/steps/file_io/_file_io.py @@ -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/.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: +# description: +# +# The note inventory lives in the body's ```` +# block: each note becomes a single line with its full frontmatter +# inlined (``- [[path]] name: ... description: ... ``), +# letting an agent scan the day at a glance. Content outside the +# auto markers is preserved verbatim across refreshes. + +_NOTES_OPEN = "" +_NOTES_CLOSE = "" + +_NOTES_BLOCK_RE = re.compile( + rf"{re.escape(_NOTES_OPEN)}(?P.*?){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 ``//*.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 ``/.md`` from the current state of its notes. + + Behaviour: + * No ``//`` 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, + } diff --git a/reme4/steps/file_io/daily_create.py b/reme4/steps/file_io/daily_create.py new file mode 100644 index 00000000..0cb15386 --- /dev/null +++ b/reme4/steps/file_io/daily_create.py @@ -0,0 +1,94 @@ +"""``daily_create`` — provision a note slug under a daily folder: ``daily//.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//.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 diff --git a/reme4/steps/daily/list.py b/reme4/steps/file_io/daily_list.py similarity index 51% rename from reme4/steps/daily/list.py rename to reme4/steps/file_io/daily_list.py index 5299375e..0552e975 100644 --- a/reme4/steps/daily/list.py +++ b/reme4/steps/file_io/daily_list.py @@ -1,23 +1,22 @@ """``daily_list`` — list the notes under a single day (pure read, no side effects). Returns one row per ``daily//.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/.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 ``//`` 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}) diff --git a/reme4/steps/daily/reindex.py b/reme4/steps/file_io/daily_reindex.py similarity index 59% rename from reme4/steps/daily/reindex.py rename to reme4/steps/file_io/daily_reindex.py index 1cbd6712..71cefb5b 100644 --- a/reme4/steps/daily/reindex.py +++ b/reme4/steps/file_io/daily_reindex.py @@ -2,18 +2,18 @@ The day index ``daily/.md`` is a derived artifact whose job is to list and describe every note file under ``daily//``. 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/.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)) diff --git a/reme4/steps/crud/delete.py b/reme4/steps/file_io/delete.py similarity index 99% rename from reme4/steps/crud/delete.py rename to reme4/steps/file_io/delete.py index a0b75ea9..7d49a152 100644 --- a/reme4/steps/crud/delete.py +++ b/reme4/steps/file_io/delete.py @@ -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: diff --git a/reme4/steps/crud/edit.py b/reme4/steps/file_io/edit.py similarity index 91% rename from reme4/steps/crud/edit.py rename to reme4/steps/file_io/edit.py index 14b51b21..fc0462cf 100644 --- a/reme4/steps/crud/edit.py +++ b/reme4/steps/file_io/edit.py @@ -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 diff --git a/reme4/steps/frontmatter/delete.py b/reme4/steps/file_io/frontmatter_delete.py similarity index 97% rename from reme4/steps/frontmatter/delete.py rename to reme4/steps/file_io/frontmatter_delete.py index 65367ce1..2091d515 100644 --- a/reme4/steps/frontmatter/delete.py +++ b/reme4/steps/file_io/frontmatter_delete.py @@ -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 diff --git a/reme4/steps/frontmatter/read.py b/reme4/steps/file_io/frontmatter_read.py similarity index 96% rename from reme4/steps/frontmatter/read.py rename to reme4/steps/file_io/frontmatter_read.py index c27fb117..c01bf6f5 100644 --- a/reme4/steps/frontmatter/read.py +++ b/reme4/steps/file_io/frontmatter_read.py @@ -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 diff --git a/reme4/steps/frontmatter/update.py b/reme4/steps/file_io/frontmatter_update.py similarity index 99% rename from reme4/steps/frontmatter/update.py rename to reme4/steps/file_io/frontmatter_update.py index c00a7f89..d1ad5ad4 100644 --- a/reme4/steps/frontmatter/update.py +++ b/reme4/steps/file_io/frontmatter_update.py @@ -17,7 +17,6 @@ from pathlib import Path import frontmatter from ..base_step import BaseStep - from ...components import R diff --git a/reme4/steps/file_io/list.py b/reme4/steps/file_io/list.py new file mode 100644 index 00000000..4506407d --- /dev/null +++ b/reme4/steps/file_io/list.py @@ -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 diff --git a/reme4/steps/crud/move.py b/reme4/steps/file_io/move.py similarity index 97% rename from reme4/steps/crud/move.py rename to reme4/steps/file_io/move.py index aaf514de..ffb33c75 100644 --- a/reme4/steps/crud/move.py +++ b/reme4/steps/file_io/move.py @@ -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" diff --git a/reme4/steps/file_io/read.py b/reme4/steps/file_io/read.py new file mode 100644 index 00000000..3263677f --- /dev/null +++ b/reme4/steps/file_io/read.py @@ -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 diff --git a/reme4/steps/crud/stat.py b/reme4/steps/file_io/stat.py similarity index 99% rename from reme4/steps/crud/stat.py rename to reme4/steps/file_io/stat.py index e6518724..97a9dcd8 100644 --- a/reme4/steps/crud/stat.py +++ b/reme4/steps/file_io/stat.py @@ -25,7 +25,6 @@ from pathlib import Path import frontmatter from ..base_step import BaseStep - from ...components import R diff --git a/reme4/steps/crud/write.py b/reme4/steps/file_io/write.py similarity index 100% rename from reme4/steps/crud/write.py rename to reme4/steps/file_io/write.py diff --git a/reme4/steps/frontmatter/__init__.py b/reme4/steps/frontmatter/__init__.py deleted file mode 100644 index 9ffbed4f..00000000 --- a/reme4/steps/frontmatter/__init__.py +++ /dev/null @@ -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") diff --git a/reme4/steps/graph/__init__.py b/reme4/steps/graph/__init__.py deleted file mode 100644 index c5147dce..00000000 --- a/reme4/steps/graph/__init__.py +++ /dev/null @@ -1,7 +0,0 @@ -"""Graph steps.""" - -from .traverse import GraphTraverseStep - -__all__ = [ - "GraphTraverseStep", -] diff --git a/reme4/steps/graph/traverse.py b/reme4/steps/graph/traverse.py deleted file mode 100644 index 25326f56..00000000 --- a/reme4/steps/graph/traverse.py +++ /dev/null @@ -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)}) diff --git a/reme4/steps/index/__init__.py b/reme4/steps/index/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/reme4/steps/index/clear_and_scan.py b/reme4/steps/index/clear_and_scan.py new file mode 100644 index 00000000..503ededc --- /dev/null +++ b/reme4/steps/index/clear_and_scan.py @@ -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 diff --git a/reme4/steps/background/update_store.py b/reme4/steps/index/scan_changes.py similarity index 79% rename from reme4/steps/background/update_store.py rename to reme4/steps/index/scan_changes.py index 9ecde745..3c1ee7d9 100644 --- a/reme4/steps/background/update_store.py +++ b/reme4/steps/index/scan_changes.py @@ -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") diff --git a/reme4/steps/common/search.py b/reme4/steps/index/search.py similarity index 53% rename from reme4/steps/common/search.py rename to reme4/steps/index/search.py index b942703e..31a8e715 100644 --- a/reme4/steps/common/search.py +++ b/reme4/steps/index/search.py @@ -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"] = [ diff --git a/reme4/steps/index/traverse.py b/reme4/steps/index/traverse.py new file mode 100644 index 00000000..ac5f1ed4 --- /dev/null +++ b/reme4/steps/index/traverse.py @@ -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 diff --git a/reme4/steps/index/update_catalog.py b/reme4/steps/index/update_catalog.py new file mode 100644 index 00000000..c94d97ad --- /dev/null +++ b/reme4/steps/index/update_catalog.py @@ -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 diff --git a/reme4/steps/background/index_changes.py b/reme4/steps/index/update_index.py similarity index 89% rename from reme4/steps/background/index_changes.py rename to reme4/steps/index/update_index.py index f6f1d135..7bcb9446 100644 --- a/reme4/steps/background/index_changes.py +++ b/reme4/steps/index/update_index.py @@ -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 diff --git a/reme4/steps/background/watch_changes.py b/reme4/steps/index/watch_changes.py similarity index 73% rename from reme4/steps/background/watch_changes.py rename to reme4/steps/index/watch_changes.py index 68d28764..dfb978a2 100644 --- a/reme4/steps/background/watch_changes.py +++ b/reme4/steps/index/watch_changes.py @@ -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 diff --git a/reme4/steps/transfer/__init__.py b/reme4/steps/transfer/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/reme4/steps/crud/download.py b/reme4/steps/transfer/download.py similarity index 100% rename from reme4/steps/crud/download.py rename to reme4/steps/transfer/download.py diff --git a/reme4/steps/crud/upload_resource.py b/reme4/steps/transfer/ingest.py similarity index 95% rename from reme4/steps/crud/upload_resource.py rename to reme4/steps/transfer/ingest.py index f732727e..1d9b31e1 100644 --- a/reme4/steps/crud/upload_resource.py +++ b/reme4/steps/transfer/ingest.py @@ -1,8 +1,8 @@ -"""``upload_resource`` — copy an externally-received asset into ``resource//``. +"""``ingest`` — capture an externally-received asset into ``resource//``. -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//`` +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//`` 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//`` and update the day's meta + index.""" +@R.register("ingest_step") +class IngestStep(BaseStep): + """Capture an external asset into ``resource//`` 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) diff --git a/reme4/steps/crud/upload.py b/reme4/steps/transfer/upload.py similarity index 99% rename from reme4/steps/crud/upload.py rename to reme4/steps/transfer/upload.py index b5a8c0c9..cbcffb32 100644 --- a/reme4/steps/crud/upload.py +++ b/reme4/steps/transfer/upload.py @@ -14,7 +14,7 @@ callers must opt in to clobber an existing destination. For the resource-bucket ingest path (channel-tagged, dated under ``resource//`` with provenance metadata) use -``upload_resource`` instead. +``ingest`` instead. """ import mimetypes diff --git a/reme4/utils/__init__.py b/reme4/utils/__init__.py index 45a5bfb9..c18d4de7 100644 --- a/reme4/utils/__init__.py +++ b/reme4/utils/__init__.py @@ -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", diff --git a/reme4/utils/common_utils.py b/reme4/utils/common_utils.py index 08fbc861..fd03eb65 100644 --- a/reme4/utils/common_utils.py +++ b/reme4/utils/common_utils.py @@ -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}", diff --git a/reme4/utils/link_expansion.py b/reme4/utils/link_expansion.py new file mode 100644 index 00000000..f3920c67 --- /dev/null +++ b/reme4/utils/link_expansion.py @@ -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 diff --git a/reme4/utils/service_utils.py b/reme4/utils/service_utils.py index aebabc2a..0bc29c78 100644 --- a/reme4/utils/service_utils.py +++ b/reme4/utils/service_utils.py @@ -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=", + f"port {port} occupied. Start on another port: reme start service.port=", file=sys.stderr, ) sys.exit(1) diff --git a/reme4/utils/wikilink_handler.py b/reme4/utils/wikilink_handler.py index fddfb148..ba93bf33 100644 --- a/reme4/utils/wikilink_handler.py +++ b/reme4/utils/wikilink_handler.py @@ -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` / diff --git a/tests4/unittest/test_background_steps.py b/tests4/unittest/test_background_steps.py index 2c6878ad..0b288a50 100644 --- a/tests4/unittest/test_background_steps.py +++ b/tests4/unittest/test_background_steps.py @@ -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() diff --git a/tests4/unittest/test_bm25_lite.py b/tests4/unittest/test_bm25_lite.py deleted file mode 100644 index 75865797..00000000 --- a/tests4/unittest/test_bm25_lite.py +++ /dev/null @@ -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所有测试通过!") diff --git a/tests4/unittest/test_common_steps.py b/tests4/unittest/test_common_steps.py index 4b573e35..136f8bb5 100644 --- a/tests4/unittest/test_common_steps.py +++ b/tests4/unittest/test_common_steps.py @@ -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() diff --git a/tests4/unittest/test_crud_steps.py b/tests4/unittest/test_crud_steps.py index 7159b287..ae53e375 100644 --- a/tests4/unittest/test_crud_steps.py +++ b/tests4/unittest/test_crud_steps.py @@ -1,27 +1,26 @@ # pylint: disable=too-many-lines """Tests for crud steps — the opaque-byte vault_dir surface plus the -text-content ops (``read`` / ``write`` / ``edit`` / ``append``). +text-content ops (``read`` / ``write`` / ``edit``). -Two surfaces share this file: +Every test drives a step directly against a freshly built +``LocalFileStore`` (embedding disabled, BM25 kept) with files seeded +on disk (and, where relevant, registered in the graph so retarget's +reverse-index lookup finds inbound edges). No app config, no HTTP +server — the step's ``vault_path`` defaults to ``cwd()`` and tests +chdir into a tmpdir to scope the vault. -* **Direct unit tests** (top half) drive each step against a freshly - built ``LocalFileStore`` (embedding disabled, BM25 kept) with files - registered in the graph so retarget's reverse-index lookup finds - inbound edges. Covers ``stat`` / ``list`` / ``download`` / ``move`` - / ``delete``. -* **HTTP/MCP E2E tests** (bottom half) spawn ``reme4 start`` via - ``mock_reme_server`` and exercise ``read`` / ``write`` / ``edit`` / - ``append`` end-to-end, including non-md degraded mode + encoding - edge cases. +Covers ``stat`` / ``list`` / ``download`` / ``move`` / ``delete`` +plus the text ops ``read`` / ``write`` / ``edit`` (including non-md +degraded mode + encoding edge cases). Frontmatter-only ops live in ``test_frontmatter_steps.py``. The ``upload`` step is a passive resource-ingest entry point with its own bucket semantics — tests for it live in ``test_resource_steps.py``. -CLI rule for the HTTP half: ``path=`` is relative-only, rooted at the -reme vault. A bare path with no suffix auto-appends ``.md``; -non-``.md`` suffix is accepted in degraded mode. Absolute paths are -accepted with a warning. +Path-shape contract (enforced by ``read`` / ``write`` / ``edit``): +``path=`` is vault-relative by default. A bare path with no suffix +auto-appends ``.md``; non-``.md`` suffix is accepted in degraded mode. +Absolute paths are accepted with a warning. """ # pylint: disable=protected-access,redefined-builtin @@ -34,14 +33,16 @@ from pathlib import Path from reme4.components.file_store import LocalFileStore from reme4.schema import FileNode -from reme4.steps.crud import ( +from reme4.steps.file_io import ( delete as crud_delete, - download as crud_download, + edit as crud_edit, list as crud_list, move as crud_move, + read as crud_read, stat as crud_stat, + write as crud_write, ) -from reme4.utils import call_action, call_and_check, mock_reme_server +from reme4.steps.transfer import download as crud_download from reme4.utils.wikilink_handler import WikilinkHandler warnings.filterwarnings("ignore", category=DeprecationWarning, module="jieba") @@ -66,7 +67,7 @@ class temp_chdir: async def _make_store(files: dict[str, str] | None = None) -> LocalFileStore: """LocalFileStore seeded with files on disk + registered in the graph.""" - store = LocalFileStore(store_name="t", embedding_model="") + store = LocalFileStore(name="t", embedding_model="") await store.start() nodes: list[FileNode] = [] for rel, content in (files or {}).items(): @@ -103,7 +104,6 @@ def _seed_md(vault_dir: Path, rel: str, body: str) -> Path: # =========================================================================== # Direct unit tests: stat / list / download / move / delete -# (LocalFileStore, no HTTP server) # =========================================================================== @@ -514,36 +514,47 @@ def test_delete_folder_empty_has_no_inbound(): # =========================================================================== -# HTTP / MCP E2E tests: read / write / edit / append -# (mock_reme_server spawns `reme4 start`, calls go over the wire) +# Direct unit tests: read / write / edit # =========================================================================== # -- read ---------------------------------------------------------------- +async def _read(store: LocalFileStore, **kwargs): + """Run a ReadStep against ``store`` and return its response.""" + step = crud_read.ReadStep(file_store=store) + await step(**kwargs) + return step.context.response + + +async def _write(store: LocalFileStore, **kwargs): + """Run a WriteStep against ``store`` and return its response.""" + step = crud_write.WriteStep(file_store=store) + await step(**kwargs) + return step.context.response + + +async def _edit(store: LocalFileStore, **kwargs): + """Run an EditStep against ``store`` and return its response.""" + step = crud_edit.EditStep(file_store=store) + await step(**kwargs) + return step.context.response + + def test_read_relative_path(): - """`reme4 read path=Templates/Recipe.md` returns the file body from vault/.""" + """`read path=Templates/Recipe.md` returns the file body from the vault.""" async def run(): with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): - working = Path(tmp) / ".reme" - working.mkdir(parents=True, exist_ok=True) body = "# Recipe\n\nMix flour and water.\n" - _seed_md(working, "Templates/Recipe.md", body) - async with mock_reme_server() as (host, port): - await call_and_check( - "read", - host=host, - port=port, - path="Templates/Recipe.md", - validator=lambda r: ( - isinstance(r, dict) - and r.get("success") is True - and "# Recipe" in str(r.get("answer", "")) - and "flour and water" in str(r.get("answer", "")) - ), - ) + _seed_md(Path(tmp), "Templates/Recipe.md", body) + store = await _make_store() + resp = await _read(store, path="Templates/Recipe.md") + assert resp.success is True + assert "# Recipe" in str(resp.answer) + assert "flour and water" in str(resp.answer) + await store.close() print("✓ test_read_relative_path passed") _run(run()) @@ -554,19 +565,12 @@ def test_read_no_suffix_autoappends_md(): async def run(): with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): - working = Path(tmp) / ".reme" - working.mkdir(parents=True, exist_ok=True) - _seed_md(working, "Templates/Recipe.md", "auto-md\n") - async with mock_reme_server() as (host, port): - await call_and_check( - "read", - host=host, - port=port, - path="Templates/Recipe", - validator=lambda r: ( - isinstance(r, dict) and r.get("success") is True and "auto-md" in str(r.get("answer", "")) - ), - ) + _seed_md(Path(tmp), "Templates/Recipe.md", "auto-md\n") + store = await _make_store() + resp = await _read(store, path="Templates/Recipe") + assert resp.success is True + assert "auto-md" in str(resp.answer) + await store.close() print("✓ test_read_no_suffix_autoappends_md passed") _run(run()) @@ -577,27 +581,16 @@ def test_read_line_range(): async def run(): with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): - working = Path(tmp) / ".reme" - working.mkdir(parents=True, exist_ok=True) - _seed_md(working, "Notes.md", "L1\nL2\nL3\nL4\nL5\n") - async with mock_reme_server() as (host, port): - await call_and_check( - "read", - host=host, - port=port, - path="Notes.md", - start_line=2, - end_line=4, - validator=lambda r: ( - isinstance(r, dict) - and r.get("success") is True - and "L2" in str(r["answer"]) - and "L3" in str(r["answer"]) - and "L4" in str(r["answer"]) - and "L1" not in str(r["answer"]) - and "L5" not in str(r["answer"]) - ), - ) + _seed_md(Path(tmp), "Notes.md", "L1\nL2\nL3\nL4\nL5\n") + store = await _make_store() + resp = await _read(store, path="Notes.md", start_line=2, end_line=4) + assert resp.success is True + assert "L2" in str(resp.answer) + assert "L3" in str(resp.answer) + assert "L4" in str(resp.answer) + assert "L1" not in str(resp.answer) + assert "L5" not in str(resp.answer) + await store.close() print("✓ test_read_line_range passed") _run(run()) @@ -608,20 +601,12 @@ def test_read_absolute_path_accepted(): async def run(): with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): - working = Path(tmp) / ".reme" - working.mkdir(parents=True, exist_ok=True) - target = _seed_md(working, "Abs.md", "x\n") - async with mock_reme_server() as (host, port): - result = await call_action( - "read", - host=host, - port=port, - path=str(target.resolve()), - ) - if not ( - isinstance(result, dict) and result.get("success") is True and "x" in str(result.get("answer", "")) - ): - raise AssertionError(f"expected absolute-path read to succeed, got {result!r}") + target = _seed_md(Path(tmp), "Abs.md", "x\n") + store = await _make_store() + resp = await _read(store, path=str(target.resolve())) + assert resp.success is True + assert "x" in str(resp.answer) + await store.close() print("✓ test_read_absolute_path_accepted passed") _run(run()) @@ -632,21 +617,12 @@ def test_read_non_md_degraded(): async def run(): with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): - working = Path(tmp) / ".reme" - working.mkdir(parents=True, exist_ok=True) - _seed_md(working, "data/foo.txt", "plain-text body\n") - async with mock_reme_server() as (host, port): - await call_and_check( - "read", - host=host, - port=port, - path="data/foo.txt", - validator=lambda r: ( - isinstance(r, dict) - and r.get("success") is True - and "plain-text body" in str(r.get("answer", "")) - ), - ) + _seed_md(Path(tmp), "data/foo.txt", "plain-text body\n") + store = await _make_store() + resp = await _read(store, path="data/foo.txt") + assert resp.success is True + assert "plain-text body" in str(resp.answer) + await store.close() print("✓ test_read_non_md_degraded passed") _run(run()) @@ -657,21 +633,11 @@ def test_read_missing_file(): async def run(): with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): - working = Path(tmp) / ".reme" - working.mkdir(parents=True, exist_ok=True) - async with mock_reme_server() as (host, port): - result = await call_action( - "read", - host=host, - port=port, - path="NotThere.md", - ) - if not ( - isinstance(result, dict) - and result.get("success") is False - and "does not exist" in str(result.get("answer", "")).lower() - ): - raise AssertionError(f"expected missing-file rejection, got {result!r}") + store = await _make_store() + resp = await _read(store, path="NotThere.md") + assert resp.success is False + assert "does not exist" in str(resp.answer).lower() + await store.close() print("✓ test_read_missing_file passed") _run(run()) @@ -682,24 +648,12 @@ def test_read_start_after_end(): async def run(): with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): - working = Path(tmp) / ".reme" - working.mkdir(parents=True, exist_ok=True) - _seed_md(working, "Range.md", "a\nb\nc\n") - async with mock_reme_server() as (host, port): - result = await call_action( - "read", - host=host, - port=port, - path="Range.md", - start_line=3, - end_line=1, - ) - if not ( - isinstance(result, dict) - and result.get("success") is False - and "start_line" in str(result.get("answer", "")) - ): - raise AssertionError(f"expected start>end rejection, got {result!r}") + _seed_md(Path(tmp), "Range.md", "a\nb\nc\n") + store = await _make_store() + resp = await _read(store, path="Range.md", start_line=3, end_line=1) + assert resp.success is False + assert "start_line" in str(resp.answer) + await store.close() print("✓ test_read_start_after_end passed") _run(run()) @@ -710,23 +664,12 @@ def test_read_start_line_exceeds_total(): async def run(): with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): - working = Path(tmp) / ".reme" - working.mkdir(parents=True, exist_ok=True) - _seed_md(working, "Short.md", "only-one-line\n") - async with mock_reme_server() as (host, port): - result = await call_action( - "read", - host=host, - port=port, - path="Short.md", - start_line=99, - ) - if not ( - isinstance(result, dict) - and result.get("success") is False - and "exceeds" in str(result.get("answer", "")).lower() - ): - raise AssertionError(f"expected exceeds-length rejection, got {result!r}") + _seed_md(Path(tmp), "Short.md", "only-one-line\n") + store = await _make_store() + resp = await _read(store, path="Short.md", start_line=99) + assert resp.success is False + assert "exceeds" in str(resp.answer).lower() + await store.close() print("✓ test_read_start_line_exceeds_total passed") _run(run()) @@ -737,24 +680,15 @@ def test_read_truncation(): async def run(): with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): - working = Path(tmp) / ".reme" - working.mkdir(parents=True, exist_ok=True) # Seed > DEFAULT_MAX_BYTES (50 KiB) so the default truncation kicks in. body = "\n".join(f"line {i}" for i in range(8000)) + "\n" - _seed_md(working, "Big.md", body) - async with mock_reme_server() as (host, port): - await call_and_check( - "read", - host=host, - port=port, - path="Big.md", - validator=lambda r: ( - isinstance(r, dict) - and r.get("success") is True - and "truncated" in str(r["answer"]) - and "start_line=" in str(r["answer"]) - ), - ) + _seed_md(Path(tmp), "Big.md", body) + store = await _make_store() + resp = await _read(store, path="Big.md") + assert resp.success is True + assert "truncated" in str(resp.answer) + assert "start_line=" in str(resp.answer) + await store.close() print("✓ test_read_truncation passed") _run(run()) @@ -765,71 +699,80 @@ def test_read_empty_path_rejected(): async def run(): with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): - working = Path(tmp) / ".reme" - working.mkdir(parents=True, exist_ok=True) - async with mock_reme_server() as (host, port): - result = await call_action("read", host=host, port=port, path="") - if not ( - isinstance(result, dict) - and result.get("success") is False - and "required" in str(result.get("answer", "")).lower() - ): - raise AssertionError(f"expected `path` required rejection, got {result!r}") + store = await _make_store() + resp = await _read(store, path="") + assert resp.success is False + assert "required" in str(resp.answer).lower() + await store.close() print("✓ test_read_empty_path_rejected passed") _run(run()) -# -- write / edit / append ----------------------------------------------- +# -- write / edit -------------------------------------------------------- def test_write_basic_with_frontmatter(): - """`reme4 write path=... name=... description=... content=...` writes a YAML front matter block.""" + """`write path=... name=... description=... content=...` writes a YAML front matter block.""" async def run(): with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): - working = Path(tmp) / ".reme" - working.mkdir(parents=True, exist_ok=True) - async with mock_reme_server() as (host, port): - await call_and_check( - "write", - host=host, - port=port, - path="Notes/A.md", - name="Greetings", - description="a friendly hello note", - content="# Hello", - validator=lambda r: ( - isinstance(r, dict) and r.get("success") is True and "Wrote" in str(r.get("answer", "")) - ), - ) - on_disk = (working / "Notes/A.md").read_text(encoding="utf-8") + store = await _make_store() + resp = await _write( + store, + path="Notes/A.md", + name="Greetings", + description="a friendly hello note", + content="# Hello", + ) + assert resp.success is True + assert "Wrote" in str(resp.answer) + on_disk = (Path(tmp) / "Notes/A.md").read_text(encoding="utf-8") assert on_disk.startswith("---\n"), on_disk assert "name: Greetings" in on_disk assert "description: a friendly hello note" in on_disk assert "# Hello" in on_disk + await store.close() print("✓ test_write_basic_with_frontmatter passed") _run(run()) +def test_write_rejects_invalid_path_components(): + """`resolve_path` validates each segment with the same rules as daily-note slugs: + Windows reserved chars / device names / trailing-dot (also blocks `..` traversal).""" + + async def run(): + with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): + store = await _make_store() + for bad in ( + "CON.md", # Windows-reserved device name (with extension) + "Notes/AUX", # device name in a sub-segment + "Notes/foo/.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//.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!") diff --git a/tests4/unittest/test_file_catalog.py b/tests4/unittest/test_file_catalog.py new file mode 100644 index 00000000..68e50b5d --- /dev/null +++ b/tests4/unittest/test_file_catalog.py @@ -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所有测试通过!") diff --git a/tests4/unittest/test_file_store.py b/tests4/unittest/test_file_store.py index e022a149..f27234f8 100644 --- a/tests4/unittest/test_file_store.py +++ b/tests4/unittest/test_file_store.py @@ -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")]) diff --git a/tests4/unittest/test_keyword_index.py b/tests4/unittest/test_keyword_index.py new file mode 100644 index 00000000..311d550b --- /dev/null +++ b/tests4/unittest/test_keyword_index.py @@ -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)} 项测试通过!") diff --git a/tests4/unittest/test_link_expansion.py b/tests4/unittest/test_link_expansion.py new file mode 100644 index 00000000..464af763 --- /dev/null +++ b/tests4/unittest/test_link_expansion.py @@ -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() diff --git a/tests4/unittest/test_resource_steps.py b/tests4/unittest/test_resource_steps.py index aa44b622..dbdf77b8 100644 --- a/tests4/unittest/test_resource_steps.py +++ b/tests4/unittest/test_resource_steps.py @@ -1,6 +1,6 @@ -"""Tests for the resource ingest path: ``UploadResourceStep`` + helpers. +"""Tests for the resource ingest path: ``IngestStep`` + helpers. -``upload_resource`` is the **passive** ingest entry point — external channels +``ingest`` is the **passive** ingest entry point — external channels push assets into ``resource//``, where each call appends a :class:`FileNode` row to ``meta.json`` (provenance on ``front_matter``) and regenerates the day's ``.md`` view from @@ -24,8 +24,8 @@ from pathlib import Path from reme4.components.file_store import LocalFileStore from reme4.schema import FileFrontMatter, FileNode -from reme4.steps.crud import upload_resource as crud_upload -from reme4.steps.crud.upload_resource import _assemble_day_md +from reme4.steps.transfer import ingest as crud_ingest +from reme4.steps.transfer.ingest import _assemble_day_md warnings.filterwarnings("ignore", category=DeprecationWarning, module="jieba") warnings.filterwarnings("ignore", category=DeprecationWarning, module="pkg_resources") @@ -49,7 +49,7 @@ class temp_chdir: async def _make_store() -> LocalFileStore: """Minimal LocalFileStore (embedding disabled). vault_path resolves to CWD.""" - store = LocalFileStore(store_name="t", embedding_model="") + store = LocalFileStore(name="t", embedding_model="") await store.start() return store @@ -120,7 +120,7 @@ def test_validate_basename_rejects_path_separators(): """Path-separator basenames are rejected even if the public API can no longer reach this code path (Path(...).name strips them) — defense in depth.""" for bad in ("evil/payload.pdf", "..\\winpath.pdf", "../escape.pdf"): - err = crud_upload._validate_basename(bad) + err = crud_ingest._validate_basename(bad) assert "path separators" in err or "reserved" in err, (bad, err) print("✓ test_validate_basename_rejects_path_separators passed") @@ -128,7 +128,7 @@ def test_validate_basename_rejects_path_separators(): def test_validate_basename_rejects_dot_segments(): """`.` and `..` are explicitly reserved.""" for bad in (".", ".."): - err = crud_upload._validate_basename(bad) + err = crud_ingest._validate_basename(bad) assert "reserved" in err or "start with '.'" in err, (bad, err) print("✓ test_validate_basename_rejects_dot_segments passed") @@ -139,19 +139,19 @@ def test_validate_basename_rejects_dot_segments(): def test_validate_channel_accepts_safe_identifiers(): """Lowercase letters / digits / dashes, starting alnum — all accepted.""" for ok in ("wechat", "email", "api", "browser", "slack-1", "ch1"): - assert crud_upload._validate_channel(ok) == "", ok + assert crud_ingest._validate_channel(ok) == "", ok print("✓ test_validate_channel_accepts_safe_identifiers passed") def test_validate_channel_rejects_unsafe_identifiers(): """Uppercase, underscores, leading dash, empty, special chars — rejected.""" for bad in ("", "WeChat", "we_chat", "-leading", "we chat", "we/chat", "我"): - err = crud_upload._validate_channel(bad) + err = crud_ingest._validate_channel(bad) assert err, bad print("✓ test_validate_channel_rejects_unsafe_identifiers passed") -# -- UploadResourceStep end-to-end -------------------------------------- +# -- IngestStep end-to-end -------------------------------------- def test_upload_first_call_creates_bucket(): @@ -163,7 +163,7 @@ def test_upload_first_call_creates_bucket(): src = Path(tmp) / "incoming.pdf" src.write_bytes(b"%PDF-fake") - step = crud_upload.UploadResourceStep(file_store=store) + step = crud_ingest.IngestStep(file_store=store) await step( path=str(src), channel="wechat", @@ -209,7 +209,7 @@ def test_upload_metadata_optional(): store = await _make_store() src = Path(tmp) / "small.txt" src.write_text("x") - step = crud_upload.UploadResourceStep(file_store=store) + step = crud_ingest.IngestStep(file_store=store) await step(path=str(src), channel="api", description="minimal") payload = _metadata(step) assert "error" not in payload, payload @@ -233,7 +233,7 @@ def test_upload_appends_to_existing_meta(): for i, suffix in enumerate(("first", "second"), start=1): src = Path(tmp) / f"{suffix}.txt" src.write_text(f"payload-{i}") - step = crud_upload.UploadResourceStep(file_store=store) + step = crud_ingest.IngestStep(file_store=store) await step(path=str(src), channel="email", description=f"item {i}") payload = _metadata(step) assert "error" not in payload, payload @@ -267,12 +267,12 @@ def test_upload_errors_on_duplicate_same_second(monkeypatch): def now(cls, tz=None): # pylint: disable=unused-argument return fixed - monkeypatch.setattr(crud_upload.datetime, "datetime", _FrozenDT) + monkeypatch.setattr(crud_ingest.datetime, "datetime", _FrozenDT) for i, body in enumerate((b"alpha", b"beta")): src = Path(tmp) / "incoming.pdf" src.write_bytes(body) - step = crud_upload.UploadResourceStep(file_store=store) + step = crud_ingest.IngestStep(file_store=store) await step(path=str(src), channel="wechat", description="dup test") payload = _metadata(step) if i == 0: @@ -306,7 +306,7 @@ def test_upload_errors_on_duplicate_against_on_disk_stray(monkeypatch): def now(cls, tz=None): # pylint: disable=unused-argument return fixed - monkeypatch.setattr(crud_upload.datetime, "datetime", _FrozenDT) + monkeypatch.setattr(crud_ingest.datetime, "datetime", _FrozenDT) bucket = Path(tmp) / "resource" / "2026-05-22" bucket.mkdir(parents=True) @@ -315,7 +315,7 @@ def test_upload_errors_on_duplicate_against_on_disk_stray(monkeypatch): src = Path(tmp) / "report.pdf" src.write_bytes(b"fresh") - step = crud_upload.UploadResourceStep(file_store=store) + step = crud_ingest.IngestStep(file_store=store) await step(path=str(src), channel="api", description="fresh copy") payload = _metadata(step) assert "duplicate" in payload.get("error", "").lower(), payload @@ -333,7 +333,7 @@ def test_upload_rejects_missing_source(): async def run(): with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp): store = await _make_store() - step = crud_upload.UploadResourceStep(file_store=store) + step = crud_ingest.IngestStep(file_store=store) await step( path=str(Path(tmp) / "ghost.txt"), channel="email", @@ -357,7 +357,7 @@ def test_upload_requires_channel(): src = Path(tmp) / "x.txt" src.write_text("x") for bad in (" ", "WeChat", "we_chat"): - step = crud_upload.UploadResourceStep(file_store=store) + step = crud_ingest.IngestStep(file_store=store) await step(path=str(src), channel=bad, description="x") payload = _metadata(step) assert "channel" in payload.get("error", ""), (bad, payload) @@ -375,7 +375,7 @@ def test_upload_requires_description(): store = await _make_store() src = Path(tmp) / "x.txt" src.write_text("x") - step = crud_upload.UploadResourceStep(file_store=store) + step = crud_ingest.IngestStep(file_store=store) await step(path=str(src), channel="api", description=" ") payload = _metadata(step) assert "description" in payload.get("error", "") @@ -393,7 +393,7 @@ def test_upload_rejects_non_dict_metadata(): store = await _make_store() src = Path(tmp) / "x.txt" src.write_text("x") - step = crud_upload.UploadResourceStep(file_store=store) + step = crud_ingest.IngestStep(file_store=store) await step( path=str(src), channel="api", @@ -418,7 +418,7 @@ def test_upload_rejects_reserved_metadata_keys(): src = Path(tmp) / "x.txt" src.write_text("x") for bad in ("name", "channel", "received_at", "description"): - step = crud_upload.UploadResourceStep(file_store=store) + step = crud_ingest.IngestStep(file_store=store) await step( path=str(src), channel="api", @@ -441,7 +441,7 @@ def test_upload_preserves_extra_metadata_keys(): store = await _make_store() src = Path(tmp) / "x.txt" src.write_text("x") - step = crud_upload.UploadResourceStep(file_store=store) + step = crud_ingest.IngestStep(file_store=store) await step( path=str(src), channel="api", @@ -477,7 +477,7 @@ def test_upload_rejects_dotfile_source(): for bad in (".hidden", ".lock", ".env"): src = Path(tmp) / bad src.write_text("x") - step = crud_upload.UploadResourceStep(file_store=store) + step = crud_ingest.IngestStep(file_store=store) await step( path=str(src), channel="api", @@ -504,7 +504,7 @@ def test_upload_records_received_at_internally(): store = await _make_store() src = Path(tmp) / "doc.pdf" src.write_bytes(b"%PDF") - step = crud_upload.UploadResourceStep(file_store=store) + step = crud_ingest.IngestStep(file_store=store) await step(path=str(src), channel="api", description="x") payload = _metadata(step) assert "error" not in payload @@ -533,7 +533,7 @@ def test_upload_preserves_description_verbatim_in_meta(): src = Path(tmp) / "doc.pdf" src.write_bytes(b"%PDF") multi = "wechat group screenshot\nfrom design-group at 14:30\nlikely a Q1 KPI table — extract numbers" - step = crud_upload.UploadResourceStep(file_store=store) + step = crud_ingest.IngestStep(file_store=store) await step( path=str(src), channel="api", diff --git a/tests4/unittest/test_wikilink_utils.py b/tests4/unittest/test_wikilink_utils.py index 1454afd2..c43aa510 100644 --- a/tests4/unittest/test_wikilink_utils.py +++ b/tests4/unittest/test_wikilink_utils.py @@ -51,7 +51,7 @@ async def _store_with(files: dict[str, str]) -> LocalFileStore: Without the parsed links the reverse-index lookup yields nothing and retarget becomes a no-op. """ - store = LocalFileStore(store_name="t", embedding_model="") + store = LocalFileStore(name="t", embedding_model="") await store.start() nodes: list[FileNode] = [] root = Path.cwd() @@ -72,7 +72,7 @@ async def _store_with(files: dict[str, str]) -> LocalFileStore: async def _empty_store() -> LocalFileStore: - store = LocalFileStore(store_name="t", embedding_model="") + store = LocalFileStore(name="t", embedding_model="") await store.start() return store