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..969fcd20 100644 --- a/reme4/application.py +++ b/reme4/application.py @@ -20,10 +20,14 @@ class Application(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) + if self.config.metadata_dir: + (vault_path / self.config.metadata_dir).mkdir(parents=True, exist_ok=True) + if self.config.daily_dir: + (vault_path / self.config.daily_dir).mkdir(parents=True, exist_ok=True) + if self.config.digest_dir: + (vault_path / self.config.digest_dir).mkdir(parents=True, exist_ok=True) + if self.config.resource_dir: + (vault_path / self.config.resource_dir).mkdir(parents=True, exist_ok=True) if self.config.enable_logo: print_logo(self.config) @@ -60,15 +64,16 @@ class Application(BaseComponent): self.context.components[component_type][name] = backend_cls(**params) # Jobs - for job_config in self.config.jobs: + for name, job_config in self.config.jobs.items(): if not job_config.backend: - raise ValueError(f"Job '{job_config.name}' is missing the required 'backend' field") + raise ValueError(f"Job '{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}'") + raise ValueError(f"Unregistered backend '{job_config.backend}' for job '{name}'") params = job_config.model_dump() + params.setdefault("name", name) params["app_context"] = self.context - self.context.jobs[job_config.name] = job_cls(**params) + self.context.jobs[name] = job_cls(**params) @property def config(self): 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 ba238cb0..9d8bf50e 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,10 @@ 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 +63,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,87 +88,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: - """Resolved component metadata directory: vault_metadata_path / component_type.""" + """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 @@ -173,7 +202,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 @@ -183,7 +212,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..3e59b88a 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/file_store/local_file_store.py b/reme4/components/file_store/local_file_store.py index 171aa66f..c3941373 100644 --- a/reme4/components/file_store/local_file_store.py +++ b/reme4/components/file_store/local_file_store.py @@ -111,40 +111,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 diff --git a/reme4/components/job/background_job.py b/reme4/components/job/background_job.py index 12251122..c72de54f 100644 --- a/reme4/components/job/background_job.py +++ b/reme4/components/job/background_job.py @@ -35,9 +35,10 @@ class BackgroundJob(BaseJob): 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 @@ -54,18 +55,38 @@ class BackgroundJob(BaseJob): async def _close(self) -> None: if self._stop_event is not None: self._stop_event.set() - if self._task is not None: - try: - 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 + 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 @@ -77,16 +98,13 @@ class BackgroundJob(BaseJob): except Exception as e: if not self.supervisor: raise + # 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 - capped = min(self.backoff_base * (2**attempt), self.backoff_cap) - delay = min(capped * (0.5 + random.random()), self.backoff_cap) + 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 steps in order; errors propagate to supervisor.""" diff --git a/reme4/components/job/base_job.py b/reme4/components/job/base_job.py index 860849f7..3a414bee 100644 --- a/reme4/components/job/base_job.py +++ b/reme4/components/job/base_job.py @@ -23,12 +23,14 @@ 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 [] + 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]] = [] diff --git a/reme4/components/job/stream_job.py b/reme4/components/job/stream_job.py index 193bf5e2..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._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..20e7e9da 100644 --- a/reme4/components/keyword_index/base_keyword_index.py +++ b/reme4/components/keyword_index/base_keyword_index.py @@ -1,7 +1,4 @@ -"""Abstract base class for keyword index implementations.""" - from abc import abstractmethod -from pathlib import Path from ..base_component import BaseComponent from ..tokenizer import BaseTokenizer @@ -9,62 +6,48 @@ from ...enumeration import ComponentEnum class BaseKeywordIndex(BaseComponent): - """Abstract base class for keyword index implementations.""" + """关键词索引基类:定义增、删、查、清的统一接口,由具体实现(如 BM25)继承。""" 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 + # 绑定分词器,未显式指定时回落到 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.""" + """对单段文本调用分词器,返回 token 列表。""" 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.""" + async def add_docs(self, docs_dict: dict[str, str]) -> None: ... @abstractmethod - async def delete_docs(self, doc_ids: list[str]) -> None: - """Remove documents by their IDs.""" + async def delete_docs(self, doc_ids: list[str]) -> None: ... @abstractmethod - async def retrieve(self, query: str, limit: int = 3) -> dict[str, float]: - """Search documents. Returns {doc_id: score} sorted descending.""" + async def retrieve(self, query: str, limit: int = 3) -> dict[str, float]: ... @abstractmethod - async def clear(self) -> None: - """Reset index to empty state.""" + async def clear(self) -> None: ... async def reset_index(self, docs_dict: dict[str, str]) -> None: - """Clear index, re-add all documents, and persist.""" + """清空索引后重新构建,并立即落盘。""" await self.clear() await self.add_docs(docs_dict) await self.dump() async def optimize_index(self) -> None: - """Optimize index for performance. Override in subclass if needed.""" + """对索引进行物理压缩或重建;基类默认无操作,由子类按需重载。""" + pass diff --git a/reme4/components/keyword_index/bm25_index.py b/reme4/components/keyword_index/bm25_index.py index 31d978a9..d5058f0e 100644 --- a/reme4/components/keyword_index/bm25_index.py +++ b/reme4/components/keyword_index/bm25_index.py @@ -1,27 +1,24 @@ -"""BM25 search engine with persistent index support. +"""基于 BM25 的倒排索引实现,支持持久化。 -Implements Okapi BM25 ranking with a numpy-vectorized inverted index for -efficient document lookup, incremental updates, and pickle-based persistence. - -Storage layout (source of truth): +核心存储结构(落盘真相源): vocab : dict[token, token_id] - _doc_ids : list[doc_id] indexed by doc_idx + _doc_ids : list[doc_id],按 doc_idx 索引 _doc_id_to_idx : dict[doc_id, doc_idx] - _doc_lens : np.ndarray[int32] indexed by doc_idx - _deleted : np.ndarray[bool] indexed by doc_idx (lazy deletion) - _doc_token_ids : list[np.ndarray[int32]] indexed by doc_idx (unique tids per doc) - _posting_doc_idxs : dict[token_id, np.ndarray[int32]] posting list doc_idxs - _posting_tfs : dict[token_id, np.ndarray[int32]] posting list tfs (parallel) + _doc_lens : np.ndarray[int32],按 doc_idx 索引 + _deleted : np.ndarray[bool],按 doc_idx 索引(懒删除标记) + _doc_token_ids : list[np.ndarray[int32]],每篇文档去重后的 token_id + _posting_doc_idxs : dict[token_id, np.ndarray[int32]],倒排表的 doc_idx + _posting_tfs : dict[token_id, np.ndarray[int32]],与上方一一对应的词频 -Deletion is lazy: ``_remove_doc`` only flips ``_deleted[idx]``; posting entries -pointing at the dead idx are masked at query time and physically dropped by -``optimize_index``. Updating an existing doc_id marks the old slot deleted and -allocates a fresh idx for the new content. +删除采用懒标记:_deleted[idx] = True 即视为删除,倒排表中的物理回收由 +optimize_index 统一完成;更新已存在的 doc_id 时,先把旧槽位标记删除,再 +分配新的 idx。 """ import math import pickle from collections import Counter +from pathlib import Path import numpy as np @@ -31,82 +28,77 @@ from ..component_registry import R @R.register("bm25") class BM25Index(BaseKeywordIndex): - """BM25 search engine with numpy-vectorized scoring and file-based persistence. - Args: - k1: Term frequency saturation parameter (default 1.5). - b: Document length normalization parameter (default 0.75). - """ - - def __init__(self, k1: float = 1.5, b: float = 0.75, **kwargs): + 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.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] = [] + + # 倒排表:token_id -> (doc_idxs, tfs) self._posting_doc_idxs: dict[int, np.ndarray] = {} self._posting_tfs: dict[int, np.ndarray] = {} + + # IDF 缓存,对增删与重建索引时失效 self._idf_cache: dict[int, float] = {} # -- Properties ----------------------------------------------------------- + @property + def index_file(self) -> Path: + """落盘文件路径,包含分词器名与索引版本,便于区分不同配置。""" + 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 (non-deleted) documents.""" - if self._deleted.size == 0: - return 0 - return int((~self._deleted).sum()) + """当前存活文档数(排除懒删除)。""" + return 0 if self._deleted.size == 0 else int((~self._deleted).sum()) @property def total_len(self) -> int: - """Total tokens across non-deleted documents.""" - if self._deleted.size == 0: - return 0 - return int(self._doc_lens[~self._deleted].sum()) + """所有存活文档的 token 总数。""" + 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 (non-deleted only).""" + """存活文档的平均长度,用于 BM25 长度归一化。""" n = self.n_docs return self.total_len / n if n > 0 else 0.0 @property def doc_meta(self) -> dict[str, dict]: - """Dict-view of {doc_id: {"len", "token_ids"}} for non-deleted docs. - - Built on demand from the numpy-backed storage; kept for backward - compatibility with callers (and tests) that read this shape. - """ - out: dict[str, dict] = {} - for idx, doc_id in enumerate(self._doc_ids): - if self._deleted[idx]: - continue - out[doc_id] = { + """对外暴露每篇存活文档的长度与去重后的 token_id 集合。""" + return { + self._doc_ids[idx]: { "len": int(self._doc_lens[idx]), "token_ids": {int(t) for t in self._doc_token_ids[idx]}, } - return out + for idx in range(len(self._doc_ids)) + if not self._deleted[idx] + } @property def inverted_index(self) -> dict[int, dict[str, int]]: - """Dict-view of {token_id: {doc_id: tf}} excluding deleted docs. - - Built on demand from the numpy-backed storage; kept for backward - compatibility. Empty posting lists (all entries deleted) are omitted. - """ + """重建可读形式的倒排表:token_id -> {doc_id: tf},跳过已删除文档。""" out: dict[int, dict[str, int]] = {} for tid, doc_idxs in self._posting_doc_idxs.items(): tfs = self._posting_tfs[tid] - posting: dict[str, int] = {} - for i, tf in zip(doc_idxs, tfs): - i = int(i) - if self._deleted[i]: - continue - posting[self._doc_ids[i]] = int(tf) + 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 @@ -114,9 +106,9 @@ class BM25Index(BaseKeywordIndex): # -- Internal helpers ----------------------------------------------------- def _tokens_to_ids(self, tokens: list[str]) -> list[int]: - """Map tokens to integer IDs, assigning new IDs on first encounter.""" + """将 token 转为 id;遇到新词时自动分配新的 token_id。""" vocab = self.vocab - ids = [] + ids: list[int] = [] for token in tokens: token = token.strip() if not token: @@ -129,7 +121,7 @@ class BM25Index(BaseKeywordIndex): return ids def _remove_doc(self, doc_id: str) -> None: - """Mark a document deleted. Posting cleanup deferred to ``optimize_index``.""" + """懒删除:仅置 _deleted 位并解除 doc_id 映射,不动倒排表。""" idx = self._doc_id_to_idx.get(doc_id) if idx is None or self._deleted[idx]: return @@ -138,7 +130,7 @@ class BM25Index(BaseKeywordIndex): self._idf_cache = {} def _get_idf(self, token_id: int, n_docs: int | None = None) -> float: - """Compute and cache IDF for a token ID against current active doc set.""" + """计算并缓存 token 的 IDF;存活文档数发生变化时缓存会被清空。""" if token_id in self._idf_cache: return self._idf_cache[token_id] doc_idxs = self._posting_doc_idxs.get(token_id) @@ -146,68 +138,36 @@ class BM25Index(BaseKeywordIndex): self._idf_cache[token_id] = 0.0 return 0.0 df = int((~self._deleted[doc_idxs]).sum()) - if df == 0: - self._idf_cache[token_id] = 0.0 - return 0.0 if n_docs is None: n_docs = self.n_docs - self._idf_cache[token_id] = math.log(1 + (n_docs - df + 0.5) / (df + 0.5)) - return self._idf_cache[token_id] + 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 - # -- Public API ----------------------------------------------------------- + def _prepare_doc(self, doc_id: str, content: str) -> tuple[np.ndarray, int, Counter] | None: + """分词并统计词频;若 doc_id 已存在则先标记旧版为删除。空文档返回 None。""" + 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 - async def add_docs(self, docs_dict: dict[str, str]) -> None: - """Index or update multiple documents. Mapping of doc_id to content. - - Updating an existing doc_id marks the old slot deleted and allocates - a new doc_idx, so the next ``optimize_index`` reclaims its postings. - """ - if not docs_dict: + def _append_doc_arrays( + self, new_doc_ids: list[str], new_doc_lens: list[int], new_doc_token_ids: list[np.ndarray] + ) -> None: + """把一批新文档的元数据一次性追加到文档数组中。""" + 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)]) - new_doc_ids: list[str] = [] - new_doc_lens: list[int] = [] - new_doc_token_ids: list[np.ndarray] = [] - pending_postings: dict[int, list[tuple[int, int]]] = {} - - next_idx = len(self._doc_ids) - - for doc_id, content in docs_dict.items(): - old_idx = self._doc_id_to_idx.get(doc_id) - if old_idx is not None and not self._deleted[old_idx]: - self._deleted[old_idx] = True - self._doc_id_to_idx.pop(doc_id, None) - - token_ids = self._tokens_to_ids(self._tokenize(content)) - if not token_ids: - continue - - token_counts = Counter(token_ids) - unique_tids = np.fromiter( - token_counts.keys(), dtype=np.int32, count=len(token_counts) - ) - - idx = next_idx - next_idx += 1 - new_doc_ids.append(doc_id) - new_doc_lens.append(len(token_ids)) - new_doc_token_ids.append(unique_tids) - self._doc_id_to_idx[doc_id] = idx - - for tid, tf in token_counts.items(): - pending_postings.setdefault(tid, []).append((idx, tf)) - - if new_doc_ids: - 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)] - ) - - for tid, items in pending_postings.items(): + def _extend_postings(self, pending: dict[int, list[tuple[int, int]]]) -> None: + """把待写入的 (doc_idx, tf) 增量按 token 追加到倒排表。""" + 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) @@ -218,113 +178,150 @@ class BM25Index(BaseKeywordIndex): self._posting_doc_idxs[tid] = new_idxs self._posting_tfs[tid] = new_tfs + def _encode_query(self, query: str) -> list[int]: + """切词、过滤未登录词并去重,返回查询的 token_id 列表。""" + 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: + """挑出得分前 limit 名(且严格大于 0)的索引,按得分降序排列。""" + 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: + """批量加入文档;已存在的 doc_id 会被替换为新版本。""" + 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(): + prepared = self._prepare_doc(doc_id, content) + if prepared is None: + continue + 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(): + 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.""" + """批量懒删除;倒排表中的物理回收由 optimize_index 完成。""" for doc_id in doc_ids: self._remove_doc(doc_id) self._idf_cache = {} - async def retrieve(self, query: str, limit: int = 3) -> dict[str, float]: - """Search documents. Returns {doc_id: score} sorted descending.""" - n_slots = self._doc_lens.size - if n_slots == 0: - return {} - - vocab = self.vocab - query_ids = list(dict.fromkeys(vocab[t] for t in self._tokenize(query) if t in vocab)) - if not query_ids: - return {} - - n_docs = self.n_docs - if n_docs == 0: - return {} - + def _score_query(self, query_ids: list[int], n_docs: int) -> np.ndarray: + """对所有文档计算 BM25 得分,已删除文档置 0。""" 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(n_slots, dtype=np.float32) - + 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=n_docs) + 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) - tf_score = tfs * (k1 + 1.0) / (tfs + denom_base + denom_norm * d_lens) - # Each doc_idx appears at most once per posting list (Counter dedups - # within a doc, and updates allocate a fresh idx), so direct - # advanced-indexing assignment-add is safe. - scores[doc_idxs] += idf * tf_score + # 同一倒排表中每个 doc_idx 至多出现一次:Counter 已在文档内去重, + # 文档更新也会分配新的 idx,因此可安全使用花式索引累加。 + 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 - positive_count = int((scores > 0).sum()) - if positive_count == 0: + async def retrieve(self, query: str, limit: int = 3) -> dict[str, float]: + """对查询做 BM25 召回,返回 {doc_id: score},按得分降序。""" + n_docs = self.n_docs + if n_docs == 0: + return {} + query_ids = self._encode_query(query) + if not query_ids: return {} - k = min(limit, positive_count) - if k >= n_slots: - top_idxs = np.argsort(-scores)[:k] - else: - top_idxs = np.argpartition(-scores, k - 1)[:k] - top_idxs = top_idxs[np.argsort(-scores[top_idxs])] + 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} + + # -- Persistence ---------------------------------------------------------- + + def _snapshot(self) -> dict: + """收集需要落盘的全部字段,集中在一处以便与 _restore 对齐。""" return { - self._doc_ids[int(i)]: float(scores[int(i)]) - for i in top_idxs - if scores[int(i)] > 0 + "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: + """从 _snapshot 产生的字典还原索引内部状态。""" + 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).""" + """通过临时文件 + 原子替换的方式持久化索引,避免半写状态。""" try: tmp = self.index_file.with_suffix(".tmp") with open(tmp, "wb") as f: - pickle.dump( - { - "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, - }, - 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.""" + """读取持久化文件并还原索引;文件不存在则不做事,损坏则清空。""" 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._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 = {} + 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}") @@ -332,7 +329,7 @@ class BM25Index(BaseKeywordIndex): await self.clear() async def clear(self) -> None: - """Reset index to empty state and remove persisted file.""" + """清空内存中的索引并删除持久化文件。""" self.vocab = {} self._doc_ids = [] self._doc_id_to_idx = {} @@ -344,64 +341,86 @@ class BM25Index(BaseKeywordIndex): self._idf_cache = {} self.index_file.unlink(missing_ok=True) + # -- Compaction ----------------------------------------------------------- + + def _build_idx_remap(self, active_mask: np.ndarray) -> tuple[np.ndarray, int]: + """构造 old_idx → new_idx 的映射数组(被删槽位为 -1),并返回存活数量。""" + 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]]: + """只保留仍被任意存活文档引用的 token,重排成连续的新 token_id。""" + 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_tids: + new_tid = len(new_vocab) + new_vocab[token] = new_tid + old_to_new[old_tid] = new_tid + return new_vocab, 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]]: + """剔除删除文档并按新 idx/tid 重写倒排表。""" + 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]]: + """在压缩后的词表下重建存活文档的 doc_id 列表与去重 token_id 数组。""" + 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: - """Compact: drop deleted docs, reassign doc_idx, prune unused tokens.""" + """物理回收懒删除的文档与未被引用的词表项,重建紧凑索引。""" if self._deleted.size == 0: return - active_mask = ~self._deleted if not active_mask.any(): await self.clear() return - active_old_idxs = np.where(active_mask)[0] - n_active = int(active_old_idxs.size) - old_to_new_idx = -np.ones(self._deleted.size, dtype=np.int32) - old_to_new_idx[active_old_idxs] = np.arange(n_active, dtype=np.int32) - - new_doc_ids = [self._doc_ids[int(i)] for i in active_old_idxs] - new_doc_lens = self._doc_lens[active_mask].astype(np.int32, copy=True) - new_doc_token_ids_pre = [self._doc_token_ids[int(i)] for i in active_old_idxs] - new_doc_id_to_idx = {doc_id: i for i, doc_id in enumerate(new_doc_ids)} - - used_tids: set[int] = set() - for tid, doc_idxs in self._posting_doc_idxs.items(): - if active_mask[doc_idxs].any(): - used_tids.add(tid) - - old_tid_to_new: dict[int, int] = {} - new_vocab: dict[str, int] = {} - for token, old_tid in self.vocab.items(): - if old_tid in used_tids: - new_tid = len(new_vocab) - new_vocab[token] = new_tid - old_tid_to_new[old_tid] = new_tid - - new_posting_doc_idxs: dict[int, np.ndarray] = {} - new_posting_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] - kept_idxs = old_to_new_idx[doc_idxs[mask]].astype(np.int32, copy=False) - kept_tfs = self._posting_tfs[tid][mask].astype(np.int32, copy=False) - new_posting_doc_idxs[old_tid_to_new[tid]] = kept_idxs - new_posting_tfs[old_tid_to_new[tid]] = kept_tfs - - new_doc_token_ids: list[np.ndarray] = [] - for arr in new_doc_token_ids_pre: - remapped = np.fromiter( - (old_tid_to_new[int(t)] for t in arr if int(t) in old_tid_to_new), - dtype=np.int32, - ) - new_doc_token_ids.append(remapped) + 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._doc_ids = new_doc_ids - self._doc_id_to_idx = new_doc_id_to_idx - self._doc_lens = new_doc_lens + 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_doc_idxs + 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..db97a963 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 8d4d2402..2384fe71 100644 --- a/reme4/components/service/base_service.py +++ b/reme4/components/service/base_service.py @@ -60,9 +60,9 @@ class BaseService(BaseComponent): return lifespan def add_jobs(self, app: "Application") -> None: - """Register every job from the app context except background-only ones.""" + """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) diff --git a/reme4/components/tokenizer/jieba_tokenizer.py b/reme4/components/tokenizer/jieba_tokenizer.py index 1dc98b13..e63c7a71 100644 --- a/reme4/components/tokenizer/jieba_tokenizer.py +++ b/reme4/components/tokenizer/jieba_tokenizer.py @@ -1,15 +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 backed by jieba for Chinese word segmentation.""" + """Tokenizer backed by jieba for Chinese word segmentation. + + `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) + 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 + + 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 + + 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]: - # Lazy import: jieba startup cost is non-trivial and only paid when used. - import jieba - - return list(jieba.cut(text)) + return list(self._cut(text)) diff --git a/reme4/config/default.yaml b/reme4/config/default.yaml index 73afdbe8..95bb8369 100644 --- a/reme4/config/default.yaml +++ b/reme4/config/default.yaml @@ -14,8 +14,8 @@ jobs: # ════════════════════════════════════════════════════════════════════ # UTILITY — service introspection # ════════════════════════════════════════════════════════════════════ - - backend: base - name: version + version: + backend: base description: "return reme4 package version" parameters: type: object @@ -23,8 +23,8 @@ jobs: steps: - backend: version_step - - backend: base - name: health_check + health_check: + backend: base description: "return a concise health-check snapshot of reme4 components" parameters: type: object @@ -32,8 +32,8 @@ jobs: steps: - backend: health_check_step - - backend: base - name: help + help: + backend: base description: "list all registered jobs with their metadata" parameters: type: object @@ -41,8 +41,8 @@ jobs: steps: - backend: help_step - - backend: base - name: reindex + reindex: + backend: base description: "wipe the file store and rebuild it from the watcher's tracked files" parameters: type: object @@ -50,8 +50,8 @@ jobs: steps: - backend: reindex_step - - backend: base - name: index_changes + index_changes: + backend: base description: "apply a batch of file changes (added/modified/deleted) into file_store" parameters: type: object @@ -82,8 +82,8 @@ jobs: # ════════════════════════════════════════════════════════════════════ # ── Retrieve ─────────────────────────────────────────────────────── - - backend: base - name: search + search: + backend: base description: "Hybrid vault search (vector + BM25, RRF-fused)." parameters: type: object @@ -108,8 +108,8 @@ jobs: expand_links: true max_links_per_direction: 10 - - backend: base - name: traverse + traverse: + backend: base description: "Walk the wikilink graph from a seed path." parameters: type: object @@ -131,8 +131,8 @@ jobs: - backend: traverse_step # ── Read Operations ─────────────────────────────────────────────────────────── - - backend: base - name: list + list: + backend: base description: "List files under a vault path." parameters: type: object @@ -152,8 +152,8 @@ jobs: steps: - backend: list_step - - backend: base - name: read + read: + backend: base description: "Read a markdown file under the vault." parameters: type: object @@ -172,8 +172,8 @@ jobs: steps: - backend: read_step - - backend: base - name: stat + stat: + backend: base description: "Stat a vault file (size, mtime, exists, is_dir, is_file)." parameters: type: object @@ -186,8 +186,8 @@ jobs: steps: - backend: stat_step - - backend: base - name: frontmatter:read + frontmatter:read: + backend: base description: "Read a file's YAML frontmatter as a dict." parameters: type: object @@ -201,8 +201,8 @@ jobs: - backend: frontmatter:read_step # ── Write Operations────────────────────────────────────────────────────────── - - backend: base - name: write + write: + backend: base description: "Write a markdown file (create or overwrite) with name/description frontmatter." parameters: type: object @@ -227,8 +227,8 @@ jobs: steps: - backend: write_step - - backend: base - name: edit + edit: + backend: base description: "Find-and-replace in a markdown file (all occurrences)." parameters: type: object @@ -250,8 +250,8 @@ jobs: steps: - backend: edit_step - - backend: base - name: append + append: + backend: base description: "Append content to a markdown file." parameters: type: object @@ -268,8 +268,8 @@ jobs: steps: - backend: append_step - - backend: base - name: frontmatter:update + frontmatter:update: + backend: base description: "Merge keys into a file's YAML frontmatter." parameters: type: object @@ -287,8 +287,8 @@ jobs: steps: - backend: frontmatter_update_step - - backend: base - name: frontmatter:delete + frontmatter:delete: + backend: base description: "Drop keys from a file's YAML frontmatter." parameters: type: object @@ -308,8 +308,8 @@ jobs: - 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 +334,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,8 +348,8 @@ jobs: steps: - backend: delete_step - - backend: base - name: upload + upload: + backend: base description: "Copy a host file into the vault at an explicit destination." parameters: type: object @@ -370,8 +370,8 @@ jobs: steps: - backend: upload_step - - backend: base - name: upload_resource + upload_resource: + backend: base description: "Ingest an external-channel asset into resource// with provenance." parameters: type: object @@ -396,8 +396,8 @@ jobs: steps: - backend: upload_resource_step - - backend: base - name: download + download: + backend: base description: "Copy a vault file out to the host filesystem." parameters: type: object @@ -419,8 +419,8 @@ jobs: - backend: download_step # ── Daily Operations (note CRUD + day-index rollup) ─────────────────── - - backend: base - name: daily:read + daily:read: + backend: base description: "Read daily//.md (body + frontmatter)." parameters: type: object @@ -437,8 +437,8 @@ jobs: steps: - backend: daily_read_step - - backend: base - name: daily:write + daily:write: + backend: base description: "Write daily//.md (body + frontmatter); refreshes the day index." parameters: type: object @@ -471,8 +471,8 @@ jobs: steps: - backend: daily_write_step - - backend: base - name: daily:list + daily:list: + backend: base description: "List notes under a single day." parameters: type: object @@ -484,8 +484,8 @@ jobs: steps: - backend: daily_list_step - - backend: base - name: daily:reindex + daily:reindex: + backend: base description: "Rebuild the day-index page daily/.md." parameters: type: object @@ -497,8 +497,8 @@ jobs: steps: - backend: daily_reindex_step - - backend: background - name: watch_file + watch_file: + backend: background watch_paths: - MEMORY.md - memory diff --git a/reme4/config/demo.yaml b/reme4/config/demo.yaml index 6fce17eb..51e6bd5d 100644 --- a/reme4/config/demo.yaml +++ b/reme4/config/demo.yaml @@ -2,8 +2,8 @@ service: backend: http jobs: - - backend: base - name: version + version: + backend: base description: "return reme4 package version" parameters: type: object @@ -11,8 +11,8 @@ jobs: steps: - backend: version_step - - backend: base - name: help + help: + backend: base description: "list all registered jobs with their metadata" parameters: type: object @@ -20,8 +20,8 @@ jobs: steps: - backend: help_step - - backend: base - name: demo + demo: + backend: base description: "demo job description" parameters: type: object @@ -39,8 +39,8 @@ jobs: - backend: demo_echo_step1 - backend: demo_echo_step2 - - backend: stream - name: stream_demo + stream_demo: + backend: stream description: "stream demo job: repeat query 10x and stream char-by-char" parameters: type: object diff --git a/reme4/config/qwenpaw.yaml b/reme4/config/qwenpaw.yaml index 73afdbe8..95bb8369 100644 --- a/reme4/config/qwenpaw.yaml +++ b/reme4/config/qwenpaw.yaml @@ -14,8 +14,8 @@ jobs: # ════════════════════════════════════════════════════════════════════ # UTILITY — service introspection # ════════════════════════════════════════════════════════════════════ - - backend: base - name: version + version: + backend: base description: "return reme4 package version" parameters: type: object @@ -23,8 +23,8 @@ jobs: steps: - backend: version_step - - backend: base - name: health_check + health_check: + backend: base description: "return a concise health-check snapshot of reme4 components" parameters: type: object @@ -32,8 +32,8 @@ jobs: steps: - backend: health_check_step - - backend: base - name: help + help: + backend: base description: "list all registered jobs with their metadata" parameters: type: object @@ -41,8 +41,8 @@ jobs: steps: - backend: help_step - - backend: base - name: reindex + reindex: + backend: base description: "wipe the file store and rebuild it from the watcher's tracked files" parameters: type: object @@ -50,8 +50,8 @@ jobs: steps: - backend: reindex_step - - backend: base - name: index_changes + index_changes: + backend: base description: "apply a batch of file changes (added/modified/deleted) into file_store" parameters: type: object @@ -82,8 +82,8 @@ jobs: # ════════════════════════════════════════════════════════════════════ # ── Retrieve ─────────────────────────────────────────────────────── - - backend: base - name: search + search: + backend: base description: "Hybrid vault search (vector + BM25, RRF-fused)." parameters: type: object @@ -108,8 +108,8 @@ jobs: expand_links: true max_links_per_direction: 10 - - backend: base - name: traverse + traverse: + backend: base description: "Walk the wikilink graph from a seed path." parameters: type: object @@ -131,8 +131,8 @@ jobs: - backend: traverse_step # ── Read Operations ─────────────────────────────────────────────────────────── - - backend: base - name: list + list: + backend: base description: "List files under a vault path." parameters: type: object @@ -152,8 +152,8 @@ jobs: steps: - backend: list_step - - backend: base - name: read + read: + backend: base description: "Read a markdown file under the vault." parameters: type: object @@ -172,8 +172,8 @@ jobs: steps: - backend: read_step - - backend: base - name: stat + stat: + backend: base description: "Stat a vault file (size, mtime, exists, is_dir, is_file)." parameters: type: object @@ -186,8 +186,8 @@ jobs: steps: - backend: stat_step - - backend: base - name: frontmatter:read + frontmatter:read: + backend: base description: "Read a file's YAML frontmatter as a dict." parameters: type: object @@ -201,8 +201,8 @@ jobs: - backend: frontmatter:read_step # ── Write Operations────────────────────────────────────────────────────────── - - backend: base - name: write + write: + backend: base description: "Write a markdown file (create or overwrite) with name/description frontmatter." parameters: type: object @@ -227,8 +227,8 @@ jobs: steps: - backend: write_step - - backend: base - name: edit + edit: + backend: base description: "Find-and-replace in a markdown file (all occurrences)." parameters: type: object @@ -250,8 +250,8 @@ jobs: steps: - backend: edit_step - - backend: base - name: append + append: + backend: base description: "Append content to a markdown file." parameters: type: object @@ -268,8 +268,8 @@ jobs: steps: - backend: append_step - - backend: base - name: frontmatter:update + frontmatter:update: + backend: base description: "Merge keys into a file's YAML frontmatter." parameters: type: object @@ -287,8 +287,8 @@ jobs: steps: - backend: frontmatter_update_step - - backend: base - name: frontmatter:delete + frontmatter:delete: + backend: base description: "Drop keys from a file's YAML frontmatter." parameters: type: object @@ -308,8 +308,8 @@ jobs: - 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 +334,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,8 +348,8 @@ jobs: steps: - backend: delete_step - - backend: base - name: upload + upload: + backend: base description: "Copy a host file into the vault at an explicit destination." parameters: type: object @@ -370,8 +370,8 @@ jobs: steps: - backend: upload_step - - backend: base - name: upload_resource + upload_resource: + backend: base description: "Ingest an external-channel asset into resource// with provenance." parameters: type: object @@ -396,8 +396,8 @@ jobs: steps: - backend: upload_resource_step - - backend: base - name: download + download: + backend: base description: "Copy a vault file out to the host filesystem." parameters: type: object @@ -419,8 +419,8 @@ jobs: - backend: download_step # ── Daily Operations (note CRUD + day-index rollup) ─────────────────── - - backend: base - name: daily:read + daily:read: + backend: base description: "Read daily//.md (body + frontmatter)." parameters: type: object @@ -437,8 +437,8 @@ jobs: steps: - backend: daily_read_step - - backend: base - name: daily:write + daily:write: + backend: base description: "Write daily//.md (body + frontmatter); refreshes the day index." parameters: type: object @@ -471,8 +471,8 @@ jobs: steps: - backend: daily_write_step - - backend: base - name: daily:list + daily:list: + backend: base description: "List notes under a single day." parameters: type: object @@ -484,8 +484,8 @@ jobs: steps: - backend: daily_list_step - - backend: base - name: daily:reindex + daily:reindex: + backend: base description: "Rebuild the day-index page daily/.md." parameters: type: object @@ -497,8 +497,8 @@ jobs: steps: - backend: daily_reindex_step - - backend: background - name: watch_file + watch_file: + backend: background watch_paths: - MEMORY.md - memory 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",